Skip to main content

elefant_client/types/
collections.rs

1use crate::protocol::FieldDescription;
2use crate::types::{EnumTypeRegistry, FromSqlBase, FromSqlBinary, FromSqlText};
3use crate::PostgresType;
4use std::error::Error;
5
6impl<'a, T> FromSqlBase<'a> for Vec<T>
7where
8    T: FromSqlBase<'a>,
9{
10    fn accepts_postgres_type(oid: i32) -> bool {
11        match PostgresType::get_by_oid(oid) {
12            None => false,
13            Some(t) => {
14                if !t.is_array {
15                    return false;
16                }
17
18                match t.element {
19                    None => false,
20                    Some(element_type) => T::accepts_postgres_type(element_type.oid),
21                }
22            }
23        }
24    }
25
26    fn accepts_with_registry(field: &FieldDescription, registry: &EnumTypeRegistry) -> bool {
27        // Static check for built-in array types
28        if Self::accepts(field) {
29            return true;
30        }
31        // Dynamic check: verify this is a registered enum array AND the element type matches T
32        if let Some(element_oid) = registry.element_oid_for_array_oid(field.data_type_oid) {
33            let element_field = FieldDescription {
34                data_type_oid: element_oid,
35                ..field.clone()
36            };
37            T::accepts_with_registry(&element_field, registry)
38        } else {
39            false
40        }
41    }
42}
43
44impl<'a, T> FromSqlBinary<'a> for Vec<T>
45where
46    T: FromSqlBinary<'a>,
47{
48    fn from_sql_binary(
49        raw: &'a [u8],
50        field: &FieldDescription,
51    ) -> Result<Self, Box<dyn Error + Sync + Send>> {
52        if raw.len() < 12 {
53            return Err(format!("Invalid length for array. Expected at least 12 bytes, got {} bytes instead. Error occurred when parsing field {:?}", raw.len(), field).into());
54        }
55        let dimensions = i32::from_be_bytes(raw[0..4].try_into().unwrap());
56
57        let has_null_bit_map = i32::from_be_bytes(raw[4..8].try_into().unwrap()) == 1;
58        let element_oid = i32::from_be_bytes(raw[8..12].try_into().unwrap());
59
60        if dimensions == 0 {
61            return Ok(Vec::new());
62        }
63
64        if raw.len() < 20 {
65            return Err(format!("Invalid length for non-empty array. Expected at least 20 bytes, got {} bytes instead. Error occurred when parsing field {:?}", raw.len(), field).into());
66        }
67
68        let size_of_first_dimension = i32::from_be_bytes(raw[12..16].try_into().unwrap());
69        // let start_index_of_first_dimension = i32::from_be_bytes(raw[16..20].try_into().unwrap());
70        let raw_data = &raw[20..];
71
72        if dimensions != 1 {
73            return Err(format!("Only one-dimensional arrays are supported. Error occurred when parsing field {field:?}").into());
74        }
75
76        let _ = element_oid; // Validated at the outer accepts_with_registry level
77
78        let mut result: Vec<T> = Vec::with_capacity(size_of_first_dimension as usize);
79
80        let mut cursor = 0;
81        for _ in 0..size_of_first_dimension {
82            let element_size = i32::from_be_bytes(raw_data[cursor..cursor + 4].try_into().unwrap());
83            cursor += 4;
84            if has_null_bit_map && element_size == -1 {
85                result.push(
86                    T::from_null(field).map_err(|e| format!("Error handling null element: {e}"))?,
87                );
88            } else {
89                let element_raw = &raw_data[cursor..cursor + element_size as usize];
90                cursor += element_size as usize;
91                result.push(T::from_sql_binary(element_raw, field)?);
92            }
93        }
94
95        Ok(result)
96    }
97}
98
99impl<'a, T> FromSqlText<'a> for Vec<T>
100where
101    T: FromSqlText<'a>,
102{
103    fn from_sql_text(
104        raw: &'a str,
105        field: &FieldDescription,
106    ) -> Result<Self, Box<dyn Error + Sync + Send>> {
107        // For enum arrays (unknown OID), default delimiter is ','
108        let delimiter_char = match PostgresType::get_by_oid(field.data_type_oid) {
109            Some(typ) => typ.array_delimiter,
110            None => ',',
111        };
112
113        let mut result = Vec::new();
114
115        // Handle explicit array bounds prefix: [lower:upper]={elements}
116        // This occurs e.g. when casting int2vector to int2[] (0-based indexing).
117        let array_body = if raw.starts_with('[') {
118            match raw.find("={") {
119                Some(pos) => &raw[pos + 1..],
120                None => raw,
121            }
122        } else {
123            raw
124        };
125
126        let narrowed = &array_body[1..array_body.len() - 1];
127
128        if narrowed.is_empty() {
129            return Ok(result);
130        }
131
132        // Parse array elements while respecting quoted boundaries
133        let mut element_start = 0;
134        let mut in_quotes = false;
135        let bytes = narrowed.as_bytes();
136
137        for (i, &byte) in bytes.iter().enumerate() {
138            let ch = byte as char;
139            match ch {
140                '"' => {
141                    in_quotes = !in_quotes;
142                }
143                c if c == delimiter_char && !in_quotes => {
144                    // End of current element
145                    if i > element_start {
146                        let element = &narrowed[element_start..i];
147                        // Remove quotes if present
148                        let clean_element = if element.starts_with('"')
149                            && element.ends_with('"')
150                            && element.len() >= 2
151                        {
152                            &element[1..element.len() - 1]
153                        } else {
154                            element
155                        };
156
157                        if clean_element == "NULL" {
158                            result.push(
159                                T::from_null(field)
160                                    .map_err(|e| format!("Error handling null element: {e}"))?,
161                            );
162                        } else {
163                            result.push(T::from_sql_text(clean_element, field)?);
164                        }
165                    }
166                    element_start = i + 1;
167                }
168                _ => {}
169            }
170        }
171
172        // Handle the last element
173        if element_start < narrowed.len() {
174            let element = &narrowed[element_start..];
175            // Remove quotes if present
176            let clean_element =
177                if element.starts_with('"') && element.ends_with('"') && element.len() >= 2 {
178                    &element[1..element.len() - 1]
179                } else {
180                    element
181                };
182
183            if clean_element == "NULL" {
184                result.push(
185                    T::from_null(field).map_err(|e| format!("Error handling null element: {e}"))?,
186                );
187            } else {
188                result.push(T::from_sql_text(clean_element, field)?);
189            }
190        }
191
192        Ok(result)
193    }
194}
195
196#[cfg(test)]
197mod tests {
198    #[cfg(feature = "tokio")]
199    mod tokio_connection {
200        use crate::test_helpers::get_settings;
201        use crate::tokio_connection::new_client;
202        use tokio::test;
203
204        #[test]
205        async fn test_array_types() {
206            let mut client = new_client(get_settings()).await.unwrap();
207
208            client
209                .execute_non_query_simple(
210                    r#"
211                drop table if exists test_array_table;
212                create table test_array_table(value int2[]);
213                "#,
214                )
215                .await
216                .unwrap();
217
218            let prepared = client
219                .prepare_query("select value from test_array_table;")
220                .await
221                .unwrap();
222
223            client
224                .execute_non_query("insert into test_array_table values ('{1,2,3}');", &[])
225                .await
226                .unwrap();
227
228            let mut value: Vec<i16> = client
229                .read_single_value("select value from test_array_table;", &[])
230                .await;
231            assert_eq!(value, vec![1, 2, 3]);
232            value = client.read_single_value(&prepared, &[]).await;
233            assert_eq!(value, vec![1, 2, 3]);
234
235            client
236                .execute_non_query("update test_array_table set value = '{}'", &[])
237                .await
238                .unwrap();
239
240            value = client
241                .read_single_value("select value from test_array_table;", &[])
242                .await;
243            assert_eq!(value, Vec::<i16>::new());
244            value = client.read_single_value(&prepared, &[]).await;
245            assert_eq!(value, Vec::<i16>::new());
246
247            client
248                .execute_non_query("update test_array_table set value = '{1,null,3}'", &[])
249                .await
250                .unwrap();
251
252            let mut value: Vec<Option<i16>> = client
253                .read_single_value("select value from test_array_table;", &[])
254                .await;
255            assert_eq!(value, vec![Some(1), None, Some(3)]);
256            value = client.read_single_value(&prepared, &[]).await;
257            assert_eq!(value, vec![Some(1), None, Some(3)]);
258
259            client
260                .execute_non_query("update test_array_table set value = '{null}'", &[])
261                .await
262                .unwrap();
263            let mut value: Vec<Option<i16>> = client
264                .read_single_value("select value from test_array_table;", &[])
265                .await;
266            assert_eq!(value, vec![None]);
267            value = client.read_single_value(&prepared, &[]).await;
268            assert_eq!(value, vec![None]);
269        }
270    }
271}