use super::naming::{literal, pascal, snake};
use super::types::{base_type, column_type};
use crate::database::introspection::{Column, Enum, Schema, Table};
use std::collections::HashSet;
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.sql_type, 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.sql_type, schema);
if column.nullable {
format!("Option<Option<{base}>>")
} else {
format!("Option<{base}>")
}
}
fn is_optional_input(column: &Column) -> bool {
column.nullable || column.default.is_some() || column.auto_increment
}