1use crate::config::TypeOverride;
2use std::collections::HashMap;
3
4#[derive(Debug, Clone)]
5pub struct ResolvedType {
6 pub rust_type: String,
9 pub borrowed_rust_type: Option<String>,
17 pub copy_cheap: bool,
18}
19
20#[derive(Debug, Clone)]
23struct OverrideEntry {
24 owned: Option<String>,
26 borrowed: Option<String>,
29 copy_cheap: bool,
30}
31
32pub struct TypeMap {
33 defaults: HashMap<&'static str, (&'static str, bool)>,
34 type_overrides: HashMap<String, OverrideEntry>,
35 custom_types: HashMap<String, (String, bool)>,
36}
37
38impl TypeMap {
39 pub fn new(overrides: &[TypeOverride], copy_cheap_types: &[String]) -> Self {
40 let mut defaults: HashMap<&'static str, (&'static str, bool)> = HashMap::new();
41
42 for n in ["bool", "boolean", "pg_catalog.bool"] {
44 defaults.insert(n, ("bool", true));
45 }
46
47 for n in [
49 "int2",
50 "smallint",
51 "pg_catalog.int2",
52 "smallserial",
53 "serial2",
54 "pg_catalog.serial2",
55 ] {
56 defaults.insert(n, ("i16", true));
57 }
58 for n in [
59 "int4",
60 "integer",
61 "int",
62 "pg_catalog.int4",
63 "serial",
64 "serial4",
65 "pg_catalog.serial4",
66 ] {
67 defaults.insert(n, ("i32", true));
68 }
69 for n in [
70 "int8",
71 "bigint",
72 "pg_catalog.int8",
73 "bigserial",
74 "serial8",
75 "pg_catalog.serial8",
76 ] {
77 defaults.insert(n, ("i64", true));
78 }
79
80 for n in ["float4", "real", "pg_catalog.float4"] {
82 defaults.insert(n, ("f32", true));
83 }
84 for n in ["float8", "float", "double precision", "pg_catalog.float8"] {
85 defaults.insert(n, ("f64", true));
86 }
87
88 for n in ["numeric", "decimal", "pg_catalog.numeric"] {
90 defaults.insert(n, ("bigdecimal::BigDecimal", false));
91 }
92
93 for n in [
95 "text",
96 "varchar",
97 "pg_catalog.varchar",
98 "pg_catalog.bpchar",
99 "bpchar",
100 "string",
101 "citext",
102 "name",
103 "pg_catalog.name",
104 ] {
105 defaults.insert(n, ("String", false));
106 }
107
108 for n in ["bytea", "blob", "pg_catalog.bytea"] {
110 defaults.insert(n, ("Vec<u8>", false));
111 }
112
113 defaults.insert("uuid", ("uuid::Uuid", true));
115
116 for n in ["json", "jsonb"] {
118 defaults.insert(n, ("serde_json::Value", false));
119 }
120
121 for n in [
123 "timestamptz",
124 "pg_catalog.timestamptz",
125 "timestamp with time zone",
126 ] {
127 defaults.insert(n, ("chrono::DateTime<chrono::Utc>", false));
128 }
129 for n in [
130 "timestamp",
131 "pg_catalog.timestamp",
132 "timestamp without time zone",
133 ] {
134 defaults.insert(n, ("chrono::NaiveDateTime", false));
135 }
136 defaults.insert("date", ("chrono::NaiveDate", true));
137 for n in ["time", "pg_catalog.time", "time without time zone"] {
138 defaults.insert(n, ("chrono::NaiveTime", false));
139 }
140
141 for n in ["inet", "cidr"] {
143 defaults.insert(n, ("ipnetwork::IpNetwork", false));
144 }
145 defaults.insert("macaddr", ("mac_address::MacAddress", true));
146
147 defaults.insert(
149 "hstore",
150 ("std::collections::HashMap<String, Option<String>>", false),
151 );
152 for n in ["interval", "pg_catalog.interval"] {
153 defaults.insert(n, ("sqlx::postgres::types::PgInterval", false));
154 }
155 defaults.insert("money", ("sqlx::postgres::types::PgMoney", true));
156 defaults.insert("oid", ("sqlx::postgres::types::Oid", true));
157 defaults.insert("pg_catalog.oid", ("sqlx::postgres::types::Oid", true));
158 for n in ["ltree", "lquery"] {
159 defaults.insert(n, ("String", false));
160 }
161
162 for n in ["int4range", "pg_catalog.int4range"] {
164 defaults.insert(n, ("sqlx::postgres::types::PgRange<i32>", false));
165 }
166 for n in ["int8range", "pg_catalog.int8range"] {
167 defaults.insert(n, ("sqlx::postgres::types::PgRange<i64>", false));
168 }
169 for n in ["numrange", "pg_catalog.numrange"] {
170 defaults.insert(
171 n,
172 (
173 "sqlx::postgres::types::PgRange<bigdecimal::BigDecimal>",
174 false,
175 ),
176 );
177 }
178 for n in ["tsrange", "pg_catalog.tsrange"] {
179 defaults.insert(
180 n,
181 (
182 "sqlx::postgres::types::PgRange<chrono::NaiveDateTime>",
183 false,
184 ),
185 );
186 }
187 for n in ["tstzrange", "pg_catalog.tstzrange"] {
188 defaults.insert(
189 n,
190 (
191 "sqlx::postgres::types::PgRange<chrono::DateTime<chrono::Utc>>",
192 false,
193 ),
194 );
195 }
196 for n in ["daterange", "pg_catalog.daterange"] {
197 defaults.insert(
198 n,
199 ("sqlx::postgres::types::PgRange<chrono::NaiveDate>", false),
200 );
201 }
202
203 for n in ["bit", "varbit", "pg_catalog.varbit"] {
205 defaults.insert(n, ("bit_vec::BitVec", false));
206 }
207
208 let mut type_overrides: HashMap<String, OverrideEntry> = HashMap::new();
209 for o in overrides {
210 if let Some(db_type) = &o.db_type {
211 type_overrides.insert(
212 db_type.to_lowercase(),
213 OverrideEntry {
214 owned: o.rs_type.clone(),
215 borrowed: o.borrowed_rs_type.clone(),
216 copy_cheap: o.copy_cheap,
217 },
218 );
219 }
220 }
221
222 for name in copy_cheap_types {
223 let key = name.to_lowercase();
224 if let Some(ovr) = type_overrides.get_mut(&key) {
225 ovr.copy_cheap = true;
226 } else if defaults.contains_key(key.as_str()) {
227 type_overrides.insert(
228 key,
229 OverrideEntry {
230 owned: None,
231 borrowed: None,
232 copy_cheap: true,
233 },
234 );
235 }
236 }
237
238 Self {
239 defaults,
240 type_overrides,
241 custom_types: HashMap::new(),
242 }
243 }
244
245 pub fn register(&mut self, pg_name: &str, rust_name: &str, copy_cheap: bool) {
249 self.custom_types
250 .insert(pg_name.to_lowercase(), (rust_name.to_string(), copy_cheap));
251 }
252
253 pub fn resolve_pg_type(
254 &self,
255 pg_type: &str,
256 nullable: bool,
257 is_array: bool,
258 ) -> Option<ResolvedType> {
259 self.resolve_pg_type_dims(pg_type, nullable, usize::from(is_array))
260 }
261
262 pub fn resolve_pg_type_dims(
263 &self,
264 pg_type: &str,
265 nullable: bool,
266 array_dims: usize,
267 ) -> Option<ResolvedType> {
268 let key = pg_type.to_lowercase();
269 let (owned_inner, borrowed_inner, copy_cheap) =
270 if let Some(ovr) = self.type_overrides.get(&key) {
271 let default = self.defaults.get(key.as_str()).map(|&(t, _)| t.to_string());
272 let owned = ovr
273 .owned
274 .clone()
275 .or(default)
276 .or_else(|| self.custom_types.get(&key).map(|(name, _)| name.clone()));
277 (owned, ovr.borrowed.clone(), ovr.copy_cheap)
278 } else if let Some(&(ty, cc)) = self.defaults.get(key.as_str()) {
279 (Some(ty.to_string()), None, cc)
280 } else if let Some((name, cc)) = self.custom_types.get(&key) {
281 (Some(name.clone()), None, *cc)
282 } else {
283 return None;
284 };
285
286 let owned = wrap_owned(owned_inner.as_deref()?, nullable, array_dims);
287 let borrowed = borrowed_inner.map(|inner| {
288 wrap_borrowed(
291 &inner,
292 owned_inner.as_deref().unwrap_or(""),
293 nullable,
294 array_dims,
295 )
296 });
297 let effective_copy_cheap = copy_cheap && !nullable && array_dims == 0;
298 Some(ResolvedType {
299 rust_type: owned,
300 borrowed_rust_type: borrowed,
301 copy_cheap: effective_copy_cheap,
302 })
303 }
304
305 pub fn resolve_column(
306 &self,
307 pg_type: &str,
308 nullable: bool,
309 is_array: bool,
310 column_key: Option<&str>,
311 column_overrides: &HashMap<String, ColumnOverride>,
312 ) -> Option<ResolvedType> {
313 self.resolve_column_dims(
314 pg_type,
315 nullable,
316 usize::from(is_array),
317 column_key,
318 column_overrides,
319 )
320 }
321
322 pub fn resolve_column_dims(
323 &self,
324 pg_type: &str,
325 nullable: bool,
326 array_dims: usize,
327 column_key: Option<&str>,
328 column_overrides: &HashMap<String, ColumnOverride>,
329 ) -> Option<ResolvedType> {
330 if let Some(key) = column_key
331 && let Some(ovr) = column_overrides.get(key)
332 {
333 let owned_inner = if let Some(owned) = &ovr.owned {
337 owned.clone()
338 } else if let Some(resolved) = self.resolve_pg_type_dims(pg_type, false, 0) {
339 resolved.rust_type
340 } else {
341 return None;
342 };
343 let owned = wrap_owned(&owned_inner, nullable, array_dims);
344 let borrowed = ovr
345 .borrowed
346 .as_ref()
347 .map(|b| wrap_borrowed(b, &owned_inner, nullable, array_dims));
348 let cc = ovr.copy_cheap && !nullable && array_dims == 0;
349 return Some(ResolvedType {
350 rust_type: owned,
351 borrowed_rust_type: borrowed,
352 copy_cheap: cc,
353 });
354 }
355 self.resolve_pg_type_dims(pg_type, nullable, array_dims)
356 }
357}
358
359#[derive(Debug, Clone)]
361pub struct ColumnOverride {
362 pub owned: Option<String>,
363 pub borrowed: Option<String>,
364 pub copy_cheap: bool,
365}
366
367impl ColumnOverride {
368 #[cfg(test)]
369 pub fn owned_form(rust_type: impl Into<String>, copy_cheap: bool) -> Self {
370 Self {
371 owned: Some(rust_type.into()),
372 borrowed: None,
373 copy_cheap,
374 }
375 }
376}
377
378fn wrap_owned(inner: &str, nullable: bool, array_dims: usize) -> String {
379 let mut t = inner.to_string();
380 for _ in 0..array_dims {
381 t = format!("Vec<{t}>");
382 }
383 if nullable { format!("Option<{t}>") } else { t }
384}
385
386fn wrap_borrowed(
387 borrowed_inner: &str,
388 owned_inner: &str,
389 nullable: bool,
390 array_dims: usize,
391) -> String {
392 let body = if array_dims == 0 {
393 borrowed_inner.to_string()
394 } else {
395 let mut t = owned_inner.to_string();
396 for _ in 0..(array_dims - 1) {
398 t = format!("Vec<{t}>");
399 }
400 format!("&[{t}]")
402 };
403 if nullable {
404 format!("Option<{body}>")
405 } else {
406 body
407 }
408}
409
410pub fn build_column_overrides(overrides: &[TypeOverride]) -> HashMap<String, ColumnOverride> {
411 overrides
412 .iter()
413 .filter_map(|o| {
414 o.column.as_ref().map(|col| {
415 (
416 col.clone(),
417 ColumnOverride {
418 owned: o.rs_type.clone(),
419 borrowed: o.borrowed_rs_type.clone(),
420 copy_cheap: o.copy_cheap,
421 },
422 )
423 })
424 })
425 .collect()
426}
427
428#[cfg(test)]
429mod tests {
430 use super::*;
431
432 fn map() -> TypeMap {
433 TypeMap::new(&[], &[])
434 }
435
436 fn owned_override(db_type: &str, rs_type: &str) -> TypeOverride {
437 TypeOverride {
438 db_type: Some(db_type.to_string()),
439 column: None,
440 rs_type: Some(rs_type.to_string()),
441 borrowed_rs_type: None,
442 copy_cheap: false,
443 }
444 }
445
446 #[test]
447 fn maps_text() {
448 let t = map().resolve_pg_type("text", false, false).unwrap();
449 assert_eq!(t.rust_type, "String");
450 assert!(t.borrowed_rust_type.is_none());
451 assert!(!t.copy_cheap);
452 }
453 #[test]
454 fn maps_int4_copy_cheap() {
455 let t = map().resolve_pg_type("int4", false, false).unwrap();
456 assert_eq!(t.rust_type, "i32");
457 assert!(t.copy_cheap);
458 }
459 #[test]
460 fn maps_bool() {
461 let t = map().resolve_pg_type("bool", false, false).unwrap();
462 assert_eq!(t.rust_type, "bool");
463 assert!(t.copy_cheap);
464 }
465 #[test]
466 fn maps_timestamptz() {
467 let t = map().resolve_pg_type("timestamptz", false, false).unwrap();
468 assert_eq!(t.rust_type, "chrono::DateTime<chrono::Utc>");
469 }
470 #[test]
471 fn maps_uuid() {
472 let t = map().resolve_pg_type("uuid", false, false).unwrap();
473 assert_eq!(t.rust_type, "uuid::Uuid");
474 assert!(t.copy_cheap);
475 }
476 #[test]
477 fn maps_jsonb() {
478 let t = map().resolve_pg_type("jsonb", false, false).unwrap();
479 assert_eq!(t.rust_type, "serde_json::Value");
480 }
481 #[test]
482 fn nullable_wraps_option() {
483 let t = map().resolve_pg_type("text", true, false).unwrap();
484 assert_eq!(t.rust_type, "Option<String>");
485 assert!(!t.copy_cheap);
486 }
487 #[test]
488 fn array_wraps_vec() {
489 let t = map().resolve_pg_type("text", false, true).unwrap();
490 assert_eq!(t.rust_type, "Vec<String>");
491 assert!(!t.copy_cheap);
492 }
493 #[test]
494 fn nullable_array() {
495 let t = map().resolve_pg_type("text", true, true).unwrap();
496 assert_eq!(t.rust_type, "Option<Vec<String>>");
497 }
498 #[test]
499 fn multidimensional_array_wraps_nested_vec() {
500 let t = map().resolve_pg_type_dims("int8", false, 2).unwrap();
501 assert_eq!(t.rust_type, "Vec<Vec<i64>>");
502 }
503 #[test]
504 fn nullable_multidimensional_array_wraps_option_nested_vec() {
505 let t = map().resolve_pg_type_dims("text", true, 3).unwrap();
506 assert_eq!(t.rust_type, "Option<Vec<Vec<Vec<String>>>>");
507 }
508 #[test]
509 fn type_override_replaces_default() {
510 let t = TypeMap::new(
511 &[owned_override("timestamptz", "time::OffsetDateTime")],
512 &[],
513 )
514 .resolve_pg_type("timestamptz", false, false)
515 .unwrap();
516 assert_eq!(t.rust_type, "time::OffsetDateTime");
517 assert!(t.borrowed_rust_type.is_none());
518 }
519 #[test]
520 fn column_override_beats_type_override() {
521 let overrides = vec![
522 owned_override("text", "TypeLevel"),
523 TypeOverride {
524 db_type: None,
525 column: Some("users.name".to_string()),
526 rs_type: Some("ColumnLevel".to_string()),
527 borrowed_rs_type: None,
528 copy_cheap: false,
529 },
530 ];
531 let col_ovrs = build_column_overrides(&overrides);
532 let map = TypeMap::new(&overrides, &[]);
533 let t = map
534 .resolve_column("text", false, false, Some("users.name"), &col_ovrs)
535 .unwrap();
536 assert_eq!(t.rust_type, "ColumnLevel");
537 }
538 #[test]
539 fn maps_numeric() {
540 let t = map().resolve_pg_type("numeric", false, false).unwrap();
541 assert_eq!(t.rust_type, "bigdecimal::BigDecimal");
542 assert!(!t.copy_cheap);
543 }
544 #[test]
545 fn maps_decimal() {
546 let t = map().resolve_pg_type("decimal", false, false).unwrap();
547 assert_eq!(t.rust_type, "bigdecimal::BigDecimal");
548 }
549 #[test]
550 fn maps_pg_catalog_numeric() {
551 let t = map()
552 .resolve_pg_type("pg_catalog.numeric", false, false)
553 .unwrap();
554 assert_eq!(t.rust_type, "bigdecimal::BigDecimal");
555 }
556 #[test]
557 fn unknown_type_returns_none() {
558 assert!(
559 map()
560 .resolve_pg_type("no_such_type", false, false)
561 .is_none()
562 );
563 }
564 #[test]
565 fn registers_custom_type() {
566 let mut map = TypeMap::new(&[], &[]);
567 map.register("my_enum", "MyEnum", false);
568 let t = map.resolve_pg_type("my_enum", false, false).unwrap();
569 assert_eq!(t.rust_type, "MyEnum");
570 assert!(!t.copy_cheap);
571 }
572 #[test]
573 fn registered_type_nullable() {
574 let mut map = TypeMap::new(&[], &[]);
575 map.register("my_enum", "MyEnum", false);
576 let t = map.resolve_pg_type("my_enum", true, false).unwrap();
577 assert_eq!(t.rust_type, "Option<MyEnum>");
578 assert!(!t.copy_cheap);
579 }
580 #[test]
581 fn type_override_beats_registered_custom() {
582 let mut map = TypeMap::new(&[owned_override("my_enum", "Override")], &[]);
583 map.register("my_enum", "MyEnum", false);
584 let t = map.resolve_pg_type("my_enum", false, false).unwrap();
585 assert_eq!(t.rust_type, "Override");
587 }
588 #[test]
589 fn registered_copy_cheap_type_is_cheap() {
590 let mut map = TypeMap::new(&[], &[]);
591 map.register("my_value_type", "MyValueType", true);
592 let t = map.resolve_pg_type("my_value_type", false, false).unwrap();
593 assert_eq!(t.rust_type, "MyValueType");
594 assert!(
595 t.copy_cheap,
596 "registered type with copy_cheap=true should resolve as copy_cheap"
597 );
598 }
599
600 #[test]
601 fn registered_copy_cheap_nullable_is_not_cheap() {
602 let mut map = TypeMap::new(&[], &[]);
603 map.register("my_value_type", "MyValueType", true);
604 let t = map.resolve_pg_type("my_value_type", true, false).unwrap();
605 assert_eq!(t.rust_type, "Option<MyValueType>");
606 assert!(
607 !t.copy_cheap,
608 "nullable type should not be copy_cheap even if base is"
609 );
610 }
611
612 #[test]
613 fn registered_copy_cheap_array_is_not_cheap() {
614 let mut map = TypeMap::new(&[], &[]);
615 map.register("my_value_type", "MyValueType", true);
616 let t = map.resolve_pg_type("my_value_type", false, true).unwrap();
617 assert_eq!(t.rust_type, "Vec<MyValueType>");
618 assert!(
619 !t.copy_cheap,
620 "array type should not be copy_cheap even if base is"
621 );
622 }
623
624 #[test]
625 fn copy_cheap_types_marks_type_as_cheap() {
626 let map = TypeMap::new(&[], &["text".to_string()]);
627 let t = map.resolve_pg_type("text", false, false).unwrap();
629 assert_eq!(t.rust_type, "String");
630 assert!(
631 t.copy_cheap,
632 "text should be copy_cheap after config promotion"
633 );
634 }
635
636 #[test]
637 fn copy_cheap_types_promotes_default_to_override() {
638 let map = TypeMap::new(&[], &["uuid".to_string()]);
640 let t = map.resolve_pg_type("uuid", false, false).unwrap();
641 assert!(t.copy_cheap);
642 }
643
644 fn borrowed_text() -> TypeOverride {
647 TypeOverride {
648 db_type: Some("text".to_string()),
649 column: None,
650 rs_type: None,
651 borrowed_rs_type: Some("&str".to_string()),
652 copy_cheap: false,
653 }
654 }
655
656 #[test]
657 fn borrowed_only_override_keeps_default_for_owned_form() {
658 let map = TypeMap::new(&[borrowed_text()], &[]);
659 let t = map.resolve_pg_type("text", false, false).unwrap();
660 assert_eq!(t.rust_type, "String");
661 assert_eq!(t.borrowed_rust_type.as_deref(), Some("&str"));
662 }
663
664 #[test]
665 fn borrowed_nullable_wraps_inside_option() {
666 let map = TypeMap::new(&[borrowed_text()], &[]);
667 let t = map.resolve_pg_type("text", true, false).unwrap();
668 assert_eq!(t.rust_type, "Option<String>");
669 assert_eq!(t.borrowed_rust_type.as_deref(), Some("Option<&str>"));
670 }
671
672 #[test]
673 fn borrowed_array_uses_owned_inner_in_slice() {
674 let map = TypeMap::new(&[borrowed_text()], &[]);
675 let t = map.resolve_pg_type("text", false, true).unwrap();
676 assert_eq!(t.rust_type, "Vec<String>");
677 assert_eq!(t.borrowed_rust_type.as_deref(), Some("&[String]"));
678 }
679
680 #[test]
681 fn borrowed_nullable_array() {
682 let map = TypeMap::new(&[borrowed_text()], &[]);
683 let t = map.resolve_pg_type("text", true, true).unwrap();
684 assert_eq!(t.rust_type, "Option<Vec<String>>");
685 assert_eq!(t.borrowed_rust_type.as_deref(), Some("Option<&[String]>"));
686 }
687
688 #[test]
689 fn borrowed_multidim_array_inner_stays_vec() {
690 let map = TypeMap::new(&[borrowed_text()], &[]);
691 let t = map.resolve_pg_type_dims("text", false, 2).unwrap();
692 assert_eq!(t.rust_type, "Vec<Vec<String>>");
693 assert_eq!(t.borrowed_rust_type.as_deref(), Some("&[Vec<String>]"));
694 }
695
696 #[test]
697 fn owned_and_borrowed_pair() {
698 let ovr = TypeOverride {
699 db_type: Some("text".to_string()),
700 column: None,
701 rs_type: Some("MyStr".to_string()),
702 borrowed_rs_type: Some("&MyStr".to_string()),
703 copy_cheap: false,
704 };
705 let map = TypeMap::new(&[ovr], &[]);
706 let t = map.resolve_pg_type("text", false, false).unwrap();
707 assert_eq!(t.rust_type, "MyStr");
708 assert_eq!(t.borrowed_rust_type.as_deref(), Some("&MyStr"));
709
710 let t = map.resolve_pg_type("text", false, true).unwrap();
711 assert_eq!(t.rust_type, "Vec<MyStr>");
713 assert_eq!(t.borrowed_rust_type.as_deref(), Some("&[MyStr]"));
714 }
715
716 #[test]
717 fn borrowed_column_override() {
718 let overrides = vec![TypeOverride {
719 db_type: None,
720 column: Some("users.name".to_string()),
721 rs_type: None,
722 borrowed_rs_type: Some("&str".to_string()),
723 copy_cheap: false,
724 }];
725 let col_ovrs = build_column_overrides(&overrides);
726 let map = TypeMap::new(&overrides, &[]);
727 let t = map
728 .resolve_column("text", false, false, Some("users.name"), &col_ovrs)
729 .unwrap();
730 assert_eq!(t.rust_type, "String");
731 assert_eq!(t.borrowed_rust_type.as_deref(), Some("&str"));
732 }
733}