use std::collections::BTreeMap;
use std::fmt;
use std::path::Path;
use serde::Deserialize;
use crate::error::{GenError, Result};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum Dialect {
Psql,
Sqlite,
Mysql,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum Hook {
BeforeInsert,
AfterInsert,
BeforeUpdate,
AfterUpdate,
BeforeDelete,
AfterDelete,
AfterSelect,
}
impl Hook {
pub fn method(self) -> &'static str {
match self {
Hook::BeforeInsert => "before_insert",
Hook::AfterInsert => "after_insert",
Hook::BeforeUpdate => "before_update",
Hook::AfterUpdate => "after_update",
Hook::BeforeDelete => "before_delete",
Hook::AfterDelete => "after_delete",
Hook::AfterSelect => "after_select",
}
}
}
impl fmt::Display for Hook {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(self.method())
}
}
#[derive(Debug, Clone, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct Config {
pub dialect: Dialect,
#[serde(default)]
pub url: Option<String>,
#[serde(default)]
pub out: Option<String>,
#[serde(default = "default_schema")]
pub schema: String,
#[serde(default)]
pub no_back_referencing: bool,
#[serde(default)]
pub only: Vec<String>,
#[serde(default)]
pub except: Vec<String>,
#[serde(default)]
pub output: Output,
#[serde(default)]
pub hooks: Hooks,
#[serde(default)]
pub inflections: BTreeMap<String, String>,
#[serde(default)]
pub tables: BTreeMap<String, TableConfig>,
#[serde(default)]
pub aliases: BTreeMap<String, TableAliases>,
#[serde(default)]
pub relationships: Vec<ManualRelationship>,
#[serde(default)]
pub types: Types,
#[serde(default)]
pub queries: Option<crate::queries::QueriesConfig>,
}
fn default_schema() -> String {
"public".to_owned()
}
#[derive(Debug, Clone, Default, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct Output {
#[serde(default)]
pub serde: bool,
#[serde(default)]
pub factories: bool,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct Hooks {
#[serde(default = "default_hooks_module")]
pub module: String,
}
impl Default for Hooks {
fn default() -> Self {
Hooks {
module: default_hooks_module(),
}
}
}
fn default_hooks_module() -> String {
"crate::hooks".to_owned()
}
#[derive(Debug, Clone, Default, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct TableConfig {
#[serde(default)]
pub only_columns: Vec<String>,
#[serde(default)]
pub except_columns: Vec<String>,
#[serde(default)]
pub hooks: Vec<Hook>,
#[serde(default)]
pub key: Vec<String>,
}
#[derive(Debug, Clone, Default, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct TableAliases {
#[serde(default)]
pub singular: Option<String>,
#[serde(default)]
pub plural: Option<String>,
#[serde(default)]
pub columns: BTreeMap<String, String>,
#[serde(default)]
pub relationships: BTreeMap<String, String>,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum Cardinality {
#[default]
ManyToOne,
OneToOne,
}
impl fmt::Display for Cardinality {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(match self {
Cardinality::ManyToOne => "many_to_one",
Cardinality::OneToOne => "one_to_one",
})
}
}
#[derive(Debug, Clone, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct ManualRelationship {
pub table: String,
pub column: String,
pub ref_table: String,
pub ref_column: String,
#[serde(default)]
pub name: Option<String>,
#[serde(default)]
pub no_back_reference: bool,
#[serde(default)]
pub cardinality: Option<Cardinality>,
}
#[derive(Debug, Clone, Default, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct Types {
#[serde(default)]
pub map: BTreeMap<String, String>,
#[serde(default, rename = "override")]
pub overrides: Vec<TypeOverride>,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct TypeOverride {
#[serde(default)]
pub tables: Vec<String>,
#[serde(rename = "match", default)]
pub matcher: Matcher,
pub rust_type: String,
}
#[derive(Debug, Clone, Default, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct Matcher {
#[serde(default)]
pub name: Option<String>,
#[serde(default)]
pub db_type: Option<String>,
#[serde(default)]
pub nullable: Option<bool>,
#[serde(default)]
pub default: Option<String>,
#[serde(default)]
pub autoincrement: Option<bool>,
#[serde(default)]
pub comment: Option<String>,
}
impl Config {
pub fn from_toml(text: &str) -> Result<Config> {
toml::from_str(text).map_err(|e| GenError::Config(e.to_string()))
}
pub fn load(path: impl AsRef<Path>) -> Result<Config> {
let text = std::fs::read_to_string(path.as_ref())?;
Config::from_toml(&text)
}
pub fn includes_table(&self, table: &str) -> bool {
if self.except.iter().any(|t| t == table) {
return false;
}
self.only.is_empty() || self.only.iter().any(|t| t == table)
}
pub fn includes_column(&self, table: &str, column: &str) -> bool {
let Some(tc) = self.tables.get(table) else {
return true;
};
if tc.except_columns.iter().any(|c| c == column) {
return false;
}
tc.only_columns.is_empty() || tc.only_columns.iter().any(|c| c == column)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_minimal_config_parses_with_defaults() {
let c = Config::from_toml("dialect = \"sqlite\"").unwrap();
assert_eq!(c.dialect, Dialect::Sqlite);
assert_eq!(c.schema, "public");
assert_eq!(c.hooks.module, "crate::hooks");
assert!(!c.no_back_referencing);
assert!(c.includes_table("anything"));
}
#[test]
fn the_full_inventory_parses() {
let c = Config::from_toml(
r#"
dialect = "psql"
url = "postgres://localhost/app"
out = "src/models"
schema = "app"
no_back_referencing = true
only = ["users", "posts"]
except = ["schema_migrations"]
[output]
serde = true
[hooks]
module = "crate::model_hooks"
[inflections]
people = "person"
[tables.users]
except_columns = ["password_digest"]
hooks = ["before_insert", "after_select"]
[aliases.users]
singular = "member"
plural = "membership"
[aliases.users.columns]
created_at = "created"
[aliases.users.relationships]
posts = "articles"
[[relationships]]
table = "posts"
column = "author_name"
ref_table = "users"
ref_column = "name"
name = "author"
no_back_reference = true
[types.map]
citext = "String"
[[types.override]]
tables = ["users"]
rust_type = "crate::types::UserId"
[types.override.match]
name = "id"
db_type = "integer"
nullable = false
"#,
)
.unwrap();
assert_eq!(c.dialect, Dialect::Psql);
assert!(c.includes_table("users"));
assert!(!c.includes_table("schema_migrations"));
assert!(!c.includes_table("tags"), "only wins");
assert!(!c.includes_column("users", "password_digest"));
assert_eq!(
c.tables["users"].hooks,
vec![Hook::BeforeInsert, Hook::AfterSelect]
);
assert_eq!(c.aliases["users"].columns["created_at"], "created");
assert_eq!(c.relationships[0].name.as_deref(), Some("author"));
assert_eq!(c.types.overrides[0].matcher.name.as_deref(), Some("id"));
}
#[test]
fn unknown_keys_are_config_errors_not_silent_noise() {
let err = Config::from_toml("dialect = \"sqlite\"\ntypo_key = 1").unwrap_err();
assert!(matches!(err, GenError::Config(_)), "{err}");
}
}