Skip to main content

radixdb_orm/
codegen.rs

1//! Deterministic Rust facade generation from saved schema descriptors.
2
3use 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}