1#![cfg_attr(
173 not(any(feature = "sqlite", feature = "postgres", feature = "mysql")),
174 allow(dead_code)
175)]
176
177pub(crate) mod batch;
178pub(crate) mod config;
179pub(crate) mod datasets;
180mod error;
181pub(crate) mod generator;
182pub(crate) mod identity;
183pub(crate) mod inference;
184#[cfg(any(feature = "sqlite", feature = "postgres", feature = "mysql"))]
185mod literal;
186#[cfg(feature = "mysql")]
187mod mysql_seed;
188pub(crate) mod rng;
189pub(crate) mod topology;
190
191pub mod generators;
192pub mod schema;
193
194pub use config::SeedConfig;
195pub use error::SeedError;
196pub use generator::{Generator, GeneratorKind, RngCore, SeedValue};
197pub use rand::Rng;
200
201use drizzle_core::{ColumnRef, TableRef};
202use rand::rngs::StdRng;
203use std::collections::{HashMap, HashSet};
204use std::sync::Arc;
205
206#[cfg(any(feature = "sqlite", feature = "postgres", feature = "mysql"))]
207use drizzle_core::{OwnedSQL, SQL, SQLChunk, Token, param::Param, traits::ToSQL};
208
209#[cfg(any(
210 feature = "postgres",
211 all(test, any(feature = "sqlite", feature = "mysql"))
212))]
213use drizzle_core::ColumnDialect;
214
215#[cfg(any(feature = "sqlite", feature = "postgres", feature = "mysql"))]
216use std::borrow::Cow;
217
218use identity::{ColumnId, TableId};
219
220#[cfg(feature = "sqlite")]
221pub use statement::SQLiteResetStatement;
222#[cfg(feature = "sqlite")]
223pub use statement::SQLiteSeedStatement;
224
225#[cfg(feature = "postgres")]
226pub use statement::PostgresResetStatement;
227#[cfg(feature = "postgres")]
228pub use statement::PostgresSeedStatement;
229
230#[cfg(feature = "mysql")]
231pub use statement::MySQLResetStatement;
232#[cfg(feature = "mysql")]
233pub use statement::MySQLSeedStatement;
234
235#[cfg(feature = "sqlite")]
236use drizzle_sqlite::values::{OwnedSQLiteValue, SQLiteValue};
237
238#[cfg(feature = "postgres")]
239use drizzle_postgres::values::{OwnedPostgresValue, PostgresValue};
240
241#[cfg(feature = "mysql")]
242use drizzle_mysql::values::{MySQLValue, OwnedMySQLValue};
243
244#[cfg(all(feature = "postgres", feature = "chrono"))]
245use chrono::{DateTime, NaiveDate, NaiveDateTime, NaiveTime, Utc};
246
247#[cfg(feature = "sqlite")]
253pub struct Sqlite;
254
255#[cfg(feature = "postgres")]
257pub struct Postgres;
258
259#[cfg(feature = "mysql")]
261pub struct MySql;
262
263mod statement {
268 #[cfg(any(feature = "sqlite", feature = "postgres", feature = "mysql"))]
269 use super::{Cow, OwnedSQL, Param, SQL, SQLChunk, ToSQL};
270
271 #[cfg(feature = "sqlite")]
272 use super::{OwnedSQLiteValue, SQLiteValue};
273
274 #[cfg(feature = "postgres")]
275 use super::{OwnedPostgresValue, PostgresValue};
276
277 #[cfg(feature = "mysql")]
278 use super::{MySQLValue, OwnedMySQLValue};
279
280 #[cfg(any(feature = "sqlite", feature = "postgres", feature = "mysql"))]
282 fn convert_to_sql<'a, Owned, Borrowed>(owned: &OwnedSQL<Owned>) -> SQL<'a, Borrowed>
283 where
284 Owned: drizzle_core::SQLParam,
285 Borrowed: drizzle_core::SQLParam + From<Owned>,
286 {
287 let chunks = owned
288 .chunks
289 .iter()
290 .map(|chunk| match chunk {
291 drizzle_core::OwnedSQLChunk::Token(t) => SQLChunk::Token(*t),
292 drizzle_core::OwnedSQLChunk::Ident(s) => SQLChunk::Ident(Cow::Owned(s.to_string())),
293 drizzle_core::OwnedSQLChunk::Raw(s) => SQLChunk::Raw(Cow::Owned(s.to_string())),
294 drizzle_core::OwnedSQLChunk::Number(v) => SQLChunk::Number(*v),
295 drizzle_core::OwnedSQLChunk::Param(p) => SQLChunk::Param(Param {
296 placeholder: p.placeholder,
297 value: p
298 .value
299 .as_ref()
300 .map(|v| Cow::Owned(Borrowed::from(v.clone()))),
301 }),
302 drizzle_core::OwnedSQLChunk::Table(t) => SQLChunk::Table(*t),
303 drizzle_core::OwnedSQLChunk::Column(c) => SQLChunk::Column(*c),
304 })
305 .collect();
306 SQL { chunks }
307 }
308
309 #[cfg(any(feature = "sqlite", feature = "postgres", feature = "mysql"))]
311 fn convert_into_sql<'a, Owned, Borrowed>(owned: OwnedSQL<Owned>) -> SQL<'a, Borrowed>
312 where
313 Owned: drizzle_core::SQLParam,
314 Borrowed: drizzle_core::SQLParam + From<Owned>,
315 {
316 let chunks = owned
317 .chunks
318 .into_iter()
319 .map(|chunk| match chunk {
320 drizzle_core::OwnedSQLChunk::Token(t) => SQLChunk::Token(t),
321 drizzle_core::OwnedSQLChunk::Ident(s) => {
322 SQLChunk::Ident(Cow::Owned(String::from(s)))
323 }
324 drizzle_core::OwnedSQLChunk::Raw(s) => SQLChunk::Raw(Cow::Owned(String::from(s))),
325 drizzle_core::OwnedSQLChunk::Number(v) => SQLChunk::Number(v),
326 drizzle_core::OwnedSQLChunk::Param(p) => SQLChunk::Param(Param {
327 placeholder: p.placeholder,
328 value: p.value.map(|v| Cow::Owned(Borrowed::from(v))),
329 }),
330 drizzle_core::OwnedSQLChunk::Table(t) => SQLChunk::Table(t),
331 drizzle_core::OwnedSQLChunk::Column(c) => SQLChunk::Column(c),
332 })
333 .collect();
334 SQL { chunks }
335 }
336
337 #[cfg(any(feature = "sqlite", feature = "postgres", feature = "mysql"))]
339 fn inline_sql<Owned>(
340 owned: &OwnedSQL<Owned>,
341 literal: fn(&Owned) -> Result<String, crate::literal::LiteralError>,
342 ) -> Result<String, crate::SeedError>
343 where
344 Owned: drizzle_core::SQLParam,
345 {
346 let mut sql: SQL<'_, Owned> = convert_to_sql(owned);
347 for (chunk, source) in sql.chunks.iter_mut().zip(owned.chunks.iter()) {
348 if let drizzle_core::OwnedSQLChunk::Param(param) = source {
349 let value = param
350 .value
351 .as_ref()
352 .ok_or_else(|| crate::SeedError::NoLiteral {
353 reason: "a placeholder has no bound value".to_owned(),
354 })?;
355 let text =
356 literal(value).map_err(|reason| crate::SeedError::NoLiteral { reason })?;
357 *chunk = SQLChunk::Raw(Cow::Owned(text));
358 }
359 }
360 Ok(sql.build().0)
361 }
362
363 macro_rules! seed_statement {
364 ($name:ident, $owned:ty, $borrowed:ty, $feature:literal, $literal:path) => {
365 #[cfg(feature = $feature)]
366 #[derive(Debug, Clone)]
367 pub struct $name {
373 pub(crate) inner: OwnedSQL<$owned>,
374 pub(crate) table: &'static str,
375 pub(crate) rows: usize,
376 }
377
378 #[cfg(feature = $feature)]
379 impl $name {
380 #[must_use]
382 pub const fn table(&self) -> &'static str {
383 self.table
384 }
385
386 #[must_use]
390 pub const fn rows(&self) -> usize {
391 self.rows
392 }
393
394 pub fn sql(&self) -> String {
396 self.inner.to_sql().build().0
397 }
398
399 pub fn build(&self) -> (String, Vec<$owned>) {
401 let sql = self.inner.to_sql();
402 let (text, params) = sql.build();
403 (text, params.into_iter().cloned().collect())
404 }
405
406 pub fn inline_sql(&self) -> Result<String, crate::SeedError> {
420 inline_sql(&self.inner, $literal)
421 }
422 }
423
424 #[cfg(feature = $feature)]
425 impl std::fmt::Display for $name {
426 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
427 f.write_str(&self.sql())
428 }
429 }
430
431 #[cfg(feature = $feature)]
432 impl<'a> ToSQL<'a, $borrowed> for $name {
433 fn to_sql(&self) -> SQL<'a, $borrowed> {
434 convert_to_sql(&self.inner)
435 }
436
437 fn into_sql(self) -> SQL<'a, $borrowed> {
438 convert_into_sql(self.inner)
439 }
440 }
441 };
442 }
443
444 seed_statement!(
445 SQLiteSeedStatement,
446 OwnedSQLiteValue,
447 SQLiteValue<'a>,
448 "sqlite",
449 crate::literal::sqlite
450 );
451 seed_statement!(
452 SQLiteResetStatement,
453 OwnedSQLiteValue,
454 SQLiteValue<'a>,
455 "sqlite",
456 crate::literal::sqlite
457 );
458 seed_statement!(
459 PostgresSeedStatement,
460 OwnedPostgresValue,
461 PostgresValue<'a>,
462 "postgres",
463 crate::literal::postgres
464 );
465 seed_statement!(
466 PostgresResetStatement,
467 OwnedPostgresValue,
468 PostgresValue<'a>,
469 "postgres",
470 crate::literal::postgres
471 );
472 seed_statement!(
473 MySQLSeedStatement,
474 OwnedMySQLValue,
475 MySQLValue<'a>,
476 "mysql",
477 crate::literal::mysql
478 );
479 seed_statement!(
480 MySQLResetStatement,
481 OwnedMySQLValue,
482 MySQLValue<'a>,
483 "mysql",
484 crate::literal::mysql
485 );
486}
487
488#[derive(Debug, Clone, PartialEq)]
499pub struct SeedRows {
500 pub schema: Option<&'static str>,
502 pub table: &'static str,
504 pub columns: Vec<&'static str>,
506 pub rows: Vec<Vec<SeedValue>>,
508}
509
510struct GeneratedChunk<'a> {
515 table: &'a TableRef,
516 rows: Vec<Vec<SeedValue>>,
517}
518
519#[derive(Clone)]
520struct RelationSpec {
521 target_table: TableId,
522 fk_columns: &'static [&'static str],
523 ref_columns: &'static [&'static str],
524 children_per_parent: usize,
525}
526
527struct RelationContext<'plan, 'schema> {
528 source_table: &'schema TableRef,
529 column_indexes: &'plan HashMap<&'static str, usize>,
530 specs: &'plan [RelationSpec],
531 generated_values: &'plan HashMap<ColumnId, Vec<SeedValue>>,
532 generated_counts: &'plan HashMap<TableId, usize>,
533 active_tables: &'plan HashMap<TableId, &'schema TableRef>,
534}
535
536struct Seeder<'a, D, S> {
541 config: &'a SeedConfig<'a, D, S>,
542}
543
544impl<'a, D, S> Seeder<'a, D, S>
545where
546 S: drizzle_core::SQLSchemaImpl,
547{
548 const fn new(config: &'a SeedConfig<'a, D, S>) -> Self {
549 Self { config }
550 }
551
552 fn generate_chunks(
553 &self,
554 dialect_max_params: usize,
555 ) -> Result<Vec<GeneratedChunk<'a>>, SeedError> {
556 self.config.check_names()?;
557 let active_tables = self.config.active_tables();
558 let order = topology::seeding_order(&active_tables).map_err(|error| {
559 SeedError::CyclicForeignKeys {
560 tables: error
561 .tables
562 .into_iter()
563 .map(|table| table.to_string())
564 .collect(),
565 }
566 })?;
567 let table_map: HashMap<TableId, &TableRef> = active_tables
568 .iter()
569 .map(|table| (TableId::from_ref(table), *table))
570 .collect();
571 let mut table_name_counts: HashMap<&'static str, usize> = HashMap::new();
572 for table in &active_tables {
573 *table_name_counts.entry(table.name).or_default() += 1;
574 }
575
576 let mut generated_values: HashMap<ColumnId, Vec<SeedValue>> = HashMap::new();
577 let mut generated_counts: HashMap<TableId, usize> = HashMap::new();
578 let mut chunks_out = Vec::new();
579
580 for table_id in order {
581 let Some(&table) = table_map.get(&table_id) else {
582 continue;
583 };
584
585 let columns = table.columns;
586 if columns.is_empty() {
587 continue;
588 }
589
590 let count = self.derived_count_for(table, &generated_counts);
591 if count == 0 {
592 generated_counts.insert(table_id, 0);
593 continue;
594 }
595
596 let generators = self.build_generators(table);
597 let col_index_map: HashMap<&'static str, usize> = columns
598 .iter()
599 .enumerate()
600 .map(|(idx, col)| (col.name, idx))
601 .collect();
602 let relation_specs = self.relation_specs_for(table);
603
604 let mut all_rows: Vec<Vec<SeedValue>> = Vec::with_capacity(count);
605 let mut col_rngs: Vec<StdRng> = columns
606 .iter()
607 .map(|column| {
608 rng::table_column_rng(
609 table_id,
610 column.name,
611 self.config.seed,
612 table_name_counts.get(table.name).copied().unwrap_or(0) > 1,
613 )
614 })
615 .collect();
616
617 let mut unique_seen: Vec<Option<HashSet<String>>> = columns
618 .iter()
619 .map(|column| unique_column(table, column).then(HashSet::new))
620 .collect();
621
622 for row_idx in 0..count {
623 let mut row = Vec::with_capacity(columns.len());
624 for (col_idx, generator) in generators.iter().enumerate() {
625 let column = &columns[col_idx];
626 let rng = &mut col_rngs[col_idx];
627 let mut val = generator.generate(rng, row_idx, column.sql_type);
628 if let Some(seen) = unique_seen[col_idx].as_mut() {
629 val = unique_value(
630 val,
631 seen,
632 row_idx,
633 column,
634 |rng| generator.generate(rng, row_idx, column.sql_type),
635 rng,
636 );
637 }
638 row.push(val);
639 }
640
641 Self::apply_many_to_one_relations(
642 &mut row,
643 row_idx,
644 &RelationContext {
645 source_table: table,
646 column_indexes: &col_index_map,
647 specs: &relation_specs,
648 generated_values: &generated_values,
649 generated_counts: &generated_counts,
650 active_tables: &table_map,
651 },
652 )?;
653
654 all_rows.push(row);
655 }
656
657 drop_composite_key_repeats(table, &col_index_map, &mut all_rows);
661 let count = all_rows.len();
662
663 for (col_idx, col) in columns.iter().enumerate() {
665 let vals: Vec<SeedValue> =
666 all_rows.iter().map(|row| row[col_idx].clone()).collect();
667 generated_values.insert(ColumnId::new(table_id, col.name), vals);
668 }
669
670 generated_counts.insert(table_id, count);
671
672 let param_limit = self
673 .config
674 .max_params_per_batch
675 .unwrap_or(dialect_max_params)
676 .max(1);
677
678 for (start, end) in
679 batch_ranges_by_param_limit(&all_rows, param_limit).map_err(|required| {
680 SeedError::ParameterLimitTooLow {
681 table: table_id.to_string(),
682 required,
683 limit: param_limit,
684 }
685 })?
686 {
687 chunks_out.push(GeneratedChunk {
688 table,
689 rows: all_rows[start..end].to_vec(),
690 });
691 }
692 }
693
694 Ok(chunks_out)
695 }
696
697 fn generate_rows(&self) -> Result<Vec<SeedRows>, SeedError> {
698 let mut out: Vec<SeedRows> = Vec::new();
699 for chunk in self.generate_chunks(usize::MAX)? {
700 let kept: Vec<usize> = chunk
701 .table
702 .columns
703 .iter()
704 .enumerate()
705 .filter(|(_, column)| generated_expression(column).is_none())
706 .map(|(index, _)| index)
707 .collect();
708 let rows = chunk
709 .rows
710 .into_iter()
711 .map(|row| kept.iter().map(|&index| row[index].clone()).collect());
712 match out.last_mut() {
713 Some(last)
714 if last.table == chunk.table.name && last.schema == chunk.table.schema =>
715 {
716 last.rows.extend(rows);
717 }
718 _ => out.push(SeedRows {
719 schema: chunk.table.schema,
720 table: chunk.table.name,
721 columns: kept
722 .iter()
723 .map(|&index| chunk.table.columns[index].name)
724 .collect(),
725 rows: rows.collect(),
726 }),
727 }
728 }
729 Ok(out)
730 }
731
732 fn reset_tables(&self) -> Result<Vec<&'static TableRef>, SeedError> {
733 self.config.check_names()?;
734 let all_tables = self.config.schema.table_refs();
735 let active_tables = self.config.active_tables();
736 let active_ids: HashSet<_> = active_tables
737 .iter()
738 .map(|table| TableId::from_ref(table))
739 .collect();
740
741 for child in all_tables {
742 let child_id = TableId::from_ref(child);
743 if active_ids.contains(&child_id) {
744 continue;
745 }
746 for foreign_key in child.foreign_keys {
747 let parent_id = TableId::foreign_target(child, foreign_key);
748 if active_ids.contains(&parent_id) {
749 return Err(SeedError::UnsafeResetSelection {
750 parent: parent_id.to_string(),
751 skipped_child: child_id.to_string(),
752 });
753 }
754 }
755 }
756
757 let order = topology::seeding_order(&active_tables).map_err(|error| {
758 SeedError::CyclicForeignKeys {
759 tables: error
760 .tables
761 .into_iter()
762 .map(|table| table.to_string())
763 .collect(),
764 }
765 })?;
766 let table_map: HashMap<_, _> = active_tables
767 .into_iter()
768 .map(|table| (TableId::from_ref(table), table))
769 .collect();
770 Ok(order
771 .into_iter()
772 .rev()
773 .filter_map(|table| table_map.get(&table).copied())
774 .collect())
775 }
776
777 fn derived_count_for(
778 &self,
779 table: &TableRef,
780 generated_counts: &HashMap<TableId, usize>,
781 ) -> usize {
782 let table_id = TableId::from_ref(table);
783 if let Some(&count) = self.config.table_counts.get(&table_id) {
784 return count;
785 }
786
787 let mut derived: Option<usize> = None;
788 for parent_id in Self::parent_table_ids(table) {
789 if let Some(&parent_count) = generated_counts.get(&parent_id) {
790 let children_per_parent = self
791 .config
792 .relation_counts
793 .get(&(parent_id, table_id))
794 .copied()
795 .unwrap_or(1);
796 let child_count = parent_count.saturating_mul(children_per_parent);
797 derived = Some(derived.map_or(child_count, |current| current.max(child_count)));
798 }
799 }
800
801 derived.unwrap_or_else(|| self.config.count_for(table_id))
802 }
803
804 fn parent_table_ids(table: &TableRef) -> Vec<TableId> {
805 let mut seen = HashSet::new();
806 let mut parent_ids = Vec::new();
807 let table_id = TableId::from_ref(table);
808
809 for fk in table.foreign_keys {
810 let parent = TableId::foreign_target(table, fk);
811 if parent != table_id && seen.insert(parent) {
812 parent_ids.push(parent);
813 }
814 }
815
816 parent_ids
817 }
818
819 fn build_generators(&self, table: &TableRef) -> Vec<Box<dyn Generator>> {
820 let table_id = TableId::from_ref(table);
821 table
822 .columns
823 .iter()
824 .map(|col| {
825 let col_name = col.name;
826 let key = ColumnId::new(table_id, col_name);
827
828 if let Some(custom) = self.config.column_generators.get(&key) {
829 return Box::new(Arc::clone(custom)) as Box<dyn Generator>;
830 }
831
832 if let Some(&kind) = self.config.column_kinds.get(&key) {
833 return kind.into_generator();
834 }
835
836 if (col.has_default() || is_postgres_identity(col)) && !col.primary_key() {
837 return Box::new(DefaultGen);
838 }
839
840 inference::infer_generator(col)
841 })
842 .collect()
843 }
844
845 fn relation_specs_for(&self, source_table: &TableRef) -> Vec<RelationSpec> {
846 let source_id = TableId::from_ref(source_table);
847 source_table
848 .foreign_keys
849 .iter()
850 .map(|fk| {
851 let target_id = TableId::foreign_target(source_table, fk);
852 let children_per_parent = self
853 .config
854 .relation_counts
855 .get(&(target_id, source_id))
856 .copied()
857 .unwrap_or(1);
858
859 RelationSpec {
860 target_table: target_id,
861 fk_columns: fk.source_columns,
862 ref_columns: fk.target_columns,
863 children_per_parent,
864 }
865 })
866 .collect()
867 }
868
869 fn apply_many_to_one_relations(
870 row: &mut [SeedValue],
871 row_idx: usize,
872 context: &RelationContext<'_, '_>,
873 ) -> Result<(), SeedError> {
874 for rel in context.specs {
875 if rel.fk_columns.len() != rel.ref_columns.len() {
876 continue;
877 }
878
879 if !context.active_tables.contains_key(&rel.target_table) {
883 continue;
884 }
885
886 let parent_count = rel
887 .ref_columns
888 .first()
889 .and_then(|first_ref| {
890 context
891 .generated_values
892 .get(&ColumnId::new(rel.target_table, first_ref))
893 .map(std::vec::Vec::len)
894 })
895 .or_else(|| context.generated_counts.get(&rel.target_table).copied())
896 .unwrap_or(0);
897
898 if parent_count == 0 || rel.children_per_parent == 0 {
899 let nullable_columns = rel
900 .fk_columns
901 .iter()
902 .filter(|fk_column| {
903 let fk_column = **fk_column;
904 context
905 .source_table
906 .columns
907 .iter()
908 .find(|column| column.name == fk_column)
909 .is_some_and(|column| !column.not_null())
910 })
911 .copied()
912 .collect::<Vec<_>>();
913 if nullable_columns.is_empty() {
914 return Err(SeedError::MissingParentRows {
915 child: TableId::from_ref(context.source_table).to_string(),
916 parent: rel.target_table.to_string(),
917 });
918 }
919 for fk_col in nullable_columns {
920 if let Some(&fk_idx) = context.column_indexes.get(fk_col) {
921 row[fk_idx] = SeedValue::Null;
922 }
923 }
924 continue;
925 }
926
927 let parent_idx = (row_idx / rel.children_per_parent) % parent_count;
928 for (fk_col, ref_col) in rel.fk_columns.iter().zip(rel.ref_columns.iter()) {
929 let Some(&fk_idx) = context.column_indexes.get(fk_col) else {
930 continue;
931 };
932
933 if let Some(parent_vals) = context
934 .generated_values
935 .get(&ColumnId::new(rel.target_table, ref_col))
936 && let Some(parent_value) = parent_vals.get(parent_idx)
937 {
938 row[fk_idx] = parent_value.clone();
939 } else {
940 return Err(SeedError::MissingParentRows {
941 child: TableId::from_ref(context.source_table).to_string(),
942 parent: rel.target_table.to_string(),
943 });
944 }
945 }
946 }
947 Ok(())
948 }
949}
950
951#[cfg(feature = "sqlite")]
952impl<S> Seeder<'_, Sqlite, S>
953where
954 S: drizzle_core::SQLSchemaImpl,
955{
956 fn generate_sqlite(&self) -> Result<Vec<SQLiteSeedStatement>, SeedError> {
957 let chunks = self.generate_chunks(batch::SQLITE_MAX_PARAMS)?;
958 let mut statements = Vec::with_capacity(chunks.len());
959 for chunk in &chunks {
960 build_sqlite_statements(chunk, &mut statements);
961 }
962 Ok(statements)
963 }
964
965 fn reset_sqlite(&self) -> Result<Vec<SQLiteResetStatement>, SeedError> {
966 Ok(build_reset_sql(&self.reset_tables()?)
967 .into_iter()
968 .map(|(table, inner)| SQLiteResetStatement {
969 inner,
970 table,
971 rows: 0,
972 })
973 .collect())
974 }
975}
976
977#[cfg(feature = "postgres")]
978impl<S> Seeder<'_, Postgres, S>
979where
980 S: drizzle_core::SQLSchemaImpl,
981{
982 fn generate_postgres(&self) -> Result<Vec<PostgresSeedStatement>, SeedError> {
983 let chunks = self.generate_chunks(batch::POSTGRES_MAX_PARAMS)?;
984 let mut statements = Vec::with_capacity(chunks.len());
985 let mut table_chunks: Vec<&GeneratedChunk<'_>> = Vec::new();
986 for chunk in &chunks {
987 if table_chunks
988 .first()
989 .is_some_and(|first| !std::ptr::eq(first.table, chunk.table))
990 {
991 statements.extend(build_postgres_sequence_sync(&table_chunks));
992 table_chunks.clear();
993 }
994 statements.push(build_postgres_statement(chunk));
995 table_chunks.push(chunk);
996 }
997 statements.extend(build_postgres_sequence_sync(&table_chunks));
998 Ok(statements)
999 }
1000
1001 fn reset_postgres(&self) -> Result<Vec<PostgresResetStatement>, SeedError> {
1002 Ok(build_reset_sql(&self.reset_tables()?)
1003 .into_iter()
1004 .map(|(table, inner)| PostgresResetStatement {
1005 inner,
1006 table,
1007 rows: 0,
1008 })
1009 .collect())
1010 }
1011}
1012
1013#[cfg(feature = "mysql")]
1014impl<S> Seeder<'_, MySql, S>
1015where
1016 S: drizzle_core::SQLSchemaImpl,
1017{
1018 fn generate_mysql(&self) -> Result<Vec<MySQLSeedStatement>, SeedError> {
1019 self.generate_chunks(batch::MYSQL_MAX_PARAMS)?
1020 .iter()
1021 .map(mysql_seed::build_statement)
1022 .collect()
1023 }
1024
1025 fn reset_mysql(&self) -> Result<Vec<MySQLResetStatement>, SeedError> {
1026 let tables = self.reset_tables()?;
1027 let mut statements = build_reset_sql(&tables)
1028 .into_iter()
1029 .map(|(table, inner)| MySQLResetStatement {
1030 inner,
1031 table,
1032 rows: 0,
1033 })
1034 .collect::<Vec<_>>();
1035 for table in tables.into_iter().rev() {
1036 if table.columns.iter().any(|column| {
1037 matches!(
1038 column.dialect,
1039 drizzle_core::ColumnDialect::MySQL {
1040 auto_increment: true,
1041 ..
1042 }
1043 )
1044 }) {
1045 statements.push(MySQLResetStatement {
1046 inner: build_mysql_auto_increment_reset_sql(table),
1047 table: table.name,
1048 rows: 0,
1049 });
1050 }
1051 }
1052 Ok(statements)
1053 }
1054}
1055
1056fn statement_table(table: &TableRef) -> TableRef {
1064 let mut table = *table;
1065 if table.schema == Some("public") {
1066 table.schema = None;
1067 }
1068 table
1069}
1070
1071const fn generated_expression(column: &ColumnRef) -> Option<&'static str> {
1073 match column.dialect {
1074 drizzle_core::ColumnDialect::SQLite {
1075 generated_expression,
1076 ..
1077 }
1078 | drizzle_core::ColumnDialect::PostgreSQL {
1079 generated_expression,
1080 ..
1081 }
1082 | drizzle_core::ColumnDialect::MySQL {
1083 generated_expression,
1084 ..
1085 } => generated_expression,
1086 }
1087}
1088
1089const fn is_postgres_identity(column: &ColumnRef) -> bool {
1091 matches!(
1092 column.dialect,
1093 drizzle_core::ColumnDialect::PostgreSQL {
1094 is_generated_identity: true,
1095 ..
1096 }
1097 )
1098}
1099
1100fn unique_column(table: &TableRef, column: &ColumnRef) -> bool {
1105 let is_foreign_key = table
1106 .foreign_keys
1107 .iter()
1108 .any(|fk| fk.source_columns.contains(&column.name));
1109 if is_foreign_key {
1110 return false;
1111 }
1112 let single_primary_key = table.primary_key.as_ref().map_or_else(
1113 || column.primary_key() && table.columns.iter().filter(|c| c.primary_key()).count() == 1,
1114 |pk| pk.columns == [column.name],
1115 );
1116 column.unique()
1117 || single_primary_key
1118 || table.constraints.iter().any(|constraint| {
1119 constraint.kind == drizzle_core::SQLConstraintKind::Unique
1120 && constraint.columns == [column.name]
1121 })
1122}
1123
1124fn drop_composite_key_repeats(
1128 table: &TableRef,
1129 column_indexes: &HashMap<&'static str, usize>,
1130 rows: &mut Vec<Vec<SeedValue>>,
1131) {
1132 let primary_key = table
1133 .primary_key
1134 .as_ref()
1135 .map(|pk| pk.columns)
1136 .unwrap_or_default();
1137 let keys: Vec<Vec<usize>> = std::iter::once(primary_key)
1138 .chain(
1139 table
1140 .constraints
1141 .iter()
1142 .filter(|constraint| constraint.kind == drizzle_core::SQLConstraintKind::Unique)
1143 .map(|constraint| constraint.columns),
1144 )
1145 .filter(|columns| columns.len() > 1)
1146 .filter_map(|columns| {
1147 columns
1148 .iter()
1149 .map(|column| column_indexes.get(column).copied())
1150 .collect()
1151 })
1152 .collect();
1153 if keys.is_empty() {
1154 return;
1155 }
1156 let mut seen: Vec<HashSet<String>> = vec![HashSet::new(); keys.len()];
1157 rows.retain(|row| {
1158 let tuples: Vec<Option<String>> = keys
1159 .iter()
1160 .map(|key| {
1161 let values: Vec<&SeedValue> = key.iter().map(|&index| &row[index]).collect();
1162 values
1163 .iter()
1164 .all(|value| !matches!(value, SeedValue::Null | SeedValue::Default))
1165 .then(|| format!("{values:?}"))
1166 })
1167 .collect();
1168 let repeats = tuples
1169 .iter()
1170 .zip(&seen)
1171 .any(|(tuple, seen)| tuple.as_ref().is_some_and(|tuple| seen.contains(tuple)));
1172 if !repeats {
1173 for (tuple, seen) in tuples.into_iter().zip(&mut seen) {
1174 if let Some(tuple) = tuple {
1175 seen.insert(tuple);
1176 }
1177 }
1178 }
1179 !repeats
1180 });
1181}
1182
1183fn unique_value<R: generator::RngCore + ?Sized>(
1187 value: SeedValue,
1188 seen: &mut HashSet<String>,
1189 row_idx: usize,
1190 column: &ColumnRef,
1191 mut regenerate: impl FnMut(&mut R) -> SeedValue,
1192 rng: &mut R,
1193) -> SeedValue {
1194 const ATTEMPTS: usize = 16;
1195 let key = |value: &SeedValue| format!("{value:?}");
1196 if matches!(
1197 value,
1198 SeedValue::Default | SeedValue::Null | SeedValue::CurrentTime
1199 ) {
1200 return value;
1201 }
1202 let mut value = value;
1203 for _ in 0..ATTEMPTS {
1204 if seen.insert(key(&value)) {
1205 return value;
1206 }
1207 value = regenerate(rng);
1208 }
1209
1210 let max_chars = inference::declared_char_length(&column.sql_type.to_uppercase());
1211 let mut suffix = row_idx;
1212 loop {
1213 let candidate = match &value {
1214 SeedValue::Integer(number) => {
1215 SeedValue::Integer(number.wrapping_add(i64::try_from(suffix).unwrap_or(0) + 1))
1216 }
1217 SeedValue::Float(number) => SeedValue::Float(number + suffix as f64 + 1.0),
1218 SeedValue::Text(text) => {
1219 let tag = format!("-{suffix}");
1220 let keep =
1221 max_chars.map_or(usize::MAX, |max| max.saturating_sub(tag.chars().count()));
1222 SeedValue::Text(text.chars().take(keep).chain(tag.chars()).collect())
1223 }
1224 SeedValue::Blob(bytes) => {
1225 let mut bytes = bytes.clone();
1226 bytes.extend_from_slice(&(suffix as u64).to_be_bytes());
1227 SeedValue::Blob(bytes)
1228 }
1229 other => return other.clone(),
1231 };
1232 if seen.insert(key(&candidate)) {
1233 return candidate;
1234 }
1235 suffix = suffix.wrapping_add(1);
1236 }
1237}
1238
1239fn row_param_count(row: &[SeedValue]) -> usize {
1244 row.iter()
1245 .filter(|v| !matches!(v, SeedValue::Default | SeedValue::CurrentTime))
1246 .count()
1247}
1248
1249fn batch_ranges_by_param_limit(
1250 rows: &[Vec<SeedValue>],
1251 param_limit: usize,
1252) -> Result<Vec<(usize, usize)>, usize> {
1253 if rows.is_empty() {
1254 return Ok(Vec::new());
1255 }
1256
1257 let mut ranges = Vec::new();
1258 let mut start = 0usize;
1259 let mut current_params = 0usize;
1260
1261 for (idx, row) in rows.iter().enumerate() {
1262 let row_params = row_param_count(row);
1263 if row_params > param_limit {
1264 return Err(row_params);
1265 }
1266 if idx > start && current_params.saturating_add(row_params) > param_limit {
1267 ranges.push((start, idx));
1268 start = idx;
1269 current_params = 0;
1270 }
1271
1272 current_params = current_params.saturating_add(row_params);
1273 }
1274
1275 if start < rows.len() {
1276 ranges.push((start, rows.len()));
1277 }
1278
1279 Ok(ranges)
1280}
1281
1282#[cfg(any(
1287 feature = "mysql",
1288 all(test, any(feature = "sqlite", feature = "postgres"))
1289))]
1290fn build_insert_sql<V>(table: &TableRef, rows: &[Vec<SQL<'static, V>>]) -> OwnedSQL<V>
1291where
1292 V: drizzle_core::SQLParam + Clone + ToOwned<Owned = V> + 'static,
1293{
1294 build_insert_sql_with(table, rows, false)
1295}
1296
1297#[cfg(any(feature = "postgres", feature = "mysql", all(test, feature = "sqlite")))]
1300fn build_insert_sql_with<V>(
1301 table: &TableRef,
1302 rows: &[Vec<SQL<'static, V>>],
1303 overriding_system_value: bool,
1304) -> OwnedSQL<V>
1305where
1306 V: drizzle_core::SQLParam + Clone + ToOwned<Owned = V> + 'static,
1307{
1308 let columns: Vec<usize> = table
1309 .columns
1310 .iter()
1311 .enumerate()
1312 .filter(|(_, column)| generated_expression(column).is_none())
1313 .map(|(index, _)| index)
1314 .collect();
1315 build_insert_sql_columns(table, &columns, rows, overriding_system_value)
1316}
1317
1318#[cfg(any(feature = "sqlite", feature = "postgres", feature = "mysql"))]
1321fn build_insert_sql_columns<V>(
1322 table: &TableRef,
1323 columns: &[usize],
1324 rows: &[Vec<SQL<'static, V>>],
1325 overriding_system_value: bool,
1326) -> OwnedSQL<V>
1327where
1328 V: drizzle_core::SQLParam + Clone + ToOwned<Owned = V> + 'static,
1329{
1330 let column_idents = SQL::join(
1331 columns
1332 .iter()
1333 .map(|&index| SQL::<'static, V>::ident(table.columns[index].name.to_string())),
1334 Token::COMMA,
1335 );
1336
1337 let mut sql = SQL::<'static, V>::token(Token::INSERT)
1338 .push(Token::INTO)
1339 .append(SQL::<'static, V>::table(statement_table(table)))
1340 .append(column_idents.parens());
1341 if overriding_system_value {
1342 sql = sql.append(SQL::raw("OVERRIDING SYSTEM VALUE"));
1343 }
1344 let sql = sql.push(Token::VALUES);
1345
1346 let mut values_sql = SQL::<'static, V>::empty();
1347 for (row_idx, row) in rows.iter().enumerate() {
1348 if row_idx > 0 {
1349 values_sql = values_sql.push(Token::COMMA);
1350 }
1351 debug_assert_eq!(row.len(), table.columns.len());
1352 let row_sql = SQL::join(
1353 columns.iter().map(|&index| row[index].clone()),
1354 Token::COMMA,
1355 );
1356 values_sql = values_sql.append(row_sql.parens());
1357 }
1358
1359 sql.append(values_sql).into_owned()
1360}
1361
1362#[cfg(any(feature = "sqlite", feature = "postgres", feature = "mysql"))]
1363fn build_delete_sql<V>(table: &TableRef) -> OwnedSQL<V>
1364where
1365 V: drizzle_core::SQLParam + Clone + ToOwned<Owned = V> + 'static,
1366{
1367 SQL::<'static, V>::token(Token::DELETE)
1368 .push(Token::FROM)
1369 .append(SQL::table(statement_table(table)))
1370 .into_owned()
1371}
1372
1373#[cfg(any(feature = "sqlite", feature = "postgres", feature = "mysql"))]
1374fn build_reset_sql<V>(tables: &[&TableRef]) -> Vec<(&'static str, OwnedSQL<V>)>
1375where
1376 V: drizzle_core::SQLParam + Clone + ToOwned<Owned = V> + 'static,
1377{
1378 let mut statements = Vec::new();
1379 for table in tables {
1380 let self_reference_columns = topology::nullable_self_reference_columns(table);
1381 if !self_reference_columns.is_empty() {
1382 let assignments = SQL::join(
1383 self_reference_columns.into_iter().map(|column| {
1384 SQL::<'static, V>::ident(column.to_string())
1385 .push(Token::EQ)
1386 .push(Token::NULL)
1387 }),
1388 Token::COMMA,
1389 );
1390 statements.push((
1391 table.name,
1392 SQL::<'static, V>::token(Token::UPDATE)
1393 .append(SQL::table(statement_table(table)))
1394 .push(Token::SET)
1395 .append(assignments)
1396 .into_owned(),
1397 ));
1398 }
1399 statements.push((table.name, build_delete_sql(table)));
1400 }
1401 statements
1402}
1403
1404#[cfg(feature = "mysql")]
1405fn build_mysql_auto_increment_reset_sql(table: &TableRef) -> OwnedSQL<OwnedMySQLValue> {
1406 SQL::<'static, OwnedMySQLValue>::token(Token::ALTER)
1407 .push(Token::TABLE)
1408 .append(SQL::table(statement_table(table)))
1409 .append(SQL::raw(" AUTO_INCREMENT = 1"))
1410 .into_owned()
1411}
1412
1413#[cfg(feature = "sqlite")]
1414fn seed_value_to_sqlite_sql(value: &SeedValue) -> SQL<'static, OwnedSQLiteValue> {
1415 match value {
1416 SeedValue::Default => SQL::token(Token::DEFAULT),
1417 SeedValue::Null => SQL::param(Cow::Owned(OwnedSQLiteValue::Null)),
1418 SeedValue::Integer(v) => SQL::param(Cow::Owned(OwnedSQLiteValue::Integer(*v))),
1419 SeedValue::Float(v) => SQL::param(Cow::Owned(OwnedSQLiteValue::Real(*v))),
1420 SeedValue::Text(v) => SQL::param(Cow::Owned(OwnedSQLiteValue::Text(v.clone()))),
1421 SeedValue::Bool(v) => SQL::param(Cow::Owned(OwnedSQLiteValue::Integer(i64::from(*v)))),
1422 SeedValue::Blob(v) => SQL::param(Cow::Owned(OwnedSQLiteValue::Blob(
1423 v.clone().into_boxed_slice(),
1424 ))),
1425 SeedValue::CurrentTime => SQL::raw("CURRENT_TIMESTAMP"),
1426 }
1427}
1428
1429#[cfg(feature = "sqlite")]
1434fn build_sqlite_statements(chunk: &GeneratedChunk<'_>, out: &mut Vec<SQLiteSeedStatement>) {
1435 let insertable: Vec<usize> = chunk
1436 .table
1437 .columns
1438 .iter()
1439 .enumerate()
1440 .filter(|(_, column)| generated_expression(column).is_none())
1441 .map(|(index, _)| index)
1442 .collect();
1443 let given = |row: &[SeedValue]| -> Vec<usize> {
1444 insertable
1445 .iter()
1446 .copied()
1447 .filter(|&index| !matches!(row[index], SeedValue::Default))
1448 .collect()
1449 };
1450
1451 let mut start = 0;
1452 while start < chunk.rows.len() {
1453 let columns = given(&chunk.rows[start]);
1454 let mut end = start + 1;
1455 while end < chunk.rows.len() && given(&chunk.rows[end]) == columns {
1456 end += 1;
1457 }
1458 let run = &chunk.rows[start..end];
1459 if columns.is_empty() {
1460 for _ in run {
1461 out.push(SQLiteSeedStatement {
1462 inner: SQL::<'static, OwnedSQLiteValue>::token(Token::INSERT)
1463 .push(Token::INTO)
1464 .append(SQL::table(statement_table(chunk.table)))
1465 .append(SQL::raw("DEFAULT VALUES"))
1466 .into_owned(),
1467 table: chunk.table.name,
1468 rows: 1,
1469 });
1470 }
1471 } else {
1472 let rows: Vec<Vec<SQL<'static, OwnedSQLiteValue>>> = run
1473 .iter()
1474 .map(|row| row.iter().map(seed_value_to_sqlite_sql).collect())
1475 .collect();
1476 out.push(SQLiteSeedStatement {
1477 inner: build_insert_sql_columns(chunk.table, &columns, &rows, false),
1478 table: chunk.table.name,
1479 rows: rows.len(),
1480 });
1481 }
1482 start = end;
1483 }
1484}
1485
1486#[cfg(feature = "postgres")]
1487fn seed_value_to_postgres_sql(
1488 value: &SeedValue,
1489 col: &ColumnRef,
1490) -> SQL<'static, OwnedPostgresValue> {
1491 match value {
1492 SeedValue::Default => SQL::token(Token::DEFAULT),
1493 SeedValue::Null => SQL::param(Cow::Owned(OwnedPostgresValue::Null)),
1494 SeedValue::Integer(v) => {
1495 let owned = match normalize_pg_type(col.sql_type).as_str() {
1498 "SMALLINT" | "INT2" | "SMALLSERIAL" | "SERIAL2" => {
1499 let clamped = (*v).clamp(i64::from(i16::MIN), i64::from(i16::MAX));
1500 OwnedPostgresValue::Smallint(i16::try_from(clamped).unwrap_or(0))
1502 }
1503 "INTEGER" | "INT" | "INT4" | "SERIAL" | "SERIAL4" => {
1504 let clamped = (*v).clamp(i64::from(i32::MIN), i64::from(i32::MAX));
1505 OwnedPostgresValue::Integer(i32::try_from(clamped).unwrap_or(0))
1507 }
1508 _ => OwnedPostgresValue::Bigint(*v),
1509 };
1510 SQL::param(Cow::Owned(owned))
1511 }
1512 SeedValue::Float(v) => SQL::param(Cow::Owned(OwnedPostgresValue::DoublePrecision(*v))),
1513 SeedValue::Text(v) => {
1514 #[cfg(feature = "chrono")]
1515 if let Some(value) = text_to_typed_postgres_value(v, col) {
1516 return SQL::param(Cow::Owned(value));
1517 }
1518
1519 let param = SQL::param(Cow::Owned(OwnedPostgresValue::Text(v.clone())));
1520 match postgres_cast_type(col) {
1524 Some(cast_type) => SQL::raw("CAST(")
1525 .append(param)
1526 .append(SQL::raw(format!(" AS {cast_type})"))),
1527 None => param,
1528 }
1529 }
1530 SeedValue::Bool(v) => SQL::param(Cow::Owned(OwnedPostgresValue::Boolean(*v))),
1531 SeedValue::Blob(v) => SQL::param(Cow::Owned(OwnedPostgresValue::Bytea(v.clone()))),
1532 SeedValue::CurrentTime => SQL::raw("now()"),
1533 }
1534}
1535
1536#[cfg(feature = "postgres")]
1539fn postgres_cast_type(col: &ColumnRef) -> Option<String> {
1540 let dimensions = match col.dialect {
1541 ColumnDialect::PostgreSQL { dimensions, .. } => dimensions.unwrap_or(0),
1542 _ => 0,
1543 };
1544 let ty = normalize_pg_type(col.sql_type);
1545 let base = ty.split('(').next().unwrap_or_default().trim();
1546 let is_text = matches!(
1547 base,
1548 "TEXT" | "VARCHAR" | "CHARACTER VARYING" | "CHAR" | "CHARACTER" | "BPCHAR" | "NAME" | ""
1549 );
1550 if is_text && dimensions == 0 {
1551 return None;
1552 }
1553 let brackets = "[]".repeat(usize::try_from(dimensions).unwrap_or(0));
1554 Some(format!("{}{brackets}", col.sql_type))
1555}
1556
1557#[cfg(feature = "postgres")]
1558fn normalize_pg_type(sql_type: &str) -> String {
1559 let mut out = String::new();
1560 let mut last_was_space = false;
1561 for ch in sql_type.trim().chars() {
1562 if ch.is_whitespace() {
1563 if !last_was_space {
1564 out.push(' ');
1565 last_was_space = true;
1566 }
1567 } else {
1568 out.push(ch.to_ascii_uppercase());
1569 last_was_space = false;
1570 }
1571 }
1572 out
1573}
1574
1575#[cfg(all(feature = "postgres", feature = "chrono"))]
1576fn text_to_typed_postgres_value(value: &str, col: &ColumnRef) -> Option<OwnedPostgresValue> {
1577 let ty = normalize_pg_type(col.sql_type);
1578
1579 if ty.contains("DATE") && !ty.contains("TIME") {
1580 return NaiveDate::parse_from_str(value, "%Y-%m-%d")
1581 .ok()
1582 .map(OwnedPostgresValue::Date);
1583 }
1584
1585 if ty.contains("TIMESTAMP") {
1586 let timestamp = NaiveDateTime::parse_from_str(value, "%Y-%m-%d %H:%M:%S").ok()?;
1587 if ty.contains("TIME ZONE") || ty.contains("TIMESTAMPTZ") {
1588 let utc = DateTime::<Utc>::from_naive_utc_and_offset(timestamp, Utc);
1589 return Some(OwnedPostgresValue::TimestampTz(utc.fixed_offset()));
1590 }
1591 return Some(OwnedPostgresValue::Timestamp(timestamp));
1592 }
1593
1594 if ty == "TIME" || ty.starts_with("TIME(") || ty.starts_with("TIME ") {
1595 return NaiveTime::parse_from_str(value, "%H:%M:%S")
1596 .ok()
1597 .map(OwnedPostgresValue::Time);
1598 }
1599
1600 None
1601}
1602
1603#[cfg(feature = "postgres")]
1604fn build_postgres_statement(chunk: &GeneratedChunk<'_>) -> PostgresSeedStatement {
1605 let columns = chunk.table.columns;
1606 let rows: Vec<Vec<SQL<'static, OwnedPostgresValue>>> = chunk
1607 .rows
1608 .iter()
1609 .map(|row| {
1610 row.iter()
1611 .enumerate()
1612 .map(|(idx, value)| seed_value_to_postgres_sql(value, &columns[idx]))
1613 .collect()
1614 })
1615 .collect();
1616
1617 let explicit_identity_always = columns.iter().enumerate().any(|(idx, column)| {
1618 matches!(
1619 column.dialect,
1620 ColumnDialect::PostgreSQL {
1621 is_identity_always: true,
1622 ..
1623 }
1624 ) && chunk
1625 .rows
1626 .iter()
1627 .any(|row| !matches!(row[idx], SeedValue::Default))
1628 });
1629
1630 PostgresSeedStatement {
1631 inner: build_insert_sql_with(chunk.table, &rows, explicit_identity_always),
1632 table: chunk.table.name,
1633 rows: rows.len(),
1634 }
1635}
1636
1637#[cfg(feature = "postgres")]
1641fn build_postgres_sequence_sync(chunks: &[&GeneratedChunk<'_>]) -> Vec<PostgresSeedStatement> {
1642 let Some(table) = chunks.first().map(|chunk| chunk.table) else {
1643 return Vec::new();
1644 };
1645 let qualified_table = match statement_table(table).schema {
1646 Some(schema) => format!("{}.{}", quote_pg_ident(schema), quote_pg_ident(table.name)),
1647 None => quote_pg_ident(table.name),
1648 };
1649 table
1650 .columns
1651 .iter()
1652 .enumerate()
1653 .filter(|(idx, column)| {
1654 let uses_sequence = matches!(
1655 column.dialect,
1656 ColumnDialect::PostgreSQL {
1657 is_serial: true,
1658 ..
1659 } | ColumnDialect::PostgreSQL {
1660 is_bigserial: true,
1661 ..
1662 } | ColumnDialect::PostgreSQL {
1663 is_generated_identity: true,
1664 ..
1665 }
1666 );
1667 uses_sequence
1668 && chunks.iter().any(|chunk| {
1669 chunk
1670 .rows
1671 .iter()
1672 .any(|row| matches!(row[*idx], SeedValue::Integer(_)))
1673 })
1674 })
1675 .map(|(_, column)| {
1676 let sql =
1677 SQL::<'static, OwnedPostgresValue>::raw("SELECT setval(pg_get_serial_sequence(")
1678 .append(SQL::param(Cow::Owned(OwnedPostgresValue::Text(
1679 qualified_table.clone(),
1680 ))))
1681 .push(Token::COMMA)
1682 .append(SQL::param(Cow::Owned(OwnedPostgresValue::Text(
1683 column.name.to_string(),
1684 ))))
1685 .append(SQL::raw("), (SELECT MAX("))
1686 .append(SQL::ident(column.name.to_string()))
1687 .append(SQL::raw(") FROM"))
1688 .append(SQL::table(statement_table(table)))
1689 .append(SQL::raw("))"));
1690 PostgresSeedStatement {
1691 inner: sql.into_owned(),
1692 table: table.name,
1693 rows: 0,
1694 }
1695 })
1696 .collect()
1697}
1698
1699#[cfg(feature = "postgres")]
1700fn quote_pg_ident(name: &str) -> String {
1701 format!("\"{}\"", name.replace('"', "\"\""))
1702}
1703
1704#[cfg(test)]
1709struct FkGen {
1710 parent_values: Vec<SeedValue>,
1711 children_per_parent: usize,
1712}
1713
1714#[cfg(test)]
1715impl Generator for FkGen {
1716 fn generate(
1717 &self,
1718 _rng: &mut dyn generator::RngCore,
1719 index: usize,
1720 _sql_type: &str,
1721 ) -> SeedValue {
1722 if self.parent_values.is_empty() || self.children_per_parent == 0 {
1723 return SeedValue::Null;
1724 }
1725 let idx = (index / self.children_per_parent) % self.parent_values.len();
1726 self.parent_values[idx].clone()
1727 }
1728 fn name(&self) -> &'static str {
1729 "ForeignKey"
1730 }
1731}
1732
1733struct DefaultGen;
1734
1735impl Generator for DefaultGen {
1736 fn generate(
1737 &self,
1738 _rng: &mut dyn generator::RngCore,
1739 _index: usize,
1740 _sql_type: &str,
1741 ) -> SeedValue {
1742 SeedValue::Default
1743 }
1744 fn name(&self) -> &'static str {
1745 "Default"
1746 }
1747}
1748
1749impl<C> Generator for &'static C
1755where
1756 C: drizzle_core::SQLColumnInfo,
1757{
1758 fn generate(
1759 &self,
1760 rng: &mut dyn generator::RngCore,
1761 index: usize,
1762 sql_type: &str,
1763 ) -> SeedValue {
1764 let mut flags = drizzle_core::ColumnFlags::empty();
1766 if self.is_primary_key() {
1767 flags |= drizzle_core::ColumnFlags::PRIMARY_KEY;
1768 }
1769 if self.has_default() {
1770 flags |= drizzle_core::ColumnFlags::HAS_DEFAULT;
1771 }
1772 let col_ref = ColumnRef {
1773 table: "",
1774 name: self.name(),
1775 sql_type: self.r#type(),
1776 flags,
1777 dialect: drizzle_core::ColumnDialect::SQLite {
1778 autoincrement: false,
1779 default: None,
1780 generated_expression: None,
1781 generated_stored: false,
1782 collate: None,
1783 enum_variants: None,
1784 },
1785 };
1786 inference::infer_generator(&col_ref).generate(rng, index, sql_type)
1787 }
1788
1789 fn name(&self) -> &'static str {
1790 "Column"
1791 }
1792}
1793
1794#[cfg(test)]
1795mod tests {
1796 use super::*;
1797
1798 #[cfg(feature = "sqlite")]
1799 type SeedTestValue = OwnedSQLiteValue;
1800 #[cfg(all(not(feature = "sqlite"), feature = "postgres"))]
1801 type SeedTestValue = OwnedPostgresValue;
1802 #[cfg(all(not(feature = "sqlite"), not(feature = "postgres"), feature = "mysql"))]
1803 type SeedTestValue = OwnedMySQLValue;
1804
1805 #[test]
1806 fn arc_generator_delegation() {
1807 use rand::SeedableRng;
1808 use rand::rngs::StdRng;
1809
1810 let g: Arc<dyn Generator> = Arc::new(generator::numeric::IntPrimaryKeyGen);
1811 let mut rng = StdRng::seed_from_u64(42);
1812
1813 assert_eq!(g.generate(&mut rng, 0, "INTEGER"), SeedValue::Integer(1));
1814 assert_eq!(g.generate(&mut rng, 4, "INTEGER"), SeedValue::Integer(5));
1815 assert_eq!(g.name(), "IntPrimaryKey");
1816 }
1817
1818 #[test]
1819 fn fk_gen_picks_from_parent_values() {
1820 use rand::SeedableRng;
1821 use rand::rngs::StdRng;
1822
1823 let parent_vals = vec![
1824 SeedValue::Integer(10),
1825 SeedValue::Integer(20),
1826 SeedValue::Integer(30),
1827 ];
1828 let g = FkGen {
1829 parent_values: parent_vals.clone(),
1830 children_per_parent: 1,
1831 };
1832 let mut rng = StdRng::seed_from_u64(42);
1833
1834 for i in 0..6 {
1835 let val = g.generate(&mut rng, i, "INTEGER");
1836 assert!(
1837 parent_vals.contains(&val),
1838 "FK value {:?} not in parent set",
1839 val
1840 );
1841 }
1842 }
1843
1844 #[test]
1845 fn fk_gen_empty_parent_returns_null() {
1846 use rand::SeedableRng;
1847 use rand::rngs::StdRng;
1848
1849 let g = FkGen {
1850 parent_values: vec![],
1851 children_per_parent: 1,
1852 };
1853 let mut rng = StdRng::seed_from_u64(42);
1854 assert_eq!(g.generate(&mut rng, 0, "INTEGER"), SeedValue::Null);
1855 }
1856
1857 #[test]
1858 fn default_gen_returns_default_keyword() {
1859 use rand::SeedableRng;
1860 use rand::rngs::StdRng;
1861
1862 let g = DefaultGen;
1863 let mut rng = StdRng::seed_from_u64(42);
1864 assert_eq!(g.generate(&mut rng, 0, "TEXT"), SeedValue::Default);
1865 }
1866
1867 #[test]
1868 fn fk_gen_with_relation_count_is_deterministic() {
1869 use rand::SeedableRng;
1870 use rand::rngs::StdRng;
1871
1872 let g = FkGen {
1873 parent_values: vec![SeedValue::Integer(1), SeedValue::Integer(2)],
1874 children_per_parent: 3,
1875 };
1876 let mut rng = StdRng::seed_from_u64(42);
1877
1878 let generated: Vec<SeedValue> =
1879 (0..6).map(|i| g.generate(&mut rng, i, "INTEGER")).collect();
1880 assert_eq!(
1881 generated,
1882 vec![
1883 SeedValue::Integer(1),
1884 SeedValue::Integer(1),
1885 SeedValue::Integer(1),
1886 SeedValue::Integer(2),
1887 SeedValue::Integer(2),
1888 SeedValue::Integer(2),
1889 ]
1890 );
1891 }
1892
1893 #[test]
1894 fn batch_ranges_split_on_param_limit() {
1895 let rows = vec![
1896 vec![SeedValue::Integer(1), SeedValue::Text("a".to_string())],
1897 vec![SeedValue::Integer(2), SeedValue::Text("b".to_string())],
1898 vec![SeedValue::Integer(3), SeedValue::Text("c".to_string())],
1899 vec![SeedValue::Integer(4), SeedValue::Text("d".to_string())],
1900 vec![SeedValue::Integer(5), SeedValue::Text("e".to_string())],
1901 ];
1902
1903 let ranges = batch_ranges_by_param_limit(&rows, 4);
1904 assert_eq!(ranges.unwrap(), vec![(0, 2), (2, 4), (4, 5)]);
1905 }
1906
1907 #[test]
1908 fn batch_ranges_counts_default_as_zero_params() {
1909 let rows = vec![
1910 vec![SeedValue::Default, SeedValue::Integer(1)],
1911 vec![SeedValue::Default, SeedValue::Integer(2)],
1912 vec![SeedValue::Default, SeedValue::Integer(3)],
1913 ];
1914
1915 let ranges = batch_ranges_by_param_limit(&rows, 2);
1916 assert_eq!(ranges.unwrap(), vec![(0, 2), (2, 3)]);
1917 }
1918
1919 #[test]
1920 fn batch_ranges_current_time_counts_as_zero_params() {
1921 let rows = vec![
1922 vec![SeedValue::Integer(1), SeedValue::CurrentTime],
1923 vec![SeedValue::Integer(2), SeedValue::CurrentTime],
1924 vec![SeedValue::Integer(3), SeedValue::CurrentTime],
1925 ];
1926
1927 let ranges = batch_ranges_by_param_limit(&rows, 2);
1930 assert_eq!(ranges.unwrap(), vec![(0, 2), (2, 3)]);
1931 }
1932
1933 #[cfg(any(feature = "sqlite", feature = "postgres", feature = "mysql"))]
1934 #[test]
1935 fn insert_sql_omits_generated_columns_for_every_dialect() {
1936 let generated_dialects = [
1937 ColumnDialect::SQLite {
1938 autoincrement: false,
1939 default: None,
1940 generated_expression: Some("LENGTH(app_default)"),
1941 generated_stored: true,
1942 collate: None,
1943 enum_variants: None,
1944 },
1945 ColumnDialect::PostgreSQL {
1946 postgres_type: "INTEGER",
1947 dimensions: None,
1948 is_serial: false,
1949 is_bigserial: false,
1950 is_generated_identity: false,
1951 is_identity_always: false,
1952 default: None,
1953 generated_expression: Some("LENGTH(app_default)"),
1954 generated_stored: true,
1955 collate: None,
1956 comment: None,
1957 enum_variants: None,
1958 },
1959 ColumnDialect::MySQL {
1960 auto_increment: false,
1961 default: None,
1962 generated_expression: Some("CHAR_LENGTH(app_default)"),
1963 generated_stored: true,
1964 charset: None,
1965 collate: None,
1966 on_update: None,
1967 },
1968 ];
1969
1970 for generated_dialect in generated_dialects {
1971 let columns = Box::leak(Box::new([
1972 ColumnRef::sql("seed_values", "db_default"),
1973 ColumnRef::sql("seed_values", "app_default"),
1974 ColumnRef {
1975 table: "seed_values",
1976 name: "computed",
1977 sql_type: "INTEGER",
1978 flags: drizzle_core::ColumnFlags::empty(),
1979 dialect: generated_dialect,
1980 },
1981 ]));
1982 let mut table =
1983 TableRef::sql("seed_values", &["db_default", "app_default", "computed"]);
1984 table.columns = columns;
1985 let rows = [vec![
1986 SQL::<'static, SeedTestValue>::token(Token::DEFAULT),
1987 SQL::raw("'application-default'"),
1988 SQL::raw("'generated-value'"),
1989 ]];
1990
1991 let sql = build_insert_sql(&table, &rows).to_sql().sql();
1992
1993 assert!(sql.contains("db_default"), "{generated_dialect:?}");
1994 assert!(sql.contains("DEFAULT"), "{generated_dialect:?}");
1995 assert!(sql.contains("app_default"), "{generated_dialect:?}");
1996 assert!(sql.contains("application-default"), "{generated_dialect:?}");
1997 assert!(!sql.contains("computed"), "{generated_dialect:?}");
1998 assert!(!sql.contains("generated-value"), "{generated_dialect:?}");
1999 }
2000 }
2001
2002 #[cfg(all(feature = "postgres", feature = "chrono"))]
2003 #[test]
2004 fn postgres_date_text_binds_as_date_param() {
2005 use drizzle_core::{ColumnDialect, ColumnFlags};
2006
2007 let col = ColumnRef {
2008 table: "employees",
2009 name: "birth_date",
2010 sql_type: "DATE",
2011 flags: ColumnFlags::empty(),
2012 dialect: ColumnDialect::PostgreSQL {
2013 postgres_type: "DATE",
2014 dimensions: None,
2015 is_serial: false,
2016 is_bigserial: false,
2017 is_generated_identity: false,
2018 is_identity_always: false,
2019 default: None,
2020 generated_expression: None,
2021 generated_stored: false,
2022 collate: None,
2023 comment: None,
2024 enum_variants: None,
2025 },
2026 };
2027
2028 let sql = seed_value_to_postgres_sql(&SeedValue::Text("2024-03-09".to_string()), &col);
2029 let (_, params) = sql.build();
2030
2031 assert!(matches!(params[0], OwnedPostgresValue::Date(_)));
2032 }
2033
2034 #[cfg(feature = "postgres")]
2035 #[test]
2036 fn postgres_integers_bind_at_the_column_width() {
2037 use drizzle_core::{ColumnDialect, ColumnFlags};
2038
2039 fn bind(sql_type: &'static str, value: i64) -> OwnedPostgresValue {
2040 let col = ColumnRef {
2041 table: "t",
2042 name: "c",
2043 sql_type,
2044 flags: ColumnFlags::empty(),
2045 dialect: ColumnDialect::PostgreSQL {
2046 postgres_type: sql_type,
2047 dimensions: None,
2048 is_serial: false,
2049 is_bigserial: false,
2050 is_generated_identity: false,
2051 is_identity_always: false,
2052 default: None,
2053 generated_expression: None,
2054 generated_stored: false,
2055 collate: None,
2056 comment: None,
2057 enum_variants: None,
2058 },
2059 };
2060 let sql = seed_value_to_postgres_sql(&SeedValue::Integer(value), &col);
2061 let (_, params) = sql.build();
2062 params[0].clone()
2063 }
2064
2065 let big = i64::from(i32::MAX) + 1;
2066 for ty in ["BIGINT", "bigint", "INT8", "BIGSERIAL", "SERIAL8"] {
2067 assert_eq!(bind(ty, big), OwnedPostgresValue::Bigint(big), "{ty}");
2068 }
2069 for ty in ["INTEGER", "int", "INT4", "SERIAL", "SERIAL4"] {
2070 assert_eq!(bind(ty, 7), OwnedPostgresValue::Integer(7), "{ty}");
2071 }
2072 for ty in ["SMALLINT", "INT2", "SMALLSERIAL", "SERIAL2"] {
2073 assert_eq!(bind(ty, 7), OwnedPostgresValue::Smallint(7), "{ty}");
2074 }
2075 }
2076
2077 #[test]
2078 fn fk_gen_zero_children_per_parent_returns_null() {
2079 use rand::SeedableRng;
2080 use rand::rngs::StdRng;
2081
2082 let g = FkGen {
2083 parent_values: vec![SeedValue::Integer(1)],
2084 children_per_parent: 0,
2085 };
2086 let mut rng = StdRng::seed_from_u64(42);
2087 assert_eq!(g.generate(&mut rng, 0, "INTEGER"), SeedValue::Null);
2088 }
2089}