Skip to main content

elefant_client/types/
enum_type.rs

1use std::collections::HashMap;
2use std::error::Error;
3
4/// Trait implemented by Rust enums that map to PostgreSQL enum types.
5///
6/// Can be implemented manually or via the `#[derive(PostgresEnum)]` macro.
7pub trait PostgresEnum: Sized {
8    /// The PostgreSQL type name (e.g., "mood").
9    /// Schema-qualify if needed (e.g., "public.mood").
10    const PG_TYPE_NAME: &'static str;
11
12    /// Convert this variant to its PostgreSQL label string.
13    fn to_label(&self) -> &'static str;
14
15    /// Parse a PostgreSQL label string into a variant.
16    fn from_label(label: &str) -> Result<Self, Box<dyn Error + Sync + Send>>;
17}
18
19/// Registration data for an enum type, used when inserting into the registry.
20#[derive(Debug, Clone)]
21pub struct EnumOidEntry {
22    pub oid: i32,
23    pub array_oid: i32,
24    pub schema: String,
25}
26
27/// Runtime registry mapping PostgreSQL enum OIDs to their type names.
28///
29/// Keyed by OID for O(1) lookups on the hot path (`accepts_with_registry`).
30/// Supports multiple schemas having enums with the same name (different OIDs).
31#[derive(Debug, Clone, Default)]
32pub struct EnumTypeRegistry {
33    /// element OID -> list of type names this OID was registered under
34    by_oid: HashMap<i32, Vec<String>>,
35    /// array OID -> element OID
36    array_oids: HashMap<i32, i32>,
37}
38
39impl EnumTypeRegistry {
40    pub fn new() -> Self {
41        Self::default()
42    }
43
44    /// Insert an OID entry for the given type name.
45    pub fn insert(&mut self, name: String, entry: EnumOidEntry) {
46        self.by_oid.entry(entry.oid).or_default().push(name);
47        if entry.array_oid != 0 {
48            self.array_oids.insert(entry.array_oid, entry.oid);
49        }
50    }
51
52    /// Check if the given OID matches any entry for the named type.
53    pub fn has_oid_for_type(&self, name: &str, oid: i32) -> bool {
54        self.by_oid
55            .get(&oid)
56            .is_some_and(|names| names.iter().any(|n| n == name))
57    }
58
59    /// Check if the given OID matches any registered enum array OID.
60    pub fn is_enum_array_oid(&self, oid: i32) -> bool {
61        self.array_oids.contains_key(&oid)
62    }
63
64    /// Look up the element OID for a registered enum array OID.
65    pub fn element_oid_for_array_oid(&self, array_oid: i32) -> Option<i32> {
66        self.array_oids.get(&array_oid).copied()
67    }
68
69    pub fn is_empty(&self) -> bool {
70        self.by_oid.is_empty()
71    }
72}
73
74#[cfg(test)]
75mod tests {
76    use super::*;
77
78    #[test]
79    fn test_registry_has_oid_for_type() {
80        let mut registry = EnumTypeRegistry::new();
81        registry.insert(
82            "mood".to_string(),
83            EnumOidEntry {
84                oid: 16384,
85                array_oid: 16385,
86                schema: "public".to_string(),
87            },
88        );
89
90        assert!(registry.has_oid_for_type("mood", 16384));
91        assert!(!registry.has_oid_for_type("mood", 99999));
92        assert!(!registry.has_oid_for_type("other", 16384));
93    }
94
95    #[test]
96    fn test_registry_multiple_schemas() {
97        let mut registry = EnumTypeRegistry::new();
98        registry.insert(
99            "mood".to_string(),
100            EnumOidEntry {
101                oid: 16384,
102                array_oid: 16385,
103                schema: "public".to_string(),
104            },
105        );
106        registry.insert(
107            "mood".to_string(),
108            EnumOidEntry {
109                oid: 16400,
110                array_oid: 16401,
111                schema: "other_schema".to_string(),
112            },
113        );
114
115        assert!(registry.has_oid_for_type("mood", 16384));
116        assert!(registry.has_oid_for_type("mood", 16400));
117        assert!(!registry.has_oid_for_type("mood", 99999));
118    }
119
120    #[test]
121    fn test_registry_array_oid() {
122        let mut registry = EnumTypeRegistry::new();
123        registry.insert(
124            "mood".to_string(),
125            EnumOidEntry {
126                oid: 16384,
127                array_oid: 16385,
128                schema: "public".to_string(),
129            },
130        );
131
132        assert!(registry.is_enum_array_oid(16385));
133        assert!(!registry.is_enum_array_oid(16384));
134        assert!(!registry.is_enum_array_oid(99999));
135    }
136
137    #[cfg(all(feature = "tokio", feature = "derive"))]
138    mod tokio_connection {
139        use crate::test_helpers::get_settings;
140        use crate::tokio_connection::TokioConnectionFactory;
141        use crate::{PostgresEnum, PostgresPool};
142        use tokio::test;
143
144        #[derive(Debug, Clone, PartialEq, PostgresEnum)]
145        #[postgres(name = "mood")]
146        enum Mood {
147            Happy,
148            Sad,
149            Neutral,
150        }
151
152        #[derive(Debug, Clone, PartialEq, PostgresEnum)]
153        enum UserRole {
154            Admin,
155            #[postgres(label = "regular_user")]
156            Regular,
157            Guest,
158        }
159
160        async fn setup_enum_client() -> crate::pool::PoolableClient<TokioConnectionFactory> {
161            // First create the enum type using a plain connection
162            let mut plain_client = crate::tokio_connection::new_client(get_settings())
163                .await
164                .unwrap();
165            plain_client
166                .execute_non_query_simple(
167                    "DROP TYPE IF EXISTS mood CASCADE;
168                     CREATE TYPE mood AS ENUM ('happy', 'sad', 'neutral');
169                     DROP TYPE IF EXISTS user_role CASCADE;
170                     CREATE TYPE user_role AS ENUM ('admin', 'regular_user', 'guest');
171                     DROP TABLE IF EXISTS enum_test;
172                     CREATE TABLE enum_test (id serial PRIMARY KEY, m mood, r user_role);",
173                )
174                .await
175                .unwrap();
176            drop(plain_client);
177
178            // Now create a pool with registered enums
179            let settings = get_settings()
180                .register_enum::<Mood>()
181                .register_enum::<UserRole>();
182
183            let pool = PostgresPool::new(TokioConnectionFactory, settings)
184                .await
185                .unwrap();
186
187            // Verify registry was populated
188            assert!(
189                !pool.enum_registry().is_empty(),
190                "Enum registry should not be empty after registering enums"
191            );
192
193            pool.get_client().await.unwrap()
194        }
195
196        #[test]
197        async fn test_enum_text_mode() {
198            let mut client = setup_enum_client().await;
199
200            // Read enum via simple query (text mode)
201            let value: Mood = client
202                .read_single_value_simple("SELECT 'happy'::mood")
203                .await;
204            assert_eq!(value, Mood::Happy);
205
206            let value: Mood = client.read_single_value_simple("SELECT 'sad'::mood").await;
207            assert_eq!(value, Mood::Sad);
208        }
209
210        #[test]
211        async fn test_enum_binary_mode() {
212            let mut client = setup_enum_client().await;
213
214            // Read enum via prepared statement (binary mode)
215            let value: Mood = client
216                .read_single_value("SELECT 'neutral'::mood", &[])
217                .await;
218            assert_eq!(value, Mood::Neutral);
219        }
220
221        #[test]
222        async fn test_enum_as_parameter() {
223            let mut client = setup_enum_client().await;
224
225            client
226                .execute_non_query_simple(
227                    "DELETE FROM enum_test;
228                     INSERT INTO enum_test (m, r) VALUES ('happy', 'admin');",
229                )
230                .await
231                .unwrap();
232
233            // Send enum as parameter and read it back
234            client
235                .execute_non_query(
236                    "UPDATE enum_test SET m = $1 WHERE r = $2",
237                    &[&Mood::Sad, &UserRole::Admin],
238                )
239                .await
240                .unwrap();
241
242            let value: Mood = client
243                .read_single_value("SELECT m FROM enum_test WHERE r = 'admin'::user_role", &[])
244                .await;
245            assert_eq!(value, Mood::Sad);
246        }
247
248        #[test]
249        async fn test_enum_nullable() {
250            let mut client = setup_enum_client().await;
251
252            client
253                .execute_non_query_simple(
254                    "DELETE FROM enum_test;
255                     INSERT INTO enum_test (m, r) VALUES (NULL, 'guest');",
256                )
257                .await
258                .unwrap();
259
260            let value: Option<Mood> = client
261                .read_single_value("SELECT m FROM enum_test WHERE r = 'guest'::user_role", &[])
262                .await;
263            assert_eq!(value, None);
264
265            // Now with a non-null value
266            client
267                .execute_non_query_simple(
268                    "UPDATE enum_test SET m = 'happy' WHERE r = 'guest'::user_role",
269                )
270                .await
271                .unwrap();
272
273            let value: Option<Mood> = client
274                .read_single_value("SELECT m FROM enum_test WHERE r = 'guest'::user_role", &[])
275                .await;
276            assert_eq!(value, Some(Mood::Happy));
277        }
278
279        #[test]
280        async fn test_enum_array_text_mode() {
281            let mut client = setup_enum_client().await;
282
283            let value: Vec<Mood> = client
284                .read_single_value_simple("SELECT ARRAY['happy', 'sad', 'neutral']::mood[]")
285                .await;
286            assert_eq!(value, vec![Mood::Happy, Mood::Sad, Mood::Neutral]);
287        }
288
289        #[test]
290        async fn test_enum_array_binary_mode() {
291            let mut client = setup_enum_client().await;
292
293            let value: Vec<Mood> = client
294                .read_single_value("SELECT ARRAY['happy', 'sad']::mood[]", &[])
295                .await;
296            assert_eq!(value, vec![Mood::Happy, Mood::Sad]);
297        }
298
299        #[test]
300        async fn test_enum_empty_array() {
301            let mut client = setup_enum_client().await;
302
303            let value: Vec<Mood> = client
304                .read_single_value_simple("SELECT ARRAY[]::mood[]")
305                .await;
306            assert_eq!(value, Vec::<Mood>::new());
307
308            let value: Vec<Mood> = client
309                .read_single_value("SELECT ARRAY[]::mood[]", &[])
310                .await;
311            assert_eq!(value, Vec::<Mood>::new());
312        }
313
314        #[test]
315        async fn test_enum_custom_labels() {
316            let mut client = setup_enum_client().await;
317
318            let value: UserRole = client
319                .read_single_value_simple("SELECT 'regular_user'::user_role")
320                .await;
321            assert_eq!(value, UserRole::Regular);
322
323            let value: UserRole = client
324                .read_single_value("SELECT 'admin'::user_role", &[])
325                .await;
326            assert_eq!(value, UserRole::Admin);
327        }
328
329        #[test]
330        async fn test_pool_no_enums_no_query() {
331            // A pool with no registered enums should not query for OIDs
332            let settings = get_settings();
333            let pool = PostgresPool::new(TokioConnectionFactory, settings)
334                .await
335                .unwrap();
336            assert!(pool.enum_registry().is_empty());
337        }
338
339        #[derive(Debug, Clone, PartialEq, PostgresEnum)]
340        #[postgres(name = "nonexistent_enum_type_12345")]
341        enum NonexistentEnum {
342            A,
343        }
344
345        #[test]
346        async fn test_unregistered_enum_silently_skipped() {
347            // Register a name that doesn't exist in the database
348            let settings = get_settings().register_enum::<NonexistentEnum>();
349
350            // Should not error
351            let pool = PostgresPool::new(TokioConnectionFactory, settings)
352                .await
353                .unwrap();
354            // Registry will be empty since the type doesn't exist
355            assert!(pool.enum_registry().is_empty());
356        }
357    }
358}