rustyroad 1.8.1

Rusty Road is a framework written in Rust that is based on Ruby on Rails. It is designed to provide the familiar conventions and ease of use of Ruby on Rails, while also taking advantage of the performance and efficiency of Rust.
Documentation
//! Rust model rendering.

use super::naming::{literal, pascal, snake};
use super::types;
use crate::database::introspection::{Column, Enum, Schema, Table};
use std::collections::HashSet;

/// Renders `models.rs`.
pub(super) fn render(schema: &Schema) -> String {
    let mut output = String::from(
        "// Generated by `rustyroad pull --language rust`; do not edit.\n\n\
         /// JSON-friendly representation of PostgreSQL's interval type.\n\
         #[derive(Debug, Clone, Copy, PartialEq, Eq, sqlx::Type)]\n\
         #[sqlx(transparent)]\n\
         pub struct Interval(pub sqlx::postgres::types::PgInterval);\n\n\
         #[derive(serde::Serialize, serde::Deserialize)]\n\
         struct IntervalWire {\n\
             months: i32,\n\
             days: i32,\n\
             microseconds: i64,\n\
         }\n\n\
         impl serde::Serialize for Interval {\n\
             fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>\n\
             where\n\
                 S: serde::Serializer,\n\
             {\n\
                 serde::Serialize::serialize(&IntervalWire {\n\
                     months: self.0.months,\n\
                     days: self.0.days,\n\
                     microseconds: self.0.microseconds,\n\
                 }, serializer)\n\
             }\n\
         }\n\n\
         impl<'de> serde::Deserialize<'de> for Interval {\n\
             fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>\n\
             where\n\
                 D: serde::Deserializer<'de>,\n\
             {\n\
                 let wire = <IntervalWire as serde::Deserialize>::deserialize(deserializer)?;\n\
                 Ok(Self(sqlx::postgres::types::PgInterval {\n\
                     months: wire.months,\n\
                     days: wire.days,\n\
                     microseconds: wire.microseconds,\n\
                 }))\n\
             }\n\
         }\n\n",
    );

    for item in &schema.enums {
        output.push_str(&render_enum(item));
    }
    for table in &schema.tables {
        output.push_str(&render_table(table, schema));
    }
    output
}

fn render_enum(item: &Enum) -> String {
    let mut used = HashSet::new();
    let variants = item
        .values
        .iter()
        .map(|value| {
            let variant = unique_variant(value, &mut used);
            format!(
                "    #[sqlx(rename = {})]\n    #[serde(rename = {})]\n    {},",
                literal(value),
                literal(value),
                variant
            )
        })
        .collect::<Vec<_>>()
        .join("\n");

    format!(
        "#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize, sqlx::Type)]\n\
         #[sqlx(type_name = {})]\n\
         pub enum {} {{\n{}\n}}\n\n",
        literal(&item.name),
        pascal(&item.name),
        variants
    )
}

fn unique_variant(value: &str, used: &mut HashSet<String>) -> String {
    let base = pascal(value);
    if used.insert(base.clone()) {
        return base;
    }

    let mut sequence = 2;
    loop {
        let candidate = format!("{base}Variant{sequence}");
        if used.insert(candidate.clone()) {
            return candidate;
        }
        sequence += 1;
    }
}

fn render_table(table: &Table, schema: &Schema) -> String {
    let name = pascal(&table.name);
    let row_fields = fields(
        &table.columns,
        |column| column_type(column, schema),
        true,
        false,
    );
    let create_columns = table
        .columns
        .iter()
        .filter(|column| !column.auto_increment)
        .cloned()
        .collect::<Vec<_>>();
    let create_fields = fields(
        &create_columns,
        |column| create_type(column, schema),
        false,
        false,
    );
    let patch_columns = table
        .columns
        .iter()
        .filter(|column| !table.primary_key.contains(&column.name))
        .cloned()
        .collect::<Vec<_>>();
    let patch_fields = fields(
        &patch_columns,
        |column| patch_type(column, schema),
        false,
        true,
    );

    format!(
        "/// A row from the `{table_name}` table.\n\
         #[derive(Debug, Clone, serde::Serialize, serde::Deserialize, sqlx::FromRow)]\n\
         pub struct {name} {{\n{row_fields}\n}}\n\n\
         /// Values accepted when creating a `{table_name}` row.\n\
         #[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]\n\
         pub struct New{name} {{\n{create_fields}\n}}\n\n\
         /// Values accepted when partially updating a `{table_name}` row.\n\
         #[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)]\n\
         pub struct Patch{name} {{\n{patch_fields}\n}}\n\n",
        table_name = table.name,
    )
}

fn fields<F>(columns: &[Column], type_for: F, sqlx: bool, all_optional: bool) -> String
where
    F: Fn(&Column) -> String,
{
    columns
        .iter()
        .map(|column| {
            let field = snake(&column.name);
            let mut attributes = Vec::new();
            if field != column.name {
                if sqlx {
                    attributes.push(format!("    #[sqlx(rename = {})]", literal(&column.name)));
                }
                attributes.push(format!("    #[serde(rename = {})]", literal(&column.name)));
            }
            if !sqlx && (all_optional || is_optional_input(column)) {
                attributes.push(
                    "    #[serde(default, skip_serializing_if = \"Option::is_none\")]".to_string(),
                );
            }
            let prefix = if attributes.is_empty() {
                String::new()
            } else {
                format!("{}\n", attributes.join("\n"))
            };
            format!("{prefix}    pub {field}: {},", type_for(column))
        })
        .collect::<Vec<_>>()
        .join("\n")
}

fn create_type(column: &Column, schema: &Schema) -> String {
    let base = base_type(column, schema);
    if column.nullable && column.default.is_some() {
        format!("Option<Option<{base}>>")
    } else if is_optional_input(column) {
        format!("Option<{base}>")
    } else {
        base
    }
}

fn patch_type(column: &Column, schema: &Schema) -> String {
    let base = base_type(column, schema);
    if column.nullable {
        // Outer Option means "field omitted"; inner Option means JSON null.
        format!("Option<Option<{base}>>")
    } else {
        format!("Option<{base}>")
    }
}

fn is_optional_input(column: &Column) -> bool {
    column.nullable || column.default.is_some() || column.auto_increment
}

fn column_type(column: &Column, schema: &Schema) -> String {
    types::map(column, schema).rust
}

fn base_type(column: &Column, schema: &Schema) -> String {
    let mut required = column.clone();
    required.nullable = false;
    types::map(&required, schema).rust
}