use crate::{DjogiError, VisageError};
use tokio_postgres::Row;
use tokio_postgres::types::{FromSql, Type};
pub trait FromPgRow: Sized {
const COLUMNS: &'static [&'static str];
const COLUMN_LIST: &'static str;
fn from_pg_row(row: &tokio_postgres::Row) -> Result<Self, crate::DjogiError>;
}
#[doc(hidden)]
pub trait FromJoinedPgRow: Sized {
fn from_joined_pg_row(row: &Row, prefix: &str) -> Result<Self, DjogiError>;
}
#[doc(hidden)]
pub fn joined_alias_for_prefix(prefix: &str, idx: usize, col_name: &str) -> String {
match prefix {
"__djogi_old__" => format!("o{idx}"),
"__djogi_new__" => format!("n{idx}"),
_ => format!("{prefix}{col_name}"),
}
}
#[doc(hidden)]
pub fn decode_at<'a, T>(row: &'a Row, idx: usize, name: &'static str) -> Result<T, DjogiError>
where
T: FromSql<'a>,
{
debug_assert_eq!(
row.columns()[idx].name(),
name,
"FromPgRow column-order drift: position {} expected {:?}, got {:?}",
idx,
name,
row.columns()[idx].name(),
);
row.try_get::<_, T>(idx)
.map_err(|e| DjogiError::Decode(format!("column `{}`: {}", name, e)))
}
#[doc(hidden)]
pub fn decode_derived_at<'a, T>(
row: &'a Row,
idx: usize,
visage: &'static str,
field: &'static str,
) -> Result<T, DjogiError>
where
T: FromSql<'a>,
{
debug_assert_eq!(
row.columns()[idx].name(),
field,
"FromPgRow column-order drift on derived visage field: position {} expected {:?}, got {:?}",
idx,
field,
row.columns()[idx].name(),
);
row.try_get::<_, T>(idx).map_err(|e| {
let actual = row
.columns()
.get(idx)
.map(|column| pg_type_name(column.type_()))
.unwrap_or("unknown");
map_derived_decode_failure::<T>(std::error::Error::source(&e), actual, visage, field)
})
}
fn map_derived_decode_failure<T>(
source: Option<&(dyn std::error::Error + 'static)>,
actual: &'static str,
visage: &'static str,
field: &'static str,
) -> DjogiError {
if source
.and_then(|source| source.downcast_ref::<postgres_types::WasNull>())
.is_some()
{
return DjogiError::Visage(VisageError::DbComputedNullForNonOptional { visage, field });
}
DjogiError::Visage(VisageError::DbComputedTypeMismatch {
visage,
field,
expected: std::any::type_name::<T>(),
actual,
})
}
fn pg_type_name(ty: &Type) -> &'static str {
if *ty == Type::BOOL {
"BOOL"
} else if *ty == Type::CHAR {
"CHAR"
} else if *ty == Type::INT2 {
"INT2"
} else if *ty == Type::INT4 {
"INT4"
} else if *ty == Type::INT8 {
"INT8"
} else if *ty == Type::FLOAT4 {
"FLOAT4"
} else if *ty == Type::FLOAT8 {
"FLOAT8"
} else if *ty == Type::NUMERIC {
"NUMERIC"
} else if *ty == Type::TEXT {
"TEXT"
} else if *ty == Type::VARCHAR {
"VARCHAR"
} else if *ty == Type::BPCHAR {
"BPCHAR"
} else if *ty == Type::TIMESTAMP {
"TIMESTAMP"
} else if *ty == Type::TIMESTAMPTZ {
"TIMESTAMPTZ"
} else if *ty == Type::DATE {
"DATE"
} else if *ty == Type::TIME {
"TIME"
} else if *ty == Type::UUID {
"UUID"
} else if *ty == Type::JSON {
"JSON"
} else if *ty == Type::JSONB {
"JSONB"
} else {
"unknown"
}
}
#[doc(hidden)]
pub fn try_get_scalar<'a, T>(row: &'a Row, idx: usize) -> Result<T, DjogiError>
where
T: FromSql<'a>,
{
row.try_get(idx).map_err(DjogiError::from)
}
#[doc(hidden)]
pub fn decode_narrowed<'a, W, N>(
row: &'a Row,
idx: usize,
name: &'static str,
) -> Result<N, DjogiError>
where
W: FromSql<'a> + std::fmt::Display + Copy,
N: TryFrom<W>,
<N as TryFrom<W>>::Error: std::fmt::Display,
{
let wide: W = decode_at(row, idx, name)?;
N::try_from(wide).map_err(|e| {
DjogiError::Decode(format!("column `{name}`: value {wide} out of range: {e}",))
})
}
#[doc(hidden)]
pub fn decode_narrowed_opt<'a, W, N>(
row: &'a Row,
idx: usize,
name: &'static str,
) -> Result<Option<N>, DjogiError>
where
W: FromSql<'a> + std::fmt::Display + Copy,
N: TryFrom<W>,
<N as TryFrom<W>>::Error: std::fmt::Display,
{
let wide: Option<W> = decode_at(row, idx, name)?;
wide.map(|w| {
N::try_from(w).map_err(|e| {
DjogiError::Decode(format!("column `{name}`: value {w} out of range: {e}",))
})
})
.transpose()
}
fn decimal_to_u64(
dec: rust_decimal::Decimal,
col: impl std::fmt::Display,
) -> Result<u64, DjogiError> {
use rust_decimal::prelude::ToPrimitive as _;
if !dec.fract().is_zero() {
return Err(DjogiError::Decode(format!(
"column `{col}`: Decimal value {dec} has a fractional part and cannot be decoded as u64",
)));
}
dec.to_u64().ok_or_else(|| {
DjogiError::Decode(format!(
"column `{col}`: Decimal value {dec} out of u64 range",
))
})
}
#[doc(hidden)]
pub fn decode_u64_from_decimal(
row: &Row,
idx: usize,
name: &'static str,
) -> Result<u64, DjogiError> {
let dec: rust_decimal::Decimal = decode_at(row, idx, name)?;
decimal_to_u64(dec, name)
}
#[doc(hidden)]
pub fn decode_opt_u64_from_decimal(
row: &Row,
idx: usize,
name: &'static str,
) -> Result<Option<u64>, DjogiError> {
let dec: Option<rust_decimal::Decimal> = decode_at(row, idx, name)?;
dec.map(|d| decimal_to_u64(d, name)).transpose()
}
#[doc(hidden)]
pub fn decode_narrowed_by_name<'a, W, N>(row: &'a Row, col_name: &str) -> Result<N, DjogiError>
where
W: FromSql<'a> + std::fmt::Display + Copy,
N: TryFrom<W>,
<N as TryFrom<W>>::Error: std::fmt::Display,
{
let wide: W = row
.try_get::<_, W>(col_name)
.map_err(|e| DjogiError::Decode(format!("column `{col_name}`: {e}")))?;
N::try_from(wide).map_err(|e| {
DjogiError::Decode(format!(
"column `{col_name}`: value {wide} out of range: {e}",
))
})
}
#[doc(hidden)]
pub fn decode_narrowed_opt_by_name<'a, W, N>(
row: &'a Row,
col_name: &str,
) -> Result<Option<N>, DjogiError>
where
W: FromSql<'a> + std::fmt::Display + Copy,
N: TryFrom<W>,
<N as TryFrom<W>>::Error: std::fmt::Display,
{
let wide: Option<W> = row
.try_get::<_, Option<W>>(col_name)
.map_err(|e| DjogiError::Decode(format!("column `{col_name}`: {e}")))?;
wide.map(|w| {
N::try_from(w).map_err(|e| {
DjogiError::Decode(format!("column `{col_name}`: value {w} out of range: {e}",))
})
})
.transpose()
}
#[doc(hidden)]
pub fn decode_u64_from_decimal_by_name(row: &Row, col_name: &str) -> Result<u64, DjogiError> {
let dec: rust_decimal::Decimal = row
.try_get::<_, rust_decimal::Decimal>(col_name)
.map_err(|e| DjogiError::Decode(format!("column `{col_name}`: {e}")))?;
decimal_to_u64(dec, col_name)
}
#[doc(hidden)]
pub fn decode_opt_u64_from_decimal_by_name(
row: &Row,
col_name: &str,
) -> Result<Option<u64>, DjogiError> {
let dec: Option<rust_decimal::Decimal> = row
.try_get::<_, Option<rust_decimal::Decimal>>(col_name)
.map_err(|e| DjogiError::Decode(format!("column `{col_name}`: {e}")))?;
dec.map(|d| decimal_to_u64(d, col_name)).transpose()
}
#[cfg(test)]
mod tests {
use super::{decimal_to_u64, joined_alias_for_prefix, map_derived_decode_failure};
use crate::{DjogiError, VisageError};
use rust_decimal::Decimal;
use std::fmt;
use std::str::FromStr as _;
#[derive(Debug)]
struct NotWasNull;
impl fmt::Display for NotWasNull {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("not a null decode failure")
}
}
impl std::error::Error for NotWasNull {}
#[test]
fn derived_decode_failure_maps_was_null_to_visage_null() {
let source = postgres_types::WasNull;
let err = map_derived_decode_failure::<String>(
Some(&source),
"TEXT",
"DerivedNullRowPublic",
"computed_label",
);
match err {
DjogiError::Visage(VisageError::DbComputedNullForNonOptional { visage, field }) => {
assert_eq!(visage, "DerivedNullRowPublic");
assert_eq!(field, "computed_label");
}
other => panic!("expected derived NULL visage error, got {other:?}"),
}
}
#[test]
fn derived_decode_failure_maps_non_null_failure_to_visage_type_mismatch() {
let source = NotWasNull;
let err = map_derived_decode_failure::<String>(
Some(&source),
"INT4",
"DerivedTypeRowPublic",
"computed_label",
);
match err {
DjogiError::Visage(VisageError::DbComputedTypeMismatch {
visage,
field,
expected,
actual,
}) => {
assert_eq!(visage, "DerivedTypeRowPublic");
assert_eq!(field, "computed_label");
assert!(expected.contains("String"), "expected type was {expected}");
assert_eq!(actual, "INT4");
}
other => panic!("expected derived type-mismatch visage error, got {other:?}"),
}
}
#[test]
fn decimal_to_u64_rejects_fractional_value() {
let dec = Decimal::from_str("1.5").unwrap();
let err = decimal_to_u64(dec, "col").unwrap_err();
let msg = format!("{err:?}");
assert!(
msg.contains("fractional"),
"error for fractional decimal must mention 'fractional': {msg}"
);
}
#[test]
fn decimal_to_u64_rejects_negative_fractional() {
let dec = Decimal::from_str("-0.1").unwrap();
let err = decimal_to_u64(dec, "col").unwrap_err();
let msg = format!("{err:?}");
assert!(
msg.contains("fractional"),
"error for negative fractional must mention 'fractional': {msg}"
);
}
#[test]
fn decimal_to_u64_rejects_value_above_u64_max() {
let dec = Decimal::from_str("18446744073709551616").unwrap();
let err = decimal_to_u64(dec, "col").unwrap_err();
let msg = format!("{err:?}");
assert!(
msg.contains("out of u64 range"),
"error for out-of-range value must mention 'out of u64 range': {msg}"
);
}
#[test]
fn decimal_to_u64_rejects_negative_integer() {
let dec = Decimal::from_str("-1").unwrap();
let err = decimal_to_u64(dec, "col").unwrap_err();
let msg = format!("{err:?}");
assert!(
msg.contains("out of u64 range"),
"error for negative integer must mention 'out of u64 range': {msg}"
);
}
#[test]
fn decimal_to_u64_accepts_zero() {
let dec = Decimal::from_str("0").unwrap();
assert_eq!(decimal_to_u64(dec, "col").unwrap(), 0u64);
}
#[test]
fn decimal_to_u64_accepts_u64_max() {
let dec = Decimal::from_str("18446744073709551615").unwrap();
assert_eq!(
decimal_to_u64(dec, "col").unwrap(),
u64::MAX,
"u64::MAX must decode without error"
);
}
#[test]
fn decimal_to_u64_accepts_positive_integer() {
let dec = Decimal::from_str("42").unwrap();
assert_eq!(decimal_to_u64(dec, "col").unwrap(), 42u64);
}
#[test]
fn joined_alias_for_prefix_maps_old_and_new_ordinals() {
assert_eq!(joined_alias_for_prefix("__djogi_old__", 7, "title"), "o7");
assert_eq!(joined_alias_for_prefix("__djogi_new__", 7, "title"), "n7");
assert_eq!(
joined_alias_for_prefix("rel_owner_id.", 7, "title"),
"rel_owner_id.title"
);
}
}