1use std::fmt::Write;
4
5use crate::*;
6
7#[derive(Debug, thiserror::Error)]
8pub enum CodegenError {
9 #[error(transparent)]
10 Descriptor(#[from] DescriptorError),
11 #[error(
12 "descriptor fingerprint mismatch for '{name}': expected {expected}, computed {actual}"
13 )]
14 Fingerprint {
15 name: String,
16 expected: String,
17 actual: String,
18 },
19 #[error("identifier '{0}' cannot be represented safely in generated Rust")]
20 Identifier(String),
21 #[error("reference '{table}.{column}' must use one source and one target column")]
22 UnsupportedReference { table: String, column: String },
23 #[error("reference target table '{0}' is absent from the descriptor")]
24 MissingReferenceTarget(String),
25 #[error("reference source column '{table}.{column}' is absent from the descriptor")]
26 MissingReferenceSourceColumn { table: String, column: String },
27 #[error("reference target column '{table}.{column}' is absent from the descriptor")]
28 MissingReferenceTargetColumn { table: String, column: String },
29 #[error(
30 "reference target '{table}.{column}' must be a one-column PRIMARY KEY or UNIQUE NOT NULL key"
31 )]
32 UnsupportedReferenceTargetKey { table: String, column: String },
33 #[error(
34 "reference type mismatch: '{source_table}.{source_column}' is {source_type:?}, but '{target_table}.{target_column}' is {target_type:?}"
35 )]
36 ReferenceTypeMismatch {
37 source_table: String,
38 source_column: String,
39 source_type: DataTypeDescriptor,
40 target_table: String,
41 target_column: String,
42 target_type: DataTypeDescriptor,
43 },
44}
45
46#[derive(Debug, Clone, PartialEq, Eq)]
47pub struct GeneratedRust {
48 pub source: String,
49 pub descriptor_fingerprint: String,
50}
51
52pub fn generate_rust_database(
53 descriptor: &DescriptorEnvelope<DatabaseDescriptor>,
54) -> Result<GeneratedRust, CodegenError> {
55 if descriptor.descriptor != SCHEMA_DESCRIPTOR_VERSION
56 || descriptor.kind != DescriptorKind::Database
57 {
58 return Err(CodegenError::Descriptor(DescriptorError::KindMismatch {
59 expected: DescriptorKind::Database,
60 actual: descriptor.kind,
61 }));
62 }
63 validate_database_fingerprints(&descriptor.payload)?;
64
65 let mut output = String::new();
66 output.push_str("// @generated by radixdb-orm; DO NOT EDIT.\n");
67 output.push_str("// Source: radixdb.schema.v1 descriptor. No live database was accessed.\n\n");
68 output.push_str("#[allow(unused_imports)]\nuse radixdb_orm::{AsyncOrmGeneratedRecordSession, DataTypeDescriptor, DynamicRecord, Expr, FieldState, FieldValue, GeneratedEntity, GeneratedQuery, GeneratedRecord, GeneratedRecordError, InsertBuilder, OrmGeneratedRecordSession, QueryBuilder, Reference, Relation, RecordMutation, TableDescriptor, TypedColumn, TypedKey, TypedValue, UpdateBuilder, DeleteBuilder, decode_generated_field, decode_generated_reference_field, ensure_schema_fingerprint};\n\n");
69 writeln!(
70 output,
71 "pub const DATABASE_SCHEMA_FINGERPRINT: &str = {:?};\n",
72 descriptor.payload.fingerprint
73 )
74 .unwrap();
75
76 for table in &descriptor.payload.tables {
77 render_table(&mut output, table, &descriptor.payload.tables)?;
78 }
79
80 Ok(GeneratedRust {
81 source: output,
82 descriptor_fingerprint: descriptor.payload.fingerprint.clone(),
83 })
84}
85
86fn validate_database_fingerprints(database: &DatabaseDescriptor) -> Result<(), CodegenError> {
87 for table in &database.tables {
88 let actual = table.computed_fingerprint()?;
89 if table.fingerprint != actual {
90 return Err(CodegenError::Fingerprint {
91 name: table.name.clone(),
92 expected: table.fingerprint.clone(),
93 actual,
94 });
95 }
96 }
97 let actual = database.computed_fingerprint()?;
98 if database.fingerprint != actual {
99 return Err(CodegenError::Fingerprint {
100 name: "database".to_string(),
101 expected: database.fingerprint.clone(),
102 actual,
103 });
104 }
105 Ok(())
106}
107
108fn render_table(
109 output: &mut String,
110 table: &TableDescriptor,
111 tables: &[TableDescriptor],
112) -> Result<(), CodegenError> {
113 validate_reference_shapes(table, tables)?;
114 let entity = rust_type_name(&table.name)?;
115 let record = format!("{entity}Record");
116 writeln!(
117 output,
118 "#[derive(Debug, Clone, PartialEq)]\npub struct {record} {{"
119 )
120 .unwrap();
121 for column in &table.columns {
122 let field_type = match column_reference(table, &column.name) {
123 Some((target_table, _)) => rust_type_name(target_table)?,
124 None => rust_type(&column.data_type).to_string(),
125 };
126 writeln!(
127 output,
128 " pub {}: FieldState<{}>,",
129 rust_field_name(&column.name)?,
130 match column_reference(table, &column.name) {
131 Some(_) => format!("Reference<{field_type}>"),
132 None => field_type,
133 }
134 )
135 .unwrap();
136 }
137 output.push_str("}\n\n");
138
139 writeln!(output, "impl Default for {record} {{").unwrap();
140 output.push_str(" fn default() -> Self {\n Self {\n");
141 for column in &table.columns {
142 writeln!(
143 output,
144 " {}: FieldState::typed({}),",
145 rust_field_name(&column.name)?,
146 rust_data_type(&column.data_type)
147 )
148 .unwrap();
149 }
150 output.push_str(" }\n }\n}\n\n");
151
152 writeln!(
153 output,
154 "#[derive(Debug, Clone, Copy, PartialEq, Eq)]\npub struct {entity};"
155 )
156 .unwrap();
157 writeln!(output, "impl GeneratedEntity for {entity} {{").unwrap();
158 writeln!(output, " type Record = {record};").unwrap();
159 writeln!(output, " const TABLE: &'static str = {:?};", table.name).unwrap();
160 writeln!(
161 output,
162 " const CATALOG_ID: &'static str = {:?};",
163 table.catalog_id
164 )
165 .unwrap();
166 writeln!(
167 output,
168 " const SCHEMA_FINGERPRINT: &'static str = {:?};",
169 table.fingerprint
170 )
171 .unwrap();
172 output.push_str("}\n\n");
173
174 writeln!(output, "impl {entity} {{").unwrap();
175 for column in &table.columns {
176 let column_type = match column_reference(table, &column.name) {
177 Some((target_table, _)) => format!("Reference<{}>", rust_type_name(target_table)?),
178 None => rust_type(&column.data_type).to_string(),
179 };
180 writeln!(
181 output,
182 " pub const {}: TypedColumn<{}, Self> = TypedColumn::new({:?}, {:?}, {}, {});",
183 rust_const_name(&column.name)?,
184 column_type,
185 table.name,
186 column.name,
187 rust_data_type(&column.data_type),
188 column.nullable
189 )
190 .unwrap();
191 }
192 writeln!(
193 output,
194 " pub fn new() -> {record} {{ {record}::default() }}"
195 )
196 .unwrap();
197 writeln!(output, " pub fn table() -> Relation {{ Relation::Table {{ name: {:?}.to_string(), alias: None }} }}", table.name).unwrap();
198 output.push_str(
199 " pub fn query() -> QueryBuilder { QueryBuilder::from_relation(Self::table()) }\n",
200 );
201 writeln!(
202 output,
203 " pub fn insert() -> InsertBuilder {{ InsertBuilder::new({:?}) }}",
204 table.name
205 )
206 .unwrap();
207 writeln!(
208 output,
209 " pub fn upsert() -> InsertBuilder {{ InsertBuilder::new({:?}) }}",
210 table.name
211 )
212 .unwrap();
213 writeln!(
214 output,
215 " pub fn update() -> UpdateBuilder {{ UpdateBuilder::new({:?}) }}",
216 table.name
217 )
218 .unwrap();
219 writeln!(
220 output,
221 " pub fn delete() -> DeleteBuilder {{ DeleteBuilder::new({:?}) }}",
222 table.name
223 )
224 .unwrap();
225 if let Some(primary_key) = primary_reference_key(table) {
226 let column = table
227 .columns
228 .iter()
229 .find(|candidate| candidate.name == primary_key)
230 .expect("constraint column exists");
231 writeln!(output, " pub fn get(key: {}) -> GeneratedQuery<{record}> {{ GeneratedQuery::new(Self::query().select([Expr::star()]).filter(Self::{}.eq(Expr::value({})))) }}", rust_type(&column.data_type), rust_const_name(primary_key)?, owned_typed_value_expression("key", &column.data_type)).unwrap();
232 }
233 for (key, is_primary) in reference_keys(table) {
234 let column = table
235 .columns
236 .iter()
237 .find(|candidate| candidate.name == key)
238 .expect("constraint column exists");
239 let method = if is_primary {
240 "reference".to_string()
241 } else {
242 format!("reference_by_{}", rust_field_name(key)?)
243 };
244 let encoder = format!("encode_{}_key", rust_field_name(key)?);
245 writeln!(
246 output,
247 " fn {encoder}(key: {}) -> TypedValue {{ {} }}",
248 rust_type(&column.data_type),
249 owned_typed_value_expression("key", &column.data_type)
250 )
251 .unwrap();
252 writeln!(
253 output,
254 " pub fn {method}(key: {}) -> Reference<Self> {{ {}Keys::{}.reference(key) }}",
255 rust_type(&column.data_type),
256 entity,
257 rust_const_name(key)?
258 )
259 .unwrap();
260 }
261 output.push_str("}\n\n");
262
263 if !reference_keys(table).is_empty() {
264 writeln!(
265 output,
266 "#[derive(Debug, Clone, Copy, PartialEq, Eq)]\npub struct {entity}Keys;"
267 )
268 .unwrap();
269 writeln!(output, "impl {entity}Keys {{").unwrap();
270 for (key, is_primary) in reference_keys(table) {
271 let column = table
272 .columns
273 .iter()
274 .find(|candidate| candidate.name == key)
275 .expect("constraint column exists");
276 writeln!(
277 output,
278 " pub const {}: TypedKey<{entity}, {}> = TypedKey::new({entity}::{}, {}, {entity}::encode_{}_key);",
279 rust_const_name(key)?,
280 rust_type(&column.data_type),
281 rust_const_name(key)?,
282 is_primary,
283 rust_field_name(key)?,
284 )
285 .unwrap();
286 }
287 output.push_str("}\n\n");
288 }
289
290 writeln!(output, "impl {record} {{").unwrap();
291 writeln!(output, " pub fn to_dynamic(&self, descriptor: &TableDescriptor) -> Result<DynamicRecord, GeneratedRecordError> {{").unwrap();
292 writeln!(output, " ensure_schema_fingerprint({entity}::SCHEMA_FINGERPRINT, &descriptor.fingerprint)?;").unwrap();
293 output.push_str(" let mut record = DynamicRecord::new(descriptor.clone());\n");
294 for column in &table.columns {
295 let field = rust_field_name(&column.name)?;
296 let typed_value = match column_reference(table, &column.name) {
297 Some(_) => "value.key().clone()".to_string(),
298 None => typed_value_expression("value", &column.data_type),
299 };
300 writeln!(output, " match self.{field}.value() {{").unwrap();
301 output.push_str(" FieldValue::Omitted => {}\n");
302 writeln!(output, " FieldValue::Null {{ data_type }} => if self.{field}.is_dirty() {{ record.set_null({:?})? }} else {{ record.hydrate({:?}, TypedValue::Null(data_type.clone()))? }},", column.name, column.name).unwrap();
303 writeln!(output, " FieldValue::Value {{ value }} => if self.{field}.is_dirty() {{ record.set({:?}, {typed_value})? }} else {{ record.hydrate({:?}, {typed_value})? }},", column.name, column.name).unwrap();
304 output.push_str(" }\n");
305 }
306 output.push_str(" Ok(record)\n }\n");
307 writeln!(output, " pub fn insert<S: OrmGeneratedRecordSession>(&mut self, session: S) -> Result<(), S::Error> {{ session.mutate_generated_record(self, RecordMutation::Insert) }}").unwrap();
308 writeln!(output, " pub fn save<S: OrmGeneratedRecordSession>(&mut self, session: S) -> Result<(), S::Error> {{ session.mutate_generated_record(self, RecordMutation::Save) }}").unwrap();
309 writeln!(output, " pub fn update<S: OrmGeneratedRecordSession>(&mut self, session: S) -> Result<(), S::Error> {{ session.mutate_generated_record(self, RecordMutation::Update) }}").unwrap();
310 writeln!(output, " pub fn delete<S: OrmGeneratedRecordSession>(&mut self, session: S) -> Result<(), S::Error> {{ session.mutate_generated_record(self, RecordMutation::Delete) }}").unwrap();
311 writeln!(output, " pub async fn insert_async<S: AsyncOrmGeneratedRecordSession>(&mut self, session: &mut S) -> Result<(), S::Error> {{ session.mutate_generated_record_async(self, RecordMutation::Insert).await }}").unwrap();
312 writeln!(output, " pub async fn save_async<S: AsyncOrmGeneratedRecordSession>(&mut self, session: &mut S) -> Result<(), S::Error> {{ session.mutate_generated_record_async(self, RecordMutation::Save).await }}").unwrap();
313 writeln!(output, " pub async fn update_async<S: AsyncOrmGeneratedRecordSession>(&mut self, session: &mut S) -> Result<(), S::Error> {{ session.mutate_generated_record_async(self, RecordMutation::Update).await }}").unwrap();
314 writeln!(output, " pub async fn delete_async<S: AsyncOrmGeneratedRecordSession>(&mut self, session: &mut S) -> Result<(), S::Error> {{ session.mutate_generated_record_async(self, RecordMutation::Delete).await }}").unwrap();
315 output.push_str("}\n\n");
316
317 writeln!(output, "impl GeneratedRecord for {record} {{").unwrap();
318 writeln!(output, " type Entity = {entity};").unwrap();
319 writeln!(output, " fn to_dynamic(&self, descriptor: &TableDescriptor) -> Result<DynamicRecord, GeneratedRecordError> {{ {record}::to_dynamic(self, descriptor) }}").unwrap();
320 output.push_str(" fn apply_dynamic(&mut self, record: &DynamicRecord) -> Result<(), GeneratedRecordError> {\n");
321 writeln!(output, " ensure_schema_fingerprint({entity}::SCHEMA_FINGERPRINT, &record.descriptor().fingerprint)?;").unwrap();
322 for column in &table.columns {
323 let field = rust_field_name(&column.name)?;
324 if let Some((target_table, target_column)) = column_reference(table, &column.name) {
325 let target_entity = rust_type_name(target_table)?;
326 writeln!(output, " self.{field}.hydrate(decode_generated_reference_field::<{target_entity}>({:?}, record.field({:?})?, &{}, {:?}, {:?})?);", column.name, column.name, rust_data_type(&column.data_type), target_table, target_column).unwrap();
327 } else {
328 writeln!(output, " self.{field}.hydrate(decode_generated_field({:?}, record.field({:?})?, &{})?);", column.name, column.name, rust_data_type(&column.data_type)).unwrap();
329 }
330 }
331 output.push_str(" Ok(())\n }\n}\n\n");
332 Ok(())
333}
334
335fn validate_reference_shapes(
336 table: &TableDescriptor,
337 tables: &[TableDescriptor],
338) -> Result<(), CodegenError> {
339 for constraint in &table.constraints {
340 let ConstraintDefinition::ForeignKey {
341 columns,
342 referenced_table,
343 referenced_columns,
344 ..
345 } = &constraint.definition
346 else {
347 continue;
348 };
349 if columns.len() != 1 || referenced_columns.len() != 1 {
350 return Err(CodegenError::UnsupportedReference {
351 table: table.name.clone(),
352 column: columns.join(","),
353 });
354 }
355 if !tables.iter().any(|table| table.name == *referenced_table) {
356 return Err(CodegenError::MissingReferenceTarget(
357 referenced_table.clone(),
358 ));
359 }
360 let source_column = table
361 .columns
362 .iter()
363 .find(|column| column.name == columns[0])
364 .ok_or_else(|| CodegenError::MissingReferenceSourceColumn {
365 table: table.name.clone(),
366 column: columns[0].clone(),
367 })?;
368 let target_table = tables
369 .iter()
370 .find(|table| table.name == *referenced_table)
371 .expect("target table presence checked");
372 let target_column = target_table
373 .columns
374 .iter()
375 .find(|column| column.name == referenced_columns[0])
376 .ok_or_else(|| CodegenError::MissingReferenceTargetColumn {
377 table: referenced_table.clone(),
378 column: referenced_columns[0].clone(),
379 })?;
380 if !reference_keys(target_table)
381 .iter()
382 .any(|(column, _)| *column == target_column.name)
383 {
384 return Err(CodegenError::UnsupportedReferenceTargetKey {
385 table: referenced_table.clone(),
386 column: referenced_columns[0].clone(),
387 });
388 }
389 if source_column.data_type != target_column.data_type {
390 return Err(CodegenError::ReferenceTypeMismatch {
391 source_table: table.name.clone(),
392 source_column: source_column.name.clone(),
393 source_type: source_column.data_type.clone(),
394 target_table: target_table.name.clone(),
395 target_column: target_column.name.clone(),
396 target_type: target_column.data_type.clone(),
397 });
398 }
399 }
400 Ok(())
401}
402
403fn column_reference<'a>(table: &'a TableDescriptor, column: &str) -> Option<(&'a str, &'a str)> {
404 table.constraints.iter().find_map(|constraint| {
405 let ConstraintDefinition::ForeignKey {
406 columns,
407 referenced_table,
408 referenced_columns,
409 ..
410 } = &constraint.definition
411 else {
412 return None;
413 };
414 (columns.as_slice() == [column] && referenced_columns.len() == 1)
415 .then(|| (referenced_table.as_str(), referenced_columns[0].as_str()))
416 })
417}
418
419fn primary_reference_key(table: &TableDescriptor) -> Option<&str> {
420 table.constraints.iter().find_map(|constraint| {
421 let ConstraintDefinition::PrimaryKey { columns } = &constraint.definition else {
422 return None;
423 };
424 (columns.len() == 1).then(|| columns[0].as_str())
425 })
426}
427
428fn reference_keys(table: &TableDescriptor) -> Vec<(&str, bool)> {
429 let mut keys = Vec::new();
430 for constraint in &table.constraints {
431 let (columns, is_primary) = match &constraint.definition {
432 ConstraintDefinition::PrimaryKey { columns } => (columns, true),
433 ConstraintDefinition::Unique { columns, .. } => (columns, false),
434 _ => continue,
435 };
436 if columns.len() == 1 {
437 if let Some(column) = table
438 .columns
439 .iter()
440 .find(|column| column.name == columns[0] && !column.nullable)
441 {
442 if !keys.iter().any(|(name, _)| *name == column.name) {
443 keys.push((column.name.as_str(), is_primary));
444 }
445 }
446 }
447 }
448 keys
449}
450
451fn rust_type(data_type: &DataTypeDescriptor) -> &'static str {
452 match data_type {
453 DataTypeDescriptor::Integer => "i64",
454 DataTypeDescriptor::Float => "f64",
455 DataTypeDescriptor::Boolean => "bool",
456 DataTypeDescriptor::Json => "serde_json::Value",
457 DataTypeDescriptor::Bytes => "String",
458 DataTypeDescriptor::Vector { .. } => "Vec<f32>",
459 DataTypeDescriptor::Null
460 | DataTypeDescriptor::Text
461 | DataTypeDescriptor::Timestamp
462 | DataTypeDescriptor::Date
463 | DataTypeDescriptor::Uuid
464 | DataTypeDescriptor::Decimal { .. } => "String",
465 }
466}
467
468fn rust_data_type(data_type: &DataTypeDescriptor) -> String {
469 match data_type {
470 DataTypeDescriptor::Null => "DataTypeDescriptor::Null".to_string(),
471 DataTypeDescriptor::Integer => "DataTypeDescriptor::Integer".to_string(),
472 DataTypeDescriptor::Float => "DataTypeDescriptor::Float".to_string(),
473 DataTypeDescriptor::Text => "DataTypeDescriptor::Text".to_string(),
474 DataTypeDescriptor::Boolean => "DataTypeDescriptor::Boolean".to_string(),
475 DataTypeDescriptor::Timestamp => "DataTypeDescriptor::Timestamp".to_string(),
476 DataTypeDescriptor::Date => "DataTypeDescriptor::Date".to_string(),
477 DataTypeDescriptor::Json => "DataTypeDescriptor::Json".to_string(),
478 DataTypeDescriptor::Uuid => "DataTypeDescriptor::Uuid".to_string(),
479 DataTypeDescriptor::Bytes => "DataTypeDescriptor::Bytes".to_string(),
480 DataTypeDescriptor::Decimal { precision, scale } => {
481 format!("DataTypeDescriptor::Decimal {{ precision: {precision:?}, scale: {scale:?} }}")
482 }
483 DataTypeDescriptor::Vector { dimensions } => {
484 format!("DataTypeDescriptor::Vector {{ dimensions: {dimensions} }}")
485 }
486 }
487}
488
489fn typed_value_expression(value: &str, data_type: &DataTypeDescriptor) -> String {
490 match data_type {
491 DataTypeDescriptor::Integer => format!("TypedValue::Integer(*{value})"),
492 DataTypeDescriptor::Float => format!("TypedValue::Float((*{value}).into())"),
493 DataTypeDescriptor::Text => format!("TypedValue::Text({value}.clone())"),
494 DataTypeDescriptor::Boolean => format!("TypedValue::Boolean(*{value})"),
495 DataTypeDescriptor::Timestamp => format!("TypedValue::Timestamp({value}.clone())"),
496 DataTypeDescriptor::Date => format!("TypedValue::Date({value}.clone())"),
497 DataTypeDescriptor::Json => format!("TypedValue::Json({value}.clone())"),
498 DataTypeDescriptor::Uuid => format!("TypedValue::Uuid({value}.clone())"),
499 DataTypeDescriptor::Bytes => format!("TypedValue::Bytes({value}.clone())"),
500 DataTypeDescriptor::Decimal { .. } => format!("TypedValue::Decimal({value}.clone())"),
501 DataTypeDescriptor::Vector { .. } => format!("TypedValue::Vector({value}.clone())"),
502 DataTypeDescriptor::Null => "TypedValue::Null(DataTypeDescriptor::Null)".to_string(),
503 }
504}
505
506fn owned_typed_value_expression(value: &str, data_type: &DataTypeDescriptor) -> String {
507 match data_type {
508 DataTypeDescriptor::Integer => format!("TypedValue::Integer({value})"),
509 DataTypeDescriptor::Float => format!("TypedValue::Float({value}.into())"),
510 DataTypeDescriptor::Text => format!("TypedValue::Text({value})"),
511 DataTypeDescriptor::Boolean => format!("TypedValue::Boolean({value})"),
512 DataTypeDescriptor::Timestamp => format!("TypedValue::Timestamp({value})"),
513 DataTypeDescriptor::Date => format!("TypedValue::Date({value})"),
514 DataTypeDescriptor::Json => format!("TypedValue::Json({value})"),
515 DataTypeDescriptor::Uuid => format!("TypedValue::Uuid({value})"),
516 DataTypeDescriptor::Bytes => format!("TypedValue::Bytes({value})"),
517 DataTypeDescriptor::Decimal { .. } => format!("TypedValue::Decimal({value})"),
518 DataTypeDescriptor::Vector { .. } => format!("TypedValue::Vector({value})"),
519 DataTypeDescriptor::Null => "TypedValue::Null(DataTypeDescriptor::Null)".to_string(),
520 }
521}
522
523fn rust_type_name(identifier: &str) -> Result<String, CodegenError> {
524 let parts = identifier_parts(identifier)?;
525 Ok(parts.into_iter().map(capitalize).collect())
526}
527
528fn rust_const_name(identifier: &str) -> Result<String, CodegenError> {
529 Ok(identifier_parts(identifier)?.join("_").to_ascii_uppercase())
530}
531
532fn rust_field_name(identifier: &str) -> Result<String, CodegenError> {
533 let mut name = identifier_parts(identifier)?.join("_").to_ascii_lowercase();
534 if is_rust_keyword(&name) {
535 name.insert_str(0, "r#");
536 }
537 Ok(name)
538}
539
540fn identifier_parts(identifier: &str) -> Result<Vec<&str>, CodegenError> {
541 let parts: Vec<_> = identifier
542 .split(|character: char| !character.is_ascii_alphanumeric())
543 .filter(|part| !part.is_empty())
544 .collect();
545 if parts.is_empty() || parts[0].as_bytes().first().is_some_and(u8::is_ascii_digit) {
546 return Err(CodegenError::Identifier(identifier.to_string()));
547 }
548 Ok(parts)
549}
550
551fn capitalize(part: &str) -> String {
552 let mut chars = part.chars();
553 match chars.next() {
554 Some(first) => first.to_ascii_uppercase().to_string() + chars.as_str(),
555 None => String::new(),
556 }
557}
558
559fn is_rust_keyword(value: &str) -> bool {
560 matches!(
561 value,
562 "as" | "break"
563 | "const"
564 | "continue"
565 | "crate"
566 | "else"
567 | "enum"
568 | "extern"
569 | "false"
570 | "fn"
571 | "for"
572 | "if"
573 | "impl"
574 | "in"
575 | "let"
576 | "loop"
577 | "match"
578 | "mod"
579 | "move"
580 | "mut"
581 | "pub"
582 | "ref"
583 | "return"
584 | "self"
585 | "Self"
586 | "static"
587 | "struct"
588 | "super"
589 | "trait"
590 | "true"
591 | "type"
592 | "unsafe"
593 | "use"
594 | "where"
595 | "while"
596 | "async"
597 | "await"
598 | "dyn"
599 )
600}
601
602#[cfg(test)]
603mod tests {
604 use std::collections::BTreeMap;
605 use std::fs;
606 use std::process::Command;
607
608 use super::*;
609
610 #[test]
611 fn codegen_is_deterministic_offline_and_fingerprint_bound() {
612 let mut table = TableDescriptor {
613 catalog_id: "018f2b34-7a10-7cc2-8f3a-9d4b5c6d7e01".to_string(),
614 name: "people".to_string(),
615 schema_generation: 1,
616 fingerprint: String::new(),
617 created_at: "x".to_string(),
618 updated_at: "x".to_string(),
619 columns: vec![ColumnDescriptor {
620 ordinal: 0,
621 name: "id".to_string(),
622 data_type: DataTypeDescriptor::Integer,
623 nullable: false,
624 auto_increment: false,
625 default_expression: None,
626 extensions: BTreeMap::new(),
627 }],
628 constraints: vec![ConstraintDescriptor {
629 id: 1,
630 name: "pk_people".to_string(),
631 definition: ConstraintDefinition::PrimaryKey {
632 columns: vec!["id".to_string()],
633 },
634 }],
635 indexes: Vec::new(),
636 extensions: BTreeMap::new(),
637 };
638 table.refresh_fingerprint().unwrap();
639 let mut database = DatabaseDescriptor {
640 schema_generation: 1,
641 fingerprint: String::new(),
642 tables: vec![table],
643 views: Vec::new(),
644 extensions: BTreeMap::new(),
645 };
646 database.refresh_fingerprint().unwrap();
647 let envelope = DescriptorEnvelope::new(DescriptorKind::Database, database);
648 let first = generate_rust_database(&envelope).unwrap();
649 let second = generate_rust_database(&envelope).unwrap();
650 assert_eq!(first, second);
651 assert!(first
652 .source
653 .contains("pub const ID: TypedColumn<i64, Self>"));
654 assert!(first.source.contains("pub const ID: TypedKey<People, i64>"));
655 assert!(first.source.contains("pub fn reference(key: i64)"));
656
657 let mut stale = envelope.clone();
658 stale.payload.tables[0].columns[0].nullable = true;
659 assert!(matches!(
660 generate_rust_database(&stale),
661 Err(CodegenError::Fingerprint { .. })
662 ));
663 }
664
665 #[test]
666 fn codegen_rejects_unusable_or_type_mismatched_reference_targets() {
667 fn descriptor(
668 target_type: DataTypeDescriptor,
669 target_nullable: bool,
670 target_constraint: ConstraintDefinition,
671 target_column: &str,
672 ) -> DescriptorEnvelope<DatabaseDescriptor> {
673 let mut target = TableDescriptor {
674 catalog_id: "target-id".to_string(),
675 name: "targets".to_string(),
676 schema_generation: 1,
677 fingerprint: String::new(),
678 created_at: "x".to_string(),
679 updated_at: "x".to_string(),
680 columns: vec![ColumnDescriptor {
681 ordinal: 0,
682 name: "id".to_string(),
683 data_type: target_type,
684 nullable: target_nullable,
685 auto_increment: false,
686 default_expression: None,
687 extensions: BTreeMap::new(),
688 }],
689 constraints: vec![ConstraintDescriptor {
690 id: 1,
691 name: "target_key".to_string(),
692 definition: target_constraint,
693 }],
694 indexes: Vec::new(),
695 extensions: BTreeMap::new(),
696 };
697 target.refresh_fingerprint().unwrap();
698 let mut source = TableDescriptor {
699 catalog_id: "source-id".to_string(),
700 name: "sources".to_string(),
701 schema_generation: 1,
702 fingerprint: String::new(),
703 created_at: "x".to_string(),
704 updated_at: "x".to_string(),
705 columns: vec![ColumnDescriptor {
706 ordinal: 0,
707 name: "target_id".to_string(),
708 data_type: DataTypeDescriptor::Integer,
709 nullable: false,
710 auto_increment: false,
711 default_expression: None,
712 extensions: BTreeMap::new(),
713 }],
714 constraints: vec![ConstraintDescriptor {
715 id: 1,
716 name: "fk_sources_target_id___targets".to_string(),
717 definition: ConstraintDefinition::ForeignKey {
718 columns: vec!["target_id".to_string()],
719 referenced_table: "targets".to_string(),
720 referenced_columns: vec![target_column.to_string()],
721 on_delete: ForeignKeyActionDescriptor::Restrict,
722 on_update: ForeignKeyActionDescriptor::Restrict,
723 },
724 }],
725 indexes: Vec::new(),
726 extensions: BTreeMap::new(),
727 };
728 source.refresh_fingerprint().unwrap();
729 let mut database = DatabaseDescriptor {
730 schema_generation: 1,
731 fingerprint: String::new(),
732 tables: vec![target, source],
733 views: Vec::new(),
734 extensions: BTreeMap::new(),
735 };
736 database.refresh_fingerprint().unwrap();
737 DescriptorEnvelope::new(DescriptorKind::Database, database)
738 }
739
740 let missing = descriptor(
741 DataTypeDescriptor::Integer,
742 false,
743 ConstraintDefinition::PrimaryKey {
744 columns: vec!["id".to_string()],
745 },
746 "absent",
747 );
748 assert!(matches!(
749 generate_rust_database(&missing),
750 Err(CodegenError::MissingReferenceTargetColumn { .. })
751 ));
752
753 let ordinary = descriptor(
754 DataTypeDescriptor::Integer,
755 false,
756 ConstraintDefinition::Check {
757 column: Some("id".to_string()),
758 expression: "id > 0".to_string(),
759 ordinal: 1,
760 },
761 "id",
762 );
763 assert!(matches!(
764 generate_rust_database(&ordinary),
765 Err(CodegenError::UnsupportedReferenceTargetKey { .. })
766 ));
767
768 let nullable_unique = descriptor(
769 DataTypeDescriptor::Integer,
770 true,
771 ConstraintDefinition::Unique {
772 columns: vec!["id".to_string()],
773 owned_index: "uq_targets_id".to_string(),
774 },
775 "id",
776 );
777 assert!(matches!(
778 generate_rust_database(&nullable_unique),
779 Err(CodegenError::UnsupportedReferenceTargetKey { .. })
780 ));
781
782 let wrong_type = descriptor(
783 DataTypeDescriptor::Text,
784 false,
785 ConstraintDefinition::PrimaryKey {
786 columns: vec!["id".to_string()],
787 },
788 "id",
789 );
790 assert!(matches!(
791 generate_rust_database(&wrong_type),
792 Err(CodegenError::ReferenceTypeMismatch { .. })
793 ));
794 }
795
796 #[test]
797 fn generated_source_compiles_as_an_independent_offline_crate() {
798 let mut table = TableDescriptor {
799 catalog_id: "018f2b34-7a10-7cc2-8f3a-9d4b5c6d7e01".to_string(),
800 name: "people".to_string(),
801 schema_generation: 1,
802 fingerprint: String::new(),
803 created_at: "x".to_string(),
804 updated_at: "x".to_string(),
805 columns: vec![
806 ColumnDescriptor {
807 ordinal: 0,
808 name: "id".to_string(),
809 data_type: DataTypeDescriptor::Integer,
810 nullable: false,
811 auto_increment: false,
812 default_expression: None,
813 extensions: BTreeMap::new(),
814 },
815 ColumnDescriptor {
816 ordinal: 1,
817 name: "metadata".to_string(),
818 data_type: DataTypeDescriptor::Json,
819 nullable: true,
820 auto_increment: false,
821 default_expression: None,
822 extensions: BTreeMap::new(),
823 },
824 ColumnDescriptor {
825 ordinal: 2,
826 name: "fio_id".to_string(),
827 data_type: DataTypeDescriptor::Integer,
828 nullable: true,
829 auto_increment: false,
830 default_expression: None,
831 extensions: BTreeMap::new(),
832 },
833 ],
834 constraints: vec![
835 ConstraintDescriptor {
836 id: 1,
837 name: "pk_people".to_string(),
838 definition: ConstraintDefinition::PrimaryKey {
839 columns: vec!["id".to_string()],
840 },
841 },
842 ConstraintDescriptor {
843 id: 2,
844 name: "fk_people_fio_id___fio".to_string(),
845 definition: ConstraintDefinition::ForeignKey {
846 columns: vec!["fio_id".to_string()],
847 referenced_table: "fio".to_string(),
848 referenced_columns: vec!["id".to_string()],
849 on_delete: ForeignKeyActionDescriptor::Restrict,
850 on_update: ForeignKeyActionDescriptor::Restrict,
851 },
852 },
853 ],
854 indexes: Vec::new(),
855 extensions: BTreeMap::new(),
856 };
857 table.refresh_fingerprint().unwrap();
858 let mut fio = TableDescriptor {
859 catalog_id: "018f2b34-7a10-7cc2-8f3a-9d4b5c6d7e02".to_string(),
860 name: "fio".to_string(),
861 schema_generation: 1,
862 fingerprint: String::new(),
863 created_at: "x".to_string(),
864 updated_at: "x".to_string(),
865 columns: vec![ColumnDescriptor {
866 ordinal: 0,
867 name: "id".to_string(),
868 data_type: DataTypeDescriptor::Integer,
869 nullable: false,
870 auto_increment: false,
871 default_expression: None,
872 extensions: BTreeMap::new(),
873 }],
874 constraints: vec![ConstraintDescriptor {
875 id: 1,
876 name: "pk_fio".to_string(),
877 definition: ConstraintDefinition::PrimaryKey {
878 columns: vec!["id".to_string()],
879 },
880 }],
881 indexes: Vec::new(),
882 extensions: BTreeMap::new(),
883 };
884 fio.refresh_fingerprint().unwrap();
885 let mut database = DatabaseDescriptor {
886 schema_generation: 1,
887 fingerprint: String::new(),
888 tables: vec![fio, table],
889 views: Vec::new(),
890 extensions: BTreeMap::new(),
891 };
892 database.refresh_fingerprint().unwrap();
893 let mut generated =
894 generate_rust_database(&DescriptorEnvelope::new(DescriptorKind::Database, database))
895 .unwrap();
896 generated.source.push_str(
897 "\n#[allow(dead_code)]\nfn compile_generated_insert<S: radixdb_orm::OrmGeneratedRecordSession>(mut record: PeopleRecord, session: S) -> Result<(), S::Error> { record.insert(session) }\n",
898 );
899 assert!(generated.source.contains("FieldState<Reference<Fio>>"));
900 generated.source.push_str(
901 "\n#[allow(dead_code)]\nfn compile_generated_get<S>(session: S) -> Result<Option<PeopleRecord>, S::Error> where S: radixdb_orm::OrmGeneratedQuerySession, S::Error: From<radixdb_orm::RecordError> + From<radixdb_orm::BuilderError> { People::get(1).optional(session) }\n",
902 );
903 generated.source.push_str(
904 "\n#[allow(dead_code)]\nfn compile_generated_null(mut record: PeopleRecord) { record.metadata.set_null(); record.metadata.unset(); }\n",
905 );
906 generated.source.push_str(
907 "\n#[allow(dead_code)]\nasync fn compile_generated_async<S>(mut record: PeopleRecord, session: &mut S) -> Result<Option<PeopleRecord>, <S as radixdb_orm::AsyncOrmGeneratedRecordSession>::Error> where S: radixdb_orm::AsyncOrmGeneratedRecordSession + radixdb_orm::AsyncOrmGeneratedQuerySession<Error = <S as radixdb_orm::AsyncOrmGeneratedRecordSession>::Error>, <S as radixdb_orm::AsyncOrmGeneratedRecordSession>::Error: From<radixdb_orm::RecordError> + From<radixdb_orm::BuilderError> { record.insert_async(session).await?; People::get(1).optional_async(session).await }\n",
908 );
909
910 let temp = std::env::temp_dir().join(format!("radixdb-orm-codegen-{}", std::process::id()));
911 let _ = fs::remove_dir_all(&temp);
912 fs::create_dir_all(temp.join("src")).unwrap();
913 let manifest_dir = env!("CARGO_MANIFEST_DIR").replace('\\', "\\\\");
914 fs::write(
915 temp.join("Cargo.toml"),
916 format!(
917 "[package]\nname = \"radixdb-orm-generated-probe\"\nversion = \"0.0.0\"\nedition = \"2021\"\n\n[dependencies]\nradixdb-orm = {{ path = {:?} }}\nserde_json = \"1\"\n",
918 manifest_dir
919 ),
920 )
921 .unwrap();
922 fs::write(temp.join("src/lib.rs"), generated.source).unwrap();
923 let status = Command::new(std::env::var_os("CARGO").unwrap_or_else(|| "cargo".into()))
924 .args(["check", "--offline", "--quiet"])
925 .current_dir(&temp)
926 .env("CARGO_TARGET_DIR", temp.join("target"))
927 .status()
928 .unwrap();
929 let _ = fs::remove_dir_all(&temp);
930 assert!(status.success(), "generated Rust facade did not compile");
931 }
932}