1use crate::{
3 dbl::model::*,
4 one::{
5 Path,
6 graph::FinGraph,
7 graph_algorithms::{ToposortData, toposort_lenient},
8 },
9 zero::{QualifiedLabel, QualifiedName, name},
10};
11use derive_more::Constructor;
12use indexmap::IndexMap;
13use itertools::Itertools;
14use nonempty::nonempty;
15use sea_query::SchemaBuilder;
16use sea_query::{
17 Alias, ColumnDef, ForeignKey, ForeignKeyCreateStatement, Iden, MysqlQueryBuilder,
18 PostgresQueryBuilder, SqliteQueryBuilder, Table, TableCreateStatement, prepare::Write,
19};
20use sqlformat::{Dialect, format};
21use std::fmt;
22
23impl Iden for QualifiedName {
24 fn unquoted(&self, s: &mut dyn Write) {
25 Iden::unquoted(&format!("{self}").as_str(), s)
26 }
27}
28
29impl Iden for QualifiedLabel {
30 fn unquoted(&self, s: &mut dyn Write) {
31 Iden::unquoted(&format!("{self}").as_str(), s)
32 }
33}
34
35impl Iden for &QualifiedLabel {
36 fn unquoted(&self, s: &mut dyn Write) {
37 Iden::unquoted(&format!("{self}").as_str(), s)
38 }
39}
40
41#[derive(Debug, Clone, PartialEq)]
44pub enum ColumnType {
45 Ordinary {
47 mor: QualifiedName,
49 tgt: QualifiedName,
51 },
52 Deferrable {
54 mor: QualifiedName,
56 tgt: QualifiedName,
58 },
59 Attribute {
61 mor: QualifiedName,
63 tgt: QualifiedName,
65 },
66}
67
68impl ColumnType {
69 fn build(
70 model: &DiscreteDblModel,
71 cycles: &IndexMap<QualifiedName, Vec<QualifiedName>>,
72 src: &QualifiedName,
73 mor: QualifiedName,
74 ) -> Self {
75 let tgt = model.get_cod(&mor).unwrap();
76 match model.mor_generator_type(&mor) {
77 t if t == Path::Seq(nonempty![name("Attr")]) => {
78 ColumnType::Attribute { mor, tgt: tgt.clone() }
79 }
80 _ => {
81 if cycles.contains_key(src) || cycles.contains_key(&tgt.clone()) {
82 ColumnType::Deferrable { mor, tgt: tgt.clone() }
83 } else {
84 ColumnType::Ordinary { mor, tgt: tgt.clone() }
85 }
86 }
87 }
88 }
89
90 fn mor(&self) -> &QualifiedName {
91 match self {
92 ColumnType::Ordinary { mor, tgt: _ }
93 | ColumnType::Deferrable { mor, tgt: _ }
94 | ColumnType::Attribute { mor, tgt: _ } => mor,
95 }
96 }
97
98 fn tgt(&self) -> &QualifiedName {
99 match self {
100 ColumnType::Ordinary { mor: _, tgt }
101 | ColumnType::Deferrable { mor: _, tgt }
102 | ColumnType::Attribute { mor: _, tgt } => tgt,
103 }
104 }
105
106 fn render_postgres_fk(
109 &self,
110 src: &QualifiedName,
111 ob_label: impl Fn(&QualifiedName) -> String,
112 mor_label: impl Fn(&QualifiedName) -> String,
113 ) -> String {
114 let fk = |src: String, mor: &String, tgt: &String| -> String {
115 format!(
116 r#"ALTER TABLE "{src}"
117 ADD CONSTRAINT fk_{mor}_{src}_{tgt}
118 FOREIGN KEY ({mor}) REFERENCES "{tgt}" (id)"#
119 )
120 };
121 match self {
122 ColumnType::Ordinary { mor, tgt } => {
123 fk(ob_label(src), &mor_label(mor), &ob_label(tgt)) + ";"
124 }
125 ColumnType::Deferrable { mor, tgt } => {
126 fk(ob_label(src), &mor_label(mor), &ob_label(tgt))
127 + "\n"
128 + r#"DEFERRABLE INITIALLY DEFERRED;"#
129 }
130 ColumnType::Attribute { mor: _, tgt: _ } => unreachable!(),
132 }
133 }
134}
135
136#[derive(Clone, Debug)]
139pub struct ForeignKeyConstraints {
140 fks: IndexMap<QualifiedName, Vec<ColumnType>>,
142}
143
144impl ForeignKeyConstraints {
145 fn new(model: &DiscreteDblModel) -> Self {
146 let g = model.generating_graph();
147 let toposort: ToposortData<QualifiedName> = toposort_lenient(g);
148 let cycles = toposort.cycles;
149 let fks = IndexMap::from_iter(toposort.stack.into_iter().rev().filter_map(|v| {
150 (name("Entity") == model.ob_generator_type(&v)).then_some((
151 v.clone(),
152 g.out_edges(&v)
153 .map(|e| ColumnType::build(model, &cycles, &v, e))
154 .collect::<Vec<ColumnType>>(),
155 ))
156 }));
157 Self { fks }
158 }
159
160 fn any_deferrable(&self) -> bool {
161 self.fks
162 .values()
163 .flatten()
164 .into_iter()
165 .any(|s| matches!(s, ColumnType::Deferrable { mor: _, tgt: _ }))
166 }
167}
168
169#[derive(Clone, Debug, PartialEq)]
171pub enum SQLAnalysisError {
172 CyclicForeignKeyError {
174 backend: SQLBackend,
177 cycles: Vec<(QualifiedName, ColumnType)>,
179 },
180}
181
182impl std::fmt::Display for SQLAnalysisError {
183 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> fmt::Result {
184 match self {
185 SQLAnalysisError::CyclicForeignKeyError { backend, cycles } => write!(
186 f,
187 "Cycle detected at tables {:#?}. {backend} cannot support cyclic foreign keys.",
188 cycles
189 ),
190 }
191 }
192}
193
194#[derive(Constructor)]
196pub struct SQLAnalysis {
197 backend: SQLBackend,
198}
199
200impl SQLAnalysis {
201 pub fn format(&self, output: &str) -> String {
203 format(
204 output,
205 &sqlformat::QueryParams::None,
206 &sqlformat::FormatOptions {
207 lines_between_queries: 2,
208 dialect: self.backend.clone().into(),
209 ..Default::default()
210 },
211 )
212 }
213
214 fn build(
216 &self,
217 tables: Vec<TableCreateStatement>,
218 constraints: ForeignKeyConstraints,
219 ob_label: impl Fn(&QualifiedName) -> String,
220 mor_label: impl Fn(&QualifiedName) -> String,
221 ) -> String {
222 let table_def: String = tables
223 .iter()
224 .map(|table| match self.backend {
225 SQLBackend::MySQL => table.to_string(MysqlQueryBuilder),
226 SQLBackend::SQLite => table.to_string(SqliteQueryBuilder),
227 SQLBackend::PostgresSQL => table.to_string(PostgresQueryBuilder),
228 })
229 .join(";\n")
230 + ";";
231
232 let deferrable_fks: String = constraints
234 .fks
235 .iter()
236 .flat_map(|(ob, mors)| {
237 mors.iter()
238 .filter(|fkb| matches!(fkb, ColumnType::Deferrable { mor: _, tgt: _ }))
239 .map(|fkb| fkb.render_postgres_fk(ob, &ob_label, &mor_label))
240 .collect::<Vec<String>>()
241 })
242 .join("\n");
243
244 table_def + &deferrable_fks
245 }
246
247 fn validate_toposort(
248 &self,
249 constraints: ForeignKeyConstraints,
250 ) -> Result<ForeignKeyConstraints, SQLAnalysisError> {
251 if (self.backend == SQLBackend::MySQL || self.backend == SQLBackend::SQLite)
253 && constraints.any_deferrable()
254 {
255 let cycles = constraints
256 .fks
257 .into_iter()
258 .flat_map(|(k, v)| v.into_iter().map(move |e| (k.clone(), e)))
259 .filter(|(_, e)| matches!(e, ColumnType::Deferrable { mor: _, tgt: _ }))
260 .collect::<Vec<_>>();
261 Err(SQLAnalysisError::CyclicForeignKeyError { backend: self.backend.clone(), cycles })
262 } else {
263 Ok(constraints)
264 }
265 }
266
267 fn toposort_morphisms(
268 &self,
269 model: &DiscreteDblModel,
270 ) -> Result<ForeignKeyConstraints, SQLAnalysisError> {
271 let constraints = ForeignKeyConstraints::new(model);
273 self.validate_toposort(constraints)
274 }
275
276 pub fn render(
278 &self,
279 model: &DiscreteDblModel,
280 ob_label: impl Fn(&QualifiedName) -> String,
281 mor_label: impl Fn(&QualifiedName) -> String,
282 ) -> Result<String, SQLAnalysisError> {
283 let constraints = self.toposort_morphisms(model);
284 let tables = self.make_tables(model, constraints.clone()?, &ob_label, &mor_label);
285 let output: String = self.build(tables, constraints.clone()?, ob_label, mor_label);
286 let formatted_output = self.format(&output);
287 match self.backend {
289 SQLBackend::SQLite => Ok(["PRAGMA foreign_keys = ON", &formatted_output].join(";\n\n")),
290 _ => Ok(formatted_output),
291 }
292 }
293
294 fn fk(&self, src: &str, tgt: &str, mor: &str) -> ForeignKeyCreateStatement {
295 ForeignKey::create()
296 .name(format!("FK_{}_{}_{}", mor, src, tgt))
297 .from(Alias::new(src), Alias::new(mor))
298 .to(Alias::new(tgt), "id")
299 .to_owned()
300 }
301
302 fn make_tables(
303 &self,
304 model: &DiscreteDblModel,
305 constraints: ForeignKeyConstraints,
306 ob_label: impl Fn(&QualifiedName) -> String,
307 mor_label: impl Fn(&QualifiedName) -> String,
308 ) -> Vec<TableCreateStatement> {
309 constraints
310 .fks
311 .into_iter()
312 .map(|(ob, mors)| {
313 let mut tbl = Table::create();
314
315 let table_column_defs = mors.iter().fold(
317 tbl.table(Alias::new(ob_label(&ob))).if_not_exists().col(
318 ColumnDef::new("id").integer().not_null().auto_increment().primary_key(),
319 ),
320 |acc, mor| {
321 let mor_tgt = mor.tgt();
322 let ob_name = ob_label(mor_tgt);
323 let mor_name = mor_label(mor.mor());
324 if model.mor_generator_type(mor.mor()) == Path::Id(name("Entity")) {
327 acc.col(
328 ColumnDef::new(Alias::new(mor_name.as_str())).integer().not_null(),
329 )
330 } else {
331 let mut col = ColumnDef::new(Alias::new(mor_name.as_str()));
332 col.not_null();
333 add_column_type(&mut col, ob_name.as_str());
334 acc.col(col)
335 }
336 },
337 );
338
339 mors.iter()
340 .filter(|mor| {
341 (model.mor_generator_type(mor.mor()) == Path::Id(name("Entity")))
342 && (if self.backend == SQLBackend::PostgresSQL {
343 matches!(mor, ColumnType::Ordinary { mor: _, tgt: _ })
344 } else {
345 true
346 })
347 })
348 .fold(
349 table_column_defs,
351 |acc, mor| {
352 acc.foreign_key(&mut self.fk(
354 ob_label(&ob).as_str(),
355 ob_label(mor.tgt()).as_str(),
356 mor_label(mor.mor()).as_str(),
357 ))
358 },
359 )
360 .to_owned()
361 })
362 .collect()
363 }
364}
365
366#[derive(Debug, Clone, PartialEq)]
370pub enum SQLBackend {
371 MySQL,
373
374 SQLite,
376
377 PostgresSQL,
379}
380
381impl SQLBackend {
382 pub fn as_type(&self) -> Box<dyn SchemaBuilder> {
384 match self {
385 SQLBackend::MySQL => Box::new(MysqlQueryBuilder),
386 SQLBackend::SQLite => Box::new(SqliteQueryBuilder),
387 SQLBackend::PostgresSQL => Box::new(PostgresQueryBuilder),
388 }
389 }
390}
391
392impl From<SQLBackend> for Dialect {
393 fn from(backend: SQLBackend) -> sqlformat::Dialect {
394 match backend {
395 SQLBackend::PostgresSQL => Dialect::PostgreSql,
396 _ => Dialect::Generic,
397 }
398 }
399}
400
401impl TryFrom<&str> for SQLBackend {
402 type Error = String;
403 fn try_from(backend: &str) -> Result<Self, Self::Error> {
404 match backend {
405 "MySQL" => Ok(SQLBackend::MySQL),
406 "SQLite" => Ok(SQLBackend::SQLite),
407 "PostgresSQL" => Ok(SQLBackend::PostgresSQL),
408 _ => Err(String::from("Invalid backend")),
409 }
410 }
411}
412
413impl fmt::Display for SQLBackend {
414 fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
415 let string = match self {
416 SQLBackend::MySQL => "MySQL",
417 SQLBackend::SQLite => "SQLite",
418 SQLBackend::PostgresSQL => "PostgresSQL",
419 };
420 write!(f, "{}", string)
421 }
422}
423
424fn add_column_type(col: &mut ColumnDef, label: &str) {
425 match label {
426 "Int" => col.integer(),
427 "TinyInt" => col.tiny_integer(),
428 "Bool" => col.boolean(),
429 "Float" => col.float(),
430 "Time" => col.timestamp(),
431 "Date" => col.date(),
432 "DateTime" => col.date_time(),
433 _ => col.custom(Alias::new(label)),
434 };
435}
436
437#[cfg(test)]
438mod tests {
439 use expect_test::expect;
440 use std::rc::Rc;
441
442 use super::*;
443 use crate::{stdlib::th_schema, tt};
444
445 #[test]
446 fn sql_schema() {
447 let th = Rc::new(th_schema());
448 let source = "[
449 Person : Entity,
450 Dog : Entity,
451 walks : (Hom Entity)[Person, Dog],
452 Hair : AttrType,
453 has : Attr[Person, Hair],
454 ]";
455 let model = tt::modelgen::Model::from_text(&th.clone().into(), source)
456 .ok()
457 .and_then(|m| m.as_discrete())
458 .unwrap();
459
460 let expected = expect![[
461 r#"CREATE TABLE IF NOT EXISTS `Dog` (`id` int NOT NULL AUTO_INCREMENT PRIMARY KEY);
462
463CREATE TABLE IF NOT EXISTS `Person` (
464 `id` int NOT NULL AUTO_INCREMENT PRIMARY KEY,
465 `walks` int NOT NULL,
466 `has` Hair NOT NULL,
467 CONSTRAINT `FK_walks_Person_Dog` FOREIGN KEY (`walks`) REFERENCES `Dog` (`id`)
468);"#
469 ]];
470 let ddl = SQLAnalysis::new(SQLBackend::MySQL)
471 .render(
472 &model,
473 |id| format!("{id}").as_str().into(),
474 |id| format!("{id}").as_str().into(),
475 )
476 .expect("SQL should render");
477 expected.assert_eq(&ddl);
478 }
479
480 #[test]
481 fn sql_postgres_cycles() {
482 let th = Rc::new(th_schema());
483 let source = "[
484 Refs : Entity,
485 Snapshots : Entity,
486 head : (Hom Entity)[Refs, Snapshots],
487 for_ref: (Hom Entity)[Snapshots, Refs],
488 Timestamp : AttrType,
489 created : Attr[Refs, Timestamp],
490 last_updated: Attr[Snapshots, Timestamp],
491 ]";
492 let model = tt::modelgen::Model::from_text(&th.into(), source)
493 .ok()
494 .and_then(|m| m.as_discrete())
495 .unwrap();
496
497 let expected = expect![[r#"CREATE TABLE IF NOT EXISTS "Snapshots" (
498 "id" serial NOT NULL PRIMARY KEY,
499 "for_ref" integer NOT NULL,
500 "last_updated" Timestamp NOT NULL
501);
502
503CREATE TABLE IF NOT EXISTS "Refs" (
504 "id" serial NOT NULL PRIMARY KEY,
505 "head" integer NOT NULL,
506 "created" Timestamp NOT NULL
507);
508
509ALTER TABLE
510 "Snapshots"
511ADD
512 CONSTRAINT fk_for_ref_Snapshots_Refs FOREIGN KEY (for_ref) REFERENCES "Refs" (id) DEFERRABLE INITIALLY DEFERRED;
513
514ALTER TABLE
515 "Refs"
516ADD
517 CONSTRAINT fk_head_Refs_Snapshots FOREIGN KEY (head) REFERENCES "Snapshots" (id) DEFERRABLE INITIALLY DEFERRED;"#]];
518 let ddl = SQLAnalysis::new(SQLBackend::PostgresSQL)
519 .render(
520 &model,
521 |id| format!("{id}").as_str().into(),
522 |id| format!("{id}").as_str().into(),
523 )
524 .expect("SQL should render");
525 expected.assert_eq(&ddl);
526 }
527
528 #[test]
529 fn sql_mysql_cycles() {
530 let th = Rc::new(th_schema());
531 let source = "[
532 Refs : Entity,
533 Snapshots : Entity,
534 head : (Hom Entity)[Refs, Snapshots],
535 for_ref: (Hom Entity)[Snapshots, Refs],
536 Timestamp : AttrType,
537 created : Attr[Refs, Timestamp],
538 last_updated: Attr[Snapshots, Timestamp],
539 ]";
540 let model = tt::modelgen::Model::from_text(&th.into(), source)
541 .ok()
542 .and_then(|m| m.as_discrete())
543 .unwrap();
544
545 let ddl = SQLAnalysis::new(SQLBackend::MySQL).render(
546 &model,
547 |id| format!("{id}").as_str().into(),
548 |id| format!("{id}").as_str().into(),
549 );
550 let e = ddl.unwrap_err();
551 assert_eq!(
552 e,
553 SQLAnalysisError::CyclicForeignKeyError {
554 backend: SQLBackend::MySQL,
555 cycles: vec![
556 (
557 name("Snapshots"),
558 ColumnType::Deferrable { mor: name("for_ref"), tgt: name("Refs") }
559 ),
560 (
561 name("Refs"),
562 ColumnType::Deferrable {
563 mor: name("head"),
564 tgt: name("Snapshots")
565 }
566 )
567 ]
568 }
569 );
570 }
571}