use crate::SeedError;
use crate::generator::{Generator, GeneratorKind};
use crate::identity::{ColumnId, TableId};
use drizzle_core::{Relation, SQLSchemaImpl, SQLTableInfo, SchemaHasTable, TableRef};
#[cfg(any(feature = "sqlite", feature = "postgres", feature = "mysql"))]
use drizzle_core::{SQLColumn, SQLColumnInfo};
use std::collections::{HashMap, HashSet};
use std::marker::PhantomData;
use std::sync::Arc;
#[cfg(feature = "sqlite")]
use crate::Sqlite;
#[cfg(feature = "sqlite")]
use drizzle_sqlite::traits::SQLiteColumn;
#[cfg(feature = "sqlite")]
use drizzle_sqlite::values::SQLiteValue;
#[cfg(feature = "postgres")]
use crate::Postgres;
#[cfg(feature = "postgres")]
use drizzle_postgres::traits::PostgresColumn;
#[cfg(feature = "postgres")]
use drizzle_postgres::values::PostgresValue;
#[cfg(feature = "mysql")]
use crate::MySql;
#[cfg(feature = "mysql")]
use drizzle_mysql::traits::MySQLColumn;
#[cfg(feature = "mysql")]
use drizzle_mysql::values::MySQLValue;
pub struct SeedConfig<'a, D, S> {
pub(crate) schema: &'a S,
pub(crate) skipped_tables: HashSet<TableId>,
pub(crate) seed: u64,
pub(crate) default_count: usize,
pub(crate) table_counts: HashMap<TableId, usize>,
pub(crate) column_generators: HashMap<ColumnId, Arc<dyn Generator>>,
pub(crate) column_kinds: HashMap<ColumnId, GeneratorKind>,
pub(crate) relation_counts: HashMap<(TableId, TableId), usize>,
pub(crate) max_params_per_batch: Option<usize>,
pub(crate) name_errors: Vec<SeedError>,
_dialect: PhantomData<D>,
_schema: PhantomData<&'a S>,
}
impl<'a, D, S> SeedConfig<'a, D, S> {
fn with_defaults(schema: &'a S) -> Self {
Self {
schema,
skipped_tables: HashSet::new(),
seed: 0,
default_count: 10,
table_counts: HashMap::new(),
column_generators: HashMap::new(),
column_kinds: HashMap::new(),
relation_counts: HashMap::new(),
max_params_per_batch: None,
name_errors: Vec::new(),
_dialect: PhantomData,
_schema: PhantomData,
}
}
#[must_use]
pub const fn seed(mut self, seed: u64) -> Self {
self.seed = seed;
self
}
#[must_use]
pub const fn default_count(mut self, count: usize) -> Self {
self.default_count = count;
self
}
#[must_use]
pub fn max_params(mut self, limit: usize) -> Self {
assert!(limit > 0, "max_params must be > 0");
self.max_params_per_batch = Some(limit);
self
}
pub(crate) fn count_for(&self, table: TableId) -> usize {
self.table_counts
.get(&table)
.copied()
.unwrap_or(self.default_count)
}
}
impl<D, S> SeedConfig<'_, D, S>
where
S: SQLSchemaImpl,
{
pub(crate) fn active_tables(&self) -> Vec<&'static TableRef> {
self.schema
.table_refs()
.iter()
.copied()
.filter(|t| !self.skipped_tables.contains(&TableId::from_ref(t)))
.collect()
}
#[must_use]
pub fn count<T>(mut self, table: &T, count: usize) -> Self
where
T: SQLTableInfo,
S: SchemaHasTable<T>,
{
self.table_counts.insert(TableId::from_info(table), count);
self
}
#[must_use]
pub fn relation<P, C>(mut self, parent: &P, child: &C, count: usize) -> Self
where
P: SQLTableInfo,
C: SQLTableInfo + Relation<P>,
S: SchemaHasTable<P> + SchemaHasTable<C>,
{
self.relation_counts.insert(
(TableId::from_info(parent), TableId::from_info(child)),
count,
);
self
}
}
impl<D, S> SeedConfig<'_, D, S> {
#[must_use]
pub fn skip<T>(mut self, table: &T) -> Self
where
T: SQLTableInfo,
S: SchemaHasTable<T>,
{
self.skipped_tables.insert(TableId::from_info(table));
self
}
}
impl<D, S> SeedConfig<'_, D, S>
where
S: SQLSchemaImpl,
{
#[must_use]
pub fn count_by_name(mut self, table: &str, count: usize) -> Self {
if let Some(table) = self.resolve_table(table) {
self.table_counts.insert(TableId::from_ref(table), count);
}
self
}
#[must_use]
pub fn relation_by_name(mut self, parent: &str, child: &str, count: usize) -> Self {
let (Some(parent_ref), Some(child_ref)) =
(self.resolve_table(parent), self.resolve_table(child))
else {
return self;
};
let parent_id = TableId::from_ref(parent_ref);
let related = child_ref
.foreign_keys
.iter()
.any(|fk| TableId::foreign_target(child_ref, fk) == parent_id);
if related {
self.relation_counts
.insert((parent_id, TableId::from_ref(child_ref)), count);
} else {
self.name_errors.push(SeedError::NotRelated {
parent: parent_id.to_string(),
child: TableId::from_ref(child_ref).to_string(),
});
}
self
}
#[must_use]
pub fn skip_by_name(mut self, table: &str) -> Self {
if let Some(table) = self.resolve_table(table) {
self.skipped_tables.insert(TableId::from_ref(table));
}
self
}
#[must_use]
pub fn kind_by_name(mut self, table: &str, column: &str, kind: GeneratorKind) -> Self {
if let Some(column) = self.resolve_column(table, column) {
self.column_kinds.insert(column, kind);
}
self
}
#[must_use]
pub fn generator_by_name(
mut self,
table: &str,
column: &str,
generator: impl Generator + 'static,
) -> Self {
if let Some(column) = self.resolve_column(table, column) {
self.column_generators.insert(column, Arc::new(generator));
}
self
}
fn resolve_table(&mut self, name: &str) -> Option<&'static TableRef> {
let tables = self.schema.table_refs();
let exact: Vec<&'static TableRef> = tables
.iter()
.copied()
.filter(|table| table.name == name)
.collect();
let candidates = if exact.is_empty() {
name.split_once('.')
.map_or_else(Vec::new, |(schema, table_name)| {
tables
.iter()
.copied()
.filter(|table| {
table.name == table_name
&& (table.schema == Some(schema)
|| (schema == "public" && table.schema.is_none()))
})
.collect()
})
} else {
exact
};
match candidates.as_slice() {
[table] => Some(table),
[] => {
self.name_errors.push(SeedError::UnknownTable {
table: name.to_string(),
known: tables
.iter()
.map(|table| TableId::from_ref(table).to_string())
.collect(),
});
None
}
_ => {
self.name_errors.push(SeedError::AmbiguousTable {
table: name.to_string(),
candidates: candidates
.iter()
.map(|table| TableId::from_ref(table).to_string())
.collect(),
});
None
}
}
}
fn resolve_column(&mut self, table: &str, column: &str) -> Option<ColumnId> {
let table = self.resolve_table(table)?;
let table_id = TableId::from_ref(table);
if let Some(found) = table
.columns
.iter()
.find(|candidate| candidate.name == column)
{
return Some(ColumnId::new(table_id, found.name));
}
self.name_errors.push(SeedError::UnknownColumn {
table: table_id.to_string(),
column: column.to_string(),
known: table
.columns
.iter()
.map(|column| column.name.to_string())
.collect(),
});
None
}
pub(crate) fn check_names(&self) -> Result<(), SeedError> {
self.name_errors.first().cloned().map_or(Ok(()), Err)
}
}
macro_rules! dialect_seed_config {
(
feature = $feature:literal,
dialect = $dialect:literal,
marker = $marker:ty,
constructor = $constructor:ident,
column = $column_trait:path,
value = $value:ty,
seed_statement = $seed_statement:path,
reset_statement = $reset_statement:path,
generate = $generate:ident,
reset = $reset:ident,
reset_note = $reset_note:literal
) => {
#[cfg(feature = $feature)]
impl<'a> SeedConfig<'a, $marker, ()> {
#[doc = concat!("Creates a ", $dialect, " seed config for `schema` (usually a `#[derive(...Schema)]` struct).")]
pub fn $constructor<Schema>(schema: &'a Schema) -> SeedConfig<'a, $marker, Schema>
where
Schema: SQLSchemaImpl,
{
SeedConfig::<'a, $marker, Schema>::with_defaults(schema)
}
}
#[cfg(feature = $feature)]
impl<S> SeedConfig<'_, $marker, S>
where
S: SQLSchemaImpl,
{
#[must_use]
pub fn kind<C>(mut self, column: &C, kind: GeneratorKind) -> Self
where
C: SQLColumnInfo + $column_trait,
S: SchemaHasTable<<C as SQLColumn<'static, $value>>::Table>,
{
self.column_kinds.insert(ColumnId::from_info(column), kind);
self
}
#[must_use]
pub fn generator<C>(mut self, column: &C, generator: impl Generator + 'static) -> Self
where
C: SQLColumnInfo + $column_trait,
S: SchemaHasTable<<C as SQLColumn<'static, $value>>::Table>,
{
self.column_generators
.insert(ColumnId::from_info(column), Arc::new(generator));
self
}
#[must_use]
pub fn generate(&self) -> Vec<$seed_statement> {
self.try_generate()
.unwrap_or_else(|error| panic!("invalid seed plan: {error}"))
}
pub fn try_generate(&self) -> Result<Vec<$seed_statement>, crate::SeedError> {
crate::Seeder::new(self).$generate()
}
pub fn try_generate_script(&self) -> Result<String, crate::SeedError> {
let mut script = String::new();
for statement in self.try_generate()? {
script.push_str(&statement.inline_sql()?);
script.push_str(";\n");
}
Ok(script)
}
pub fn try_generate_rows(&self) -> Result<Vec<crate::SeedRows>, crate::SeedError> {
crate::Seeder::new(self).generate_rows()
}
#[doc = $reset_note]
pub fn reset_plan(&self) -> Result<Vec<$reset_statement>, crate::SeedError> {
crate::Seeder::new(self).$reset()
}
}
};
}
dialect_seed_config!(
feature = "sqlite",
dialect = "SQLite",
marker = Sqlite,
constructor = sqlite,
column = SQLiteColumn<'static>,
value = SQLiteValue<'static>,
seed_statement = crate::SQLiteSeedStatement,
reset_statement = crate::SQLiteResetStatement,
generate = generate_sqlite,
reset = reset_sqlite,
reset_note = "Wrap execution in a transaction when the connection API supports one."
);
dialect_seed_config!(
feature = "postgres",
dialect = "PostgreSQL",
marker = Postgres,
constructor = postgres,
column = PostgresColumn<'static>,
value = PostgresValue<'static>,
seed_statement = crate::PostgresSeedStatement,
reset_statement = crate::PostgresResetStatement,
generate = generate_postgres,
reset = reset_postgres,
reset_note = "Wrap execution in a transaction when the connection API supports one."
);
dialect_seed_config!(
feature = "mysql",
dialect = "MySQL",
marker = MySql,
constructor = mysql,
column = MySQLColumn<'static>,
value = MySQLValue<'static>,
seed_statement = crate::MySQLSeedStatement,
reset_statement = crate::MySQLResetStatement,
generate = generate_mysql,
reset = reset_mysql,
reset_note = "The plan appends `ALTER TABLE ... AUTO_INCREMENT = 1` for auto-increment tables. Those statements implicitly commit in MySQL, so callers must not assume the whole reset is transactional."
);