Skip to main content

rusqlite_facet/
lib.rs

1use facet::Facet;
2use facet_core::{Def, Shape, StructKind, Type, UserType};
3use facet_reflect::{AllocError, HasFields, Partial, Peek, ReflectError, ShapeMismatchError};
4use rusqlite::types::{Type as SqlType, Value as SqlValue, ValueRef};
5use rusqlite::{Connection, Row, Rows, Statement};
6
7#[derive(Debug)]
8pub enum Error {
9    Sql(rusqlite::Error),
10    Reflect(ReflectError),
11    Alloc(AllocError),
12    ShapeMismatch(ShapeMismatchError),
13    NotAStruct {
14        shape: &'static Shape,
15    },
16    UnsupportedParamType {
17        field: String,
18        shape: &'static Shape,
19    },
20    UnsupportedRowType {
21        field: String,
22        shape: &'static Shape,
23    },
24    MissingNamedParam {
25        parameter: String,
26    },
27    MissingColumn {
28        column: String,
29    },
30    UnnamedParameter {
31        index: usize,
32    },
33    UnusedParamFields {
34        fields: Vec<String>,
35    },
36    UnusedPositionalParams {
37        provided: usize,
38        used: usize,
39    },
40    TooManyRows {
41        expected: usize,
42        actual_at_least: usize,
43    },
44    WithSqlContext {
45        sql: String,
46        source: Box<Error>,
47    },
48    OutOfRange {
49        field: String,
50        source: i128,
51        target: &'static str,
52    },
53    TypeMismatch {
54        field: String,
55        expected: &'static Shape,
56        actual: SqlType,
57    },
58}
59
60impl core::fmt::Display for Error {
61    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
62        match self {
63            Error::Sql(e) => write!(f, "sqlite error: {e}"),
64            Error::Reflect(e) => write!(f, "reflection error: {e}"),
65            Error::Alloc(e) => write!(f, "allocation error: {e}"),
66            Error::ShapeMismatch(e) => write!(f, "shape mismatch: {e}"),
67            Error::NotAStruct { shape } => write!(f, "expected a struct shape, got {shape}"),
68            Error::UnsupportedParamType { field, shape } => {
69                write!(f, "unsupported parameter type for field '{field}': {shape}")
70            }
71            Error::UnsupportedRowType { field, shape } => {
72                write!(f, "unsupported row type for field '{field}': {shape}")
73            }
74            Error::MissingNamedParam { parameter } => {
75                write!(f, "missing named parameter for SQL binding: {parameter}")
76            }
77            Error::MissingColumn { column } => write!(f, "missing required column: {column}"),
78            Error::UnnamedParameter { index } => {
79                write!(f, "statement parameter #{index} is unnamed")
80            }
81            Error::UnusedParamFields { fields } => {
82                write!(f, "unused parameter fields: {}", fields.join(", "))
83            }
84            Error::UnusedPositionalParams { provided, used } => {
85                write!(
86                    f,
87                    "unused positional parameters: provided {provided}, used {used}"
88                )
89            }
90            Error::TooManyRows {
91                expected,
92                actual_at_least,
93            } => {
94                write!(
95                    f,
96                    "query returned too many rows: expected {expected}, got at least {actual_at_least}"
97                )
98            }
99            Error::WithSqlContext { sql, source } => {
100                write!(f, "{source} (sql: {sql})")
101            }
102            Error::OutOfRange {
103                field,
104                source,
105                target,
106            } => {
107                write!(
108                    f,
109                    "out-of-range conversion for field '{field}': {source} cannot fit in {target}"
110                )
111            }
112            Error::TypeMismatch {
113                field,
114                expected,
115                actual,
116            } => {
117                write!(
118                    f,
119                    "type mismatch for field '{field}': expected {expected}, got {actual:?}"
120                )
121            }
122        }
123    }
124}
125
126impl std::error::Error for Error {
127    fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
128        match self {
129            Error::Sql(err) => Some(err),
130            Error::Reflect(err) => Some(err),
131            Error::Alloc(err) => Some(err),
132            Error::ShapeMismatch(err) => Some(err),
133            Error::WithSqlContext { source, .. } => Some(source),
134            _ => None,
135        }
136    }
137}
138
139impl From<rusqlite::Error> for Error {
140    fn from(value: rusqlite::Error) -> Self {
141        Self::Sql(value)
142    }
143}
144
145impl From<ReflectError> for Error {
146    fn from(value: ReflectError) -> Self {
147        Self::Reflect(value)
148    }
149}
150
151impl From<AllocError> for Error {
152    fn from(value: AllocError) -> Self {
153        Self::Alloc(value)
154    }
155}
156
157impl From<ShapeMismatchError> for Error {
158    fn from(value: ShapeMismatchError) -> Self {
159        Self::ShapeMismatch(value)
160    }
161}
162
163pub type Result<T> = core::result::Result<T, Error>;
164
165pub struct FacetRows<'stmt, T> {
166    rows: Rows<'stmt>,
167    _marker: core::marker::PhantomData<T>,
168}
169
170impl<T: Facet<'static>> Iterator for FacetRows<'_, T> {
171    type Item = Result<T>;
172
173    fn next(&mut self) -> Option<Self::Item> {
174        match self.rows.next() {
175            Ok(Some(row)) => Some(from_row::<T>(row)),
176            Ok(None) => None,
177            Err(err) => Some(Err(Error::Sql(err))),
178        }
179    }
180}
181
182pub trait StatementFacetExt {
183    fn facet_execute_ref<'p, P: Facet<'p> + ?Sized>(&mut self, params: &'p P) -> Result<usize>;
184    fn facet_query_iter_ref<'stmt, 'p, T: Facet<'static>, P: Facet<'p> + ?Sized>(
185        &'stmt mut self,
186        params: &'p P,
187    ) -> Result<FacetRows<'stmt, T>>;
188    fn facet_query_ref<'p, T: Facet<'static>, P: Facet<'p> + ?Sized>(
189        &mut self,
190        params: &'p P,
191    ) -> Result<Vec<T>>;
192    fn facet_query_optional_ref<'p, T: Facet<'static>, P: Facet<'p> + ?Sized>(
193        &mut self,
194        params: &'p P,
195    ) -> Result<Option<T>>;
196    fn facet_query_one_ref<'p, T: Facet<'static>, P: Facet<'p> + ?Sized>(
197        &mut self,
198        params: &'p P,
199    ) -> Result<T>;
200    fn facet_query_row_ref<'p, T: Facet<'static>, P: Facet<'p> + ?Sized>(
201        &mut self,
202        params: &'p P,
203    ) -> Result<T>;
204    fn facet_execute<P: Facet<'static>>(&mut self, params: P) -> Result<usize>;
205    fn facet_query_iter<'stmt, T: Facet<'static>, P: Facet<'static>>(
206        &'stmt mut self,
207        params: P,
208    ) -> Result<FacetRows<'stmt, T>>;
209    fn facet_query<T: Facet<'static>, P: Facet<'static>>(&mut self, params: P) -> Result<Vec<T>>;
210    fn facet_query_optional<T: Facet<'static>, P: Facet<'static>>(
211        &mut self,
212        params: P,
213    ) -> Result<Option<T>>;
214    fn facet_query_one<T: Facet<'static>, P: Facet<'static>>(&mut self, params: P) -> Result<T>;
215    fn facet_query_row<T: Facet<'static>, P: Facet<'static>>(&mut self, params: P) -> Result<T>;
216}
217
218pub trait ConnectionFacetExt {
219    fn facet_prepare_cached(&self, sql: &str) -> rusqlite::Result<rusqlite::CachedStatement<'_>>;
220    fn facet_execute_ref<'p, P: Facet<'p> + ?Sized>(
221        &self,
222        sql: &str,
223        params: &'p P,
224    ) -> Result<usize>;
225    fn facet_query_ref<'p, T: Facet<'static>, P: Facet<'p> + ?Sized>(
226        &self,
227        sql: &str,
228        params: &'p P,
229    ) -> Result<Vec<T>>;
230    fn facet_query_optional_ref<'p, T: Facet<'static>, P: Facet<'p> + ?Sized>(
231        &self,
232        sql: &str,
233        params: &'p P,
234    ) -> Result<Option<T>>;
235    fn facet_query_one_ref<'p, T: Facet<'static>, P: Facet<'p> + ?Sized>(
236        &self,
237        sql: &str,
238        params: &'p P,
239    ) -> Result<T>;
240    fn facet_execute<P: Facet<'static>>(&self, sql: &str, params: P) -> Result<usize>;
241    fn facet_query<T: Facet<'static>, P: Facet<'static>>(
242        &self,
243        sql: &str,
244        params: P,
245    ) -> Result<Vec<T>>;
246    fn facet_query_optional<T: Facet<'static>, P: Facet<'static>>(
247        &self,
248        sql: &str,
249        params: P,
250    ) -> Result<Option<T>>;
251    fn facet_query_one<T: Facet<'static>, P: Facet<'static>>(
252        &self,
253        sql: &str,
254        params: P,
255    ) -> Result<T>;
256}
257
258impl StatementFacetExt for Statement<'_> {
259    fn facet_execute_ref<'p, P: Facet<'p> + ?Sized>(&mut self, params: &'p P) -> Result<usize> {
260        bind_facet_params_ref(self, params)?;
261        Ok(self.raw_execute()?)
262    }
263
264    fn facet_query_iter_ref<'stmt, 'p, T: Facet<'static>, P: Facet<'p> + ?Sized>(
265        &'stmt mut self,
266        params: &'p P,
267    ) -> Result<FacetRows<'stmt, T>> {
268        bind_facet_params_ref(self, params)?;
269        Ok(FacetRows {
270            rows: self.raw_query(),
271            _marker: core::marker::PhantomData,
272        })
273    }
274
275    fn facet_query_ref<'p, T: Facet<'static>, P: Facet<'p> + ?Sized>(
276        &mut self,
277        params: &'p P,
278    ) -> Result<Vec<T>> {
279        let mut out = Vec::new();
280        for row in self.facet_query_iter_ref::<T, P>(params)? {
281            out.push(row?);
282        }
283        Ok(out)
284    }
285
286    fn facet_query_optional_ref<'p, T: Facet<'static>, P: Facet<'p> + ?Sized>(
287        &mut self,
288        params: &'p P,
289    ) -> Result<Option<T>> {
290        bind_facet_params_ref(self, params)?;
291        let mut rows = self.raw_query();
292        let Some(first_row) = rows.next()? else {
293            return Ok(None);
294        };
295        let first = from_row::<T>(first_row)?;
296        if rows.next()?.is_some() {
297            return Err(Error::TooManyRows {
298                expected: 1,
299                actual_at_least: 2,
300            });
301        }
302        Ok(Some(first))
303    }
304
305    fn facet_query_one_ref<'p, T: Facet<'static>, P: Facet<'p> + ?Sized>(
306        &mut self,
307        params: &'p P,
308    ) -> Result<T> {
309        match self.facet_query_optional_ref::<T, P>(params)? {
310            Some(row) => Ok(row),
311            None => Err(Error::Sql(rusqlite::Error::QueryReturnedNoRows)),
312        }
313    }
314
315    fn facet_query_row_ref<'p, T: Facet<'static>, P: Facet<'p> + ?Sized>(
316        &mut self,
317        params: &'p P,
318    ) -> Result<T> {
319        self.facet_query_one_ref::<T, P>(params)
320    }
321
322    fn facet_execute<P: Facet<'static>>(&mut self, params: P) -> Result<usize> {
323        bind_facet_params_static(self, &params)?;
324        Ok(self.raw_execute()?)
325    }
326
327    fn facet_query_iter<'stmt, T: Facet<'static>, P: Facet<'static>>(
328        &'stmt mut self,
329        params: P,
330    ) -> Result<FacetRows<'stmt, T>> {
331        bind_facet_params_static(self, &params)?;
332        Ok(FacetRows {
333            rows: self.raw_query(),
334            _marker: core::marker::PhantomData,
335        })
336    }
337
338    fn facet_query<T: Facet<'static>, P: Facet<'static>>(&mut self, params: P) -> Result<Vec<T>> {
339        let mut out = Vec::new();
340        for row in self.facet_query_iter::<T, P>(params)? {
341            out.push(row?);
342        }
343        Ok(out)
344    }
345
346    fn facet_query_optional<T: Facet<'static>, P: Facet<'static>>(
347        &mut self,
348        params: P,
349    ) -> Result<Option<T>> {
350        bind_facet_params_static(self, &params)?;
351        let mut rows = self.raw_query();
352        let Some(first_row) = rows.next()? else {
353            return Ok(None);
354        };
355        let first = from_row::<T>(first_row)?;
356        if rows.next()?.is_some() {
357            return Err(Error::TooManyRows {
358                expected: 1,
359                actual_at_least: 2,
360            });
361        }
362        Ok(Some(first))
363    }
364
365    fn facet_query_one<T: Facet<'static>, P: Facet<'static>>(&mut self, params: P) -> Result<T> {
366        match self.facet_query_optional::<T, P>(params)? {
367            Some(row) => Ok(row),
368            None => Err(Error::Sql(rusqlite::Error::QueryReturnedNoRows)),
369        }
370    }
371
372    fn facet_query_row<T: Facet<'static>, P: Facet<'static>>(&mut self, params: P) -> Result<T> {
373        self.facet_query_one::<T, P>(params)
374    }
375}
376
377impl ConnectionFacetExt for Connection {
378    fn facet_prepare_cached(&self, sql: &str) -> rusqlite::Result<rusqlite::CachedStatement<'_>> {
379        self.prepare_cached(sql)
380    }
381
382    fn facet_execute_ref<'p, P: Facet<'p> + ?Sized>(
383        &self,
384        sql: &str,
385        params: &'p P,
386    ) -> Result<usize> {
387        let mut stmt = self.prepare(sql)?;
388        with_sql_context(sql, stmt.facet_execute_ref(params))
389    }
390
391    fn facet_query_ref<'p, T: Facet<'static>, P: Facet<'p> + ?Sized>(
392        &self,
393        sql: &str,
394        params: &'p P,
395    ) -> Result<Vec<T>> {
396        let mut stmt = self.prepare(sql)?;
397        with_sql_context(sql, stmt.facet_query_ref::<T, P>(params))
398    }
399
400    fn facet_query_optional_ref<'p, T: Facet<'static>, P: Facet<'p> + ?Sized>(
401        &self,
402        sql: &str,
403        params: &'p P,
404    ) -> Result<Option<T>> {
405        let mut stmt = self.prepare(sql)?;
406        with_sql_context(sql, stmt.facet_query_optional_ref::<T, P>(params))
407    }
408
409    fn facet_query_one_ref<'p, T: Facet<'static>, P: Facet<'p> + ?Sized>(
410        &self,
411        sql: &str,
412        params: &'p P,
413    ) -> Result<T> {
414        let mut stmt = self.prepare(sql)?;
415        with_sql_context(sql, stmt.facet_query_one_ref::<T, P>(params))
416    }
417
418    fn facet_execute<P: Facet<'static>>(&self, sql: &str, params: P) -> Result<usize> {
419        let mut stmt = self.prepare(sql)?;
420        with_sql_context(sql, stmt.facet_execute(params))
421    }
422
423    fn facet_query<T: Facet<'static>, P: Facet<'static>>(
424        &self,
425        sql: &str,
426        params: P,
427    ) -> Result<Vec<T>> {
428        let mut stmt = self.prepare(sql)?;
429        with_sql_context(sql, stmt.facet_query::<T, P>(params))
430    }
431
432    fn facet_query_optional<T: Facet<'static>, P: Facet<'static>>(
433        &self,
434        sql: &str,
435        params: P,
436    ) -> Result<Option<T>> {
437        let mut stmt = self.prepare(sql)?;
438        with_sql_context(sql, stmt.facet_query_optional::<T, P>(params))
439    }
440
441    fn facet_query_one<T: Facet<'static>, P: Facet<'static>>(
442        &self,
443        sql: &str,
444        params: P,
445    ) -> Result<T> {
446        let mut stmt = self.prepare(sql)?;
447        with_sql_context(sql, stmt.facet_query_one::<T, P>(params))
448    }
449}
450
451fn with_sql_context<T>(sql: &str, result: Result<T>) -> Result<T> {
452    result.map_err(|source| Error::WithSqlContext {
453        sql: sql.to_string(),
454        source: Box::new(source),
455    })
456}
457
458pub fn from_row<T: Facet<'static>>(row: &Row<'_>) -> Result<T> {
459    let partial = Partial::alloc_owned::<T>()?;
460    let partial = deserialize_row_into(row, partial, T::SHAPE)?;
461    let heap_value = partial.build()?;
462    Ok(heap_value.materialize()?)
463}
464
465fn bind_facet_params_static<P: Facet<'static>>(stmt: &mut Statement<'_>, params: &P) -> Result<()> {
466    bind_facet_params_impl(stmt, Peek::new(params), P::SHAPE)
467}
468
469fn bind_facet_params_ref<'p, P: Facet<'p> + ?Sized>(
470    stmt: &mut Statement<'_>,
471    params: &'p P,
472) -> Result<()> {
473    bind_facet_params_impl(stmt, Peek::new(params), P::SHAPE)
474}
475
476fn bind_facet_params_impl(
477    stmt: &mut Statement<'_>,
478    peek: Peek<'_, '_>,
479    shape: &'static Shape,
480) -> Result<()> {
481    stmt.clear_bindings();
482
483    if matches!(
484        peek.shape().def,
485        Def::List(_) | Def::Array(_) | Def::Slice(_)
486    ) {
487        return bind_list_like_params(stmt, peek);
488    }
489
490    let struct_peek = peek
491        .into_struct()
492        .map_err(|_| Error::NotAStruct { shape })?;
493
494    let mut field_names: Vec<String> = Vec::new();
495    let mut field_values: Vec<SqlValue> = Vec::new();
496    for (field, value) in struct_peek.fields() {
497        let name = field.rename.unwrap_or(field.name).to_string();
498        field_names.push(name.clone());
499        field_values.push(peek_to_sql_value(value, &name)?);
500    }
501
502    let mut used = vec![false; field_names.len()];
503    let mut positional_cursor = 0usize;
504    for param_index in 1..=stmt.parameter_count() {
505        let field_index = if let Some(name) = stmt.parameter_name(param_index) {
506            if let Some(stripped) = name.strip_prefix(':') {
507                field_names
508                    .iter()
509                    .position(|f| f == stripped)
510                    .ok_or_else(|| Error::MissingNamedParam {
511                        parameter: name.to_string(),
512                    })?
513            } else if let Some(stripped) = name.strip_prefix('@') {
514                field_names
515                    .iter()
516                    .position(|f| f == stripped)
517                    .ok_or_else(|| Error::MissingNamedParam {
518                        parameter: name.to_string(),
519                    })?
520            } else if let Some(stripped) = name.strip_prefix('$') {
521                field_names
522                    .iter()
523                    .position(|f| f == stripped)
524                    .ok_or_else(|| Error::MissingNamedParam {
525                        parameter: name.to_string(),
526                    })?
527            } else if let Some(stripped) = name.strip_prefix('?') {
528                if stripped.is_empty() {
529                    let idx = positional_cursor;
530                    positional_cursor += 1;
531                    idx
532                } else {
533                    let raw = stripped
534                        .parse::<usize>()
535                        .map_err(|_| Error::MissingNamedParam {
536                            parameter: name.to_string(),
537                        })?;
538                    raw.saturating_sub(1)
539                }
540            } else {
541                return Err(Error::UnnamedParameter { index: param_index });
542            }
543        } else {
544            if positional_cursor >= field_values.len() {
545                return Err(Error::UnnamedParameter { index: param_index });
546            }
547            let idx = positional_cursor;
548            positional_cursor += 1;
549            idx
550        };
551
552        let value = field_values
553            .get(field_index)
554            .ok_or(Error::UnnamedParameter { index: param_index })?;
555
556        stmt.raw_bind_parameter(param_index, value)?;
557        used[field_index] = true;
558    }
559
560    let unused: Vec<String> = field_names
561        .iter()
562        .enumerate()
563        .filter_map(|(idx, name)| (!used[idx]).then_some(name.clone()))
564        .collect();
565    if !unused.is_empty() {
566        return Err(Error::UnusedParamFields { fields: unused });
567    }
568
569    Ok(())
570}
571
572fn bind_list_like_params(stmt: &mut Statement<'_>, peek: Peek<'_, '_>) -> Result<()> {
573    let list_like = peek.into_list_like().map_err(Error::Reflect)?;
574    let mut values = Vec::with_capacity(list_like.len());
575    for value in list_like.iter() {
576        values.push(peek_to_sql_value(value, "positional_param")?);
577    }
578
579    let mut used = vec![false; values.len()];
580    let mut positional_cursor = 0usize;
581    for param_index in 1..=stmt.parameter_count() {
582        let value_index = if let Some(name) = stmt.parameter_name(param_index) {
583            if let Some(stripped) = name.strip_prefix('?') {
584                if stripped.is_empty() {
585                    let idx = positional_cursor;
586                    positional_cursor += 1;
587                    idx
588                } else {
589                    let raw = stripped
590                        .parse::<usize>()
591                        .map_err(|_| Error::MissingNamedParam {
592                            parameter: name.to_string(),
593                        })?;
594                    raw.saturating_sub(1)
595                }
596            } else {
597                return Err(Error::MissingNamedParam {
598                    parameter: name.to_string(),
599                });
600            }
601        } else {
602            let idx = positional_cursor;
603            positional_cursor += 1;
604            idx
605        };
606
607        let value = values
608            .get(value_index)
609            .ok_or(Error::UnnamedParameter { index: param_index })?;
610        stmt.raw_bind_parameter(param_index, value)?;
611        used[value_index] = true;
612    }
613
614    let used_count = used.iter().filter(|v| **v).count();
615    if used_count != values.len() {
616        return Err(Error::UnusedPositionalParams {
617            provided: values.len(),
618            used: used_count,
619        });
620    }
621
622    Ok(())
623}
624
625fn peek_to_sql_value(peek: Peek<'_, '_>, field_name: &str) -> Result<SqlValue> {
626    if let Ok(option) = peek.into_option() {
627        let Some(inner) = option.value() else {
628            return Ok(SqlValue::Null);
629        };
630        return peek_to_sql_value(inner, field_name);
631    }
632
633    let peek = peek.innermost_peek();
634    if peek.shape() == bool::SHAPE {
635        return Ok(SqlValue::Integer(i64::from(*peek.get::<bool>()?)));
636    }
637    if peek.shape() == i8::SHAPE {
638        return Ok(SqlValue::Integer(i64::from(*peek.get::<i8>()?)));
639    }
640    if peek.shape() == i16::SHAPE {
641        return Ok(SqlValue::Integer(i64::from(*peek.get::<i16>()?)));
642    }
643    if peek.shape() == i32::SHAPE {
644        return Ok(SqlValue::Integer(i64::from(*peek.get::<i32>()?)));
645    }
646    if peek.shape() == i64::SHAPE {
647        return Ok(SqlValue::Integer(*peek.get::<i64>()?));
648    }
649    if peek.shape() == u8::SHAPE {
650        return Ok(SqlValue::Integer(i64::from(*peek.get::<u8>()?)));
651    }
652    if peek.shape() == u16::SHAPE {
653        return Ok(SqlValue::Integer(i64::from(*peek.get::<u16>()?)));
654    }
655    if peek.shape() == u32::SHAPE {
656        return Ok(SqlValue::Integer(i64::from(*peek.get::<u32>()?)));
657    }
658    if peek.shape() == u64::SHAPE {
659        let value = *peek.get::<u64>()?;
660        let value = i64::try_from(value).map_err(|_| Error::OutOfRange {
661            field: field_name.to_string(),
662            source: value as i128,
663            target: "i64",
664        })?;
665        return Ok(SqlValue::Integer(value));
666    }
667    if peek.shape() == f32::SHAPE {
668        return Ok(SqlValue::Real(f64::from(*peek.get::<f32>()?)));
669    }
670    if peek.shape() == f64::SHAPE {
671        return Ok(SqlValue::Real(*peek.get::<f64>()?));
672    }
673    if peek.shape() == String::SHAPE {
674        return Ok(SqlValue::Text(peek.get::<String>()?.clone()));
675    }
676    if peek.shape() == <Vec<u8>>::SHAPE {
677        return Ok(SqlValue::Blob(peek.get::<Vec<u8>>()?.clone()));
678    }
679    if let Some(text) = peek.as_str() {
680        return Ok(SqlValue::Text(text.to_string()));
681    }
682
683    Err(Error::UnsupportedParamType {
684        field: field_name.to_string(),
685        shape: peek.shape(),
686    })
687}
688
689fn deserialize_row_into(
690    row: &Row<'_>,
691    mut partial: Partial<'static, false>,
692    shape: &'static Shape,
693) -> Result<Partial<'static, false>> {
694    let struct_def = match &shape.ty {
695        Type::User(UserType::Struct(s)) if s.kind == StructKind::Struct => s,
696        _ => return Err(Error::NotAStruct { shape }),
697    };
698
699    for field in struct_def.fields {
700        let column_name = field.rename.unwrap_or(field.name);
701        let column_idx =
702            find_column_index(row, column_name).ok_or_else(|| Error::MissingColumn {
703                column: column_name.to_string(),
704            })?;
705
706        partial = partial.begin_field(field.name)?;
707        partial = deserialize_column(row, column_idx, column_name, partial, field.shape())?;
708        partial = partial.end()?;
709    }
710
711    Ok(partial)
712}
713
714fn find_column_index(row: &Row<'_>, column_name: &str) -> Option<usize> {
715    let stmt = row.as_ref();
716    (0..stmt.column_count()).find(|idx| {
717        stmt.column_name(*idx)
718            .map(|name| name == column_name)
719            .unwrap_or(false)
720    })
721}
722
723fn deserialize_column(
724    row: &Row<'_>,
725    column_idx: usize,
726    field_name: &str,
727    mut partial: Partial<'static, false>,
728    shape: &'static Shape,
729) -> Result<Partial<'static, false>> {
730    if shape.decl_id == Option::<()>::SHAPE.decl_id {
731        let value_ref = row.get_ref(column_idx)?;
732        if matches!(value_ref, ValueRef::Null) {
733            partial = partial.set_default()?;
734            return Ok(partial);
735        }
736
737        let inner = shape.inner.expect("Option shape must have inner");
738        partial = partial.begin_some()?;
739        partial = deserialize_column(row, column_idx, field_name, partial, inner)?;
740        partial = partial.end()?;
741        return Ok(partial);
742    }
743
744    if let Some(inner) = shape.inner {
745        partial = partial.begin_inner()?;
746        partial = deserialize_column(row, column_idx, field_name, partial, inner)?;
747        partial = partial.end()?;
748        return Ok(partial);
749    }
750
751    let value_ref = row.get_ref(column_idx)?;
752    if matches!(value_ref, ValueRef::Null) {
753        return Err(Error::TypeMismatch {
754            field: field_name.to_string(),
755            expected: shape,
756            actual: SqlType::Null,
757        });
758    }
759
760    if shape == bool::SHAPE {
761        partial = partial.set(row.get::<_, bool>(column_idx)?)?;
762    } else if shape == i8::SHAPE {
763        partial = partial.set(row.get::<_, i8>(column_idx)?)?;
764    } else if shape == i16::SHAPE {
765        partial = partial.set(row.get::<_, i16>(column_idx)?)?;
766    } else if shape == i32::SHAPE {
767        partial = partial.set(row.get::<_, i32>(column_idx)?)?;
768    } else if shape == i64::SHAPE {
769        partial = partial.set(row.get::<_, i64>(column_idx)?)?;
770    } else if shape == u8::SHAPE {
771        partial = partial.set(checked_unsigned::<u8>(
772            row.get::<_, i64>(column_idx)?,
773            field_name,
774        )?)?;
775    } else if shape == u16::SHAPE {
776        partial = partial.set(checked_unsigned::<u16>(
777            row.get::<_, i64>(column_idx)?,
778            field_name,
779        )?)?;
780    } else if shape == u32::SHAPE {
781        partial = partial.set(checked_unsigned::<u32>(
782            row.get::<_, i64>(column_idx)?,
783            field_name,
784        )?)?;
785    } else if shape == u64::SHAPE {
786        partial = partial.set(checked_unsigned::<u64>(
787            row.get::<_, i64>(column_idx)?,
788            field_name,
789        )?)?;
790    } else if shape == f32::SHAPE {
791        partial = partial.set(row.get::<_, f32>(column_idx)?)?;
792    } else if shape == f64::SHAPE {
793        partial = partial.set(row.get::<_, f64>(column_idx)?)?;
794    } else if shape == String::SHAPE {
795        partial = partial.set(row.get::<_, String>(column_idx)?)?;
796    } else if shape == <Vec<u8>>::SHAPE {
797        partial = partial.set(row.get::<_, Vec<u8>>(column_idx)?)?;
798    } else if shape.vtable.has_parse() {
799        let raw: String = row.get(column_idx)?;
800        partial = partial.parse_from_str(&raw)?;
801    } else {
802        return Err(Error::UnsupportedRowType {
803            field: field_name.to_string(),
804            shape,
805        });
806    }
807
808    Ok(partial)
809}
810
811fn checked_unsigned<T>(value: i64, field_name: &str) -> Result<T>
812where
813    T: TryFrom<i64>,
814{
815    T::try_from(value).map_err(|_| Error::OutOfRange {
816        field: field_name.to_string(),
817        source: value as i128,
818        target: core::any::type_name::<T>(),
819    })
820}
821
822#[cfg(test)]
823mod tests {
824    use super::{ConnectionFacetExt, Error, StatementFacetExt};
825    use facet::Facet;
826    use rusqlite::Connection;
827
828    #[derive(Debug, Facet, PartialEq)]
829    struct InsertConn {
830        conn_id: i64,
831        label: String,
832    }
833
834    #[derive(Debug, Facet, PartialEq)]
835    struct RowConn {
836        conn_id: u64,
837        label: String,
838    }
839
840    #[derive(Debug, Facet)]
841    struct QueryConn {
842        conn_id: i64,
843    }
844
845    #[derive(Debug, Facet, PartialEq)]
846    struct MaybeConn {
847        conn_id: i64,
848        label: Option<String>,
849    }
850
851    #[test]
852    fn facet_execute_and_query_named_params() {
853        let conn = Connection::open_in_memory().unwrap();
854        conn.execute(
855            "CREATE TABLE connections (conn_id INTEGER NOT NULL, label TEXT)",
856            (),
857        )
858        .unwrap();
859
860        let mut insert = conn
861            .prepare("INSERT INTO connections (conn_id, label) VALUES (:conn_id, :label)")
862            .unwrap();
863        insert
864            .facet_execute(InsertConn {
865                conn_id: 42,
866                label: "alpha".to_string(),
867            })
868            .unwrap();
869
870        let mut query = conn
871            .prepare("SELECT conn_id, label FROM connections WHERE conn_id = :conn_id")
872            .unwrap();
873        let rows = query
874            .facet_query::<RowConn, _>(QueryConn { conn_id: 42 })
875            .unwrap();
876        assert_eq!(
877            rows,
878            vec![RowConn {
879                conn_id: 42,
880                label: "alpha".to_string()
881            }]
882        );
883    }
884
885    #[test]
886    fn facet_query_positional_params_and_option() {
887        let conn = Connection::open_in_memory().unwrap();
888        conn.execute(
889            "CREATE TABLE items (conn_id INTEGER NOT NULL, label TEXT)",
890            (),
891        )
892        .unwrap();
893        conn.execute("INSERT INTO items (conn_id, label) VALUES (1, NULL)", ())
894            .unwrap();
895
896        #[derive(Facet)]
897        struct Positional {
898            conn_id: i64,
899        }
900
901        let mut stmt = conn
902            .prepare("SELECT conn_id, label FROM items WHERE conn_id = ?1")
903            .unwrap();
904        let rows = stmt
905            .facet_query::<MaybeConn, _>(Positional { conn_id: 1 })
906            .unwrap();
907        assert_eq!(
908            rows,
909            vec![MaybeConn {
910                conn_id: 1,
911                label: None
912            }]
913        );
914    }
915
916    #[test]
917    fn facet_query_accepts_array_params() {
918        let conn = Connection::open_in_memory().unwrap();
919        conn.execute(
920            "CREATE TABLE pairs (left_id INTEGER NOT NULL, right_id INTEGER NOT NULL)",
921            (),
922        )
923        .unwrap();
924        conn.execute("INSERT INTO pairs (left_id, right_id) VALUES (10, 20)", ())
925            .unwrap();
926
927        #[derive(Debug, Facet, PartialEq)]
928        struct PairRow {
929            left_id: i64,
930            right_id: i64,
931        }
932
933        let mut stmt = conn
934            .prepare("SELECT left_id, right_id FROM pairs WHERE left_id = ?1 AND right_id = ?2")
935            .unwrap();
936        let rows = stmt.facet_query::<PairRow, _>([10_i64, 20_i64]).unwrap();
937        assert_eq!(
938            rows,
939            vec![PairRow {
940                left_id: 10,
941                right_id: 20
942            }]
943        );
944    }
945
946    #[test]
947    fn facet_query_ref_accepts_slice_params() {
948        let conn = Connection::open_in_memory().unwrap();
949        conn.execute("CREATE TABLE ids (id INTEGER NOT NULL)", ())
950            .unwrap();
951        conn.execute("INSERT INTO ids (id) VALUES (7)", ()).unwrap();
952
953        #[derive(Debug, Facet, PartialEq)]
954        struct IdRow {
955            id: i64,
956        }
957
958        let values = [7_i64];
959        let mut stmt = conn.prepare("SELECT id FROM ids WHERE id = ?1").unwrap();
960        let rows = stmt.facet_query_ref::<IdRow, [i64]>(&values[..]).unwrap();
961        assert_eq!(rows, vec![IdRow { id: 7 }]);
962    }
963
964    #[test]
965    fn facet_query_iter_streams_rows() {
966        let conn = Connection::open_in_memory().unwrap();
967        conn.execute("CREATE TABLE nums (n INTEGER NOT NULL)", ())
968            .unwrap();
969        conn.execute("INSERT INTO nums (n) VALUES (1)", ()).unwrap();
970        conn.execute("INSERT INTO nums (n) VALUES (2)", ()).unwrap();
971
972        #[derive(Debug, Facet, PartialEq)]
973        struct NumRow {
974            n: i64,
975        }
976
977        let mut stmt = conn.prepare("SELECT n FROM nums ORDER BY n ASC").unwrap();
978        let mut iter = stmt.facet_query_iter::<NumRow, _>(()).unwrap();
979        assert_eq!(iter.next().unwrap().unwrap(), NumRow { n: 1 });
980        assert_eq!(iter.next().unwrap().unwrap(), NumRow { n: 2 });
981        assert!(iter.next().is_none());
982    }
983
984    #[test]
985    fn facet_query_iter_ref_streams_slice_params() {
986        let conn = Connection::open_in_memory().unwrap();
987        conn.execute("CREATE TABLE ids (id INTEGER NOT NULL)", ())
988            .unwrap();
989        conn.execute("INSERT INTO ids (id) VALUES (7)", ()).unwrap();
990
991        #[derive(Debug, Facet, PartialEq)]
992        struct IdRow {
993            id: i64,
994        }
995
996        let values = [7_i64];
997        let mut stmt = conn.prepare("SELECT id FROM ids WHERE id = ?1").unwrap();
998        let mut iter = stmt
999            .facet_query_iter_ref::<IdRow, [i64]>(&values[..])
1000            .unwrap();
1001        assert_eq!(iter.next().unwrap().unwrap(), IdRow { id: 7 });
1002        assert!(iter.next().is_none());
1003    }
1004
1005    #[test]
1006    fn facet_query_optional_returns_none_for_no_rows() {
1007        let conn = Connection::open_in_memory().unwrap();
1008        conn.execute("CREATE TABLE ids (id INTEGER NOT NULL)", ())
1009            .unwrap();
1010
1011        #[derive(Facet)]
1012        struct QueryId {
1013            id: i64,
1014        }
1015
1016        #[derive(Debug, Facet, PartialEq)]
1017        struct IdRow {
1018            id: i64,
1019        }
1020
1021        let mut stmt = conn.prepare("SELECT id FROM ids WHERE id = :id").unwrap();
1022        let row = stmt
1023            .facet_query_optional::<IdRow, _>(QueryId { id: 99 })
1024            .unwrap();
1025        assert_eq!(row, None);
1026    }
1027
1028    #[test]
1029    fn facet_query_optional_errors_on_multiple_rows() {
1030        let conn = Connection::open_in_memory().unwrap();
1031        conn.execute("CREATE TABLE ids (id INTEGER NOT NULL)", ())
1032            .unwrap();
1033        conn.execute("INSERT INTO ids (id) VALUES (1)", ()).unwrap();
1034        conn.execute("INSERT INTO ids (id) VALUES (1)", ()).unwrap();
1035
1036        #[derive(Facet)]
1037        struct QueryId {
1038            id: i64,
1039        }
1040
1041        #[derive(Debug, Facet, PartialEq)]
1042        struct IdRow {
1043            id: i64,
1044        }
1045
1046        let mut stmt = conn.prepare("SELECT id FROM ids WHERE id = :id").unwrap();
1047        let err = stmt
1048            .facet_query_optional::<IdRow, _>(QueryId { id: 1 })
1049            .unwrap_err();
1050        match err {
1051            Error::TooManyRows {
1052                expected,
1053                actual_at_least,
1054            } => {
1055                assert_eq!(expected, 1);
1056                assert_eq!(actual_at_least, 2);
1057            }
1058            _ => panic!("unexpected error: {err}"),
1059        }
1060    }
1061
1062    #[test]
1063    fn facet_query_one_errors_on_no_rows() {
1064        let conn = Connection::open_in_memory().unwrap();
1065        conn.execute("CREATE TABLE ids (id INTEGER NOT NULL)", ())
1066            .unwrap();
1067
1068        #[derive(Facet)]
1069        struct QueryId {
1070            id: i64,
1071        }
1072
1073        #[derive(Debug, Facet, PartialEq)]
1074        struct IdRow {
1075            id: i64,
1076        }
1077
1078        let mut stmt = conn.prepare("SELECT id FROM ids WHERE id = :id").unwrap();
1079        let err = stmt
1080            .facet_query_one::<IdRow, _>(QueryId { id: 1 })
1081            .unwrap_err();
1082        match err {
1083            Error::Sql(rusqlite::Error::QueryReturnedNoRows) => {}
1084            _ => panic!("unexpected error: {err}"),
1085        }
1086    }
1087
1088    #[test]
1089    fn connection_ext_execute_and_query() {
1090        let conn = Connection::open_in_memory().unwrap();
1091        conn.execute(
1092            "CREATE TABLE users (id INTEGER NOT NULL, name TEXT NOT NULL)",
1093            (),
1094        )
1095        .unwrap();
1096
1097        #[derive(Facet)]
1098        struct InsertUser {
1099            id: i64,
1100            name: String,
1101        }
1102
1103        #[derive(Facet)]
1104        struct QueryUser {
1105            id: i64,
1106        }
1107
1108        #[derive(Debug, Facet, PartialEq)]
1109        struct UserRow {
1110            id: i64,
1111            name: String,
1112        }
1113
1114        conn.facet_execute(
1115            "INSERT INTO users (id, name) VALUES (:id, :name)",
1116            InsertUser {
1117                id: 11,
1118                name: "alice".to_string(),
1119            },
1120        )
1121        .unwrap();
1122
1123        let row = conn
1124            .facet_query_one::<UserRow, _>(
1125                "SELECT id, name FROM users WHERE id = :id",
1126                QueryUser { id: 11 },
1127            )
1128            .unwrap();
1129        assert_eq!(
1130            row,
1131            UserRow {
1132                id: 11,
1133                name: "alice".to_string()
1134            }
1135        );
1136    }
1137
1138    #[test]
1139    fn connection_ext_query_ref_accepts_slice() {
1140        let conn = Connection::open_in_memory().unwrap();
1141        conn.execute("CREATE TABLE ids (id INTEGER NOT NULL)", ())
1142            .unwrap();
1143        conn.execute("INSERT INTO ids (id) VALUES (3)", ()).unwrap();
1144
1145        #[derive(Debug, Facet, PartialEq)]
1146        struct IdRow {
1147            id: i64,
1148        }
1149
1150        let values = [3_i64];
1151        let rows = conn
1152            .facet_query_ref::<IdRow, [i64]>("SELECT id FROM ids WHERE id = ?1", &values[..])
1153            .unwrap();
1154        assert_eq!(rows, vec![IdRow { id: 3 }]);
1155    }
1156
1157    #[test]
1158    fn connection_ext_errors_include_sql_context() {
1159        let conn = Connection::open_in_memory().unwrap();
1160        conn.execute("CREATE TABLE ids (id INTEGER NOT NULL)", ())
1161            .unwrap();
1162
1163        #[derive(Facet)]
1164        struct QueryId {
1165            id: i64,
1166        }
1167
1168        #[derive(Debug, Facet, PartialEq)]
1169        struct IdRow {
1170            id: i64,
1171        }
1172
1173        let sql = "SELECT id FROM ids WHERE id = :id";
1174        let err = conn
1175            .facet_query_one::<IdRow, _>(sql, QueryId { id: 1 })
1176            .unwrap_err();
1177        match err {
1178            Error::WithSqlContext {
1179                sql: actual,
1180                source,
1181            } => {
1182                assert_eq!(actual, sql.to_string());
1183                match *source {
1184                    Error::Sql(rusqlite::Error::QueryReturnedNoRows) => {}
1185                    _ => panic!("unexpected nested source"),
1186                }
1187            }
1188            _ => panic!("expected SQL context wrapper"),
1189        }
1190    }
1191
1192    #[test]
1193    fn connection_ext_prepare_cached_works_with_facet_methods() {
1194        let conn = Connection::open_in_memory().unwrap();
1195        conn.execute("CREATE TABLE ids (id INTEGER NOT NULL)", ())
1196            .unwrap();
1197        conn.execute("INSERT INTO ids (id) VALUES (5)", ()).unwrap();
1198
1199        #[derive(Facet)]
1200        struct QueryId {
1201            id: i64,
1202        }
1203
1204        #[derive(Debug, Facet, PartialEq)]
1205        struct IdRow {
1206            id: i64,
1207        }
1208
1209        let mut stmt = conn
1210            .facet_prepare_cached("SELECT id FROM ids WHERE id = :id")
1211            .unwrap();
1212        let row = stmt.facet_query_one::<IdRow, _>(QueryId { id: 5 }).unwrap();
1213        assert_eq!(row, IdRow { id: 5 });
1214    }
1215
1216    #[test]
1217    fn transparent_wrapper_works_for_params_and_rows() {
1218        let conn = Connection::open_in_memory().unwrap();
1219        conn.execute("CREATE TABLE monks (name TEXT NOT NULL)", ())
1220            .unwrap();
1221        conn.execute("INSERT INTO monks (name) VALUES ('teacup')", ())
1222            .unwrap();
1223
1224        #[derive(Debug, Facet, PartialEq, Eq)]
1225        #[facet(transparent)]
1226        struct MonkString(String);
1227
1228        #[derive(Facet)]
1229        struct QueryMonk {
1230            name: MonkString,
1231        }
1232
1233        #[derive(Debug, Facet, PartialEq, Eq)]
1234        struct MonkRow {
1235            name: MonkString,
1236        }
1237
1238        let row = conn
1239            .facet_query_one::<MonkRow, _>(
1240                "SELECT name FROM monks WHERE name = :name",
1241                QueryMonk {
1242                    name: MonkString("teacup".to_string()),
1243                },
1244            )
1245            .unwrap();
1246        assert_eq!(
1247            row,
1248            MonkRow {
1249                name: MonkString("teacup".to_string())
1250            }
1251        );
1252    }
1253}