#![allow(clippy::cast_ptr_alignment)]
use super::{
Column, Conjunct, Graph, IdentityKind, Lookup, Maps, Occurrence, OutputColumn, Piece, 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 HashMap<Oid, String>,
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,
nullable: HashSet::new(),
};
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>),
Expr(Computed),
Opaque,
}
#[derive(Debug, Clone)]
struct Computed {
sql: Sql,
element: bool,
strict: bool,
type_oid: Oid,
}
impl Resolved {
fn occs(&self) -> Vec<usize> {
match self {
Self::Col(c) => vec![c.occ],
Self::Alt(cs) => cs.iter().map(|c| c.occ).collect(),
Self::Expr(e) => sql_occs(&e.sql),
Self::Opaque => vec![],
}
}
}
fn sql_occs(sql: &Sql) -> Vec<usize> {
let mut occs = Vec::new();
for piece in &sql.0 {
if let Piece::Column { occ, .. } = piece
&& !occs.contains(occ)
{
occs.push(*occ);
}
}
occs
}
#[derive(Debug, Clone)]
enum RteInfo {
Base(usize),
Outputs(Vec<Resolved>),
Join(*mut pg_sys::List),
Tview {
entity: String,
relid: Oid,
},
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<Resolved>,
}
struct Operand {
sql: Sql,
type_oid: Oid,
column: bool,
element: bool,
}
#[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,
nullable: HashSet<usize>,
}
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)?;
self.graph.outputs = elements::<pg_sys::TargetEntry>((*query).targetList)
.iter()
.zip(&outputs)
.filter(|(tle, _)| !(***tle).resjunk)
.map(|(tle, out)| {
let column = match out {
Resolved::Col(c) => Some(c.clone()),
_ => None,
};
(cstr((**tle).resname), column)
})
.collect();
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 pass_through: Vec<bool> = tles
.iter()
.zip(skipped)
.map(|(&tle, skip)| {
!skip
&& (link == Link::Top
|| (!opaque
&& (!grouped || keyed(tle, (*query).groupClause, &group_keys))
&& (!(*query).hasDistinctOn
|| keyed(tle, (*query).distinctClause, &distinct_keys))))
})
.collect();
Ok(tles
.iter()
.zip(pass_through)
.map(|(&tle, pass)| {
if pass {
self.output((*tle).expr.cast())
} else {
Resolved::Opaque
}
})
.collect())
}
}
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::Expr(_) | 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 if (*rte).self_reference => Ok(RteInfo::Other),
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).ctequery.cast()).cast::<pg_sys::Query>();
self.cte_parent = Some(defined_at);
if !(*cte).cterecursive {
return Ok(RteInfo::Outputs(self.level(copy, flags, Link::From)?));
}
let recursive = Flags {
opaque_level: Some(format!(
"read in a recursive CTE ({})",
flags.via_view.as_deref().unwrap_or("the definition")
)),
..flags.clone()
};
let width = self.level(copy, &recursive, Link::From)?.len();
Ok(RteInfo::Outputs(vec![Resolved::Opaque; width]))
}
pg_sys::RTEKind::RTE_JOIN => Ok(RteInfo::Join((*rte).joinaliasvars)),
pg_sys::RTEKind::RTE_FUNCTION => Ok(self.function_rte(rte)),
#[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_key(&relid) => {
let entity = self.ctx.tview_tables[&relid].clone();
self.graph.tview_keys.entry(entity.clone()).or_default();
Ok(RteInfo::Tview { entity, relid })
}
b'r' | b'p' => {
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 query = unsafe { view_query(relid)? };
let outputs = unsafe { self.level(query, &inner, Link::From)? };
if let Some(entity) = self.ctx.tview_views.get(&relid) {
let key = unsafe { output_position(query, &format!("pk_{entity}")) };
let columns = match key.and_then(|i| outputs.get(i)) {
Some(Resolved::Col(c)) => vec![c.clone()],
Some(Resolved::Alt(cs)) => cs.clone(),
_ => Vec::new(),
};
self.graph
.tview_keys
.entry(entity.clone())
.or_default()
.extend(columns);
}
Ok(RteInfo::Outputs(outputs))
}
_ => Ok(RteInfo::Other),
}
}
unsafe fn function_rte(&mut self, rte: *mut pg_sys::RangeTblEntry) -> RteInfo {
unsafe {
let functions = elements::<pg_sys::RangeTblFunction>((*rte).functions);
let [function] = functions[..] else {
return RteInfo::Other;
};
if (*rte).funcordinality {
return RteInfo::Other;
}
match unnest_array((*function).funcexpr) {
Some(array) => RteInfo::Outputs(vec![self.computed(array, true)]),
None => RteInfo::Other,
}
}
}
fn cte(&self, name: &str, levelsup: usize) -> Option<(*mut pg_sys::CommonTableExpr, 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, 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)?;
match (*join).jointype {
pg_sys::JoinType::JOIN_INNER => {}
pg_sys::JoinType::JOIN_LEFT => self.nullable.extend(&right),
pg_sys::JoinType::JOIN_RIGHT => self.nullable.extend(&left),
_ => self.nullable.extend(left.iter().chain(&right)),
}
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(Resolved::occs).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 {
self.tview_key_equality(qual);
if matches!(origin, Origin::None)
|| pg_sys::contain_mutable_functions(qual)
|| has_sublink(qual)
{
return;
}
let mut sites = Vec::new();
if !self.sites(qual, &mut sites) || !self.null_safe(qual, &sites) {
return;
}
self.add_conjuncts(qual, &sites, &|_| None, origin, true);
}
}
unsafe fn tview_key_equality(&mut self, qual: *mut pg_sys::Node) {
unsafe {
if tag(qual) != Some(pg_sys::NodeTag::T_OpExpr)
|| cstr(pg_sys::get_opname((*qual.cast::<pg_sys::OpExpr>()).opno)) != "="
{
return;
}
let args = elements::<pg_sys::Node>((*qual.cast::<pg_sys::OpExpr>()).args);
let [l, r] = args[..] else { return };
let (l, r) = (strip_relabel(l), strip_relabel(r));
if tag(l) != Some(pg_sys::NodeTag::T_Var) || tag(r) != Some(pg_sys::NodeTag::T_Var) {
return;
}
for (key, other) in [(l, r), (r, l)] {
let (key, other) = (key.cast::<pg_sys::Var>(), other.cast::<pg_sys::Var>());
let Some(entity) = self.tview_key_of(key) else {
continue;
};
if let Resolved::Col(c) = self.resolve_var(other, (*other).varlevelsup as usize) {
self.graph.tview_keys.entry(entity).or_default().push(c);
}
}
}
}
unsafe fn tview_key_of(&self, var: *mut pg_sys::Var) -> Option<String> {
unsafe {
let index = self
.levels
.len()
.checked_sub(1 + (*var).varlevelsup as usize)?;
let rtindex = usize::try_from((*var).varno).ok()?;
let Some(RteInfo::Tview { entity, relid }) =
self.levels[index].rtes.get(rtindex.wrapping_sub(1))
else {
return None;
};
(cstr(pg_sys::get_attname(*relid, (*var).varattno, true)) == format!("pk_{entity}"))
.then(|| entity.clone())
}
}
unsafe fn null_safe(&self, expr: *mut pg_sys::Node, sites: &[Site]) -> bool {
let terms = || sites.iter().flat_map(|s| &s.candidates);
let strict = unsafe { !pg_sys::contain_nonstrict_functions(expr) }
&& terms().all(|t| !matches!(t, Resolved::Expr(e) if !e.strict));
strict
|| terms()
.flat_map(Resolved::occs)
.all(|occ| !self.nullable.contains(&occ))
}
unsafe fn test_predicate(
&mut self,
test: *mut pg_sys::Node,
outputs: &[Resolved],
required: bool,
) {
unsafe {
if 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<Resolved>> {
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::Alt(cs) => Some(cs.iter().cloned().map(Resolved::Col).collect()),
Resolved::Opaque => None,
term => Some(vec![term.clone()]),
}
};
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,
});
}
if !self.null_safe(test, &sites) {
return;
}
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::Alt(cs) => cs.into_iter().map(Resolved::Col).collect(),
Resolved::Opaque => return false,
term => vec![term],
};
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<&Resolved> = 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().flat_map(|c| c.occs()).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 term_of = |node: *mut pg_sys::Node| -> Option<Resolved> {
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, lookups)) =
(unsafe { self.conjunct_sql(expr, &term_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,
lookups,
});
}
}
fn directions(
&self,
sites: &[Site],
chosen: &[&Resolved],
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.occs().contains(&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 output(&mut self, node: *mut pg_sys::Node) -> Resolved {
unsafe {
if tag(strip_relabel(node)) == Some(pg_sys::NodeTag::T_Var) {
self.resolve_expr(node)
} else if pg_sys::expression_returns_set(node) {
unnest_array(node).map_or(Resolved::Opaque, |array| self.computed(array, true))
} else {
self.computed(node, false)
}
}
}
unsafe fn computed(&mut self, expr: *mut pg_sys::Node, element: bool) -> Resolved {
unsafe {
if expr.is_null()
|| pg_sys::contain_mutable_functions(expr)
|| has_sublink(expr)
|| pg_sys::expression_returns_set(expr)
{
return Resolved::Opaque;
}
let mut vars = Vec::new();
collect_vars(expr, &mut vars);
if vars.is_empty() {
return Resolved::Opaque;
}
let mut terms = Vec::new();
for var in vars {
match self.resolve_var(var, (*var).varlevelsup as usize) {
term @ (Resolved::Col(_) | Resolved::Expr(_)) => terms.push(term),
_ => return Resolved::Opaque,
}
}
let strict = !pg_sys::contain_nonstrict_functions(expr)
&& terms
.iter()
.all(|t| !matches!(t, Resolved::Expr(e) if !e.strict));
let index = std::cell::Cell::new(0_usize);
let term_of = |node: *mut pg_sys::Node| {
(tag(node) == Some(pg_sys::NodeTag::T_Var))
.then(|| terms.get(index.get()).cloned())
.flatten()
};
let mut next_var = || index.set(index.get() + 1);
match self.deparse(expr, &term_of, &mut next_var) {
Some(sql) => Resolved::Expr(Computed {
sql,
element,
strict,
type_oid: pg_sys::exprType(expr),
}),
None => Resolved::Opaque,
}
}
}
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 conjunct_sql(
&mut self,
expr: *mut pg_sys::Node,
term_of: &dyn Fn(*mut pg_sys::Node) -> Option<Resolved>,
next_var: &mut dyn FnMut(),
) -> Option<(Sql, Vec<Lookup>)> {
unsafe {
match tag(expr)? {
pg_sys::NodeTag::T_OpExpr => {
let op = expr.cast::<pg_sys::OpExpr>();
if let [l, r] = elements::<pg_sys::Node>((*op).args)[..]
&& cstr(pg_sys::get_opname((*op).opno)) == "="
{
let l = self.operand(l, term_of, next_var)?;
let r = self.operand(r, term_of, next_var)?;
return match (l.element, r.element) {
(false, false) => Some(self.comparison((*op).opno, &l, &r)),
(false, true) => self.membership((*op).opno, &l, &r),
(true, false) => {
let commutator = pg_sys::get_commutator((*op).opno);
(commutator != Oid::INVALID)
.then(|| self.membership(commutator, &r, &l))
.flatten()
}
(true, true) => None,
};
}
}
pg_sys::NodeTag::T_ScalarArrayOpExpr => {
let op = expr.cast::<pg_sys::ScalarArrayOpExpr>();
if let [l, r] = elements::<pg_sys::Node>((*op).args)[..]
&& (*op).useOr
{
let l = self.operand(l, term_of, next_var)?;
let r = self.operand(r, term_of, next_var)?;
if l.element || r.element {
return None;
}
return self.membership((*op).opno, &l, &r);
}
}
_ => {}
}
Some((self.deparse(expr, term_of, next_var)?, Vec::new()))
}
}
unsafe fn operand(
&mut self,
arg: *mut pg_sys::Node,
term_of: &dyn Fn(*mut pg_sys::Node) -> Option<Resolved>,
next_var: &mut dyn FnMut(),
) -> Option<Operand> {
unsafe {
let bare = strip_relabel(arg);
let term = match tag(bare) {
Some(pg_sys::NodeTag::T_Var | pg_sys::NodeTag::T_Param) => term_of(bare),
_ => None,
};
if let Some(Resolved::Expr(e)) = &term
&& e.element
{
if tag(bare) == Some(pg_sys::NodeTag::T_Var) {
next_var();
}
return Some(Operand {
sql: e.sql.clone(),
type_oid: e.type_oid,
column: false,
element: true,
});
}
Some(Operand {
sql: self.deparse(arg, term_of, next_var)?,
type_oid: pg_sys::exprType(arg),
column: matches!(term, Some(Resolved::Col(_))),
element: false,
})
}
}
fn comparison(&mut self, opno: Oid, l: &Operand, r: &Operand) -> (Sql, Vec<Lookup>) {
let name = self.operator_name(opno);
let mut sql = Sql::text("(");
sql.push_sql(l.sql.clone());
sql.push_text(&format!(" OPERATOR({name}) "));
sql.push_sql(r.sql.clone());
sql.push_text(")");
let mut lookups = Vec::new();
for (side, other) in [(l, r), (r, l)] {
if !side.column
&& other.column
&& let [occ] = sql_occs(&side.sql)[..]
{
lookups.push(Lookup {
occ,
expr: side.sql.clone(),
gin: false,
});
}
}
(sql, lookups)
}
fn membership(
&mut self,
opno: Oid,
scalar: &Operand,
array: &Operand,
) -> Option<(Sql, Vec<Lookup>)> {
let name = self.operator_name(opno);
let mut sql = Sql::text("(");
sql.push_sql(scalar.sql.clone());
sql.push_text(&format!(" OPERATOR({name}) ANY ("));
sql.push_sql(array.sql.clone());
sql.push_text("))");
let containment = unsafe {
let element = pg_sys::get_element_type(array.type_oid);
element != Oid::INVALID
&& element == scalar.type_oid
&& (*pg_sys::lookup_type_cache(element, pg_sys::TYPECACHE_EQ_OPR.cast_signed()))
.eq_opr
== opno
};
let mut lookups = Vec::new();
if containment {
let mut both = Sql::text("(");
both.push_sql(sql);
both.push_text(" AND (");
both.push_sql(array.sql.clone());
both.push_text(") OPERATOR(pg_catalog.@>) ARRAY[");
both.push_sql(scalar.sql.clone());
both.push_text("])");
sql = both;
if let [occ] = sql_occs(&array.sql)[..] {
lookups.push(Lookup {
occ,
expr: array.sql.clone(),
gin: true,
});
}
}
Some((sql, lookups))
}
unsafe fn deparse(
&mut self,
expr: *mut pg_sys::Node,
term_of: &dyn Fn(*mut pg_sys::Node) -> Option<Resolved>,
next_var: &mut dyn FnMut(),
) -> Option<Sql> {
unsafe {
match tag(expr)? {
pg_sys::NodeTag::T_Var => {
let term = term_of(expr)?;
next_var();
term_sql(&term)
}
pg_sys::NodeTag::T_Param => term_sql(&term_of(expr)?),
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, term_of, next_var)?);
}
[l, r] => {
sql.push_sql(self.deparse(l, term_of, next_var)?);
sql.push_text(&format!(" OPERATOR({name}) "));
sql.push_sql(self.deparse(r, term_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, term_of, next_var)?);
sql.push_text(&format!(
" OPERATOR({name}) {} (",
if (*op).useOr { "ANY" } else { "ALL" }
));
sql.push_sql(self.deparse(r, term_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, term_of, next_var)?);
}
sql.push_text(")");
} else {
let [arg, ..] = args[..] else { return None };
sql = Sql::text("(");
sql.push_sql(self.deparse(arg, term_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(), term_of, next_var)?);
sql.push_text(&format!(")::{}", self.type_name((*r).resulttype)));
Some(sql)
}
pg_sys::NodeTag::T_ArrayCoerceExpr => {
let r = expr.cast::<pg_sys::ArrayCoerceExpr>();
let mut sql = Sql::text("(");
sql.push_sql(self.deparse((*r).arg.cast(), term_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(), term_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()?, term_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, term_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).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)
.or_else(|| {
(*query)
.hasTargetSRFs
.then(|| "read under a set-returning function".to_string())
})
.map(|why| format!("{why} in the top-level SELECT"))
}
}
unsafe fn output_position(query: *mut pg_sys::Query, name: &str) -> Option<usize> {
unsafe {
elements::<pg_sys::TargetEntry>((*query).targetList)
.iter()
.position(|tle| !(**tle).resjunk && cstr((**tle).resname) == name)
}
}
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
)
}
}
fn term_sql(term: &Resolved) -> Option<Sql> {
match term {
Resolved::Col(c) => Some(c.sql()),
Resolved::Expr(e) if !e.element => {
let mut sql = Sql::text("(");
sql.push_sql(e.sql.clone());
sql.push_text(")");
Some(sql)
}
_ => None,
}
}
unsafe fn unnest_array(node: *mut pg_sys::Node) -> Option<*mut pg_sys::Node> {
unsafe {
let node = strip_relabel(node);
if tag(node) != Some(pg_sys::NodeTag::T_FuncExpr) {
return None;
}
let f = node.cast::<pg_sys::FuncExpr>();
if (*f).funcid != Oid::from(pg_sys::F_UNNEST_ANYARRAY) {
return None;
}
let [array] = elements::<pg_sys::Node>((*f).args)[..] else {
return None;
};
Some(array)
}
}
unsafe fn equality(expr: *mut pg_sys::Node, chosen: &[&Resolved]) -> Option<(Column, Column)> {
unsafe {
let [Resolved::Col(x), Resolved::Col(y)] = chosen[..] else {
return None;
};
if tag(expr) != Some(pg_sys::NodeTag::T_OpExpr) {
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(|| (x.clone(), y.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(),
);
}
}