Skip to main content

elefant_client/types/
numbers.rs

1use crate::protocol::FieldDescription;
2use crate::types::{FromSqlBase, FromSqlBinary, FromSqlText, PostgresNamedType, ToSql};
3use crate::PostgresType;
4use std::error::Error;
5
6macro_rules! impl_number {
7    ($typ: ty, $standard_type: expr $(, $also_accept: expr)*) => {
8        impl<'a> FromSqlBase<'a> for $typ {
9            fn accepts_postgres_type(oid: i32) -> bool {
10                oid == $standard_type.oid $( || oid == $also_accept.oid)*
11            }
12        }
13
14        impl<'a> FromSqlBinary<'a> for $typ {
15            fn from_sql_binary(
16                raw: &'a [u8],
17                field: &FieldDescription,
18            ) -> Result<Self, Box<dyn Error + Sync + Send>> {
19
20                const BYTE_SIZE: usize = std::mem::size_of::<$typ>();
21
22                if raw.len() != BYTE_SIZE {
23                    return Err(format!("Invalid length for {}. Expected {} bytes, got {} bytes instead. Error occurred when parsing field {:?}", std::any::type_name::<$typ>(), BYTE_SIZE, raw.len(), field).into());
24                }
25
26                Ok(<$typ>::from_be_bytes(raw.try_into().unwrap()))
27            }
28        }
29
30        impl<'a> FromSqlText<'a> for $typ {
31            fn from_sql_text(
32                raw: &'a str,
33                _field: &FieldDescription,
34            ) -> Result<Self, Box<dyn Error + Sync + Send>> {
35                Ok(raw.parse()?)
36            }
37        }
38
39        impl PostgresNamedType for $typ {
40            const PG_NAME: &'static str = $standard_type.name;
41        }
42
43        impl ToSql for $typ {
44            fn to_sql_binary(&self, target_buffer: &mut Vec<u8>) -> Result<(), Box<dyn Error + Sync + Send>> {
45                target_buffer.extend_from_slice(&self.to_be_bytes());
46                Ok(())
47            }
48        }
49    };
50}
51
52impl_number!(i16, PostgresType::INT2);
53impl_number!(i32, PostgresType::INT4, PostgresType::INT2);
54impl_number!(
55    i64,
56    PostgresType::INT8,
57    PostgresType::INT4,
58    PostgresType::INT2
59);
60impl_number!(f32, PostgresType::FLOAT4);
61impl_number!(f64, PostgresType::FLOAT8, PostgresType::FLOAT4);
62
63#[cfg(test)]
64mod tests {
65    #[cfg(feature = "tokio")]
66    mod tokio_connection {
67        use crate::test_helpers::get_settings;
68        use crate::tokio_connection::{new_client, TokioPoolableClient};
69        use crate::types::*;
70        use std::fmt::{Debug, Display};
71        use tokio::test;
72
73        struct DataReaderTest {
74            client: TokioPoolableClient,
75        }
76
77        impl DataReaderTest {
78            async fn new() -> Self {
79                let client = new_client(get_settings()).await.unwrap();
80                Self { client }
81            }
82
83            pub async fn test_read_special_cast<T>(&mut self, value: T, cast_to: &str)
84            where
85                T: FromSqlOwned + Display + PartialEq + Debug,
86            {
87                let sql = format!("select '{value}'::{cast_to}; ");
88
89                let received_value: T = self
90                    .client
91                    .read_single_column_and_row_exactly(sql.as_str(), &[])
92                    .await;
93
94                assert_eq!(received_value, value);
95
96                let prepared_query = self.client.prepare_query(&sql).await.unwrap();
97
98                let received_value: T = self
99                    .client
100                    .read_single_column_and_row_exactly(&prepared_query, &[])
101                    .await;
102
103                assert_eq!(received_value, value);
104            }
105
106            pub async fn test_read<T>(&mut self, value: T)
107            where
108                T: FromSqlOwned + Display + PartialEq + Debug + PostgresNamedType,
109            {
110                self.test_read_special_cast(value, T::PG_NAME).await
111            }
112
113            pub async fn test_round_trip<T>(&mut self, value: T)
114            where
115                T: FromSqlOwned + Display + PartialEq + Debug + ToSql + PostgresNamedType,
116            {
117                let sql = format!(
118                    "select t.f::{0} from (select b.f::text from (select $1::{0} as f) as b) as t ",
119                    T::PG_NAME
120                );
121
122                let received_value: T = self
123                    .client
124                    .read_single_column_and_row_exactly(sql.as_str(), &[&value])
125                    .await;
126
127                assert_eq!(received_value, value);
128            }
129        }
130
131        #[test]
132        async fn test_integer_types() {
133            let mut helper = DataReaderTest::new().await;
134
135            macro_rules! test_integer_values {
136                ($typ: ty) => {
137                    helper.test_read::<$typ>(1).await;
138                    helper.test_read(<$typ>::MAX).await;
139                    helper.test_read(<$typ>::MIN).await;
140                    helper.test_read::<$typ>(-1).await;
141                    helper.test_read::<$typ>(0).await;
142
143                    helper.test_round_trip::<$typ>(1).await;
144                    helper.test_round_trip::<$typ>(<$typ>::MAX).await;
145                    helper.test_round_trip::<$typ>(<$typ>::MIN).await;
146                    helper.test_round_trip::<$typ>(-1).await;
147                    helper.test_round_trip::<$typ>(0).await;
148                };
149            }
150
151            test_integer_values!(i16);
152            test_integer_values!(i32);
153            test_integer_values!(i64);
154
155            helper.test_read_special_cast(1i16, "smallint").await;
156        }
157
158        #[test]
159        async fn test_float_types() {
160            let mut helper = DataReaderTest::new().await;
161
162            macro_rules! test_float_values {
163                ($typ: ty) => {
164                    helper.test_read::<$typ>(1.0).await;
165                    helper.test_read::<$typ>(-1.0).await;
166                    helper.test_read::<$typ>(0.5).await;
167                    helper.test_read::<$typ>(-0.5).await;
168                    helper.test_read::<$typ>(0.0).await;
169                    helper.test_read::<$typ>(<$typ>::INFINITY).await;
170                    helper.test_read::<$typ>(<$typ>::NEG_INFINITY).await;
171
172                    helper.test_round_trip::<$typ>(1.0).await;
173                    helper.test_round_trip::<$typ>(-1.0).await;
174                    helper.test_round_trip::<$typ>(0.5).await;
175                    helper.test_round_trip::<$typ>(-0.5).await;
176                    helper.test_round_trip::<$typ>(0.0).await;
177                    helper.test_round_trip::<$typ>(<$typ>::INFINITY).await;
178                    helper.test_round_trip::<$typ>(<$typ>::NEG_INFINITY).await;
179
180                    let nan_text: $typ = helper
181                        .client
182                        .read_single_value_simple(&format!("select 'NaN'::{} ", <$typ>::PG_NAME))
183                        .await;
184                    assert!(nan_text.is_nan());
185
186                    let nan_binary: $typ = helper
187                        .client
188                        .read_single_value(&format!("select 'NaN'::{} ", <$typ>::PG_NAME), &[])
189                        .await;
190                    assert!(nan_binary.is_nan());
191
192                    let should_be_nan: $typ = helper
193                        .client
194                        .read_single_value(
195                            &format!("select $1::{} ", <$typ>::PG_NAME),
196                            &[&<$typ>::NAN],
197                        )
198                        .await;
199                    assert!(should_be_nan.is_nan());
200                };
201            }
202
203            test_float_values!(f32);
204            test_float_values!(f64);
205        }
206    }
207}