use crate::config::ConnectionInfo;
use crate::errors::ErdifyError;
use crate::schema::{
CheckConstraint, Column, ForeignKey, IndexInfo, SYSTEM_SCHEMAS, Table, TableKind,
UniqueConstraint,
};
use std::collections::{HashMap, HashSet};
use tokio::time::{Duration, timeout};
use tokio_postgres::types::{Oid, ToSql};
use tokio_postgres::{Client, NoTls};
const CONNECT_TIMEOUT: Duration = Duration::from_secs(10);
struct ConstraintRow {
table_oid: Oid,
name: String,
kind: String,
columns: Vec<String>,
ref_schema: Option<String>,
ref_table: Option<String>,
ref_columns: Vec<String>,
definition: String,
}
struct IndexRow {
table_oid: Oid,
name: String,
columns: Vec<String>,
is_unique: bool,
}
struct ColumnRow {
table_oid: Oid,
name: String,
data_type: String,
not_null: bool,
default_expr: Option<String>,
}
struct TableRow {
oid: Oid,
schema: String,
name: String,
relkind: String,
}
fn table_kind(relkind: &str) -> TableKind {
match relkind {
"v" => TableKind::View,
"m" => TableKind::MaterializedView,
_ => TableKind::Table,
}
}
pub async fn connect(info: &ConnectionInfo) -> Result<Client, ErdifyError> {
let mut config = tokio_postgres::Config::new();
config
.host(&info.host)
.port(info.port)
.dbname(&info.database)
.connect_timeout(CONNECT_TIMEOUT);
if !info.user.is_empty() {
config.user(&info.user);
}
if !info.password.is_empty() {
config.password(&info.password);
}
let fut = config.connect(NoTls);
match timeout(CONNECT_TIMEOUT, fut).await {
Ok(Ok((client, connection))) => {
tokio::spawn(async move {
if let Err(e) = connection.await {
eprintln!("warning: connection interrupted: {e}");
}
});
ping(&client).await?;
Ok(client)
}
Ok(Err(e)) => Err(ErdifyError::DatabaseConnection(e.to_string())),
Err(_) => Err(ErdifyError::ConnectionTimeout),
}
}
pub async fn ping(client: &Client) -> Result<(), ErdifyError> {
client
.query_one("SELECT 1", &[])
.await
.map_err(|e| ErdifyError::DatabaseConnection(e.to_string()))?;
Ok(())
}
pub async fn fetch_tables(
client: &Client,
schemas: &[&str],
tables_filter: &[&str],
ignore_tables: &[&str],
) -> Result<Vec<Table>, ErdifyError> {
let table_rows = fetch_table_list(client, schemas).await?;
warn_missing_schemas(schemas, &table_rows);
if table_rows.is_empty() {
return Ok(Vec::new());
}
let oids: Vec<Oid> = table_rows.iter().map(|t| t.oid).collect();
let columns = fetch_columns(client, &oids).await?;
let constraints = fetch_constraints(client, &oids).await?;
let indexes = fetch_indexes(client, &oids).await?;
let tables = assemble_tables(table_rows, columns, constraints, indexes);
let tables = crate::schema::filter_tables(tables, schemas, tables_filter, ignore_tables);
warn_missing_tables(tables_filter, &tables);
Ok(tables)
}
async fn fetch_table_list(client: &Client, schemas: &[&str]) -> Result<Vec<TableRow>, ErdifyError> {
let query = "\
SELECT c.oid, n.nspname AS schema_name, c.relname AS table_name, \
c.relkind::text AS relkind \
FROM pg_class c \
JOIN pg_namespace n ON n.oid = c.relnamespace \
WHERE c.relkind IN ('r', 'p', 'v', 'm') \
AND NOT c.relispartition \
AND n.nspname <> ALL($2) \
AND ($1::text[] IS NULL OR n.nspname = ANY($1)) \
ORDER BY n.nspname, c.relname";
let schemas_param: Option<Vec<&str>> = if schemas.is_empty() {
None
} else {
Some(schemas.to_vec())
};
let system: Vec<&str> = SYSTEM_SCHEMAS.to_vec();
let rows = client
.query(query, &[&schemas_param, &system])
.await
.map_err(|e| ErdifyError::QueryError(e.to_string()))?;
Ok(rows
.into_iter()
.map(|row| TableRow {
oid: row.get("oid"),
schema: row.get("schema_name"),
name: row.get("table_name"),
relkind: row.get("relkind"),
})
.collect())
}
async fn fetch_columns(client: &Client, oids: &[Oid]) -> Result<Vec<ColumnRow>, ErdifyError> {
let query = "\
SELECT a.attrelid AS table_oid, \
a.attname AS column_name, \
format_type(a.atttypid, a.atttypmod) AS data_type, \
a.attnotnull AS not_null, \
pg_get_expr(ad.adbin, ad.adrelid) AS default_expr \
FROM pg_attribute a \
LEFT JOIN pg_attrdef ad ON ad.adrelid = a.attrelid AND ad.adnum = a.attnum \
WHERE a.attrelid = ANY($1) \
AND a.attnum > 0 \
AND NOT a.attisdropped \
ORDER BY a.attrelid, a.attnum";
let rows = client
.query(query, &[&oids])
.await
.map_err(|e| ErdifyError::QueryError(e.to_string()))?;
Ok(rows
.into_iter()
.map(|row| ColumnRow {
table_oid: row.get("table_oid"),
name: row.get("column_name"),
data_type: row.get("data_type"),
not_null: row.get("not_null"),
default_expr: row.get("default_expr"),
})
.collect())
}
async fn fetch_constraints(
client: &Client,
oids: &[Oid],
) -> Result<Vec<ConstraintRow>, ErdifyError> {
let query = "\
SELECT co.conrelid AS table_oid, \
co.conname AS constraint_name, \
co.contype::text AS constraint_type, \
ARRAY( \
SELECT a.attname \
FROM unnest(co.conkey) WITH ORDINALITY AS k(attnum, ord) \
JOIN pg_attribute a ON a.attrelid = co.conrelid AND a.attnum = k.attnum \
ORDER BY k.ord \
) AS columns, \
rn.nspname AS ref_schema, \
rc.relname AS ref_table, \
ARRAY( \
SELECT a.attname \
FROM unnest(co.confkey) WITH ORDINALITY AS k(attnum, ord) \
JOIN pg_attribute a ON a.attrelid = co.confrelid AND a.attnum = k.attnum \
ORDER BY k.ord \
) AS ref_columns, \
pg_get_constraintdef(co.oid) AS definition \
FROM pg_constraint co \
LEFT JOIN pg_class rc ON rc.oid = co.confrelid \
LEFT JOIN pg_namespace rn ON rn.oid = rc.relnamespace \
WHERE co.conrelid = ANY($1) \
AND co.contype IN ('p', 'f', 'u', 'c') \
ORDER BY co.conrelid, co.conname";
let rows = client
.query(query, &[&oids])
.await
.map_err(|e| ErdifyError::QueryError(e.to_string()))?;
Ok(rows
.into_iter()
.map(|row| ConstraintRow {
table_oid: row.get("table_oid"),
name: row.get("constraint_name"),
kind: row.get("constraint_type"),
columns: row.get("columns"),
ref_schema: row.get("ref_schema"),
ref_table: row.get("ref_table"),
ref_columns: row.get("ref_columns"),
definition: row.get("definition"),
})
.collect())
}
async fn fetch_indexes(client: &Client, oids: &[Oid]) -> Result<Vec<IndexRow>, ErdifyError> {
let query = "\
SELECT ix.indrelid AS table_oid, \
i.relname AS index_name, \
ix.indisunique AS is_unique, \
ARRAY( \
SELECT a.attname \
FROM unnest(ix.indkey::smallint[]) WITH ORDINALITY AS k(attnum, ord) \
JOIN pg_attribute a ON a.attrelid = ix.indrelid AND a.attnum = k.attnum \
ORDER BY k.ord \
) AS columns \
FROM pg_index ix \
JOIN pg_class i ON i.oid = ix.indexrelid \
WHERE ix.indrelid = ANY($1) \
AND NOT ix.indisprimary \
ORDER BY ix.indrelid, i.relname";
let rows = client
.query(query, &[&oids])
.await
.map_err(|e| ErdifyError::QueryError(e.to_string()))?;
Ok(rows
.into_iter()
.map(|row| IndexRow {
table_oid: row.get("table_oid"),
name: row.get("index_name"),
columns: row.get("columns"),
is_unique: row.get("is_unique"),
})
.collect())
}
fn assemble_tables(
table_rows: Vec<TableRow>,
columns: Vec<ColumnRow>,
constraints: Vec<ConstraintRow>,
indexes: Vec<IndexRow>,
) -> Vec<Table> {
let mut tables: Vec<Table> = Vec::with_capacity(table_rows.len());
let mut position: HashMap<Oid, usize> = HashMap::with_capacity(table_rows.len());
for row in table_rows {
position.insert(row.oid, tables.len());
tables.push(Table {
schema: row.schema,
name: row.name,
kind: table_kind(&row.relkind),
..Table::default()
});
}
for col in columns {
let Some(&idx) = position.get(&col.table_oid) else {
continue;
};
let table = &mut tables[idx];
if col.not_null {
table.not_null_cols.insert(col.name.clone());
}
table.columns.push(Column {
name: col.name,
data_type: col.data_type,
default: col.default_expr,
});
}
for c in constraints {
let Some(&idx) = position.get(&c.table_oid) else {
continue;
};
let table = &mut tables[idx];
match c.kind.as_str() {
"p" => table.primary_keys = c.columns,
"f" => {
if let (Some(to_schema), Some(to_table)) = (c.ref_schema, c.ref_table) {
table.foreign_keys.push(ForeignKey {
name: c.name,
from_columns: c.columns,
to_schema,
to_table,
to_columns: c.ref_columns,
});
}
}
"u" => table.unique_constraints.push(UniqueConstraint {
name: c.name,
columns: c.columns,
}),
"c" => table.check_constraints.push(CheckConstraint {
name: c.name,
definition: c.definition,
}),
_ => {}
}
}
for idx_row in indexes {
let Some(&idx) = position.get(&idx_row.table_oid) else {
continue;
};
tables[idx].indexes.push(IndexInfo {
name: idx_row.name,
columns: idx_row.columns,
is_unique: idx_row.is_unique,
});
}
tables
}
fn warn_missing_schemas(requested: &[&str], found: &[TableRow]) {
if requested.is_empty() {
return;
}
let present: HashSet<&str> = found.iter().map(|t| t.schema.as_str()).collect();
for schema in requested {
if !present.contains(schema) {
eprintln!("warning: no table found in schema \"{schema}\"");
}
}
}
fn warn_missing_tables(requested: &[&str], found: &[Table]) {
if requested.is_empty() {
return;
}
let present: HashSet<&str> = found.iter().map(|t| t.name.as_str()).collect();
let missing: Vec<&&str> = requested
.iter()
.filter(|t| !present.contains(**t))
.collect();
if !missing.is_empty() {
let list = missing
.iter()
.map(|t| format!("\"{t}\""))
.collect::<Vec<_>>()
.join(", ");
eprintln!("warning: table(s) not found: {list}");
}
}
const _: fn() = || {
fn assert_to_sql<T: ToSql + Sync>() {}
assert_to_sql::<Vec<Oid>>();
assert_to_sql::<Option<Vec<&str>>>();
};
#[cfg(test)]
mod tests {
use super::*;
fn table_row(oid: Oid, schema: &str, name: &str) -> TableRow {
table_row_with_kind(oid, schema, name, "r")
}
fn table_row_with_kind(oid: Oid, schema: &str, name: &str, relkind: &str) -> TableRow {
TableRow {
oid,
schema: schema.to_string(),
name: name.to_string(),
relkind: relkind.to_string(),
}
}
#[test]
fn assemble_tables_preserves_catalog_order() {
let rows = vec![
table_row(1, "extended", "audit"),
table_row(2, "public", "orders"),
table_row(3, "public", "users"),
];
let tables = assemble_tables(rows, Vec::new(), Vec::new(), Vec::new());
let keys: Vec<_> = tables.iter().map(Table::key).collect();
assert_eq!(
keys,
vec![
("extended", "audit"),
("public", "orders"),
("public", "users"),
]
);
}
#[test]
fn assemble_tables_attaches_columns_and_not_null() {
let rows = vec![table_row(1, "public", "users")];
let columns = vec![
ColumnRow {
table_oid: 1,
name: "id".to_string(),
data_type: "integer".to_string(),
not_null: true,
default_expr: None,
},
ColumnRow {
table_oid: 1,
name: "bio".to_string(),
data_type: "text".to_string(),
not_null: false,
default_expr: None,
},
];
let tables = assemble_tables(rows, columns, Vec::new(), Vec::new());
assert_eq!(tables[0].columns.len(), 2);
assert_eq!(tables[0].columns[0].name, "id");
assert!(tables[0].not_null_cols.contains("id"));
assert!(!tables[0].not_null_cols.contains("bio"));
}
#[test]
fn assemble_tables_attaches_column_default() {
let rows = vec![table_row(1, "public", "users")];
let columns = vec![
ColumnRow {
table_oid: 1,
name: "created_at".to_string(),
data_type: "timestamp".to_string(),
not_null: true,
default_expr: Some("now()".to_string()),
},
ColumnRow {
table_oid: 1,
name: "bio".to_string(),
data_type: "text".to_string(),
not_null: false,
default_expr: None,
},
];
let tables = assemble_tables(rows, columns, Vec::new(), Vec::new());
assert_eq!(tables[0].columns[0].default, Some("now()".to_string()));
assert_eq!(tables[0].columns[1].default, None);
}
#[test]
fn assemble_tables_assigns_kind_from_relkind() {
let rows = vec![
table_row_with_kind(1, "public", "users", "r"),
table_row_with_kind(2, "public", "orders_p1", "p"),
table_row_with_kind(3, "public", "orders_summary", "v"),
table_row_with_kind(4, "public", "orders_summary_mat", "m"),
];
let tables = assemble_tables(rows, Vec::new(), Vec::new(), Vec::new());
assert_eq!(tables[0].kind, TableKind::Table);
assert_eq!(tables[1].kind, TableKind::Table);
assert_eq!(tables[2].kind, TableKind::View);
assert_eq!(tables[3].kind, TableKind::MaterializedView);
}
#[test]
fn assemble_tables_dispatches_constraints_by_type() {
let rows = vec![table_row(1, "public", "orders")];
let constraints = vec![
ConstraintRow {
table_oid: 1,
name: "orders_pkey".to_string(),
kind: "p".to_string(),
columns: vec!["id".to_string()],
ref_schema: None,
ref_table: None,
ref_columns: Vec::new(),
definition: "PRIMARY KEY (id)".to_string(),
},
ConstraintRow {
table_oid: 1,
name: "orders_user_fkey".to_string(),
kind: "f".to_string(),
columns: vec!["user_id".to_string()],
ref_schema: Some("public".to_string()),
ref_table: Some("users".to_string()),
ref_columns: vec!["id".to_string()],
definition: "FOREIGN KEY (user_id) REFERENCES users(id)".to_string(),
},
ConstraintRow {
table_oid: 1,
name: "orders_ref_key".to_string(),
kind: "u".to_string(),
columns: vec!["reference".to_string()],
ref_schema: None,
ref_table: None,
ref_columns: Vec::new(),
definition: "UNIQUE (reference)".to_string(),
},
ConstraintRow {
table_oid: 1,
name: "orders_total_check".to_string(),
kind: "c".to_string(),
columns: vec!["total".to_string()],
ref_schema: None,
ref_table: None,
ref_columns: Vec::new(),
definition: "CHECK ((total > 0))".to_string(),
},
];
let tables = assemble_tables(rows, Vec::new(), constraints, Vec::new());
let table = &tables[0];
assert_eq!(table.primary_keys, vec!["id".to_string()]);
assert_eq!(table.foreign_keys.len(), 1);
assert_eq!(table.foreign_keys[0].to_schema, "public");
assert_eq!(table.foreign_keys[0].to_table, "users");
assert_eq!(table.unique_constraints[0].name, "orders_ref_key");
assert_eq!(table.check_constraints[0].definition, "CHECK ((total > 0))");
}
#[test]
fn assemble_tables_drops_foreign_key_without_target() {
let rows = vec![table_row(1, "public", "orders")];
let constraints = vec![ConstraintRow {
table_oid: 1,
name: "dangling".to_string(),
kind: "f".to_string(),
columns: vec!["user_id".to_string()],
ref_schema: None,
ref_table: None,
ref_columns: Vec::new(),
definition: String::new(),
}];
let tables = assemble_tables(rows, Vec::new(), constraints, Vec::new());
assert!(tables[0].foreign_keys.is_empty());
}
#[test]
fn assemble_tables_ignores_rows_of_unknown_tables() {
let rows = vec![table_row(1, "public", "users")];
let columns = vec![ColumnRow {
table_oid: 999,
name: "ghost".to_string(),
data_type: "text".to_string(),
not_null: false,
default_expr: None,
}];
let tables = assemble_tables(rows, columns, Vec::new(), Vec::new());
assert!(tables[0].columns.is_empty());
}
#[test]
fn assemble_tables_attaches_indexes() {
let rows = vec![table_row(1, "public", "users")];
let indexes = vec![IndexRow {
table_oid: 1,
name: "idx_users_email".to_string(),
columns: vec!["email".to_string()],
is_unique: true,
}];
let tables = assemble_tables(rows, Vec::new(), Vec::new(), indexes);
assert_eq!(tables[0].indexes.len(), 1);
assert_eq!(tables[0].indexes[0].name, "idx_users_email");
assert!(tables[0].indexes[0].is_unique);
}
#[test]
fn assemble_tables_ignores_indexes_of_unknown_tables() {
let rows = vec![table_row(1, "public", "users")];
let indexes = vec![IndexRow {
table_oid: 999,
name: "ghost_idx".to_string(),
columns: Vec::new(),
is_unique: false,
}];
let tables = assemble_tables(rows, Vec::new(), Vec::new(), indexes);
assert!(tables[0].indexes.is_empty());
}
#[test]
fn warn_missing_schemas_does_nothing_when_no_schema_requested() {
warn_missing_schemas(&[], &[]);
}
#[test]
fn warn_missing_schemas_does_nothing_when_all_present() {
let found = vec![table_row(1, "public", "users")];
warn_missing_schemas(&["public"], &found);
}
#[test]
fn warn_missing_schemas_warns_on_absent_schema() {
let found = vec![table_row(1, "public", "users")];
warn_missing_schemas(&["public", "extended"], &found);
}
#[test]
fn warn_missing_tables_does_nothing_when_no_table_requested() {
warn_missing_tables(&[], &[]);
}
#[test]
fn warn_missing_tables_does_nothing_when_all_present() {
let table = Table {
schema: "public".to_string(),
name: "users".to_string(),
..Table::default()
};
warn_missing_tables(&["users"], &[table]);
}
#[test]
fn warn_missing_tables_warns_on_absent_table() {
let table = Table {
schema: "public".to_string(),
name: "users".to_string(),
..Table::default()
};
warn_missing_tables(&["users", "ghost"], &[table]);
}
}