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, ¶ms)?;
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, ¶ms)?;
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, ¶ms)?;
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}