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, ("sqlx::types::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 {
339 let resolved = self.resolve_pg_type_dims(pg_type, false, 0)?;
340 resolved.rust_type
341 };
342 let owned = wrap_owned(&owned_inner, nullable, array_dims);
343 let borrowed = ovr
344 .borrowed
345 .as_ref()
346 .map(|b| wrap_borrowed(b, &owned_inner, nullable, array_dims));
347 let cc = ovr.copy_cheap && !nullable && array_dims == 0;
348 return Some(ResolvedType {
349 rust_type: owned,
350 borrowed_rust_type: borrowed,
351 copy_cheap: cc,
352 });
353 }
354 self.resolve_pg_type_dims(pg_type, nullable, array_dims)
355 }
356}
357
358#[derive(Debug, Clone)]
360pub struct ColumnOverride {
361 pub owned: Option<String>,
362 pub borrowed: Option<String>,
363 pub copy_cheap: bool,
364}
365
366impl ColumnOverride {
367 #[cfg(test)]
368 pub fn owned_form(rust_type: impl Into<String>, copy_cheap: bool) -> Self {
369 Self {
370 owned: Some(rust_type.into()),
371 borrowed: None,
372 copy_cheap,
373 }
374 }
375}
376
377fn wrap_owned(inner: &str, nullable: bool, array_dims: usize) -> String {
378 let mut t = inner.to_string();
379 for _ in 0..array_dims {
380 t = format!("Vec<{t}>");
381 }
382 if nullable { format!("Option<{t}>") } else { t }
383}
384
385fn wrap_borrowed(
386 borrowed_inner: &str,
387 owned_inner: &str,
388 nullable: bool,
389 array_dims: usize,
390) -> String {
391 let body = if array_dims == 0 {
392 borrowed_inner.to_string()
393 } else {
394 let mut t = owned_inner.to_string();
395 for _ in 0..(array_dims - 1) {
397 t = format!("Vec<{t}>");
398 }
399 format!("&[{t}]")
401 };
402 if nullable {
403 format!("Option<{body}>")
404 } else {
405 body
406 }
407}
408
409pub fn build_column_overrides(overrides: &[TypeOverride]) -> HashMap<String, ColumnOverride> {
410 overrides
411 .iter()
412 .filter_map(|o| {
413 o.column.as_ref().map(|col| {
414 (
415 col.clone(),
416 ColumnOverride {
417 owned: o.rs_type.clone(),
418 borrowed: o.borrowed_rs_type.clone(),
419 copy_cheap: o.copy_cheap,
420 },
421 )
422 })
423 })
424 .collect()
425}
426
427#[cfg(test)]
428mod tests {
429 use super::*;
430
431 fn map() -> TypeMap {
432 TypeMap::new(&[], &[])
433 }
434
435 fn owned_override(db_type: &str, rs_type: &str) -> TypeOverride {
436 TypeOverride {
437 db_type: Some(db_type.to_string()),
438 column: None,
439 rs_type: Some(rs_type.to_string()),
440 borrowed_rs_type: None,
441 copy_cheap: false,
442 }
443 }
444
445 #[test]
446 fn maps_text() {
447 let t = map().resolve_pg_type("text", false, false).unwrap();
448 assert_eq!(t.rust_type, "String");
449 assert!(t.borrowed_rust_type.is_none());
450 assert!(!t.copy_cheap);
451 }
452 #[test]
453 fn maps_int4_copy_cheap() {
454 let t = map().resolve_pg_type("int4", false, false).unwrap();
455 assert_eq!(t.rust_type, "i32");
456 assert!(t.copy_cheap);
457 }
458 #[test]
459 fn maps_bool() {
460 let t = map().resolve_pg_type("bool", false, false).unwrap();
461 assert_eq!(t.rust_type, "bool");
462 assert!(t.copy_cheap);
463 }
464 #[test]
465 fn maps_timestamptz() {
466 let t = map().resolve_pg_type("timestamptz", false, false).unwrap();
467 assert_eq!(t.rust_type, "chrono::DateTime<chrono::Utc>");
468 }
469 #[test]
470 fn maps_uuid() {
471 let t = map().resolve_pg_type("uuid", false, false).unwrap();
472 assert_eq!(t.rust_type, "uuid::Uuid");
473 assert!(t.copy_cheap);
474 }
475 #[test]
476 fn maps_jsonb() {
477 let t = map().resolve_pg_type("jsonb", false, false).unwrap();
478 assert_eq!(t.rust_type, "serde_json::Value");
479 }
480 #[test]
481 fn nullable_wraps_option() {
482 let t = map().resolve_pg_type("text", true, false).unwrap();
483 assert_eq!(t.rust_type, "Option<String>");
484 assert!(!t.copy_cheap);
485 }
486 #[test]
487 fn array_wraps_vec() {
488 let t = map().resolve_pg_type("text", false, true).unwrap();
489 assert_eq!(t.rust_type, "Vec<String>");
490 assert!(!t.copy_cheap);
491 }
492 #[test]
493 fn nullable_array() {
494 let t = map().resolve_pg_type("text", true, true).unwrap();
495 assert_eq!(t.rust_type, "Option<Vec<String>>");
496 }
497 #[test]
498 fn multidimensional_array_wraps_nested_vec() {
499 let t = map().resolve_pg_type_dims("int8", false, 2).unwrap();
500 assert_eq!(t.rust_type, "Vec<Vec<i64>>");
501 }
502 #[test]
503 fn nullable_multidimensional_array_wraps_option_nested_vec() {
504 let t = map().resolve_pg_type_dims("text", true, 3).unwrap();
505 assert_eq!(t.rust_type, "Option<Vec<Vec<Vec<String>>>>");
506 }
507 #[test]
508 fn type_override_replaces_default() {
509 let t = TypeMap::new(
510 &[owned_override("timestamptz", "time::OffsetDateTime")],
511 &[],
512 )
513 .resolve_pg_type("timestamptz", false, false)
514 .unwrap();
515 assert_eq!(t.rust_type, "time::OffsetDateTime");
516 assert!(t.borrowed_rust_type.is_none());
517 }
518 #[test]
519 fn column_override_beats_type_override() {
520 let overrides = vec![
521 owned_override("text", "TypeLevel"),
522 TypeOverride {
523 db_type: None,
524 column: Some("users.name".to_string()),
525 rs_type: Some("ColumnLevel".to_string()),
526 borrowed_rs_type: None,
527 copy_cheap: false,
528 },
529 ];
530 let col_ovrs = build_column_overrides(&overrides);
531 let map = TypeMap::new(&overrides, &[]);
532 let t = map
533 .resolve_column("text", false, false, Some("users.name"), &col_ovrs)
534 .unwrap();
535 assert_eq!(t.rust_type, "ColumnLevel");
536 }
537 #[test]
538 fn maps_numeric() {
539 let t = map().resolve_pg_type("numeric", false, false).unwrap();
540 assert_eq!(t.rust_type, "bigdecimal::BigDecimal");
541 assert!(!t.copy_cheap);
542 }
543 #[test]
544 fn maps_decimal() {
545 let t = map().resolve_pg_type("decimal", false, false).unwrap();
546 assert_eq!(t.rust_type, "bigdecimal::BigDecimal");
547 }
548 #[test]
549 fn maps_pg_catalog_numeric() {
550 let t = map()
551 .resolve_pg_type("pg_catalog.numeric", false, false)
552 .unwrap();
553 assert_eq!(t.rust_type, "bigdecimal::BigDecimal");
554 }
555 #[test]
556 fn unknown_type_returns_none() {
557 assert!(
558 map()
559 .resolve_pg_type("no_such_type", false, false)
560 .is_none()
561 );
562 }
563 #[test]
564 fn registers_custom_type() {
565 let mut map = TypeMap::new(&[], &[]);
566 map.register("my_enum", "MyEnum", false);
567 let t = map.resolve_pg_type("my_enum", false, false).unwrap();
568 assert_eq!(t.rust_type, "MyEnum");
569 assert!(!t.copy_cheap);
570 }
571 #[test]
572 fn registered_type_nullable() {
573 let mut map = TypeMap::new(&[], &[]);
574 map.register("my_enum", "MyEnum", false);
575 let t = map.resolve_pg_type("my_enum", true, false).unwrap();
576 assert_eq!(t.rust_type, "Option<MyEnum>");
577 assert!(!t.copy_cheap);
578 }
579 #[test]
580 fn type_override_beats_registered_custom() {
581 let mut map = TypeMap::new(&[owned_override("my_enum", "Override")], &[]);
582 map.register("my_enum", "MyEnum", false);
583 let t = map.resolve_pg_type("my_enum", false, false).unwrap();
584 assert_eq!(t.rust_type, "Override");
586 }
587 #[test]
588 fn registered_copy_cheap_type_is_cheap() {
589 let mut map = TypeMap::new(&[], &[]);
590 map.register("my_value_type", "MyValueType", true);
591 let t = map.resolve_pg_type("my_value_type", false, false).unwrap();
592 assert_eq!(t.rust_type, "MyValueType");
593 assert!(
594 t.copy_cheap,
595 "registered type with copy_cheap=true should resolve as copy_cheap"
596 );
597 }
598
599 #[test]
600 fn registered_copy_cheap_nullable_is_not_cheap() {
601 let mut map = TypeMap::new(&[], &[]);
602 map.register("my_value_type", "MyValueType", true);
603 let t = map.resolve_pg_type("my_value_type", true, false).unwrap();
604 assert_eq!(t.rust_type, "Option<MyValueType>");
605 assert!(
606 !t.copy_cheap,
607 "nullable type should not be copy_cheap even if base is"
608 );
609 }
610
611 #[test]
612 fn registered_copy_cheap_array_is_not_cheap() {
613 let mut map = TypeMap::new(&[], &[]);
614 map.register("my_value_type", "MyValueType", true);
615 let t = map.resolve_pg_type("my_value_type", false, true).unwrap();
616 assert_eq!(t.rust_type, "Vec<MyValueType>");
617 assert!(
618 !t.copy_cheap,
619 "array type should not be copy_cheap even if base is"
620 );
621 }
622
623 #[test]
624 fn copy_cheap_types_marks_type_as_cheap() {
625 let map = TypeMap::new(&[], &["text".to_string()]);
626 let t = map.resolve_pg_type("text", false, false).unwrap();
628 assert_eq!(t.rust_type, "String");
629 assert!(
630 t.copy_cheap,
631 "text should be copy_cheap after config promotion"
632 );
633 }
634
635 #[test]
636 fn copy_cheap_types_promotes_default_to_override() {
637 let map = TypeMap::new(&[], &["uuid".to_string()]);
639 let t = map.resolve_pg_type("uuid", false, false).unwrap();
640 assert!(t.copy_cheap);
641 }
642
643 fn borrowed_text() -> TypeOverride {
646 TypeOverride {
647 db_type: Some("text".to_string()),
648 column: None,
649 rs_type: None,
650 borrowed_rs_type: Some("&str".to_string()),
651 copy_cheap: false,
652 }
653 }
654
655 #[test]
656 fn borrowed_only_override_keeps_default_for_owned_form() {
657 let map = TypeMap::new(&[borrowed_text()], &[]);
658 let t = map.resolve_pg_type("text", false, false).unwrap();
659 assert_eq!(t.rust_type, "String");
660 assert_eq!(t.borrowed_rust_type.as_deref(), Some("&str"));
661 }
662
663 #[test]
664 fn borrowed_nullable_wraps_inside_option() {
665 let map = TypeMap::new(&[borrowed_text()], &[]);
666 let t = map.resolve_pg_type("text", true, false).unwrap();
667 assert_eq!(t.rust_type, "Option<String>");
668 assert_eq!(t.borrowed_rust_type.as_deref(), Some("Option<&str>"));
669 }
670
671 #[test]
672 fn borrowed_array_uses_owned_inner_in_slice() {
673 let map = TypeMap::new(&[borrowed_text()], &[]);
674 let t = map.resolve_pg_type("text", false, true).unwrap();
675 assert_eq!(t.rust_type, "Vec<String>");
676 assert_eq!(t.borrowed_rust_type.as_deref(), Some("&[String]"));
677 }
678
679 #[test]
680 fn borrowed_nullable_array() {
681 let map = TypeMap::new(&[borrowed_text()], &[]);
682 let t = map.resolve_pg_type("text", true, true).unwrap();
683 assert_eq!(t.rust_type, "Option<Vec<String>>");
684 assert_eq!(t.borrowed_rust_type.as_deref(), Some("Option<&[String]>"));
685 }
686
687 #[test]
688 fn borrowed_multidim_array_inner_stays_vec() {
689 let map = TypeMap::new(&[borrowed_text()], &[]);
690 let t = map.resolve_pg_type_dims("text", false, 2).unwrap();
691 assert_eq!(t.rust_type, "Vec<Vec<String>>");
692 assert_eq!(t.borrowed_rust_type.as_deref(), Some("&[Vec<String>]"));
693 }
694
695 #[test]
696 fn owned_and_borrowed_pair() {
697 let ovr = TypeOverride {
698 db_type: Some("text".to_string()),
699 column: None,
700 rs_type: Some("MyStr".to_string()),
701 borrowed_rs_type: Some("&MyStr".to_string()),
702 copy_cheap: false,
703 };
704 let map = TypeMap::new(&[ovr], &[]);
705 let t = map.resolve_pg_type("text", false, false).unwrap();
706 assert_eq!(t.rust_type, "MyStr");
707 assert_eq!(t.borrowed_rust_type.as_deref(), Some("&MyStr"));
708
709 let t = map.resolve_pg_type("text", false, true).unwrap();
710 assert_eq!(t.rust_type, "Vec<MyStr>");
712 assert_eq!(t.borrowed_rust_type.as_deref(), Some("&[MyStr]"));
713 }
714
715 #[test]
716 fn borrowed_column_override() {
717 let overrides = vec![TypeOverride {
718 db_type: None,
719 column: Some("users.name".to_string()),
720 rs_type: None,
721 borrowed_rs_type: Some("&str".to_string()),
722 copy_cheap: false,
723 }];
724 let col_ovrs = build_column_overrides(&overrides);
725 let map = TypeMap::new(&overrides, &[]);
726 let t = map
727 .resolve_column("text", false, false, Some("users.name"), &col_ovrs)
728 .unwrap();
729 assert_eq!(t.rust_type, "String");
730 assert_eq!(t.borrowed_rust_type.as_deref(), Some("&str"));
731 }
732}