1use crate::diff::{MigrationStep, OpKey};
38use crate::schema::{EnumDescriptor, FunctionDescriptor, ScalarDescriptor, SchemaDescriptor, TypeDescriptor};
39
40pub fn python_snippet_for_step(step: &MigrationStep, schema: &SchemaDescriptor) -> Option<String> {
43 match &step.op_key {
44 OpKey::Table(module, table) | OpKey::ForeignKey(module, table) => {
49 let td = schema.types.iter().find(|t| &t.module == module && &t.table == table)?;
50 Some(python_snippet_for_class(td, "@pylon.type"))
51 }
52 OpKey::View(module, name) => {
53 let td = schema.types.iter().find(|t| &t.module == module && &t.name == name)?;
54 Some(python_snippet_for_class(td, "@pylon.interface"))
55 }
56 OpKey::Scalar(module, name) => {
57 if let Some(e) = schema.enums.iter().find(|e| &e.module == module && &e.name == name) {
58 return Some(python_snippet_for_enum(e));
59 }
60 let s = schema.scalars.iter().find(|s| &s.module == module && &s.name == name)?;
61 Some(python_snippet_for_scalar(s))
62 }
63 OpKey::Function(module, name) => {
64 let f = schema
65 .functions
66 .iter()
67 .find(|f| &f.module == module && &f.name == name)?;
68 Some(python_snippet_for_function(f))
69 }
70 OpKey::Module(_) => None,
71 }
72}
73
74fn python_snippet_for_class(td: &TypeDescriptor, decorator: &str) -> String {
77 let mut lines = vec![decorator.to_string(), format!("class {}:", td.name)];
78 let mut body: Vec<String> = Vec::new();
79
80 for p in &td.properties {
81 if p.name == "id" {
82 continue; }
84 let hint = python_type_hint(&p.pg_type, p.column_type.as_deref());
85 let hint = if p.nullable { format!("{hint} | None") } else { hint };
86 body.push(format!(" {}: {}", p.name, hint));
87 }
88 for l in &td.links {
89 let hint = format!("Link[{}]", bare_name(&l.target));
90 let hint = if l.nullable { format!("{hint} | None") } else { hint };
91 body.push(format!(" {}: {}", l.name, hint));
92 }
93 for ml in &td.multilinks {
94 body.push(format!(" {}: MultiLink[{}]", ml.name, bare_name(&ml.target)));
95 }
96
97 if body.is_empty() {
98 lines.push(" pass".to_string());
99 } else {
100 lines.extend(body);
101 }
102 lines.join("\n")
103}
104
105fn python_snippet_for_enum(e: &EnumDescriptor) -> String {
106 let members: Vec<String> = e.members.iter().map(|m| format!("\"{m}\"")).collect();
107 format!(
108 "@pylon.enum({})\nclass {}(pylon.Enum):\n pass",
109 members.join(", "),
110 e.name
111 )
112}
113
114fn python_snippet_for_scalar(s: &ScalarDescriptor) -> String {
115 format!(
116 "@pylon.scalar(pylon.{})\nclass {}(pylon.Scalar):\n pass",
117 s.base, s.name
118 )
119}
120
121fn python_snippet_for_function(f: &FunctionDescriptor) -> String {
122 let params: Vec<String> = f
123 .params
124 .iter()
125 .map(|p| format!("{}: {}", p.name, python_type_hint(&p.pg_type, None)))
126 .collect();
127
128 let mut return_hint = if f.return_is_object {
129 bare_name(&f.return_pg_type).to_string()
130 } else {
131 python_type_hint(&f.return_pg_type, None)
132 };
133 if f.return_is_set {
134 return_hint = format!("set[{return_hint}]");
135 }
136
137 format!(
138 "@pylon.function\ndef {}({}) -> {}:\n \"\"\"\n{}\n \"\"\"",
139 f.name,
140 params.join(", "),
141 return_hint,
142 indent_block(f.body.trim(), " "),
143 )
144}
145
146fn indent_block(text: &str, indent: &str) -> String {
152 text.lines()
153 .map(|line| {
154 if line.is_empty() {
155 line.to_string()
156 } else {
157 format!("{indent}{line}")
158 }
159 })
160 .collect::<Vec<_>>()
161 .join("\n")
162}
163
164fn bare_name(qualified: &str) -> &str {
168 qualified.rsplit("::").next().unwrap_or(qualified)
169}
170
171fn python_type_hint(pg_type: &str, column_type: Option<&str>) -> String {
175 if let Some(ct) = column_type {
176 let dequoted = ct.replace("\".\"", "::").replace('"', "");
179 return bare_name(&dequoted).to_string();
180 }
181 if let Some(elem) = pg_type.strip_suffix("[]") {
182 return format!("list[{}]", python_type_hint(elem, None));
183 }
184 if let Some(nt_name) = pg_type.strip_prefix("__nt__:") {
185 return bare_name(nt_name).to_string();
186 }
187 if pg_type.starts_with('"') {
188 let dequoted = pg_type.replace("\".\"", "::").replace('"', "");
191 return bare_name(&dequoted).to_string();
192 }
193 match pg_type {
194 "text" | "varchar" => "str".to_string(),
195 "boolean" => "bool".to_string(),
196 "float8" => "float".to_string(),
197 "int4" => "int".to_string(),
198 "int2" => "pylon.Int16".to_string(),
199 "int8" => "pylon.Int64".to_string(),
200 "float4" => "pylon.Float32".to_string(),
201 "numeric" => "pylon.Decimal".to_string(),
202 "timestamptz" => "pylon.DateTime".to_string(),
203 "timestamp" => "pylon.LocalDateTime".to_string(),
204 "date" => "pylon.LocalDate".to_string(),
205 "time" => "pylon.LocalTime".to_string(),
206 "uuid" => "pylon.UUID".to_string(),
207 "bytea" => "pylon.Bytes".to_string(),
208 "jsonb" => "pylon.Json".to_string(),
209 "interval" => "pylon.Duration".to_string(),
210 other => other.to_string(),
211 }
212}
213
214#[cfg(test)]
215mod tests {
216 use super::*;
217 use crate::diff::Verb;
218
219 fn type_step(td: TypeDescriptor, verb: Verb) -> MigrationStep {
220 MigrationStep {
221 prompt: String::new(),
222 verb,
223 object_desc: String::new(),
224 ddl: vec![],
225 required_input: vec![],
226 op_key: OpKey::Table(td.module.clone(), td.table.clone()),
227 }
228 }
229
230 fn empty_type(module: &str, name: &str, table: &str) -> TypeDescriptor {
231 TypeDescriptor {
232 name: name.into(),
233 module: module.into(),
234 table: table.into(),
235 abstract_: false,
236 materialized: false,
237 description: None,
238 parents: vec![],
239 interfaces: vec![],
240 bases: vec![],
241 properties: vec![],
242 links: vec![],
243 multilinks: vec![],
244 computed: vec![],
245 constraints: vec![],
246 indexes: vec![],
247 partition: None,
248 vector_indexes: vec![],
249 search_indexes: vec![],
250 triggers: vec![],
251 junction: false,
252 signals: vec![],
253 }
254 }
255
256 fn prop(name: &str, pg_type: &str, nullable: bool) -> crate::schema::PropertyDescriptor {
257 crate::schema::PropertyDescriptor {
258 name: name.into(),
259 pg_type: pg_type.into(),
260 nullable,
261 default_sql: None,
262 default_pyql: None,
263 description: None,
264 check_constraints: vec![],
265 is_exclusive: false,
266 is_pk: name == "id",
267 is_readonly: name == "id",
268 rewrites: vec![],
269 tuple_members: None,
270 column_type: None,
271 }
272 }
273
274 #[test]
275 fn renders_a_simple_object_type() {
276 let mut td = empty_type("blog", "Post", "Post");
277 td.properties.push(prop("id", "uuid", false));
278 td.properties.push(prop("title", "text", false));
279 td.properties.push(prop("views", "int8", true));
280 let schema = SchemaDescriptor {
281 types: vec![td],
282 scalars: vec![],
283 enums: vec![],
284 named_tuples: vec![],
285 globals: vec![],
286 functions: vec![],
287 aliases: vec![],
288 channels: vec![],
289 ..Default::default()
290 };
291 let step = type_step(schema.types[0].clone(), Verb::Create);
292 let snippet = python_snippet_for_step(&step, &schema).unwrap();
293 assert_eq!(
294 snippet,
295 "@pylon.type\nclass Post:\n title: str\n views: pylon.Int64 | None"
296 );
297 }
298
299 #[test]
300 fn renders_an_interface_type_as_pylon_interface_not_the_view_sql() {
301 let mut td = empty_type("default", "Account", "Account");
302 td.abstract_ = true;
303 td.materialized = true;
304 td.properties.push(prop("id", "uuid", false));
305 td.properties.push(prop("email", "text", false));
306 let schema = SchemaDescriptor {
307 types: vec![td],
308 scalars: vec![],
309 enums: vec![],
310 named_tuples: vec![],
311 globals: vec![],
312 functions: vec![],
313 aliases: vec![],
314 channels: vec![],
315 ..Default::default()
316 };
317 let step = MigrationStep {
318 prompt: String::new(),
319 verb: Verb::Create,
320 object_desc: String::new(),
321 ddl: vec![],
322 required_input: vec![],
323 op_key: OpKey::View("default".into(), "Account".into()),
324 };
325 let snippet = python_snippet_for_step(&step, &schema).unwrap();
326 assert_eq!(snippet, "@pylon.interface\nclass Account:\n email: str");
327 }
328
329 #[test]
330 fn renders_an_enum() {
331 let schema = SchemaDescriptor {
332 types: vec![],
333 scalars: vec![],
334 enums: vec![EnumDescriptor {
335 name: "Status".into(),
336 module: "default".into(),
337 members: vec!["Active".into(), "Inactive".into()],
338 }],
339 named_tuples: vec![],
340 globals: vec![],
341 functions: vec![],
342 aliases: vec![],
343 channels: vec![],
344 ..Default::default()
345 };
346 let step = MigrationStep {
347 prompt: String::new(),
348 verb: Verb::Create,
349 object_desc: String::new(),
350 ddl: vec![],
351 required_input: vec![],
352 op_key: OpKey::Scalar("default".into(), "Status".into()),
353 };
354 let snippet = python_snippet_for_step(&step, &schema).unwrap();
355 assert_eq!(
356 snippet,
357 "@pylon.enum(\"Active\", \"Inactive\")\nclass Status(pylon.Enum):\n pass"
358 );
359 }
360
361 #[test]
362 fn renders_a_scalar_function_as_its_pyql_body_not_compiled_sql() {
363 let schema = SchemaDescriptor {
364 types: vec![],
365 scalars: vec![],
366 enums: vec![],
367 named_tuples: vec![],
368 globals: vec![],
369 aliases: vec![],
370 channels: vec![],
371 functions: vec![FunctionDescriptor {
372 name: "get_content_type".into(),
373 module: "default".into(),
374 params: vec![crate::schema::FunctionParamDescriptor {
375 name: "uuid_val".into(),
376 pg_type: "uuid".into(),
377 }],
378 return_pg_type: "int2".into(),
379 return_is_object: false,
380 return_is_set: false,
381 return_is_polymorphic: false,
382 volatility: "immutable".into(),
383 body: "select 1".into(),
384 }],
385 ..Default::default()
386 };
387 let step = MigrationStep {
388 prompt: String::new(),
389 verb: Verb::Create,
390 object_desc: String::new(),
391 ddl: vec![],
392 required_input: vec![],
393 op_key: OpKey::Function("default".into(), "get_content_type".into()),
394 };
395 let snippet = python_snippet_for_step(&step, &schema).unwrap();
396 assert_eq!(
397 snippet,
398 "@pylon.function\ndef get_content_type(uuid_val: pylon.UUID) -> pylon.Int16:\n \"\"\"\n select 1\n \"\"\""
399 );
400 }
401
402 #[test]
403 fn indents_every_line_of_a_multiline_pyql_body() {
404 let schema = SchemaDescriptor {
405 types: vec![], scalars: vec![], enums: vec![], named_tuples: vec![], globals: vec![], aliases: vec![], channels: vec![],
406 functions: vec![FunctionDescriptor {
407 name: "get_content_type".into(),
408 module: "default".into(),
409 params: vec![crate::schema::FunctionParamDescriptor { name: "uuid_val".into(), pg_type: "uuid".into() }],
410 return_pg_type: "int8".into(),
411 return_is_object: false,
412 return_is_set: false,
413 return_is_polymorphic: false,
414 volatility: "immutable".into(),
415 body: "with\n h := str_replace(<str>uuid_val, '-', ''),\n ct_bytes := std::from_hex(h[20:24])\nselect ct_bytes".into(),
416 }],
417 ..Default::default()
418 };
419 let step = MigrationStep {
420 prompt: String::new(),
421 verb: Verb::Create,
422 object_desc: String::new(),
423 ddl: vec![],
424 required_input: vec![],
425 op_key: OpKey::Function("default".into(), "get_content_type".into()),
426 };
427 let snippet = python_snippet_for_step(&step, &schema).unwrap();
428 assert_eq!(
429 snippet,
430 "@pylon.function\ndef get_content_type(uuid_val: pylon.UUID) -> pylon.Int64:\n \"\"\"\n with\n h := str_replace(<str>uuid_val, '-', ''),\n ct_bytes := std::from_hex(h[20:24])\n select ct_bytes\n \"\"\""
431 );
432 }
433
434 #[test]
435 fn function_snippet_survives_the_real_diff_pipeline() {
436 let schema = SchemaDescriptor {
443 types: vec![],
444 scalars: vec![],
445 enums: vec![],
446 named_tuples: vec![],
447 globals: vec![],
448 aliases: vec![],
449 channels: vec![],
450 functions: vec![FunctionDescriptor {
451 name: "get_content_type".into(),
452 module: "default".into(),
453 params: vec![crate::schema::FunctionParamDescriptor {
454 name: "uuid_val".into(),
455 pg_type: "uuid".into(),
456 }],
457 return_pg_type: "int2".into(),
458 return_is_object: false,
459 return_is_set: false,
460 return_is_polymorphic: false,
461 volatility: "immutable".into(),
462 body: "select 1".into(),
463 }],
464 ..Default::default()
465 };
466 let steps = crate::diff::diff_schema_steps_with_renames_and_fills(
467 &schema,
468 &crate::diff::DbState::default(),
469 &[],
470 &[],
471 &[],
472 )
473 .unwrap();
474 let step = steps
475 .iter()
476 .find(|s| s.prompt.contains("get_content_type"))
477 .expect("expected a step for get_content_type");
478 let snippet = python_snippet_for_step(step, &schema);
479 assert!(
480 snippet.is_some(),
481 "expected a snippet for the function step, prompt was: {:?}",
482 step.prompt
483 );
484 }
485
486 #[test]
487 fn returns_none_for_a_module_step() {
488 let schema = SchemaDescriptor {
489 types: vec![],
490 scalars: vec![],
491 enums: vec![],
492 named_tuples: vec![],
493 globals: vec![],
494 functions: vec![],
495 aliases: vec![],
496 channels: vec![],
497 ..Default::default()
498 };
499 let step = MigrationStep {
500 prompt: String::new(),
501 verb: Verb::Create,
502 object_desc: String::new(),
503 ddl: vec![],
504 required_input: vec![],
505 op_key: OpKey::Module("catalog".into()),
506 };
507 assert!(python_snippet_for_step(&step, &schema).is_none());
508 }
509}