use std::collections::BTreeSet;
use anyhow::Result;
use pgevolve_core::ir::catalog::Catalog;
use pgevolve_core::parse::normalize_body::NormalizedBody;
use pgevolve_core::plan::edges::{DepEdge, DepSource, NodeId, build_create_graph};
use pgevolve_core::render::render_catalog;
use crate::shadow::ShadowBackend;
#[derive(Debug, Default)]
pub struct CrossCheckReport {
pub structural_edges_checked: usize,
pub canonical_mismatches: Vec<CanonicalMismatch>,
pub missing_ast_edges: Vec<MissingEdge>,
pub extra_ast_edges: Vec<ExtraEdge>,
}
#[derive(Debug)]
pub struct CanonicalMismatch {
pub view_qname: String,
pub source_canonical: String,
pub catalog_canonical: String,
}
#[derive(Debug)]
pub struct MissingEdge {
pub view_qname: String,
pub ref_schema: String,
pub ref_name: String,
}
#[derive(Debug)]
pub struct ExtraEdge {
pub view_qname: String,
pub dep_node: String,
}
pub async fn cross_check(
backend: &dyn ShadowBackend,
source: &Catalog,
pg_major: u32,
strict: bool,
) -> Result<CrossCheckReport> {
let mut report = CrossCheckReport::default();
let graph = build_create_graph(source);
for edge in graph.dep_edges() {
if matches!(edge.source, DepSource::Structural) {
report.structural_edges_checked += 1;
}
}
if source.views.is_empty() && source.materialized_views.is_empty() {
return Ok(report);
}
let guard = backend.checkout(pg_major).await?;
apply_source_to_shadow(guard.url(), source).await?;
let (client, conn) = tokio_postgres::connect(guard.url(), tokio_postgres::NoTls).await?;
tokio::spawn(conn);
check_views(&client, source, &mut report).await?;
check_materialized_views(&client, source, &mut report).await?;
if strict
&& (!report.canonical_mismatches.is_empty()
|| !report.extra_ast_edges.is_empty()
|| !report.missing_ast_edges.is_empty())
{
let n_canon = report.canonical_mismatches.len();
let n_extra = report.extra_ast_edges.len();
let n_missing = report.missing_ast_edges.len();
anyhow::bail!(
"shadow-strict: {n_canon} canonical mismatch(es), \
{n_extra} extra AST edge(s), \
{n_missing} missing AST edge(s)"
);
}
Ok(report)
}
async fn check_views(
client: &tokio_postgres::Client,
source: &Catalog,
report: &mut CrossCheckReport,
) -> Result<()> {
for view in &source.views {
let qname = view.qname.to_string();
let qname_sql = view.qname.render_sql();
check_body_canonical(client, &qname, &qname_sql, &view.body_canonical, report).await?;
check_dep_edges(client, &qname, &qname_sql, &view.body_dependencies, report).await?;
}
Ok(())
}
async fn check_materialized_views(
client: &tokio_postgres::Client,
source: &Catalog,
report: &mut CrossCheckReport,
) -> Result<()> {
for mv in &source.materialized_views {
let qname = mv.qname.to_string();
let qname_sql = mv.qname.render_sql();
check_body_canonical(client, &qname, &qname_sql, &mv.body_canonical, report).await?;
check_dep_edges(client, &qname, &qname_sql, &mv.body_dependencies, report).await?;
}
Ok(())
}
async fn apply_source_to_shadow(url: &str, source: &Catalog) -> Result<()> {
let sql = render_catalog(source);
if sql.trim().is_empty() {
return Ok(());
}
let (client, conn) = tokio_postgres::connect(url, tokio_postgres::NoTls).await?;
tokio::spawn(conn);
client
.batch_execute(&sql)
.await
.map_err(|e| anyhow::anyhow!("apply_source_to_shadow failed: {e}\nSQL:\n{sql}"))?;
Ok(())
}
async fn check_body_canonical(
client: &tokio_postgres::Client,
qname: &str,
qname_sql: &str,
source_body: &NormalizedBody,
report: &mut CrossCheckReport,
) -> Result<()> {
let pg_body: String = client
.query_one(
&format!("SELECT pg_get_viewdef('{qname_sql}'::regclass, true)"),
&[],
)
.await
.map(|row| row.get::<_, String>(0))
.map_err(|e| anyhow::anyhow!("pg_get_viewdef failed for {qname}: {e}"))?;
match NormalizedBody::from_sql(&pg_body) {
Ok(catalog_canonical) => {
if catalog_canonical.canonical_text() != source_body.canonical_text() {
report.canonical_mismatches.push(CanonicalMismatch {
view_qname: qname.to_string(),
source_canonical: source_body.canonical_text().to_string(),
catalog_canonical: catalog_canonical.canonical_text().to_string(),
});
}
}
Err(e) => {
report.canonical_mismatches.push(CanonicalMismatch {
view_qname: qname.to_string(),
source_canonical: source_body.canonical_text().to_string(),
catalog_canonical: format!("<parse error: {e}>"),
});
}
}
Ok(())
}
async fn check_dep_edges(
client: &tokio_postgres::Client,
qname: &str,
qname_sql: &str,
body_dependencies: &[DepEdge],
report: &mut CrossCheckReport,
) -> Result<()> {
let ast_extracted: BTreeSet<String> = body_dependencies
.iter()
.filter(|e| e.source == DepSource::AstExtracted)
.map(|e| node_id_to_key(&e.to))
.collect();
if ast_extracted.is_empty() {
return Ok(());
}
let rows = client
.query(
&format!(
"WITH view_info AS ( \
SELECT c.oid AS view_oid \
FROM pg_class c \
JOIN pg_namespace n ON n.oid = c.relnamespace \
WHERE n.nspname || '.' || c.relname = '{qname_sql}' \
) \
SELECT DISTINCT ref_n.nspname AS ref_schema, ref_c.relname AS ref_name \
FROM view_info vi \
JOIN pg_rewrite rw ON rw.ev_class = vi.view_oid \
JOIN pg_depend d ON d.objid = rw.oid \
AND d.classid = 'pg_rewrite'::regclass \
AND d.refclassid = 'pg_class'::regclass \
JOIN pg_class ref_c ON ref_c.oid = d.refobjid \
JOIN pg_namespace ref_n ON ref_n.oid = ref_c.relnamespace \
WHERE d.refobjid <> vi.view_oid"
),
&[],
)
.await
.map_err(|e| anyhow::anyhow!("pg_depend/pg_rewrite query failed for {qname}: {e}"))?;
let pg_deps: BTreeSet<String> = rows
.iter()
.map(|row| {
let schema: String = row.get("ref_schema");
let name: String = row.get("ref_name");
format!("{schema}.{name}")
})
.collect();
for key in &ast_extracted {
if !pg_deps.contains(key) {
report.extra_ast_edges.push(ExtraEdge {
view_qname: qname.to_string(),
dep_node: key.clone(),
});
}
}
for key in &pg_deps {
if !ast_extracted.contains(key) {
report.missing_ast_edges.push(MissingEdge {
view_qname: qname.to_string(),
ref_schema: key.split('.').next().unwrap_or("").to_string(),
ref_name: key.split('.').nth(1).unwrap_or("").to_string(),
});
}
}
Ok(())
}
fn node_id_to_key(node: &NodeId) -> String {
match node {
NodeId::Table(q)
| NodeId::View(q)
| NodeId::Mv(q)
| NodeId::Index(q)
| NodeId::Sequence(q)
| NodeId::Type(q)
| NodeId::Trigger(q)
| NodeId::Procedure(q)
| NodeId::Statistic(q)
| NodeId::Collation(q)
| NodeId::Function(q, _) => q.to_string(),
NodeId::Schema(id)
| NodeId::Extension(id)
| NodeId::Publication(id)
| NodeId::Subscription(id) => id.to_string(),
NodeId::Constraint { table, name } => format!("{table}.{name}"),
}
}