Skip to main content

drizzle_core/prepared/
mod.rs

1mod owned;
2pub use owned::OwnedPreparedStatement;
3
4use crate::prelude::*;
5use crate::{
6    error::DrizzleError,
7    param::{Param, ParamBind},
8    sql::{SQL, SQLChunk, SQLiteNamedParams},
9    traits::{SQLParam, ToSQL},
10};
11use compact_str::CompactString;
12use core::fmt;
13use smallvec::SmallVec;
14
15/// A statement rendered once, whose placeholders are bound each time it
16/// runs.
17///
18/// Create one with [`prepare_render`]; drivers wrap it in their own prepared
19/// statement types. The SQL is stored as text segments with one parameter
20/// between each pair: `[text, param, text, param, text]`. Parameters that
21/// already had a value when the statement was rendered keep it.
22///
23/// # Examples
24///
25/// ```
26/// use drizzle_core::prepared::prepare_render;
27/// use drizzle_core::{ParamBind, Placeholder, SQL, ToSQL};
28/// # use drizzle_core::{Dialect, SQLParam, SQLiteDialect};
29/// # use std::borrow::Cow;
30/// # #[derive(Debug, Clone, PartialEq)]
31/// # struct Value(i64);
32/// # impl SQLParam for Value {
33/// #     const DIALECT: Dialect = Dialect::SQLite;
34/// #     type DialectMarker = SQLiteDialect;
35/// # }
36/// # impl From<Value> for Cow<'_, Value> {
37/// #     fn from(value: Value) -> Self { Cow::Owned(value) }
38/// # }
39///
40/// let sql: SQL<'_, Value> = SQL::raw("SELECT * FROM users WHERE id =")
41///     .append(Placeholder::named("id").to_sql());
42/// let prepared = prepare_render(&sql);
43/// assert_eq!(prepared.sql(), "SELECT * FROM users WHERE id = :id");
44/// assert_eq!(prepared.external_param_count(), 1);
45///
46/// let (text, values) = prepared.bind([ParamBind::new("id", Value(5))])?;
47/// assert_eq!(text, "SELECT * FROM users WHERE id = :id");
48/// assert_eq!(values.collect::<Vec<_>>(), [Value(5)]);
49/// # Ok::<(), drizzle_core::error::DrizzleError>(())
50/// ```
51#[derive(Debug, Clone)]
52pub struct PreparedStatement<'a, V: SQLParam> {
53    /// Rendered SQL text between the parameters; one more than `params`.
54    pub text_segments: Box<[CompactString]>,
55    /// The parameters, in order.
56    pub params: Box<[Param<'a, V>]>,
57    /// The full SQL text, with the dialect's placeholders.
58    pub sql: CompactString,
59}
60
61impl<V: SQLParam> From<OwnedPreparedStatement<V>> for PreparedStatement<'_, V> {
62    fn from(value: OwnedPreparedStatement<V>) -> Self {
63        Self {
64            text_segments: value.text_segments,
65            params: value.params.iter().map(|v| v.clone().into()).collect(),
66            sql: value.sql,
67        }
68    }
69}
70
71impl<V: SQLParam> core::fmt::Display for PreparedStatement<'_, V> {
72    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
73        write!(f, "{}", self.sql())
74    }
75}
76
77/// Matches `param_binds` to `params` and returns the values to send, in
78/// order. See [`PreparedStatement::bind`] for the errors.
79pub(crate) fn bind_values_internal<'a, V, T, P>(
80    params: &[P],
81    param_binds: impl IntoIterator<Item = ParamBind<'a, T>>,
82    param_name_fn: impl Fn(&P) -> Option<&str>,
83    param_value_fn: impl Fn(&P) -> Option<&V>,
84) -> crate::error::Result<SmallVec<[V; 8]>>
85where
86    V: SQLParam + Clone,
87    T: SQLParam + Into<V>,
88{
89    #[cfg(feature = "profiling")]
90    crate::drizzle_profile_scope!("prepared", "bind_values_internal");
91    let param_binds = param_binds.into_iter();
92    let (binds_lower, binds_upper) = param_binds.size_hint();
93
94    let mut expected_named = HashMap::<&str, usize>::new();
95    let mut expected_positional = 0usize;
96    for param in params {
97        if param_value_fn(param).is_some() {
98            continue;
99        }
100
101        match param_name_fn(param) {
102            Some(name) if !name.is_empty() => {
103                *expected_named.entry(name).or_insert(0) += 1;
104            }
105            _ => expected_positional += 1,
106        }
107    }
108
109    let mut param_map = HashMap::<&str, V>::with_capacity(expected_named.len().max(binds_lower));
110
111    let mut positional_params: SmallVec<[V; 8]> =
112        SmallVec::with_capacity(binds_upper.unwrap_or(binds_lower));
113
114    for bind in param_binds {
115        if bind.name.is_empty() {
116            positional_params.push(bind.value.into());
117        } else if param_map.insert(bind.name, bind.value.into()).is_some() {
118            return Err(DrizzleError::ParameterError(
119                format!("Duplicate parameter binding: '{}'", bind.name).into(),
120            ));
121        }
122    }
123
124    if positional_params.len() < expected_positional {
125        return Err(DrizzleError::ParameterError(
126            format!(
127                "Missing positional parameter(s): expected {}, got {}",
128                expected_positional,
129                positional_params.len()
130            )
131            .into(),
132        ));
133    }
134    if positional_params.len() > expected_positional {
135        return Err(DrizzleError::ParameterError(
136            format!(
137                "Unexpected positional parameter(s): expected {}, got {}",
138                expected_positional,
139                positional_params.len()
140            )
141            .into(),
142        ));
143    }
144
145    let mut missing_named: SmallVec<[&str; 8]> = expected_named
146        .keys()
147        .filter(|name| !param_map.contains_key(**name))
148        .copied()
149        .collect();
150    if !missing_named.is_empty() {
151        missing_named.sort_unstable();
152        return Err(DrizzleError::ParameterError(
153            format!("Missing named parameter(s): {}", missing_named.join(", ")).into(),
154        ));
155    }
156
157    let mut extra_named: SmallVec<[&str; 8]> = param_map
158        .keys()
159        .filter(|name| !expected_named.contains_key(**name))
160        .copied()
161        .collect();
162    if !extra_named.is_empty() {
163        extra_named.sort_unstable();
164        return Err(DrizzleError::ParameterError(
165            format!("Unexpected named parameter(s): {}", extra_named.join(", ")).into(),
166        ));
167    }
168
169    let mut positional_iter = positional_params.into_iter();
170
171    let mut bound_params = SmallVec::<[V; 8]>::with_capacity(params.len());
172    let mut sqlite_names = SQLiteNamedParams::default();
173
174    for param in params {
175        // SQLite renders a named parameter as `:name` and gives each distinct
176        // name one slot, so only its first occurrence binds a value.
177        if V::DIALECT == crate::dialect::Dialect::SQLite
178            && let Some(name) = param_name_fn(param)
179            && sqlite_names.is_repeat(name)
180        {
181            continue;
182        }
183
184        // For parameters, prioritize internal values first, then external bindings
185        if let Some(value) = param_value_fn(param) {
186            // Use internal parameter value (from prepared statement)
187            bound_params.push(value.clone());
188        } else if let Some(name) = param_name_fn(param) {
189            // If no internal value, try external binding for named parameters
190            if !name.is_empty() {
191                if let Some(value) = param_map.get(name) {
192                    bound_params.push(value.clone());
193                }
194            } else if let Some(value) = positional_iter.next() {
195                bound_params.push(value);
196            }
197        } else if let Some(value) = positional_iter.next() {
198            bound_params.push(value);
199        }
200    }
201
202    Ok(bound_params)
203}
204
205impl<'a, V: SQLParam> PreparedStatement<'a, V> {
206    /// Returns how many bindings [`bind`](Self::bind) expects.
207    ///
208    /// Counts parameters without a value, with each placeholder name counted
209    /// once, since one binding fills every use of a name.
210    #[must_use]
211    pub fn external_param_count(&self) -> usize {
212        let mut named = HashSet::<&str>::new();
213        let mut positional = 0usize;
214        for param in &self.params {
215            if param.value.is_some() {
216                continue;
217            }
218            match param.placeholder.name {
219                Some(name) if !name.is_empty() => {
220                    named.insert(name);
221                }
222                _ => positional += 1,
223            }
224        }
225        named.len() + positional
226    }
227
228    /// Binds values to the placeholders and returns the SQL text with the
229    /// values to send, in order.
230    ///
231    /// Named bindings match placeholders by name; unnamed ones
232    /// ([`ParamBind::positional`]) fill unnamed placeholders in order. For
233    /// SQLite, a name used more than once is sent once.
234    ///
235    /// # Errors
236    ///
237    /// Returns [`DrizzleError::ParameterError`] when a name is bound twice,
238    /// a placeholder has no binding, or a binding matches no placeholder.
239    pub fn bind<T: SQLParam + Into<V>>(
240        &self,
241        param_binds: impl IntoIterator<Item = ParamBind<'a, T>>,
242    ) -> crate::error::Result<(&str, impl Iterator<Item = V>)> {
243        let bound_params = bind_values_internal(
244            &self.params,
245            param_binds,
246            |p| p.placeholder.name,
247            |p| p.value.as_ref().map(core::convert::AsRef::as_ref),
248        )?;
249
250        Ok((self.sql.as_str(), bound_params.into_iter()))
251    }
252
253    /// Returns the SQL text, with the dialect's placeholders.
254    #[must_use]
255    pub fn sql(&self) -> &str {
256        self.sql.as_str()
257    }
258}
259
260impl<'a, V: SQLParam> ToSQL<'a, V> for PreparedStatement<'a, V> {
261    fn to_sql(&self) -> SQL<'a, V> {
262        // Calculate exact capacity needed: text_segments.len() + params.len()
263        let capacity = self.text_segments.len() + self.params.len();
264        let mut chunks = SmallVec::with_capacity(capacity);
265
266        // Interleave text segments and params: text[0], param[0], text[1], param[1], ..., text[n]
267        // Use iterators to avoid bounds checking and minimize allocations
268        let mut param_iter = self.params.iter();
269
270        for text_segment in &self.text_segments {
271            chunks.push(SQLChunk::Raw(Cow::Owned(text_segment.to_string())));
272
273            // Add corresponding param if available
274            if let Some(param) = param_iter.next() {
275                chunks.push(SQLChunk::Param(param.clone()));
276            }
277        }
278
279        SQL { chunks }
280    }
281}
282/// Renders `sql` into a [`PreparedStatement`], splitting the text around its
283/// parameters.
284pub fn prepare_render<'a, V: SQLParam>(sql: &SQL<'a, V>) -> PreparedStatement<'a, V> {
285    use crate::dialect::{Dialect, write_placeholder};
286    use crate::sql::chunk_needs_space;
287
288    #[cfg(feature = "profiling")]
289    crate::drizzle_profile_scope!("prepared", "prepare_render");
290
291    if !sql
292        .chunks
293        .iter()
294        .any(|chunk| matches!(chunk, SQLChunk::Param(_)))
295    {
296        #[cfg(feature = "profiling")]
297        crate::drizzle_profile_scope!("prepared", "prepare_render.no_params");
298        let rendered_sql = CompactString::new(sql.sql());
299        return PreparedStatement {
300            text_segments: vec![rendered_sql.clone()].into_boxed_slice(),
301            params: Vec::new().into_boxed_slice(),
302            sql: rendered_sql,
303        };
304    }
305
306    #[cfg(feature = "profiling")]
307    crate::drizzle_profile_scope!("prepared", "prepare_render.scan");
308    let mut text_segments = Vec::new();
309    let mut params = Vec::new();
310    let mut current_text = String::new();
311    let mut rendered_sql = String::with_capacity(sql.chunks.len().saturating_mul(8).max(64));
312    let mut param_index = 1usize;
313
314    for (i, chunk) in sql.chunks.iter().enumerate() {
315        let current_text_ends_with_space = if let SQLChunk::Param(param) = chunk {
316            text_segments.push(CompactString::new(&current_text));
317            rendered_sql.push_str(&current_text);
318            current_text.clear();
319            params.push(param.clone());
320
321            if let Some(name) = param.placeholder.name
322                && V::DIALECT == Dialect::SQLite
323            {
324                rendered_sql.push(':');
325                rendered_sql.push_str(name);
326            } else {
327                write_placeholder(V::DIALECT, param_index, &mut rendered_sql);
328            }
329            param_index += 1;
330            false
331        } else {
332            sql.write_chunk_to(&mut current_text, chunk, i);
333            matches!(chunk, SQLChunk::Raw(text) if text.ends_with(' '))
334        };
335
336        // Use the canonical spacing logic, with an extra check for trailing spaces
337        // already in the accumulated text buffer
338        if let Some(next) = sql.chunks.get(i + 1)
339            && !current_text_ends_with_space
340            && chunk_needs_space(chunk, next)
341        {
342            current_text.push(' ');
343        }
344    }
345
346    text_segments.push(CompactString::new(&current_text));
347    rendered_sql.push_str(&current_text);
348
349    #[cfg(feature = "profiling")]
350    crate::drizzle_profile_scope!("prepared", "prepare_render.finalize");
351    let text_segments = text_segments.into_boxed_slice();
352    let params = params.into_boxed_slice();
353    let rendered_sql = CompactString::new(rendered_sql);
354
355    PreparedStatement {
356        text_segments,
357        params,
358        sql: rendered_sql,
359    }
360}