1use std::fmt;
37
38use crate::value::Value;
39
40pub trait Queryable: Sized {
44 fn decode(value: &Value) -> Result<Self, DecodeError>;
45}
46
47#[derive(Debug, Clone, PartialEq)]
49pub struct DecodeError {
50 type_root: Option<String>,
55 path: Vec<String>,
58 kind: DecodeErrorKind,
59}
60
61#[derive(Debug, Clone, PartialEq)]
62pub enum DecodeErrorKind {
63 MissingField(String),
66 WrongType {
67 expected: &'static str,
68 actual: &'static str,
69 },
70 UnknownVariant(String),
72 OutOfRange { value: i64, target: &'static str },
75 Json(String),
77 Invalid(String),
80}
81
82impl DecodeError {
83 fn new(kind: DecodeErrorKind) -> Self {
84 Self {
85 type_root: None,
86 path: Vec::new(),
87 kind,
88 }
89 }
90
91 fn under(mut self, segment: impl Into<String>) -> Self {
94 self.path.insert(0, segment.into());
95 self
96 }
97
98 fn rooted_at(mut self, type_name: &str) -> Self {
102 self.type_root = Some(type_name.to_string());
103 self
104 }
105
106 pub fn kind(&self) -> &DecodeErrorKind {
107 &self.kind
108 }
109
110 pub fn custom(message: impl Into<String>) -> Self {
117 Self::new(DecodeErrorKind::Invalid(message.into()))
118 }
119
120 fn wrong_type(expected: &'static str, actual: &Value) -> Self {
121 Self::new(DecodeErrorKind::WrongType {
122 expected,
123 actual: type_name(actual),
124 })
125 }
126}
127
128impl fmt::Display for DecodeError {
129 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
130 let mut location: Vec<&str> = Vec::with_capacity(self.path.len() + 1);
131 location.extend(self.type_root.as_deref());
132 location.extend(self.path.iter().map(String::as_str));
133 if location.is_empty() {
134 write!(f, "cannot decode query result: {}", self.kind)
135 } else {
136 write!(f, "cannot decode {}: {}", location.join("."), self.kind)
137 }
138 }
139}
140
141impl std::error::Error for DecodeError {}
142
143impl fmt::Display for DecodeErrorKind {
144 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
145 match self {
146 Self::MissingField(name) => write!(f, "the query's shape has no '{name}'"),
147 Self::WrongType { expected, actual } => write!(f, "expected {expected}, got {actual}"),
148 Self::UnknownVariant(label) => write!(f, "'{label}' is not a known variant"),
149 Self::OutOfRange { value, target } => write!(f, "{value} is out of range for {target}"),
150 Self::Json(message) => write!(f, "{message}"),
151 Self::Invalid(message) => write!(f, "{message}"),
152 }
153 }
154}
155
156fn type_name(value: &Value) -> &'static str {
158 match value {
159 Value::Null => "an empty set",
160 Value::Bool(_) => "a bool",
161 Value::Int64(_) => "an integer",
162 Value::Float64(_) => "a float",
163 Value::Str(_) => "a string",
164 Value::Bytes(_) => "bytes",
165 Value::Uuid(_) => "a uuid",
166 Value::Decimal(_) => "a decimal",
167 Value::Duration { .. } => "a duration",
168 Value::Date(_) => "a date",
169 Value::Time(_) => "a time",
170 Value::Timestamp(_) => "a local datetime",
171 Value::Timestamptz(_) => "a datetime",
172 Value::Range(_) => "a range",
173 Value::Array(_) => "a set",
174 Value::Tuple(_) => "a tuple",
175 Value::Object(_) => "an object",
176 Value::Enum { .. } => "an enum value",
177 Value::Group(_) => "a group",
178 Value::VectorSearch { .. } => "a vector search result",
179 Value::FtsSearch { .. } => "an FTS search result",
180 }
181}
182
183impl Queryable for Value {
187 fn decode(value: &Value) -> Result<Self, DecodeError> {
188 Ok(value.clone())
189 }
190}
191
192impl<T: Queryable> Queryable for Option<T> {
194 fn decode(value: &Value) -> Result<Self, DecodeError> {
195 match value {
196 Value::Null => Ok(None),
197 other => T::decode(other).map(Some),
198 }
199 }
200}
201
202impl<T: Queryable> Queryable for Vec<T> {
207 fn decode(value: &Value) -> Result<Self, DecodeError> {
208 match value {
209 Value::Array(items) => items
210 .iter()
211 .enumerate()
212 .map(|(index, item)| T::decode(item).map_err(|error| error.under(index.to_string())))
213 .collect(),
214 Value::Null => Ok(Vec::new()),
215 single => T::decode(single).map(|decoded| vec![decoded]),
216 }
217 }
218}
219
220impl Queryable for bool {
221 fn decode(value: &Value) -> Result<Self, DecodeError> {
222 match value {
223 Value::Bool(b) => Ok(*b),
224 other => Err(DecodeError::wrong_type("a bool", other)),
225 }
226 }
227}
228
229impl Queryable for String {
230 fn decode(value: &Value) -> Result<Self, DecodeError> {
231 match value {
232 Value::Str(s) => Ok(s.clone()),
233 Value::Enum { value, .. } => Ok(value.clone()),
238 Value::Decimal(s) => Ok(s.clone()),
239 other => Err(DecodeError::wrong_type("a string", other)),
240 }
241 }
242}
243
244impl Queryable for Vec<u8> {
245 fn decode(value: &Value) -> Result<Self, DecodeError> {
246 match value {
247 Value::Bytes(b) => Ok(b.clone()),
248 other => Err(DecodeError::wrong_type("bytes", other)),
249 }
250 }
251}
252
253impl Queryable for uuid::Uuid {
254 fn decode(value: &Value) -> Result<Self, DecodeError> {
255 match value {
256 Value::Uuid(u) => Ok(*u),
257 other => Err(DecodeError::wrong_type("a uuid", other)),
258 }
259 }
260}
261
262impl Queryable for i64 {
263 fn decode(value: &Value) -> Result<Self, DecodeError> {
264 match value {
265 Value::Int64(n) => Ok(*n),
266 other => Err(DecodeError::wrong_type("an integer", other)),
267 }
268 }
269}
270
271macro_rules! queryable_narrow_int {
275 ($($target:ty),* $(,)?) => {
276 $(
277 impl Queryable for $target {
278 fn decode(value: &Value) -> Result<Self, DecodeError> {
279 match value {
280 Value::Int64(n) => <$target>::try_from(*n).map_err(|_| {
281 DecodeError::new(DecodeErrorKind::OutOfRange {
282 value: *n,
283 target: stringify!($target),
284 })
285 }),
286 other => Err(DecodeError::wrong_type("an integer", other)),
287 }
288 }
289 }
290 )*
291 };
292}
293
294queryable_narrow_int!(i16, i32, u16, u32, u64);
295
296impl Queryable for f64 {
297 fn decode(value: &Value) -> Result<Self, DecodeError> {
298 match value {
299 Value::Float64(f) => Ok(*f),
300 Value::Int64(n) => Ok(*n as f64),
304 other => Err(DecodeError::wrong_type("a float", other)),
305 }
306 }
307}
308
309impl Queryable for f32 {
310 fn decode(value: &Value) -> Result<Self, DecodeError> {
311 f64::decode(value).map(|f| f as f32)
312 }
313}
314
315impl Queryable for chrono::DateTime<chrono::Utc> {
316 fn decode(value: &Value) -> Result<Self, DecodeError> {
317 match value {
318 Value::Timestamptz(micros) => from_pg_micros(*micros).map(|naive| naive.and_utc()),
319 other => Err(DecodeError::wrong_type("a datetime", other)),
320 }
321 }
322}
323
324impl Queryable for chrono::NaiveDateTime {
325 fn decode(value: &Value) -> Result<Self, DecodeError> {
326 match value {
327 Value::Timestamp(micros) => from_pg_micros(*micros),
328 other => Err(DecodeError::wrong_type("a local datetime", other)),
329 }
330 }
331}
332
333impl Queryable for chrono::NaiveDate {
334 fn decode(value: &Value) -> Result<Self, DecodeError> {
335 match value {
336 Value::Date(days) => pg_epoch()
337 .checked_add_signed(chrono::Duration::days(i64::from(*days)))
338 .ok_or_else(|| {
339 DecodeError::new(DecodeErrorKind::Invalid(format!(
340 "date {days} days from 2000-01-01 is outside the supported range"
341 )))
342 }),
343 other => Err(DecodeError::wrong_type("a date", other)),
344 }
345 }
346}
347
348impl Queryable for chrono::NaiveTime {
349 fn decode(value: &Value) -> Result<Self, DecodeError> {
350 match value {
351 Value::Time(micros) => chrono::NaiveTime::from_hms_opt(0, 0, 0)
352 .and_then(|midnight| {
353 midnight
354 .overflowing_add_signed(chrono::Duration::microseconds(*micros))
355 .0
356 .into()
357 })
358 .ok_or_else(|| {
359 DecodeError::new(DecodeErrorKind::Invalid(format!(
360 "time {micros}\u{b5}s after midnight is not a valid time of day"
361 )))
362 }),
363 other => Err(DecodeError::wrong_type("a time", other)),
364 }
365 }
366}
367
368fn pg_epoch() -> chrono::NaiveDate {
369 chrono::NaiveDate::from_ymd_opt(2000, 1, 1).expect("2000-01-01 is a valid date")
371}
372
373fn from_pg_micros(micros: i64) -> Result<chrono::NaiveDateTime, DecodeError> {
374 pg_epoch()
375 .and_hms_opt(0, 0, 0)
376 .and_then(|midnight| midnight.checked_add_signed(chrono::Duration::microseconds(micros)))
377 .ok_or_else(|| {
378 DecodeError::new(DecodeErrorKind::Invalid(format!(
379 "timestamp {micros}\u{b5}s from 2000-01-01 is outside the supported range"
380 )))
381 })
382}
383
384#[doc(hidden)]
388pub mod derive {
389 use super::{DecodeError, DecodeErrorKind, Queryable};
390 use crate::value::{Object, Value};
391
392 pub fn object<'v>(value: &'v Value, container: &'static str) -> Result<&'v Object, DecodeError> {
393 match value {
394 Value::Object(object) => Ok(object),
395 other => Err(DecodeError::wrong_type("an object", other).rooted_at(container)),
396 }
397 }
398
399 fn missing(container: &'static str, name: &str) -> DecodeError {
403 DecodeError::new(DecodeErrorKind::MissingField(name.to_string())).rooted_at(container)
404 }
405
406 pub fn field<T: Queryable>(object: &Object, container: &'static str, name: &str) -> Result<T, DecodeError> {
407 let value = object.get(name).ok_or_else(|| missing(container, name))?;
408 T::decode(value).map_err(|error| error.under(name).rooted_at(container))
409 }
410
411 pub fn json_field<T: serde::de::DeserializeOwned>(
412 object: &Object,
413 container: &'static str,
414 name: &str,
415 ) -> Result<T, DecodeError> {
416 let value = object.get(name).ok_or_else(|| missing(container, name))?;
417 from_json_value(value).map_err(|error| error.under(name).rooted_at(container))
418 }
419
420 pub fn from_json<T: serde::de::DeserializeOwned>(value: &Value, container: &'static str) -> Result<T, DecodeError> {
421 from_json_value(value).map_err(|error| error.rooted_at(container))
422 }
423
424 fn from_json_value<T: serde::de::DeserializeOwned>(value: &Value) -> Result<T, DecodeError> {
432 serde_json::from_str(&crate::json::to_json(value))
433 .map_err(|error| DecodeError::new(DecodeErrorKind::Json(error.to_string())))
434 }
435
436 pub fn enum_label<'v>(value: &'v Value, container: &'static str) -> Result<&'v str, DecodeError> {
439 match value {
440 Value::Enum { value, .. } => Ok(value.as_str()),
441 Value::Str(label) => Ok(label.as_str()),
442 other => Err(DecodeError::wrong_type("an enum value", other).rooted_at(container)),
443 }
444 }
445
446 pub fn unknown_variant(container: &'static str, label: &str) -> DecodeError {
447 DecodeError::new(DecodeErrorKind::UnknownVariant(label.to_string())).rooted_at(container)
448 }
449}
450
451pub(crate) fn decode_rows<R: Queryable>(values: Vec<Value>) -> crate::Result<Vec<R>> {
455 values
456 .iter()
457 .map(|value| R::decode(value).map_err(crate::Error::from))
458 .collect()
459}
460
461pub(crate) fn decode_optional_row<R: Queryable>(value: Option<Value>) -> crate::Result<Option<R>> {
462 value.as_ref().map(R::decode).transpose().map_err(crate::Error::from)
463}
464
465pub(crate) fn decode_row<R: Queryable>(value: Value) -> crate::Result<R> {
466 R::decode(&value).map_err(crate::Error::from)
467}
468
469#[cfg(test)]
470mod tests {
471 use super::*;
472 use crate::value::Object;
473 use crate::{QueryArgs, named_args};
474
475 #[derive(Debug, PartialEq, crate::Queryable)]
479 #[pylon(crate_path = crate)]
480 struct Row {
481 id: uuid::Uuid,
482 attempt: i32,
483 last_error: Option<String>,
484 created_at: chrono::DateTime<chrono::Utc>,
485 }
486
487 #[derive(Debug, PartialEq, crate::Queryable)]
488 #[pylon(crate_path = crate)]
489 enum WebhookEvent {
490 #[pylon(rename = "contact.created")]
491 ContactCreated,
492 #[pylon(rename = "contact.updated")]
493 ContactUpdated,
494 Other,
495 }
496
497 #[derive(Debug, PartialEq, crate::Queryable)]
498 #[pylon(crate_path = crate)]
499 struct Nested {
500 latest_version: Option<Inner>,
501 #[pylon(rename = "type")]
502 kind: WebhookEvent,
503 }
504
505 #[derive(Debug, PartialEq, crate::Queryable)]
506 #[pylon(crate_path = crate)]
507 struct Inner {
508 runner: String,
509 }
510
511 fn object(fields: Vec<(&str, Value)>) -> Value {
512 Value::Object(Object {
513 type_name: Some("test::Row".to_string()),
514 fields: fields.into_iter().map(|(n, v)| (n.to_string(), v)).collect(),
515 implicit_id: false,
516 })
517 }
518
519 #[test]
520 fn decodes_a_row_struct() {
521 let id = uuid::Uuid::from_u128(7);
522 let row = Row::decode(&object(vec![
523 ("id", Value::Uuid(id)),
524 ("attempt", Value::Int64(3)),
525 ("last_error", Value::Null),
526 ("created_at", Value::Timestamptz(0)),
527 ]))
528 .unwrap();
529 assert_eq!(row.id, id);
530 assert_eq!(row.attempt, 3);
531 assert_eq!(row.last_error, None);
532 assert_eq!(row.created_at.to_rfc3339(), "2000-01-01T00:00:00+00:00");
533 }
534
535 #[test]
539 fn field_order_does_not_matter() {
540 let id = uuid::Uuid::from_u128(1);
541 let row = Row::decode(&object(vec![
542 ("created_at", Value::Timestamptz(0)),
543 ("last_error", Value::Str("boom".into())),
544 ("attempt", Value::Int64(1)),
545 ("id", Value::Uuid(id)),
546 ]))
547 .unwrap();
548 assert_eq!(row.last_error.as_deref(), Some("boom"));
549 assert_eq!(row.id, id);
550 }
551
552 #[test]
555 fn a_field_the_shape_omitted_is_an_error() {
556 let error = Row::decode(&object(vec![
557 ("id", Value::Uuid(uuid::Uuid::nil())),
558 ("attempt", Value::Int64(1)),
559 ("created_at", Value::Timestamptz(0)),
560 ]))
561 .unwrap_err();
562 assert_eq!(error.kind(), &DecodeErrorKind::MissingField("last_error".to_string()));
563 assert_eq!(
564 error.to_string(),
565 "cannot decode Row: the query's shape has no 'last_error'"
566 );
567 }
568
569 #[test]
570 fn nested_failures_name_their_whole_path() {
571 let error = Nested::decode(&object(vec![
572 ("latest_version", object(vec![("runner", Value::Int64(4))])),
573 ("type", Value::Str("contact.created".into())),
574 ]))
575 .unwrap_err();
576 assert_eq!(
577 error.to_string(),
578 "cannot decode Nested.latest_version.runner: expected a string, got an integer"
579 );
580 }
581
582 #[test]
583 fn decodes_a_renamed_enum_from_either_a_real_enum_or_a_str_cast() {
584 let from_enum = WebhookEvent::decode(&Value::Enum {
585 type_name: "integration::WebhookEvent".to_string(),
586 value: "contact.updated".to_string(),
587 })
588 .unwrap();
589 assert_eq!(from_enum, WebhookEvent::ContactUpdated);
590 let from_cast = WebhookEvent::decode(&Value::Str("contact.created".into())).unwrap();
593 assert_eq!(from_cast, WebhookEvent::ContactCreated);
594 assert_eq!(
596 WebhookEvent::decode(&Value::Str("Other".into())).unwrap(),
597 WebhookEvent::Other
598 );
599 }
600
601 #[test]
602 fn an_unknown_enum_label_names_the_label_it_saw() {
603 let error = WebhookEvent::decode(&Value::Str("contact.merged".into())).unwrap_err();
604 assert_eq!(
605 error.to_string(),
606 "cannot decode WebhookEvent: 'contact.merged' is not a known variant"
607 );
608 }
609
610 #[test]
613 fn decodes_doubly_optional_nesting() {
614 let row = Nested::decode(&object(vec![
615 ("latest_version", Value::Null),
616 ("type", Value::Str("Other".into())),
617 ]))
618 .unwrap();
619 assert_eq!(row.latest_version, None);
620 }
621
622 #[test]
625 fn a_custom_decode_failure_converts_into_the_crate_error() {
626 let error: crate::Error = DecodeError::custom("failed to decode query result JSON: eof").into();
627 assert_eq!(
628 error.to_string(),
629 "cannot decode query result: failed to decode query result JSON: eof"
630 );
631 }
632
633 #[test]
634 fn narrow_integers_range_check_instead_of_wrapping() {
635 assert_eq!(i32::decode(&Value::Int64(-5)).unwrap(), -5);
636 let error = i32::decode(&Value::Int64(i64::from(i32::MAX) + 1)).unwrap_err();
637 assert_eq!(
638 error.kind(),
639 &DecodeErrorKind::OutOfRange {
640 value: 2_147_483_648,
641 target: "i32"
642 }
643 );
644 }
645
646 #[test]
647 fn a_multi_pointer_decodes_into_a_vec() {
648 let rows: Vec<Inner> = Vec::decode(&Value::Array(vec![
649 object(vec![("runner", Value::Str("lambda".into()))]),
650 object(vec![("runner", Value::Str("firecracker".into()))]),
651 ]))
652 .unwrap();
653 assert_eq!(rows.len(), 2);
654 assert_eq!(rows[1].runner, "firecracker");
655 assert_eq!(Vec::<Inner>::decode(&Value::Null).unwrap(), vec![]);
657 }
658
659 #[test]
660 fn a_json_field_deserializes_from_a_natively_decoded_document() {
661 #[derive(Debug, PartialEq, crate::Queryable)]
662 #[pylon(crate_path = crate)]
663 struct WithJson {
664 #[pylon(json)]
665 payload: Option<serde_json::Value>,
666 }
667 let row = WithJson::decode(&object(vec![(
670 "payload",
671 Value::Object(Object {
672 type_name: None,
673 fields: vec![("email".to_string(), Value::Str("a@b.test".into()))],
674 implicit_id: false,
675 }),
676 )]))
677 .unwrap();
678 assert_eq!(row.payload.unwrap()["email"], "a@b.test");
679 }
680
681 #[test]
684 fn a_duration_binds_as_a_fixed_microsecond_interval() {
685 let bound = crate::QueryArg::to_decoded(&std::time::Duration::from_secs(90));
686 assert_eq!(
687 bound,
688 pylon_value::DecodedValue::Interval {
689 months: 0,
690 days: 0,
691 microseconds: 90_000_000
692 }
693 );
694 }
695
696 #[test]
697 fn positional_arguments_bind_by_index() {
698 let id = uuid::Uuid::from_u128(9);
699 let args = (id, "urgent", 4i32);
703 let params = args.to_params();
704 assert_eq!(params[0].0, "0");
705 assert_eq!(params[1], ("1", pylon_value::DecodedValue::Str("urgent".into())));
706 assert_eq!(params[2], ("2", pylon_value::DecodedValue::I64(4)));
707 assert!(().to_params().is_empty());
708 }
709
710 #[test]
711 fn named_arguments_accept_mixed_types_including_absent_optionals() {
712 let args = named_args! {
713 "id" => uuid::Uuid::nil(),
714 "width" => Option::<i32>::None,
715 "labels" => vec!["a".to_string()],
716 };
717 let params: std::collections::HashMap<_, _> = args.to_params().into_iter().collect();
718 assert_eq!(params["width"], pylon_value::DecodedValue::Null);
719 assert_eq!(
720 params["labels"],
721 pylon_value::DecodedValue::Array(vec![pylon_value::DecodedValue::Str("a".into())])
722 );
723 assert_eq!(params["id"], pylon_value::DecodedValue::Uuid([0; 16]));
724 }
725
726 #[test]
729 fn datetimes_round_trip_through_the_pg_epoch() {
730 let when = chrono::DateTime::parse_from_rfc3339("2026-09-24T12:34:56Z")
731 .unwrap()
732 .with_timezone(&chrono::Utc);
733 let bound = crate::QueryArg::to_decoded(&when);
734 let pylon_value::DecodedValue::Timestamptz(micros) = bound else {
735 panic!("a datetime must bind as a timestamptz, got {bound:?}");
736 };
737 assert_eq!(
738 chrono::DateTime::<chrono::Utc>::decode(&Value::Timestamptz(micros)).unwrap(),
739 when
740 );
741 }
742}