#![allow(clippy::cast_ptr_alignment)]
use super::{
Column, Conjunct, Graph, IdentityKind, Maps, Occurrence, OutputColumn, Root, Sql,
WalkedIdentity,
};
use crate::error::{TViewError, TViewResult};
use pgrx::pg_sys::{self, Oid};
use pgrx::prelude::*;
use std::collections::{HashMap, HashSet};
use std::ffi::CStr;
pub const MAX_DEPTH: usize = 32;
pub struct Context<'a> {
pub tview_tables: &'a HashSet<Oid>,
pub tview_views: &'a HashMap<Oid, String>,
pub entity: &'a str,
pub key_column: &'a str,
}
pub fn analyze(view_oid: Oid, ctx: &Context<'_>) -> TViewResult<Graph> {
let mut walker = Walker {
ctx,
graph: Graph::default(),
levels: Vec::new(),
catalog: CatalogNames::default(),
cte_parent: None,
read_ctes: HashSet::new(),
wanted: None,
identity_level: false,
};
let query = unsafe { view_query(view_oid)? };
let flags = Flags::default();
unsafe { walker.top(query, &flags)? };
walker.note_virtual_columns();
Ok(walker.graph)
}
unsafe fn view_query(view_oid: Oid) -> TViewResult<*mut pg_sys::Query> {
unsafe {
let rel = pg_sys::try_relation_open(view_oid, pg_sys::AccessShareLock.cast_signed());
if rel.is_null() {
return Err(TViewError::CatalogError {
operation: format!("Open view {view_oid:?}"),
pg_error: "relation does not exist".to_string(),
});
}
let query = pg_sys::get_view_query(rel);
let copy = pg_sys::copyObjectImpl(query.cast()).cast::<pg_sys::Query>();
pg_sys::relation_close(rel, pg_sys::NoLock.cast_signed());
Ok(copy)
}
}
#[derive(Debug, Clone, Default)]
struct Flags {
branch: usize,
via_view: Option<String>,
via_tview: Option<String>,
in_sublink: bool,
opaque_level: Option<String>,
unread: bool,
}
#[derive(Debug, Clone)]
enum Resolved {
Col(Column),
Alt(Vec<Column>),
Opaque,
}
#[derive(Debug, Clone)]
enum RteInfo {
Base(usize),
Outputs(Vec<Resolved>),
Join(*mut pg_sys::List),
Other,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Link {
Top,
From,
Sublink {
required: bool,
},
}
struct Level {
query: *mut pg_sys::Query,
rtes: Vec<RteInfo>,
link: Link,
cte_parent: Option<usize>,
}
#[derive(Clone, Copy)]
enum Origin<'a> {
Required,
Outer { nullable: &'a HashSet<usize> },
None,
}
struct Site {
levelsup: usize,
candidates: Vec<Column>,
}
#[derive(Default)]
struct CatalogNames {
operators: HashMap<u32, String>,
functions: HashMap<u32, (String, bool)>,
types: HashMap<u32, String>,
}
struct Walker<'c> {
ctx: &'c Context<'c>,
graph: Graph,
levels: Vec<Level>,
catalog: CatalogNames,
cte_parent: Option<usize>,
read_ctes: HashSet<(usize, String)>,
wanted: Option<HashSet<i16>>,
identity_level: bool,
}
fn cstr(ptr: *const std::ffi::c_char) -> String {
if ptr.is_null() {
return String::new();
}
unsafe { CStr::from_ptr(ptr) }
.to_string_lossy()
.into_owned()
}
unsafe fn elements<T>(list: *mut pg_sys::List) -> Vec<*mut T> {
if list.is_null() {
return Vec::new();
}
unsafe {
(0..(*list).length)
.map(|i| pg_sys::list_nth(list, i).cast::<T>())
.collect()
}
}
unsafe fn tag(node: *const pg_sys::Node) -> Option<pg_sys::NodeTag> {
unsafe { (!node.is_null()).then(|| (*node).type_) }
}
unsafe fn conjuncts(node: *mut pg_sys::Node) -> Vec<*mut pg_sys::Node> {
unsafe {
match tag(node) {
None => Vec::new(),
Some(pg_sys::NodeTag::T_List) => elements::<pg_sys::Node>(node.cast())
.into_iter()
.flat_map(|n| conjuncts(n))
.collect(),
Some(pg_sys::NodeTag::T_BoolExpr)
if (*node.cast::<pg_sys::BoolExpr>()).boolop == pg_sys::BoolExprType::AND_EXPR =>
{
elements::<pg_sys::Node>((*node.cast::<pg_sys::BoolExpr>()).args)
.into_iter()
.flat_map(|n| conjuncts(n))
.collect()
}
Some(_) => vec![node],
}
}
}
impl Walker<'_> {
fn note_virtual_columns(&mut self) {
let graph = &self.graph;
let columns: Vec<&Column> = graph
.roots
.iter()
.map(|r| &r.key)
.chain(
graph
.conjuncts
.iter()
.filter_map(|c| c.equality.as_ref())
.flat_map(|(x, y)| [x, y]),
)
.collect();
let found: Vec<(usize, i16)> = columns
.into_iter()
.filter(|c| {
let relid = Oid::from(graph.occurrences[c.occ].relid);
unsafe { pg_sys::get_attgenerated(relid, c.attnum) as u8 == b'v' }
})
.map(|c| (c.occ, c.attnum))
.collect();
self.graph.virtual_columns.extend(found);
}
unsafe fn top(&mut self, query: *mut pg_sys::Query, flags: &Flags) -> TViewResult<()> {
unsafe {
let key_position = elements::<pg_sys::TargetEntry>((*query).targetList)
.iter()
.position(|tle| cstr((**tle).resname) == self.ctx.key_column);
if (*query).setOperations.is_null() {
let opaque = top_opaque_reason(query);
let flags = Flags {
opaque_level: opaque.clone().or_else(|| flags.opaque_level.clone()),
..flags.clone()
};
self.identity_level = true;
let outputs = self.level(query, &flags, Link::Top)?;
let root_position = match &self.graph.identity {
Some(Ok(identity)) => Some(identity.position),
_ => key_position,
};
if opaque.is_none()
&& let Some(Resolved::Col(key)) = root_position.and_then(|p| outputs.get(p))
{
self.graph.roots.push(Root {
branch: flags.branch,
key: key.clone(),
});
}
return Ok(());
}
self.graph.set_operation = true;
let whole = top_opaque_reason(query);
self.push_level(
query,
vec![RteInfo::Other; list_len((*query).rtable)],
Link::Top,
);
let mut leaves = Vec::new();
setop_leaves((*query).setOperations, &mut leaves);
let result: TViewResult<()> = (|| {
for (branch, rtindex) in leaves.into_iter().enumerate() {
let Some(&rte) = elements::<pg_sys::RangeTblEntry>((*query).rtable)
.get(rtindex.wrapping_sub(1))
else {
continue;
};
let opaque = whole.clone().or_else(|| top_opaque_reason((*rte).subquery));
let leaf_flags = Flags {
branch,
opaque_level: opaque.clone().or_else(|| flags.opaque_level.clone()),
..flags.clone()
};
let outputs = self.level((*rte).subquery, &leaf_flags, Link::Top)?;
if opaque.is_none()
&& let Some(Resolved::Col(key)) = key_position.and_then(|p| outputs.get(p))
{
self.graph.roots.push(Root {
branch,
key: key.clone(),
});
}
}
self.unread_ctes(flags)
})();
self.levels.pop();
let identity = key_position
.and_then(|p| {
elements::<pg_sys::TargetEntry>((*query).targetList)
.get(p)
.copied()
})
.zip(key_position)
.map(|(tle, position)| WalkedIdentity {
name: self.ctx.key_column.to_string(),
position,
type_oid: pg_sys::exprType((*tle).expr.cast()).to_u32(),
kind: IdentityKind::Pk,
columns: self.graph.roots.iter().map(|r| r.key.clone()).collect(),
})
.ok_or(super::IdentityError::Missing);
self.graph.identity = Some(identity);
result
}
}
unsafe fn level(
&mut self,
query: *mut pg_sys::Query,
flags: &Flags,
link: Link,
) -> TViewResult<Vec<Resolved>> {
let wanted = self.wanted.take();
if self.levels.len() >= MAX_DEPTH {
return Err(TViewError::InvalidInput {
parameter: "tview definition".to_string(),
reason: format!(
"the backing view nests views, subqueries and CTEs more than {MAX_DEPTH} levels deep"
),
});
}
unsafe {
if !(*query).setOperations.is_null() {
return self.union_outputs(query, flags);
}
let opaque = (link != Link::Top).then(|| opaque_reason(query)).flatten();
let flags = Flags {
opaque_level: opaque.clone().or_else(|| flags.opaque_level.clone()),
..flags.clone()
};
self.push_level(query, Vec::new(), link);
let result = self
.level_body(query, &flags, link, opaque.is_some(), wanted.as_ref())
.and_then(|outputs| self.unread_ctes(&flags).map(|()| outputs));
self.levels.pop();
result
}
}
unsafe fn level_body(
&mut self,
query: *mut pg_sys::Query,
flags: &Flags,
link: Link,
opaque: bool,
wanted: Option<&HashSet<i16>>,
) -> TViewResult<Vec<Resolved>> {
let identity_level = std::mem::take(&mut self.identity_level);
unsafe {
let read = referenced_columns(query);
for (i, rte) in elements::<pg_sys::RangeTblEntry>((*query).rtable)
.into_iter()
.enumerate()
{
self.wanted = match read.get(&(i + 1)) {
Some(Some(columns)) => Some(columns.clone()),
Some(None) => None,
None => Some(HashSet::new()),
};
let info = self.rte(rte, flags);
self.wanted = None;
self.current().rtes.push(info?);
}
let jointree = (*query).jointree;
if !jointree.is_null() {
self.join_item(jointree.cast(), flags)?;
}
let unread = Flags {
unread: true,
..flags.clone()
};
let skipped: Vec<bool> = elements::<pg_sys::TargetEntry>((*query).targetList)
.iter()
.map(|tle| {
wanted.is_some_and(|w| !w.contains(&(**tle).resno))
&& (**tle).ressortgroupref == 0
&& !pg_sys::expression_returns_set((**tle).expr.cast())
})
.collect();
for (tle, skip) in elements::<pg_sys::TargetEntry>((*query).targetList)
.into_iter()
.zip(&skipped)
{
let tle_flags = if *skip { &unread } else { flags };
self.sublinks((*tle).expr.cast(), tle_flags, false)?;
}
self.sublinks((*query).havingQual, flags, false)?;
self.note_functions(query.cast());
if identity_level {
let tles = elements::<pg_sys::TargetEntry>((*query).targetList);
self.graph.identity = Some(self.identity(query, &tles));
}
let grouped = (*query).hasAggs || !(*query).groupClause.is_null();
let tles = elements::<pg_sys::TargetEntry>((*query).targetList);
let key_columns = |clause: *mut pg_sys::List| -> Vec<Column> {
tles.iter()
.filter(|tle| in_clause((***tle).ressortgroupref, clause))
.filter_map(|tle| match self.resolve_expr((**tle).expr.cast()) {
Resolved::Col(c) => Some(c),
_ => None,
})
.collect()
};
let group_keys = key_columns((*query).groupClause);
let distinct_keys = key_columns((*query).distinctClause);
let keyed =
|tle: *mut pg_sys::TargetEntry, clause: *mut pg_sys::List, keys: &[Column]| {
in_clause((*tle).ressortgroupref, clause)
|| matches!(self.resolve_expr((*tle).expr.cast()),
Resolved::Col(c) if self.equal_to_key(&c, keys))
};
let mut outputs = Vec::new();
for (&tle, skip) in tles.iter().zip(skipped) {
let pass_through = !skip
&& (link == Link::Top
|| (!opaque
&& (!grouped || keyed(tle, (*query).groupClause, &group_keys))
&& (!(*query).hasDistinctOn
|| keyed(tle, (*query).distinctClause, &distinct_keys))));
outputs.push(if pass_through {
self.resolve_expr((*tle).expr.cast())
} else {
Resolved::Opaque
});
}
Ok(outputs)
}
}
fn equal_to_key(&self, column: &Column, keys: &[Column]) -> bool {
self.graph.conjuncts.iter().any(|c| {
c.equality.as_ref().is_some_and(|(x, y)| {
let (toward, key) = if x == column {
(if c.a == x.occ { c.a_to_b } else { c.b_to_a }, y)
} else if y == column {
(if c.a == y.occ { c.a_to_b } else { c.b_to_a }, x)
} else {
return false;
};
toward == Maps::Yes && keys.contains(key)
})
})
}
unsafe fn identity(
&self,
query: *mut pg_sys::Query,
tles: &[*mut pg_sys::TargetEntry],
) -> Result<WalkedIdentity, super::IdentityError> {
unsafe {
let outputs: Vec<OutputColumn> = tles
.iter()
.map(|&tle| OutputColumn {
name: cstr((*tle).resname),
junk: (*tle).resjunk,
sortgroupref: (*tle).ressortgroupref,
column: match self.resolve_expr((*tle).expr.cast()) {
Resolved::Col(c) => Some(c),
_ => None,
},
type_oid: pg_sys::exprType((*tle).expr.cast()).to_u32(),
})
.collect();
let distinct_on: Option<Vec<u32>> = (*query).hasDistinctOn.then(|| {
elements::<pg_sys::SortGroupClause>((*query).distinctClause)
.iter()
.map(|c| (**c).tleSortGroupRef)
.collect()
});
let selected = super::select_identity(
self.ctx.entity,
&outputs,
distinct_on.as_deref(),
&|column, key| self.equal_to_key(column, std::slice::from_ref(key)),
)?;
let chosen = &outputs[selected.position];
Ok(WalkedIdentity {
name: chosen.name.clone(),
position: selected.position,
type_oid: chosen.type_oid,
kind: selected.kind,
columns: chosen.column.iter().cloned().collect(),
})
}
}
unsafe fn union_outputs(
&mut self,
query: *mut pg_sys::Query,
flags: &Flags,
) -> TViewResult<Vec<Resolved>> {
unsafe {
self.push_level(
query,
vec![RteInfo::Other; list_len((*query).rtable)],
Link::From,
);
let mut leaves = Vec::new();
setop_leaves((*query).setOperations, &mut leaves);
let width = list_len((*query).targetList);
let mut columns: Vec<Vec<Column>> = vec![Vec::new(); width];
let result: TViewResult<()> = (|| {
for rtindex in leaves {
let Some(&rte) = elements::<pg_sys::RangeTblEntry>((*query).rtable)
.get(rtindex.wrapping_sub(1))
else {
continue;
};
let outputs = self.level((*rte).subquery, flags, Link::From)?;
for (i, out) in outputs.into_iter().take(width).enumerate() {
match out {
Resolved::Col(c) => columns[i].push(c),
Resolved::Alt(cs) => columns[i].extend(cs),
Resolved::Opaque => {}
}
}
}
self.unread_ctes(flags)
})();
self.levels.pop();
result?;
Ok(columns
.into_iter()
.map(|cs| {
if cs.is_empty() {
Resolved::Opaque
} else {
Resolved::Alt(cs)
}
})
.collect())
}
}
fn current(&mut self) -> &mut Level {
self.levels.last_mut().expect("inside a query level")
}
fn push_level(&mut self, query: *mut pg_sys::Query, rtes: Vec<RteInfo>, link: Link) {
let cte_parent = self
.cte_parent
.take()
.or_else(|| self.levels.len().checked_sub(1));
self.levels.push(Level {
query,
rtes,
link,
cte_parent,
});
}
unsafe fn unread_ctes(&mut self, flags: &Flags) -> TViewResult<()> {
let index = self.levels.len() - 1;
let query = self.levels[index].query;
let ctes = unsafe { elements::<pg_sys::CommonTableExpr>((*query).cteList) };
for cte in ctes {
let (name, body) = unsafe { (cstr((*cte).ctename), (*cte).ctequery) };
if self.read_ctes.contains(&(query as usize, name.clone())) {
continue;
}
self.read_ctes.insert((query as usize, name));
let unread = Flags {
unread: true,
..flags.clone()
};
self.cte_parent = Some(index);
unsafe {
let copy = pg_sys::copyObjectImpl(body.cast()).cast::<pg_sys::Query>();
self.level(copy, &unread, Link::From)?;
}
}
Ok(())
}
unsafe fn rte(
&mut self,
rte: *mut pg_sys::RangeTblEntry,
flags: &Flags,
) -> TViewResult<RteInfo> {
unsafe {
match (*rte).rtekind {
pg_sys::RTEKind::RTE_RELATION => self.relation((*rte).relid, flags),
pg_sys::RTEKind::RTE_SUBQUERY => Ok(RteInfo::Outputs(self.level(
(*rte).subquery,
flags,
Link::From,
)?)),
pg_sys::RTEKind::RTE_CTE => {
let name = cstr((*rte).ctename);
let Some((cte, defined_at)) = self.cte(&name, (*rte).ctelevelsup as usize)
else {
return Ok(RteInfo::Other);
};
self.read_ctes
.insert((self.levels[defined_at].query as usize, name));
let copy = pg_sys::copyObjectImpl(cte.cast()).cast::<pg_sys::Query>();
self.cte_parent = Some(defined_at);
Ok(RteInfo::Outputs(self.level(copy, flags, Link::From)?))
}
pg_sys::RTEKind::RTE_JOIN => Ok(RteInfo::Join((*rte).joinaliasvars)),
#[cfg(feature = "pg18")]
pg_sys::RTEKind::RTE_GROUP => Ok(RteInfo::Join((*rte).groupexprs)),
_ => Ok(RteInfo::Other),
}
}
}
fn relation(&mut self, relid: Oid, flags: &Flags) -> TViewResult<RteInfo> {
let (relkind, relname, qualified) = unsafe {
let relkind = pg_sys::get_rel_relkind(relid) as u8;
let relname = cstr(pg_sys::get_rel_name(relid));
let nsp = cstr(pg_sys::get_namespace_name(pg_sys::get_rel_namespace(relid)));
let qualified = format!("{}.{}", quote_ident(&nsp), quote_ident(&relname));
(relkind, relname, qualified)
};
match relkind {
b'r' | b'p' if flags.unread => {
self.graph.unread_tables.insert(relid.to_u32());
Ok(RteInfo::Other)
}
b'r' | b'p' if !self.ctx.tview_tables.contains(&relid) => {
self.graph.occurrences.push(Occurrence {
relid: relid.to_u32(),
relname,
qualified,
branch: flags.branch,
via_view: flags.via_view.clone(),
via_tview: flags.via_tview.clone(),
in_sublink: flags.in_sublink,
opaque_level: flags.opaque_level.clone(),
});
Ok(RteInfo::Base(self.graph.occurrences.len() - 1))
}
b'v' => {
let readable = unsafe {
pg_sys::pg_class_aclcheck(
relid,
pg_sys::GetUserId(),
pg_sys::AclMode::from(pg_sys::ACL_SELECT),
) == pg_sys::AclResult::ACLCHECK_OK
};
if !readable {
return Err(TViewError::InvalidInput {
parameter: "tview definition".to_string(),
reason: format!("permission denied to read view {qualified}"),
});
}
let inner = match self.ctx.tview_views.get(&relid) {
Some(entity) => Flags {
via_tview: flags.via_tview.clone().or_else(|| Some(entity.clone())),
..flags.clone()
},
None => Flags {
via_view: flags.via_view.clone().or(Some(qualified)),
..flags.clone()
},
};
let outputs = unsafe { self.level(view_query(relid)?, &inner, Link::From)? };
Ok(RteInfo::Outputs(outputs))
}
_ => Ok(RteInfo::Other),
}
}
fn cte(&self, name: &str, levelsup: usize) -> Option<(*mut pg_sys::Query, usize)> {
let mut level = self.levels.len().checked_sub(1)?;
for _ in 0..levelsup {
level = self.levels[level].cte_parent?;
}
let query = self.levels[level].query;
unsafe {
elements::<pg_sys::CommonTableExpr>((*query).cteList)
.into_iter()
.find(|cte| cstr((**cte).ctename) == *name)
.map(|cte| ((*cte).ctequery.cast(), level))
}
}
unsafe fn join_item(
&mut self,
node: *mut pg_sys::Node,
flags: &Flags,
) -> TViewResult<HashSet<usize>> {
unsafe {
match tag(node) {
Some(pg_sys::NodeTag::T_RangeTblRef) => {
let rtindex = (*node.cast::<pg_sys::RangeTblRef>()).rtindex as usize;
Ok(self.occurrences_of(rtindex))
}
Some(pg_sys::NodeTag::T_FromExpr) => {
let from = node.cast::<pg_sys::FromExpr>();
let mut under = HashSet::new();
for item in elements::<pg_sys::Node>((*from).fromlist) {
under.extend(self.join_item(item, flags)?);
}
for qual in conjuncts((*from).quals) {
self.predicate(qual, Origin::Required);
}
for qual in conjuncts((*from).quals) {
let required = is_required_sublink(qual);
self.sublinks(qual, flags, required)?;
}
Ok(under)
}
Some(pg_sys::NodeTag::T_JoinExpr) => {
let join = node.cast::<pg_sys::JoinExpr>();
let left = self.join_item((*join).larg, flags)?;
let right = self.join_item((*join).rarg, flags)?;
for qual in conjuncts((*join).quals) {
let origin = match (*join).jointype {
pg_sys::JoinType::JOIN_INNER => Origin::Required,
pg_sys::JoinType::JOIN_LEFT => Origin::Outer { nullable: &right },
pg_sys::JoinType::JOIN_RIGHT => Origin::Outer { nullable: &left },
_ => Origin::None,
};
self.predicate(qual, origin);
}
self.sublinks((*join).quals, flags, false)?;
Ok(left.union(&right).copied().collect())
}
_ => Ok(HashSet::new()),
}
}
}
fn occurrences_of(&self, rtindex: usize) -> HashSet<usize> {
let level = self.levels.last().expect("inside a query level");
match level.rtes.get(rtindex.wrapping_sub(1)) {
Some(RteInfo::Base(occ)) => HashSet::from([*occ]),
Some(RteInfo::Outputs(outputs)) => outputs
.iter()
.flat_map(|r| match r {
Resolved::Col(c) => vec![c.occ],
Resolved::Alt(cs) => cs.iter().map(|c| c.occ).collect(),
Resolved::Opaque => vec![],
})
.collect(),
_ => HashSet::new(),
}
}
unsafe fn sublinks(
&mut self,
node: *mut pg_sys::Node,
flags: &Flags,
required: bool,
) -> TViewResult<()> {
let mut found: Vec<*mut pg_sys::SubLink> = Vec::new();
unsafe {
collect_sublinks(node, &mut found);
for sublink in found {
self.sublink(sublink, flags, required)?;
}
}
Ok(())
}
unsafe fn sublink(
&mut self,
sublink: *mut pg_sys::SubLink,
flags: &Flags,
required: bool,
) -> TViewResult<()> {
unsafe {
let subselect = (*sublink).subselect.cast::<pg_sys::Query>();
if subselect.is_null() {
return Ok(());
}
let required = required
&& matches!(
(*sublink).subLinkType,
pg_sys::SubLinkType::EXISTS_SUBLINK | pg_sys::SubLinkType::ANY_SUBLINK
);
let inner = Flags {
in_sublink: true,
..flags.clone()
};
let outputs = self.level(subselect, &inner, Link::Sublink { required })?;
if (*sublink).subLinkType == pg_sys::SubLinkType::ANY_SUBLINK {
for test in conjuncts((*sublink).testexpr) {
self.test_predicate(test, &outputs, required);
}
}
Ok(())
}
}
unsafe fn predicate(&mut self, qual: *mut pg_sys::Node, origin: Origin<'_>) {
unsafe {
if matches!(origin, Origin::None)
|| pg_sys::contain_nonstrict_functions(qual)
|| pg_sys::contain_mutable_functions(qual)
|| has_sublink(qual)
{
return;
}
let mut sites = Vec::new();
if !self.sites(qual, &mut sites) {
return;
}
self.add_conjuncts(qual, &sites, &|_| None, origin, true);
}
}
unsafe fn test_predicate(
&mut self,
test: *mut pg_sys::Node,
outputs: &[Resolved],
required: bool,
) {
unsafe {
if pg_sys::contain_nonstrict_functions(test) || pg_sys::contain_mutable_functions(test)
{
return;
}
let mut sites = Vec::new();
if !self.sites(test, &mut sites) {
return;
}
let param_column = |param: *mut pg_sys::Param| -> Option<Vec<Column>> {
let p = &*param;
if p.paramkind != pg_sys::ParamKind::PARAM_SUBLINK {
return None;
}
match outputs.get(usize::try_from(p.paramid).ok()?.checked_sub(1)?)? {
Resolved::Col(c) => Some(vec![c.clone()]),
Resolved::Alt(cs) => Some(cs.clone()),
Resolved::Opaque => None,
}
};
let mut param_sites = Vec::new();
collect_params(test, &mut param_sites);
for param in ¶m_sites {
let Some(candidates) = param_column(*param) else {
return;
};
sites.push(Site {
levelsup: usize::MAX,
candidates,
});
}
let lookup = |node: *mut pg_sys::Node| -> Option<usize> {
param_sites
.iter()
.position(|p| p.cast::<pg_sys::Node>() == node)
};
self.add_conjuncts(test, &sites, &lookup, Origin::Required, required);
}
}
unsafe fn sites(&self, expr: *mut pg_sys::Node, sites: &mut Vec<Site>) -> bool {
let mut vars = Vec::new();
unsafe {
collect_vars(expr, &mut vars);
for var in vars {
let v = &*var;
let levelsup = v.varlevelsup as usize;
let candidates = match self.resolve_var(var, levelsup) {
Resolved::Col(c) => vec![c],
Resolved::Alt(cs) => cs,
Resolved::Opaque => return false,
};
sites.push(Site {
levelsup,
candidates,
});
}
}
true
}
unsafe fn add_conjuncts(
&mut self,
expr: *mut pg_sys::Node,
sites: &[Site],
param_site: &dyn Fn(*mut pg_sys::Node) -> Option<usize>,
origin: Origin<'_>,
outer_to_inner: bool,
) {
if sites.is_empty() {
return;
}
let combinations: usize = sites.iter().map(|s| s.candidates.len()).product();
if combinations == 0 || combinations > 16 {
return;
}
let var_count = sites.iter().filter(|s| s.levelsup != usize::MAX).count();
for n in 0..combinations {
let mut rest = n;
let chosen: Vec<&Column> = sites
.iter()
.map(|s| {
let c = &s.candidates[rest % s.candidates.len()];
rest /= s.candidates.len();
c
})
.collect();
let mut occs: Vec<usize> = chosen.iter().map(|c| c.occ).collect();
occs.sort_unstable();
occs.dedup();
let [a, b] = occs[..] else { continue };
let Some((a_to_b, b_to_a)) =
self.directions(sites, &chosen, a, b, origin, outer_to_inner)
else {
continue;
};
let var_index = std::cell::Cell::new(0_usize);
let column_of = |node: *mut pg_sys::Node| -> Option<Column> {
match unsafe { tag(node) } {
Some(pg_sys::NodeTag::T_Var) => {
chosen.get(var_index.get()).map(|c| (*c).clone())
}
Some(pg_sys::NodeTag::T_Param) => param_site(node)
.and_then(|i| chosen.get(var_count + i).map(|c| (*c).clone())),
_ => None,
}
};
let mut next_var = || var_index.set(var_index.get() + 1);
let Some(sql) = (unsafe { self.deparse(expr, &column_of, &mut next_var) }) else {
return;
};
let equality = unsafe { equality(expr, &chosen) };
let checked = |m: Maps| {
if m == Maps::IfMatched && equality.is_none() {
Maps::No
} else {
m
}
};
let (a_to_b, b_to_a) = (checked(a_to_b), checked(b_to_a));
if a_to_b == Maps::No && b_to_a == Maps::No {
continue;
}
self.graph.conjuncts.push(Conjunct {
sql,
a,
b,
a_to_b,
b_to_a,
equality,
});
}
}
fn directions(
&self,
sites: &[Site],
chosen: &[&Column],
a: usize,
b: usize,
origin: Origin<'_>,
outer_to_inner: bool,
) -> Option<(Maps, Maps)> {
let level_of = |occ: usize| {
sites
.iter()
.zip(chosen)
.filter(|(_, c)| c.occ == occ)
.map(|(s, _)| s.levelsup)
.min()
};
let (la, lb) = (level_of(a)?, level_of(b)?);
let depth = |l: usize| {
if l == usize::MAX {
-1_i64
} else {
i64::try_from(l).unwrap_or(i64::MAX)
}
};
let (da, db) = (depth(la), depth(lb));
let outer_ok = |outer_levelsup: i64| outer_to_inner && self.levels_required(outer_levelsup);
let (a_to_b, b_to_a) = match da.cmp(&db) {
std::cmp::Ordering::Equal => (true, true),
std::cmp::Ordering::Less => (true, outer_ok(db)),
std::cmp::Ordering::Greater => (outer_ok(da), true),
};
let maps = |holds: bool, from: usize, to: usize| match origin {
Origin::Outer { nullable } if holds && !nullable.contains(&from) => {
if nullable.contains(&to) {
Maps::IfMatched
} else {
Maps::No
}
}
_ if holds => Maps::Yes,
_ => Maps::No,
};
let (a_to_b, b_to_a) = (maps(a_to_b, a, b), maps(b_to_a, b, a));
(a_to_b != Maps::No || b_to_a != Maps::No).then_some((a_to_b, b_to_a))
}
fn levels_required(&self, outer: i64) -> bool {
let top = i64::try_from(self.levels.len()).unwrap_or(i64::MAX) - 1;
(0..outer).all(|up| {
usize::try_from(top - up).is_ok_and(|index| {
matches!(self.levels[index].link, Link::Sublink { required: true })
})
})
}
unsafe fn resolve_var(&self, var: *mut pg_sys::Var, levelsup: usize) -> Resolved {
unsafe {
let Some(index) = self.levels.len().checked_sub(1 + levelsup) else {
return Resolved::Opaque;
};
let level = &self.levels[index];
let (Ok(rtindex), attno) = (usize::try_from((*var).varno), (*var).varattno) else {
return Resolved::Opaque;
};
if attno <= 0 {
return Resolved::Opaque;
}
match level.rtes.get(rtindex.wrapping_sub(1)) {
Some(RteInfo::Base(occ)) => {
let relid = Oid::from(self.graph.occurrences[*occ].relid);
Resolved::Col(Column {
occ: *occ,
attnum: attno,
name: cstr(pg_sys::get_attname(relid, attno, true)),
})
}
Some(RteInfo::Outputs(outputs)) => outputs
.get(attno as usize - 1)
.cloned()
.unwrap_or(Resolved::Opaque),
Some(RteInfo::Join(aliases)) => {
let alias = elements::<pg_sys::Node>(*aliases)
.get(attno as usize - 1)
.copied()
.unwrap_or(std::ptr::null_mut());
self.resolve_alias(alias, index)
}
_ => Resolved::Opaque,
}
}
}
unsafe fn resolve_alias(&self, node: *mut pg_sys::Node, index: usize) -> Resolved {
unsafe {
let node = strip_relabel(node);
if tag(node) != Some(pg_sys::NodeTag::T_Var) {
return Resolved::Opaque;
}
let var = node.cast::<pg_sys::Var>();
let levelsup = self.levels.len() - 1 - index + (*var).varlevelsup as usize;
self.resolve_var(var, levelsup)
}
}
unsafe fn resolve_expr(&self, node: *mut pg_sys::Node) -> Resolved {
unsafe {
let node = strip_relabel(node);
if tag(node) == Some(pg_sys::NodeTag::T_Var) {
let var = node.cast::<pg_sys::Var>();
self.resolve_var(var, (*var).varlevelsup as usize)
} else {
Resolved::Opaque
}
}
}
unsafe fn deparse(
&mut self,
expr: *mut pg_sys::Node,
column_of: &dyn Fn(*mut pg_sys::Node) -> Option<Column>,
next_var: &mut dyn FnMut(),
) -> Option<Sql> {
unsafe {
match tag(expr)? {
pg_sys::NodeTag::T_Var => {
let column = column_of(expr)?;
next_var();
Some(column.sql())
}
pg_sys::NodeTag::T_Param => Some(column_of(expr)?.sql()),
pg_sys::NodeTag::T_Const => {
let c = expr.cast::<pg_sys::Const>();
let ty = self.type_name((*c).consttype);
if (*c).constisnull {
return Some(Sql::text(format!("NULL::{ty}")));
}
let mut output: Oid = Oid::INVALID;
let mut varlena = false;
pg_sys::getTypeOutputInfo((*c).consttype, &raw mut output, &raw mut varlena);
let text = cstr(pg_sys::OidOutputFunctionCall(output, (*c).constvalue));
Some(Sql::text(format!("{}::{ty}", quote_literal(&text)?)))
}
pg_sys::NodeTag::T_OpExpr => {
let op = expr.cast::<pg_sys::OpExpr>();
let name = self.operator_name((*op).opno);
let args = elements::<pg_sys::Node>((*op).args);
let mut sql = Sql::text("(");
match args[..] {
[arg] => {
sql.push_text(&format!("OPERATOR({name}) "));
sql.push_sql(self.deparse(arg, column_of, next_var)?);
}
[l, r] => {
sql.push_sql(self.deparse(l, column_of, next_var)?);
sql.push_text(&format!(" OPERATOR({name}) "));
sql.push_sql(self.deparse(r, column_of, next_var)?);
}
_ => return None,
}
sql.push_text(")");
Some(sql)
}
pg_sys::NodeTag::T_ScalarArrayOpExpr => {
let op = expr.cast::<pg_sys::ScalarArrayOpExpr>();
let name = self.operator_name((*op).opno);
let [l, r] = elements::<pg_sys::Node>((*op).args)[..] else {
return None;
};
let mut sql = Sql::text("(");
sql.push_sql(self.deparse(l, column_of, next_var)?);
sql.push_text(&format!(
" OPERATOR({name}) {} (",
if (*op).useOr { "ANY" } else { "ALL" }
));
sql.push_sql(self.deparse(r, column_of, next_var)?);
sql.push_text("))");
Some(sql)
}
pg_sys::NodeTag::T_FuncExpr => {
let f = expr.cast::<pg_sys::FuncExpr>();
let args = elements::<pg_sys::Node>((*f).args);
let mut sql;
if (*f).funcformat == pg_sys::CoercionForm::COERCE_EXPLICIT_CALL
|| (*f).funcformat == pg_sys::CoercionForm::COERCE_SQL_SYNTAX
{
let (name, _) = self.function((*f).funcid);
sql = Sql::text(format!("{name}("));
for (i, arg) in args.into_iter().enumerate() {
if i > 0 {
sql.push_text(", ");
}
sql.push_sql(self.deparse(arg, column_of, next_var)?);
}
sql.push_text(")");
} else {
let [arg, ..] = args[..] else { return None };
sql = Sql::text("(");
sql.push_sql(self.deparse(arg, column_of, next_var)?);
sql.push_text(&format!(")::{}", self.type_name((*f).funcresulttype)));
}
Some(sql)
}
pg_sys::NodeTag::T_RelabelType => {
let r = expr.cast::<pg_sys::RelabelType>();
let mut sql = Sql::text("(");
sql.push_sql(self.deparse((*r).arg.cast(), column_of, next_var)?);
sql.push_text(&format!(")::{}", self.type_name((*r).resulttype)));
Some(sql)
}
pg_sys::NodeTag::T_CoerceViaIO => {
let r = expr.cast::<pg_sys::CoerceViaIO>();
let mut sql = Sql::text("(");
sql.push_sql(self.deparse((*r).arg.cast(), column_of, next_var)?);
sql.push_text(&format!(")::{}", self.type_name((*r).resulttype)));
Some(sql)
}
pg_sys::NodeTag::T_BoolExpr => {
let b = expr.cast::<pg_sys::BoolExpr>();
let args = elements::<pg_sys::Node>((*b).args);
let mut sql = Sql::text("(");
match (*b).boolop {
pg_sys::BoolExprType::NOT_EXPR => {
sql.push_text("NOT ");
sql.push_sql(self.deparse(*args.first()?, column_of, next_var)?);
}
op => {
let joiner = if op == pg_sys::BoolExprType::AND_EXPR {
" AND "
} else {
" OR "
};
for (i, arg) in args.into_iter().enumerate() {
if i > 0 {
sql.push_text(joiner);
}
sql.push_sql(self.deparse(arg, column_of, next_var)?);
}
}
}
sql.push_text(")");
Some(sql)
}
_ => None,
}
}
}
fn type_name(&mut self, ty: Oid) -> String {
self.catalog
.types
.entry(ty.to_u32())
.or_insert_with(|| cstr(unsafe { pg_sys::format_type_be_qualified(ty) }))
.clone()
}
fn operator_name(&mut self, opno: Oid) -> String {
self.catalog
.operators
.entry(opno.to_u32())
.or_insert_with(|| {
Spi::get_one_with_args::<String>(
"SELECT pg_catalog.quote_ident(n.nspname) || '.' || o.oprname \
FROM pg_catalog.pg_operator o \
JOIN pg_catalog.pg_namespace n ON n.oid = o.oprnamespace \
WHERE o.oid = $1",
&[unsafe {
pgrx::datum::DatumWithOid::new(
opno,
PgOid::BuiltIn(PgBuiltInOids::OIDOID).value(),
)
}],
)
.ok()
.flatten()
.unwrap_or_default()
})
.clone()
}
fn function(&mut self, funcid: Oid) -> (String, bool) {
self.catalog
.functions
.entry(funcid.to_u32())
.or_insert_with(|| {
unsafe {
let name = cstr(pg_sys::get_func_name(funcid));
let nsp = cstr(pg_sys::get_namespace_name(pg_sys::get_func_namespace(
funcid,
)));
let immutable = pg_sys::func_volatile(funcid)
== pg_sys::PROVOLATILE_IMMUTABLE.cast_signed();
(
format!("{}.{}", quote_ident(&nsp), quote_ident(&name)),
immutable,
)
}
})
.clone()
}
unsafe fn note_functions(&mut self, query: *mut pg_sys::Node) {
let mut funcids = Vec::new();
unsafe { collect_functions(query, &mut funcids) };
for funcid in funcids {
let (name, immutable) = self.function(funcid);
if !immutable
&& !name.starts_with("pg_catalog.")
&& !self.graph.untracked_functions.contains(&name)
{
self.graph.untracked_functions.push(name);
}
}
}
}
impl Clone for Level {
fn clone(&self) -> Self {
Self {
query: self.query,
rtes: self.rtes.clone(),
link: self.link,
cte_parent: self.cte_parent,
}
}
}
fn list_len(list: *mut pg_sys::List) -> usize {
if list.is_null() {
0
} else {
usize::try_from(unsafe { (*list).length }).unwrap_or(0)
}
}
unsafe fn opaque_reason(query: *mut pg_sys::Query) -> Option<String> {
unsafe {
if (*query).hasWindowFuncs {
Some("read under a window function".to_string())
} else if !(*query).limitCount.is_null() || !(*query).limitOffset.is_null() {
Some("read under LIMIT/OFFSET".to_string())
} else if (*query).hasTargetSRFs {
Some("read under a set-returning function".to_string())
} else if !(*query).groupingSets.is_null() {
Some("read under GROUPING SETS".to_string())
} else {
None
}
}
}
unsafe fn top_opaque_reason(query: *mut pg_sys::Query) -> Option<String> {
unsafe { opaque_reason(query) }.map(|why| format!("{why} in the top-level SELECT"))
}
unsafe fn in_clause(sortref: pg_sys::Index, clause: *mut pg_sys::List) -> bool {
sortref != 0
&& unsafe { elements::<pg_sys::SortGroupClause>(clause) }
.iter()
.any(|c| unsafe { (**c).tleSortGroupRef } == sortref)
}
unsafe fn setop_leaves(node: *mut pg_sys::Node, leaves: &mut Vec<usize>) {
unsafe {
match tag(node) {
Some(pg_sys::NodeTag::T_RangeTblRef) => {
leaves.push((*node.cast::<pg_sys::RangeTblRef>()).rtindex as usize);
}
Some(pg_sys::NodeTag::T_SetOperationStmt) => {
let op = node.cast::<pg_sys::SetOperationStmt>();
setop_leaves((*op).larg, leaves);
setop_leaves((*op).rarg, leaves);
}
_ => {}
}
}
}
unsafe fn strip_relabel(mut node: *mut pg_sys::Node) -> *mut pg_sys::Node {
unsafe {
while tag(node) == Some(pg_sys::NodeTag::T_RelabelType) {
node = (*node.cast::<pg_sys::RelabelType>()).arg.cast();
}
}
node
}
unsafe fn is_required_sublink(node: *mut pg_sys::Node) -> bool {
unsafe {
tag(node) == Some(pg_sys::NodeTag::T_SubLink)
&& matches!(
(*node.cast::<pg_sys::SubLink>()).subLinkType,
pg_sys::SubLinkType::EXISTS_SUBLINK | pg_sys::SubLinkType::ANY_SUBLINK
)
}
}
unsafe fn equality(expr: *mut pg_sys::Node, chosen: &[&Column]) -> Option<(Column, Column)> {
unsafe {
if tag(expr) != Some(pg_sys::NodeTag::T_OpExpr) || chosen.len() != 2 {
return None;
}
let op = expr.cast::<pg_sys::OpExpr>();
let args = elements::<pg_sys::Node>((*op).args);
let plain = args.len() == 2
&& args.iter().all(|a| {
matches!(
tag(strip_relabel(*a)),
Some(pg_sys::NodeTag::T_Var | pg_sys::NodeTag::T_Param)
)
});
(plain && cstr(pg_sys::get_opname((*op).opno)) == "=")
.then(|| (chosen[0].clone(), chosen[1].clone()))
}
}
fn quote_ident(name: &str) -> String {
let Ok(c) = std::ffi::CString::new(name) else {
return crate::utils::quote_identifier(name);
};
cstr(unsafe { pg_sys::quote_identifier(c.as_ptr()) })
}
fn quote_literal(text: &str) -> Option<String> {
let c = std::ffi::CString::new(text).ok()?;
Some(cstr(unsafe { pg_sys::quote_literal_cstr(c.as_ptr()) }))
}
unsafe fn referenced_columns(query: *mut pg_sys::Query) -> HashMap<usize, Option<HashSet<i16>>> {
struct Refs {
depth: u32,
vars: Vec<(usize, i16)>,
}
unsafe extern "C-unwind" fn walker(
node: *mut pg_sys::Node,
ctx: *mut std::ffi::c_void,
) -> bool {
unsafe {
let refs = &mut *ctx.cast::<Refs>();
match tag(node) {
None => false,
Some(pg_sys::NodeTag::T_Var) => {
let var = node.cast::<pg_sys::Var>();
if (*var).varlevelsup == refs.depth {
refs.vars.push(((*var).varno as usize, (*var).varattno));
}
false
}
Some(pg_sys::NodeTag::T_Query) => {
refs.depth += 1;
let done = pg_sys::query_tree_walker(node.cast(), Some(walker), ctx, 0);
refs.depth -= 1;
done
}
Some(_) => pg_sys::expression_tree_walker(node, Some(walker), ctx),
}
}
}
let mut refs = Refs {
depth: 0,
vars: Vec::new(),
};
#[cfg(feature = "pg18")]
let flags = pg_sys::QTW_IGNORE_JOINALIASES | pg_sys::QTW_IGNORE_GROUPEXPRS;
#[cfg(not(feature = "pg18"))]
let flags = pg_sys::QTW_IGNORE_JOINALIASES;
unsafe {
pg_sys::query_tree_walker(
query,
Some(walker),
std::ptr::from_mut(&mut refs).cast(),
flags.cast_signed(),
);
}
let rtable = unsafe { elements::<pg_sys::RangeTblEntry>((*query).rtable) };
let mut read: HashMap<usize, Option<HashSet<i16>>> = HashMap::new();
let mut pending = refs.vars;
let mut seen: HashSet<(usize, i16)> = HashSet::new();
while let Some((varno, attno)) = pending.pop() {
if !seen.insert((varno, attno)) {
continue;
}
let Some(&rte) = rtable.get(varno.wrapping_sub(1)) else {
continue;
};
let behind = unsafe {
match (*rte).rtekind {
pg_sys::RTEKind::RTE_JOIN => Some((*rte).joinaliasvars),
#[cfg(feature = "pg18")]
pg_sys::RTEKind::RTE_GROUP => Some((*rte).groupexprs),
_ => None,
}
};
if let Some(list) = behind {
let exprs = unsafe { elements::<pg_sys::Node>(list) };
let chosen: Vec<*mut pg_sys::Node> = if attno == 0 {
exprs
} else {
exprs
.get(usize::try_from(attno - 1).unwrap_or(usize::MAX))
.copied()
.into_iter()
.collect()
};
for expr in chosen {
let mut vars = Vec::new();
unsafe { collect_vars(expr, &mut vars) };
for var in vars {
unsafe {
if (*var).varlevelsup == 0 {
pending.push(((*var).varno as usize, (*var).varattno));
}
}
}
}
continue;
}
let entry = read.entry(varno).or_insert_with(|| Some(HashSet::new()));
if attno == 0 {
*entry = None;
} else if let Some(columns) = entry {
columns.insert(attno);
}
}
read
}
unsafe fn collect_vars(node: *mut pg_sys::Node, out: &mut Vec<*mut pg_sys::Var>) {
unsafe extern "C-unwind" fn walker(
node: *mut pg_sys::Node,
ctx: *mut std::ffi::c_void,
) -> bool {
unsafe {
match tag(node) {
None => false,
Some(pg_sys::NodeTag::T_Var) => {
(*ctx.cast::<Vec<*mut pg_sys::Var>>()).push(node.cast());
false
}
Some(pg_sys::NodeTag::T_SubLink | pg_sys::NodeTag::T_Query) => false,
Some(_) => pg_sys::expression_tree_walker(node, Some(walker), ctx),
}
}
}
unsafe { walker(node, std::ptr::from_mut(out).cast()) };
}
unsafe fn collect_params(node: *mut pg_sys::Node, out: &mut Vec<*mut pg_sys::Param>) {
unsafe extern "C-unwind" fn walker(
node: *mut pg_sys::Node,
ctx: *mut std::ffi::c_void,
) -> bool {
unsafe {
match tag(node) {
None => false,
Some(pg_sys::NodeTag::T_Param) => {
(*ctx.cast::<Vec<*mut pg_sys::Param>>()).push(node.cast());
false
}
Some(pg_sys::NodeTag::T_SubLink | pg_sys::NodeTag::T_Query) => false,
Some(_) => pg_sys::expression_tree_walker(node, Some(walker), ctx),
}
}
}
unsafe { walker(node, std::ptr::from_mut(out).cast()) };
}
unsafe fn collect_sublinks(node: *mut pg_sys::Node, out: &mut Vec<*mut pg_sys::SubLink>) {
unsafe extern "C-unwind" fn walker(
node: *mut pg_sys::Node,
ctx: *mut std::ffi::c_void,
) -> bool {
unsafe {
match tag(node) {
None => false,
Some(pg_sys::NodeTag::T_SubLink) => {
(*ctx.cast::<Vec<*mut pg_sys::SubLink>>()).push(node.cast());
false
}
Some(pg_sys::NodeTag::T_Query) => false,
Some(_) => pg_sys::expression_tree_walker(node, Some(walker), ctx),
}
}
}
unsafe { walker(node, std::ptr::from_mut(out).cast()) };
}
unsafe fn has_sublink(node: *mut pg_sys::Node) -> bool {
let mut found = Vec::new();
unsafe { collect_sublinks(node, &mut found) };
!found.is_empty()
}
unsafe fn collect_functions(node: *mut pg_sys::Node, out: &mut Vec<Oid>) {
unsafe extern "C-unwind" fn walker(
node: *mut pg_sys::Node,
ctx: *mut std::ffi::c_void,
) -> bool {
unsafe {
match tag(node) {
None => false,
Some(pg_sys::NodeTag::T_FuncExpr) => {
(*ctx.cast::<Vec<Oid>>()).push((*node.cast::<pg_sys::FuncExpr>()).funcid);
pg_sys::expression_tree_walker(node, Some(walker), ctx)
}
Some(pg_sys::NodeTag::T_SubLink | pg_sys::NodeTag::T_Query) => false,
Some(_) => pg_sys::expression_tree_walker(node, Some(walker), ctx),
}
}
}
unsafe {
pg_sys::query_tree_walker(
node.cast(),
Some(walker),
std::ptr::from_mut(out).cast(),
(pg_sys::QTW_IGNORE_RT_SUBQUERIES | pg_sys::QTW_IGNORE_CTE_SUBQUERIES).cast_signed(),
);
}
}