1use crate::error::Error;
2
3#[derive(Debug, serde::Deserialize)]
4#[serde(default)]
5pub struct Config {
6 pub output: String,
7 pub overrides: Vec<TypeOverride>,
8 pub row_derives: Vec<String>,
9 pub enum_derives: Vec<String>,
10 pub composite_derives: Vec<String>,
11 pub copy_cheap_types: Vec<String>,
12}
13
14impl Default for Config {
15 fn default() -> Self {
16 Self {
17 output: "queries.rs".to_string(),
18 overrides: vec![],
19 row_derives: vec![],
20 enum_derives: vec![],
21 composite_derives: vec![],
22 copy_cheap_types: vec![],
23 }
24 }
25}
26
27impl Config {
28 pub fn from_bytes(bytes: &[u8]) -> Result<Self, Error> {
29 let cfg: Config = serde_json::from_slice(bytes)?;
30 for o in &cfg.overrides {
31 let target = || {
32 o.db_type
33 .as_deref()
34 .or(o.column.as_deref())
35 .unwrap_or("<unspecified>")
36 .to_string()
37 };
38 if o.rs_type.is_none() && o.borrowed_rs_type.is_none() {
39 return Err(Error::Codegen(format!(
40 "override for '{}' must set at least one of 'rs_type' or 'borrowed_rs_type'",
41 target()
42 )));
43 }
44 if let Some(borrowed) = &o.borrowed_rs_type {
45 validate_borrowed_type(borrowed, &target())?;
46 }
47 }
48 Ok(cfg)
49 }
50}
51
52fn validate_borrowed_type(rs_type: &str, target: &str) -> Result<(), Error> {
57 syn::parse_str::<syn::Type>(rs_type).map_err(|e| {
58 Error::Codegen(format!(
59 "override for '{target}': borrowed_rs_type '{rs_type}' is not a valid Rust type: {e}"
60 ))
61 })?;
62 if !rs_type.contains('&') {
63 return Err(Error::Codegen(format!(
64 "override for '{target}': borrowed_rs_type '{rs_type}' must contain a reference \
65 (e.g. '&str' or 'Option<&str>'); use 'rs_type' for owned types"
66 )));
67 }
68 Ok(())
69}
70
71#[derive(Debug, serde::Deserialize)]
72pub struct TypeOverride {
73 pub db_type: Option<String>,
74 pub column: Option<String>,
75 pub rs_type: Option<String>,
79 pub borrowed_rs_type: Option<String>,
85 #[serde(default)]
86 pub copy_cheap: bool,
87}
88
89#[cfg(test)]
90mod tests {
91 use super::*;
92
93 #[test]
94 fn parses_empty_json() {
95 let c = Config::from_bytes(b"{}").unwrap();
96 assert_eq!(c.output, "queries.rs");
97 assert!(c.overrides.is_empty());
98 assert!(c.row_derives.is_empty());
99 }
100
101 #[test]
102 fn parses_output_name() {
103 let c = Config::from_bytes(br#"{"output": "db.rs"}"#).unwrap();
104 assert_eq!(c.output, "db.rs");
105 }
106
107 #[test]
108 fn parses_type_override() {
109 let json = br#"{"overrides":[{"db_type":"timestamptz","rs_type":"chrono::DateTime<chrono::Utc>"}]}"#;
110 let c = Config::from_bytes(json).unwrap();
111 assert_eq!(c.overrides.len(), 1);
112 assert_eq!(c.overrides[0].db_type, Some("timestamptz".to_string()));
113 assert_eq!(
114 c.overrides[0].rs_type.as_deref(),
115 Some("chrono::DateTime<chrono::Utc>")
116 );
117 assert!(c.overrides[0].borrowed_rs_type.is_none());
118 }
119
120 #[test]
121 fn parses_column_override() {
122 let json = br#"{"overrides":[{"column":"users.created_at","rs_type":"chrono::DateTime<chrono::Local>","copy_cheap":true}]}"#;
123 let c = Config::from_bytes(json).unwrap();
124 assert_eq!(c.overrides[0].column, Some("users.created_at".to_string()));
125 assert!(c.overrides[0].copy_cheap);
126 }
127
128 #[test]
129 fn parses_derives() {
130 let json = br#"{"row_derives":["serde::Serialize"],"enum_derives":["serde::Serialize","serde::Deserialize"]}"#;
131 let c = Config::from_bytes(json).unwrap();
132 assert_eq!(c.row_derives, ["serde::Serialize"]);
133 assert_eq!(c.enum_derives.len(), 2);
134 }
135
136 #[test]
137 fn parses_borrowed_only_override() {
138 let json = br#"{"overrides":[{"db_type":"text","borrowed_rs_type":"&str"}]}"#;
139 let c = Config::from_bytes(json).unwrap();
140 assert_eq!(c.overrides.len(), 1);
141 assert!(c.overrides[0].rs_type.is_none());
142 assert_eq!(c.overrides[0].borrowed_rs_type.as_deref(), Some("&str"));
143 }
144
145 #[test]
146 fn parses_owned_and_borrowed_override() {
147 let json =
148 br#"{"overrides":[{"db_type":"text","rs_type":"MyStr","borrowed_rs_type":"&MyStr"}]}"#;
149 let c = Config::from_bytes(json).unwrap();
150 assert_eq!(c.overrides[0].rs_type.as_deref(), Some("MyStr"));
151 assert_eq!(c.overrides[0].borrowed_rs_type.as_deref(), Some("&MyStr"));
152 }
153
154 #[test]
155 fn rejects_non_borrowed_in_borrowed_rs_type() {
156 let json = br#"{"overrides":[{"db_type":"text","borrowed_rs_type":"Option<String>"}]}"#;
157 let err = Config::from_bytes(json).unwrap_err();
158 let msg = err.to_string();
159 assert!(
160 msg.contains("must contain a reference"),
161 "expected reference-required error, got: {msg}"
162 );
163 assert!(msg.contains("text"), "expected target name in error: {msg}");
164 }
165
166 #[test]
167 fn rejects_invalid_rust_in_borrowed_rs_type() {
168 let json = br#"{"overrides":[{"db_type":"text","borrowed_rs_type":"¬ a type!!"}]}"#;
169 let err = Config::from_bytes(json).unwrap_err();
170 let msg = err.to_string();
171 assert!(
172 msg.contains("not a valid Rust type"),
173 "expected parse error, got: {msg}"
174 );
175 }
176
177 #[test]
178 fn accepts_option_of_reference_in_borrowed_rs_type() {
179 let json = br#"{"overrides":[{"db_type":"text","borrowed_rs_type":"Option<&str>"}]}"#;
180 Config::from_bytes(json).expect("Option<&str> should be accepted as borrowed");
181 }
182}