use crate::codegen::traits::file_writer::FileInfo;
use crate::ir::types::{
IrEnum, IrEnumValueType, IrIntersection, IrObject, IrPrimitive, IrSchema, IrSchemaKind, IrSpec,
IrTaggedUnion, IrTypeExpr, IrUnion, TaggingStyle,
};
use heck::{ToPascalCase, ToSnakeCase};
use sigil_stitch::prelude::{CodeBlock, sigil_quote};
use sigil_stitch::spec::file_spec::FileSpec;
use sigil_stitch::spec::import_spec::ImportSpec;
use sigil_stitch::type_name::TypeName;
use super::config::{ExtraDeriveConfig, RustGeneratorConfig};
pub fn generate_model_files(
ir: &IrSpec,
header: &str,
config: &RustGeneratorConfig,
) -> Result<Vec<FileInfo>, String> {
let mut files = Vec::new();
let mut mod_entries = Vec::new();
for (_name, schema) in &ir.schemas {
let Some(file_spec) = emit_model_file(schema, config) else {
return Err(format!(
"unsupported schema kind for {}: {:?}",
schema.name, schema.kind
));
};
let stem = schema.name.to_snake_case();
let filename = format!("{stem}.rs");
mod_entries.push(stem);
let rendered = file_spec
.render(100)
.map_err(|e| format!("render error for {}: {e}", schema.name))?;
let mut content = String::with_capacity(header.len() + rendered.len());
content.push_str(header);
content.push_str(&rendered);
files.push(FileInfo::model(filename, content));
}
let mut mod_content = String::from(header);
for entry in &mod_entries {
mod_content.push_str(&format!("mod {entry};\npub use {entry}::*;\n"));
}
files.push(FileInfo::model("mod.rs".to_string(), mod_content));
Ok(files)
}
fn emit_model_file(schema: &IrSchema, config: &RustGeneratorConfig) -> Option<FileSpec> {
let extra = config.extra_derives.as_ref();
match &schema.kind {
IrSchemaKind::Object(obj) => {
emit_object(schema, obj, extra.and_then(|e| e.structs.as_ref()))
}
IrSchemaKind::Enum(en) => emit_enum(schema, en, extra.and_then(|e| e.enums.as_ref())),
IrSchemaKind::Alias(expr) => emit_alias(schema, expr),
IrSchemaKind::Union(u) => emit_union(schema, u, extra.and_then(|e| e.unions.as_ref())),
IrSchemaKind::Intersection(i) => {
emit_intersection(schema, i, extra.and_then(|e| e.structs.as_ref()))
}
IrSchemaKind::TaggedUnion(tu) => {
emit_tagged_union(schema, tu, extra.and_then(|e| e.unions.as_ref()))
}
}
}
fn derive_attr(base: &str, extra: Option<&ExtraDeriveConfig>) -> String {
match extra {
Some(cfg) if !cfg.derives.is_empty() => {
format!("#[derive({base}, {})]", cfg.derives.join(", "))
}
_ => format!("#[derive({base})]"),
}
}
fn emit_object(
schema: &IrSchema,
obj: &IrObject,
extra: Option<&ExtraDeriveConfig>,
) -> Option<FileSpec> {
let name = schema.name.to_pascal_case();
let stem = schema.name.to_snake_case();
let mut fsb = FileSpec::builder(&format!("{stem}.rs"));
fsb = fsb.add_import(ImportSpec::named("serde", "Deserialize"));
fsb = fsb.add_import(ImportSpec::named("serde", "Serialize"));
let body = sigil_quote!(RustLang {
$if(schema.description.is_some()) {
$L(doc_comment_block(schema.description.as_deref().unwrap()).trim_end())
}
$L(derive_attr("Debug, Clone, Serialize, Deserialize", extra))
pub struct $N(name.as_str()) {
$for((json_name, prop) in obj.properties.iter()) {
$if(prop.description.is_some()) {
$L(doc_comment_block(prop.description.as_deref().unwrap()).trim_end())
}
$if(escape_rust_keyword(&json_name.to_snake_case()) != *json_name) {
$L(format!("#[serde(rename = \"{json_name}\")]"))
}
$if(!prop.required || prop.nullable) {
#[serde(skip_serializing_if = "Option::is_none", default)]
$L(format!("pub {}: Option<{}>,", escape_rust_keyword(&json_name.to_snake_case()), rust_type_str_model(&prop.type_expr)))
} $else {
$L(format!("pub {}: {},", escape_rust_keyword(&json_name.to_snake_case()), rust_type_str_model(&prop.type_expr)))
}
}
$if(obj.additional_properties.is_some()) {
#[serde(flatten)]
$L(format!("pub additional_properties: std::collections::HashMap<String, {}>,", rust_type_str_model(obj.additional_properties.as_ref().unwrap())))
}
}
})
.ok()?;
fsb = fsb.add_code(body);
fsb.build().ok()
}
fn emit_enum(
schema: &IrSchema,
en: &IrEnum,
extra: Option<&ExtraDeriveConfig>,
) -> Option<FileSpec> {
let name = schema.name.to_pascal_case();
match en.value_type {
IrEnumValueType::Mixed | IrEnumValueType::Number => {
return emit_type_alias_file(schema, "serde_json::Value");
}
IrEnumValueType::Integer => {
return emit_integer_enum(schema, en, extra);
}
IrEnumValueType::String => {}
}
let stem = schema.name.to_snake_case();
let mut fsb = FileSpec::builder(&format!("{stem}.rs"));
fsb = fsb.add_import(ImportSpec::named("serde", "Deserialize"));
fsb = fsb.add_import(ImportSpec::named("serde", "Serialize"));
let mut variants: Vec<(String, String)> = Vec::new();
for v in &en.values {
let s = v.value.as_str()?;
variants.push((s.to_pascal_case(), s.to_string()));
}
let body = sigil_quote!(RustLang {
$if(schema.description.is_some()) {
$L(doc_comment_block(schema.description.as_deref().unwrap()).trim_end())
}
$L(derive_attr("Debug, Clone, PartialEq, Eq, Serialize, Deserialize", extra))
pub enum $N(name.as_str()) {
$for((variant, wire) in variants.iter()) {
$if(variant != wire) {
$L(format!("#[serde(rename = \"{}\")]", escape_str(wire)))
}
$L(format!("{variant},"))
}
}
})
.ok()?;
fsb = fsb.add_code(body);
if let Some(display_block) = build_string_enum_display(&name, &variants) {
fsb = fsb.add_code(display_block);
}
fsb.build().ok()
}
fn emit_integer_enum(
schema: &IrSchema,
en: &IrEnum,
extra: Option<&ExtraDeriveConfig>,
) -> Option<FileSpec> {
let name = schema.name.to_pascal_case();
let stem = schema.name.to_snake_case();
let mut fsb = FileSpec::builder(&format!("{stem}.rs"));
fsb = fsb.add_import(ImportSpec::named("serde_repr", "Deserialize_repr"));
fsb = fsb.add_import(ImportSpec::named("serde_repr", "Serialize_repr"));
let int_variants: Vec<(String, i64)> = en
.values
.iter()
.map(|v| {
let n = v.value.as_i64()?;
let variant_name = if n < 0 {
format!("Neg{}", n.unsigned_abs())
} else {
format!("N{n}")
};
Some((variant_name, n))
})
.collect::<Option<Vec<_>>>()?;
let body = sigil_quote!(RustLang {
$if(schema.description.is_some()) {
$L(doc_comment_block(schema.description.as_deref().unwrap()).trim_end())
}
$L(derive_attr("Debug, Clone, Copy, PartialEq, Eq, Serialize_repr, Deserialize_repr", extra))
#[repr(i64)]
pub enum $N(name.as_str()) {
$for((variant_name, n) in int_variants.iter()) {
$L(format!("{variant_name} = {n},"))
}
}
})
.ok()?;
fsb = fsb.add_code(body);
if let Some(display_block) = build_integer_enum_display(&name) {
fsb = fsb.add_code(display_block);
}
fsb.build().ok()
}
fn emit_alias(schema: &IrSchema, expr: &IrTypeExpr) -> Option<FileSpec> {
let name = schema.name.to_pascal_case();
let stem = schema.name.to_snake_case();
let rhs = rust_type_str_model(expr);
let rhs_type = TypeName::raw(&rhs);
let mut fsb = FileSpec::builder(&format!("{stem}.rs"));
let block = sigil_quote!(RustLang {
$if(schema.description.is_some()) {
$L(doc_comment_block(schema.description.as_deref().unwrap()).trim_end())
}
pub type $N(name.as_str()) = $T(rhs_type);
})
.ok()?;
fsb = fsb.add_code(block);
fsb.build().ok()
}
fn emit_type_alias_file(schema: &IrSchema, rhs_str: &str) -> Option<FileSpec> {
let name = schema.name.to_pascal_case();
let stem = schema.name.to_snake_case();
let mut fsb = FileSpec::builder(&format!("{stem}.rs"));
let rhs_type = TypeName::raw(rhs_str);
let block = sigil_quote!(RustLang {
$if(schema.description.is_some()) {
$L(doc_comment_block(schema.description.as_deref().unwrap()).trim_end())
}
pub type $N(name.as_str()) = $T(rhs_type);
})
.ok()?;
fsb = fsb.add_code(block);
fsb.build().ok()
}
fn emit_union(
schema: &IrSchema,
union: &IrUnion,
extra: Option<&ExtraDeriveConfig>,
) -> Option<FileSpec> {
let name = schema.name.to_pascal_case();
let stem = schema.name.to_snake_case();
let mut fsb = FileSpec::builder(&format!("{stem}.rs"));
fsb = fsb.add_import(ImportSpec::named("serde", "Deserialize"));
fsb = fsb.add_import(ImportSpec::named("serde", "Serialize"));
let variants: Vec<(String, String)> = union
.members
.iter()
.enumerate()
.map(|(i, member)| {
let variant_name = union_variant_name(member, i);
let rust_type = rust_type_str_model(member);
(variant_name, rust_type)
})
.collect();
let body = sigil_quote!(RustLang {
$if(schema.description.is_some()) {
$L(doc_comment_block(schema.description.as_deref().unwrap()).trim_end())
}
$L(derive_attr("Debug, Clone, Serialize, Deserialize", extra))
#[serde(untagged)]
pub enum $N(name.as_str()) {
$for((variant_name, rust_type) in variants.iter()) {
$L(format!("{variant_name}({rust_type}),"))
}
}
})
.ok()?;
fsb = fsb.add_code(body);
fsb.build().ok()
}
fn union_variant_name(expr: &IrTypeExpr, index: usize) -> String {
match expr {
IrTypeExpr::Named(n) => n.to_pascal_case(),
IrTypeExpr::Primitive(p) => primitive_variant_name(p),
IrTypeExpr::Array(_) => format!("Array{index}"),
IrTypeExpr::Map(_) => format!("Map{index}"),
_ => format!("Variant{index}"),
}
}
fn primitive_variant_name(p: &IrPrimitive) -> String {
match p {
IrPrimitive::String | IrPrimitive::StringWithFormat(_) => "String".to_string(),
IrPrimitive::Integer | IrPrimitive::IntegerWithFormat(_) => "Integer".to_string(),
IrPrimitive::Number | IrPrimitive::NumberWithFormat(_) => "Number".to_string(),
IrPrimitive::Boolean => "Boolean".to_string(),
IrPrimitive::Binary => "Binary".to_string(),
IrPrimitive::Date => "Date".to_string(),
IrPrimitive::DateTime => "DateTime".to_string(),
IrPrimitive::Uuid => "Uuid".to_string(),
}
}
fn emit_tagged_union(
schema: &IrSchema,
tu: &IrTaggedUnion,
extra: Option<&ExtraDeriveConfig>,
) -> Option<FileSpec> {
if tu.variants.is_empty() {
return None;
}
let name = schema.name.to_pascal_case();
let stem = schema.name.to_snake_case();
let mut fsb = FileSpec::builder(&format!("{stem}.rs"));
fsb = fsb.add_import(ImportSpec::named("serde", "Deserialize"));
fsb = fsb.add_import(ImportSpec::named("serde", "Serialize"));
let serde_tag_attr = match &tu.tagging {
TaggingStyle::Internal => {
format!(
"#[serde(tag = \"{}\")]",
escape_str(&tu.discriminator_field)
)
}
TaggingStyle::Adjacent { content_field } => {
format!(
"#[serde(tag = \"{}\", content = \"{}\")]",
escape_str(&tu.discriminator_field),
escape_str(content_field)
)
}
TaggingStyle::External => String::new(),
};
let body = sigil_quote!(RustLang {
$if(schema.description.is_some()) {
$L(doc_comment_block(schema.description.as_deref().unwrap()).trim_end())
}
$L(derive_attr("Debug, Clone, Serialize, Deserialize", extra))
$if(!serde_tag_attr.is_empty()) {
$L(serde_tag_attr.as_str())
}
pub enum $N(name.as_str()) {
$for(variant in tu.variants.iter()) {
$if(variant.discriminator_value.to_pascal_case() != variant.discriminator_value) {
$L(format!("#[serde(rename = \"{}\")]", escape_str(&variant.discriminator_value)))
}
$L(format!("{}({}),", variant.discriminator_value.to_pascal_case(), rust_type_str_model(&variant.content_type)))
}
}
})
.ok()?;
fsb = fsb.add_code(body);
fsb.build().ok()
}
fn emit_intersection(
schema: &IrSchema,
inter: &IrIntersection,
extra: Option<&ExtraDeriveConfig>,
) -> Option<FileSpec> {
let name = schema.name.to_pascal_case();
let stem = schema.name.to_snake_case();
let mut fsb = FileSpec::builder(&format!("{stem}.rs"));
fsb = fsb.add_import(ImportSpec::named("serde", "Deserialize"));
fsb = fsb.add_import(ImportSpec::named("serde", "Serialize"));
let fields: Vec<(String, String)> = inter
.members
.iter()
.enumerate()
.map(|(i, member)| {
let raw_name = match member {
IrTypeExpr::Named(n) => n.to_snake_case(),
_ => format!("member_{i}"),
};
let field_name = escape_rust_keyword(&raw_name);
let rust_type = rust_type_str_model(member);
(field_name, rust_type)
})
.collect();
let body = sigil_quote!(RustLang {
$if(schema.description.is_some()) {
$L(doc_comment_block(schema.description.as_deref().unwrap()).trim_end())
}
$L(derive_attr("Debug, Clone, Serialize, Deserialize", extra))
pub struct $N(name.as_str()) {
$for((field_name, rust_type) in fields.iter()) {
#[serde(flatten)]
$L(format!("pub {field_name}: {rust_type},"))
}
}
})
.ok()?;
fsb = fsb.add_code(body);
fsb.build().ok()
}
fn build_integer_enum_display(name: &str) -> Option<CodeBlock> {
sigil_quote!(RustLang {
impl std::fmt::Display for $N(name) {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{}", *self as i64)
}
}
})
.ok()
}
fn build_string_enum_display(name: &str, variants: &[(String, String)]) -> Option<CodeBlock> {
sigil_quote!(RustLang {
impl std::fmt::Display for $N(name) {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
$for((variant, wire_value) in variants.iter()) {
$L(format!("{name}::{variant} => write!(f, {wire_value:?}),"))
}
}
}
}
})
.ok()
}
fn doc_comment_block(doc: &str) -> String {
let mut out = String::new();
for line in doc.lines() {
out.push_str(&format!("/// {line}\n"));
}
out
}
pub fn rust_type_str(expr: &IrTypeExpr) -> String {
match expr {
IrTypeExpr::Named(name) => name.to_pascal_case(),
IrTypeExpr::Primitive(p) => rust_primitive(p).to_string(),
IrTypeExpr::Array(inner) => format!("Vec<{}>", rust_type_str(inner)),
IrTypeExpr::Map(inner) => format!(
"std::collections::HashMap<String, {}>",
rust_type_str(inner)
),
IrTypeExpr::Nullable(inner) => format!("Option<{}>", rust_type_str(inner)),
IrTypeExpr::StringLiteral(_) | IrTypeExpr::StringEnum(_) => "String".to_string(),
IrTypeExpr::Union(_) | IrTypeExpr::Any => "serde_json::Value".to_string(),
}
}
pub fn rust_type_str_qualified(expr: &IrTypeExpr) -> String {
match expr {
IrTypeExpr::Named(name) => format!("crate::models::{}", name.to_pascal_case()),
IrTypeExpr::Array(inner) => format!("Vec<{}>", rust_type_str_qualified(inner)),
IrTypeExpr::Map(inner) => format!(
"std::collections::HashMap<String, {}>",
rust_type_str_qualified(inner)
),
IrTypeExpr::Nullable(inner) => format!("Option<{}>", rust_type_str_qualified(inner)),
other => rust_type_str(other),
}
}
fn rust_type_str_model(expr: &IrTypeExpr) -> String {
match expr {
IrTypeExpr::Named(name) => format!("super::{}", name.to_pascal_case()),
IrTypeExpr::Array(inner) => format!("Vec<{}>", rust_type_str_model(inner)),
IrTypeExpr::Map(inner) => format!(
"std::collections::HashMap<String, {}>",
rust_type_str_model(inner)
),
IrTypeExpr::Nullable(inner) => format!("Option<{}>", rust_type_str_model(inner)),
other => rust_type_str(other),
}
}
fn rust_primitive(p: &IrPrimitive) -> &'static str {
match p {
IrPrimitive::String
| IrPrimitive::Date
| IrPrimitive::DateTime
| IrPrimitive::Uuid
| IrPrimitive::StringWithFormat(_) => "String",
IrPrimitive::Binary => "Vec<u8>",
IrPrimitive::Integer => "i64",
IrPrimitive::IntegerWithFormat(format) => match format.as_str() {
"int32" => "i32",
"int64" => "i64",
_ => "i64",
},
IrPrimitive::Number => "f64",
IrPrimitive::NumberWithFormat(format) => match format.as_str() {
"float" => "f32",
_ => "f64",
},
IrPrimitive::Boolean => "bool",
}
}
fn escape_str(s: &str) -> String {
s.replace('\\', "\\\\").replace('"', "\\\"")
}
fn escape_rust_keyword(name: &str) -> String {
const KEYWORDS: &[&str] = &[
"as", "async", "await", "break", "const", "continue", "crate", "dyn", "else", "enum",
"extern", "false", "fn", "for", "if", "impl", "in", "let", "loop", "match", "mod", "move",
"mut", "pub", "ref", "return", "self", "Self", "static", "struct", "super", "trait",
"true", "type", "union", "unsafe", "use", "where", "while", "yield",
];
if KEYWORDS.contains(&name) {
format!("r#{name}")
} else {
name.to_string()
}
}