use crate::protocol::FieldDescription;
use crate::types::PostgresType;
use crate::types::{FromSqlBase, FromSqlBinary, FromSqlText, ToSql};
use rust_decimal::Decimal;
use std::error::Error;
impl<'a> FromSqlBase<'a> for Decimal {
fn accepts_postgres_type(oid: i32) -> bool {
oid == PostgresType::NUMERIC.oid
}
}
impl<'a> FromSqlBinary<'a> for Decimal {
fn from_sql_binary(
raw: &'a [u8],
_field: &FieldDescription,
) -> Result<Self, Box<dyn Error + Sync + Send>> {
if raw.len() < 8 {
return Err("NUMERIC data too short".into());
}
let ndigits = i16::from_be_bytes([raw[0], raw[1]]);
let weight = i16::from_be_bytes([raw[2], raw[3]]);
let sign = i16::from_be_bytes([raw[4], raw[5]]);
let _dscale = i16::from_be_bytes([raw[6], raw[7]]);
if sign == 0xC000u16 as i16 {
return Err("NUMERIC NaN values are not supported by rust_decimal".into());
}
let is_negative = match sign {
0x0000 => false,
0x4000 => true,
_ => return Err(format!("Invalid NUMERIC sign: {sign:#x}").into()),
};
let expected_len = 8 + (ndigits as usize * 2);
if raw.len() < expected_len {
return Err(format!(
"NUMERIC data too short: expected {} bytes, got {}",
expected_len,
raw.len()
)
.into());
}
let mut digits = Vec::with_capacity(ndigits as usize);
for i in 0..ndigits {
let offset = 8 + (i as usize * 2);
let digit = i16::from_be_bytes([raw[offset], raw[offset + 1]]);
if !(0..10000).contains(&digit) {
return Err(format!("Invalid NUMERIC digit: {digit}").into());
}
digits.push(digit);
}
let mut mantissa: i128 = 0;
let mantissa_scale;
let digits_before_decimal = (weight + 1) as i32;
let total_digit_positions = ndigits as i32;
if digits_before_decimal <= 0 {
mantissa_scale = (-digits_before_decimal * 4 + total_digit_positions * 4) as u32;
for &digit in &digits {
mantissa = mantissa
.checked_mul(10000)
.and_then(|m| m.checked_add(digit as i128))
.ok_or("Numeric value exceeds rust_decimal precision (mantissa overflow)")?;
}
} else if digits_before_decimal >= total_digit_positions {
mantissa_scale = 0;
for &digit in &digits {
mantissa = mantissa
.checked_mul(10000)
.and_then(|m| m.checked_add(digit as i128))
.ok_or("Numeric value exceeds rust_decimal precision (mantissa overflow)")?;
}
let extra_zero_positions = (digits_before_decimal - total_digit_positions) * 4;
for _ in 0..extra_zero_positions {
mantissa = mantissa.checked_mul(10).ok_or(
"Numeric value exceeds rust_decimal precision (trailing zeros overflow)",
)?;
}
} else {
mantissa_scale = ((total_digit_positions - digits_before_decimal) * 4) as u32;
for &digit in &digits {
mantissa = mantissa
.checked_mul(10000)
.and_then(|m| m.checked_add(digit as i128))
.ok_or("Numeric value exceeds rust_decimal precision (mantissa overflow)")?;
}
}
if is_negative {
mantissa = -mantissa;
}
match Decimal::try_from_i128_with_scale(mantissa, mantissa_scale) {
Ok(decimal) => Ok(decimal),
Err(e) => Err(format!("Failed to create Decimal from mantissa {mantissa} with scale {mantissa_scale}: {e}").into()),
}
}
}
impl<'a> FromSqlText<'a> for Decimal {
fn from_sql_text(
raw: &'a str,
field: &FieldDescription,
) -> Result<Self, Box<dyn Error + Sync + Send>> {
match raw.parse::<Decimal>() {
Ok(decimal) => Ok(decimal),
Err(e) => Err(format!(
"Failed to parse NUMERIC from text '{raw}': {e}. Error occurred when parsing field {field:?}"
).into()),
}
}
}
impl ToSql for Decimal {
fn to_sql_binary(
&self,
target_buffer: &mut Vec<u8>,
) -> Result<(), Box<dyn Error + Sync + Send>> {
if self.is_zero() {
target_buffer.extend_from_slice(&0i16.to_be_bytes()); target_buffer.extend_from_slice(&0i16.to_be_bytes()); target_buffer.extend_from_slice(&0i16.to_be_bytes()); target_buffer.extend_from_slice(&0i16.to_be_bytes()); return Ok(());
}
let mantissa = self.mantissa().abs(); let scale = self.scale() as i16; let is_negative = self.is_sign_negative();
let mut decimal_digits = Vec::new();
let mut temp_mantissa = mantissa;
if temp_mantissa == 0 {
decimal_digits.push(0);
} else {
while temp_mantissa > 0 {
decimal_digits.insert(0, (temp_mantissa % 10) as u8);
temp_mantissa /= 10;
}
}
let total_decimal_digits = decimal_digits.len() as i16;
let digits_before_decimal = total_decimal_digits - scale;
let mut digits_10000 = Vec::new();
if digits_before_decimal > 0 {
let mut i = digits_before_decimal as usize;
while i > 0 {
let start = i.saturating_sub(4);
let mut group_value: i16 = 0;
for digit in decimal_digits.iter().take(i).skip(start) {
group_value = group_value * 10 + *digit as i16;
}
digits_10000.insert(0, group_value);
i = start;
}
}
if scale > 0 {
let fractional_start = if digits_before_decimal > 0 {
digits_before_decimal as usize
} else {
0
};
let mut i = fractional_start;
while i < decimal_digits.len() {
let end = std::cmp::min(i + 4, decimal_digits.len());
let mut group_value: i16 = 0;
for digit in decimal_digits.iter().take(end).skip(i) {
group_value = group_value * 10 + *digit as i16;
}
let digits_in_group = end - i;
for _ in digits_in_group..4 {
group_value *= 10;
}
digits_10000.push(group_value);
i += 4;
}
}
let weight = if digits_before_decimal <= 0 {
let leading_zero_groups = (-digits_before_decimal + 3) / 4;
-(leading_zero_groups + 1)
} else {
((digits_before_decimal + 3) / 4) - 1
};
while digits_10000.len() > 1 && digits_10000.last() == Some(&0) {
digits_10000.pop();
}
let ndigits = digits_10000.len() as i16;
let sign = if is_negative {
0x4000u16 as i16
} else {
0x0000i16
};
target_buffer.extend_from_slice(&ndigits.to_be_bytes());
target_buffer.extend_from_slice(&weight.to_be_bytes());
target_buffer.extend_from_slice(&sign.to_be_bytes());
target_buffer.extend_from_slice(&scale.to_be_bytes());
for digit in digits_10000 {
target_buffer.extend_from_slice(&digit.to_be_bytes());
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use rust_decimal::Decimal;
#[cfg(feature = "tokio")]
mod tokio_connection {
use super::*;
use crate::test_helpers::get_settings;
use crate::tokio_connection::new_client;
use tokio::test;
#[test]
async fn test_numeric_basic_values() {
let mut client = new_client(get_settings()).await.unwrap();
let zero: Decimal = client
.read_single_value_dual_mode("select 0::numeric")
.await;
assert_eq!(zero, Decimal::from(0));
let positive: Decimal = client
.read_single_value_dual_mode("select 12345::numeric")
.await;
assert_eq!(positive, Decimal::from(12345));
let negative: Decimal = client
.read_single_value_dual_mode("select -67890::numeric")
.await;
assert_eq!(negative, Decimal::from(-67890));
let decimal: Decimal = client
.read_single_value_dual_mode("select 123.456::numeric")
.await;
assert_eq!(decimal, "123.456".parse::<Decimal>().unwrap());
}
#[test]
async fn test_numeric_precision_scale() {
let mut client = new_client(get_settings()).await.unwrap();
let high_precision: Decimal = client
.read_single_value_dual_mode::<Decimal>("select 123456789.123456789::numeric(18,9)")
.await;
assert_eq!(
high_precision,
"123456789.123456789".parse::<Decimal>().unwrap()
);
let many_decimals: Decimal = client
.read_single_value_dual_mode::<Decimal>("select 1.000000001::numeric(10,9)")
.await;
assert_eq!(many_decimals, "1.000000001".parse::<Decimal>().unwrap());
let large_int: Decimal = client
.read_single_value_dual_mode("select 999999999999999999::numeric")
.await;
assert_eq!(large_int, "999999999999999999".parse::<Decimal>().unwrap());
}
#[test]
async fn test_numeric_postgresql_direct() {
let mut client = new_client(get_settings()).await.unwrap();
let small_decimal: Decimal = client
.read_single_value_dual_mode("select 0.000000001::numeric")
.await;
let expected = "0.000000001".parse::<Decimal>().unwrap();
assert_eq!(small_decimal, expected, "Direct PostgreSQL read failed");
}
#[test]
async fn test_numeric_small_decimal_round_trip() {
let mut client = new_client(get_settings()).await.unwrap();
client.execute_non_query_simple("drop table if exists test_numeric_debug; create table test_numeric_debug(value numeric);").await.unwrap();
let test_value = "0.000000001".parse::<Decimal>().unwrap();
client
.execute_non_query(
"insert into test_numeric_debug values ($1);",
&[&test_value],
)
.await
.unwrap();
let retrieved: Decimal = client
.read_single_value_dual_mode("select value from test_numeric_debug")
.await;
assert_eq!(
retrieved, test_value,
"Small decimal round-trip failed for {test_value}"
);
}
#[test]
async fn test_numeric_round_trip() {
let mut client = new_client(get_settings()).await.unwrap();
client.execute_non_query_simple("drop table if exists test_numeric_table; create table test_numeric_table(value numeric);").await.unwrap();
let test_values = vec![
"0".parse::<Decimal>().unwrap(),
"123.456".parse::<Decimal>().unwrap(),
"-789.012".parse::<Decimal>().unwrap(),
"999999999.999999999".parse::<Decimal>().unwrap(),
"0.000000001".parse::<Decimal>().unwrap(),
"1000000000".parse::<Decimal>().unwrap(),
];
for test_value in &test_values {
client
.execute_non_query("insert into test_numeric_table values ($1);", &[test_value])
.await
.unwrap();
let retrieved: Decimal = client
.read_single_value(
"select value from test_numeric_table order by value desc limit 1;",
&[],
)
.await;
assert_eq!(&retrieved, test_value, "Round-trip failed for {test_value}");
client
.execute_non_query("delete from test_numeric_table;", &[])
.await
.unwrap();
}
}
#[test]
async fn test_numeric_null_handling() {
let mut client = new_client(get_settings()).await.unwrap();
let null_value: Option<Decimal> = client
.read_single_value_dual_mode("select null::numeric")
.await;
assert_eq!(null_value, None);
}
#[test]
async fn test_numeric_edge_cases() {
let mut client = new_client(get_settings()).await.unwrap();
let small: Decimal = client
.read_single_value_dual_mode("select 0.0001::numeric")
.await;
assert_eq!(small, "0.0001".parse::<Decimal>().unwrap());
let trailing_zeros: Decimal = client
.read_single_value_dual_mode("select 123.4500::numeric")
.await;
assert_eq!(trailing_zeros, "123.45".parse::<Decimal>().unwrap()); }
#[test]
async fn test_numeric_nan_error() {
let mut client = new_client(get_settings()).await.unwrap();
let nan_result = client
.try_read_single_value::<Decimal>("select 'NaN'::numeric;", &[])
.await;
assert!(nan_result.is_err(), "Expected error for NaN NUMERIC");
assert!(nan_result.unwrap_err().to_string().contains("NaN"));
}
#[test]
async fn test_numeric_array_support() {
let mut client = new_client(get_settings()).await.unwrap();
let numeric_array: Vec<Decimal> = client
.read_single_value_dual_mode("select ARRAY[0::numeric, 123.456::numeric, -789.012::numeric, 0.000000001::numeric]").await;
let expected = vec![
"0".parse::<Decimal>().unwrap(),
"123.456".parse::<Decimal>().unwrap(),
"-789.012".parse::<Decimal>().unwrap(),
"0.000000001".parse::<Decimal>().unwrap(),
];
assert_eq!(numeric_array, expected);
}
}
}