Skip to main content

pg_query/
parse_result.rs

1use std::collections::HashMap;
2use std::collections::HashSet;
3use std::iter::FromIterator;
4use std::string::String;
5
6use itertools::join;
7
8use crate::*;
9
10macro_rules! cast {
11    ($target: expr, $pat: path) => {{
12        if let $pat(a) = $target {
13            // #1
14            a
15        } else {
16            panic!("mismatch variant when cast to {}", stringify!($pat)); // #2
17        }
18    }};
19}
20
21impl protobuf::ParseResult {
22    pub fn deparse(&self) -> Result<String> {
23        crate::deparse(self)
24    }
25
26    // Note: this doesn't iterate over every possible node type, since we only care about a subset of nodes.
27    pub fn nodes(&self) -> Vec<(NodeRef<'_>, i32, Context, bool)> {
28        self.stmts
29            .iter()
30            .filter_map(|s|
31            // RawStmt  ->  Node   ->    NodeEnum           ->              NodeRef
32            s.stmt.as_ref().and_then(|s| s.node.as_ref()).map(|n| n.nodes()))
33            .flatten()
34            .collect()
35    }
36
37    /// Returns a mutable reference to nested nodes.
38    ///
39    /// # Safety
40    ///
41    /// The caller may have to deal with dangling pointers, and passing an
42    /// invalid tree back to libpg_query may cause it to panic.
43    pub unsafe fn nodes_mut(&mut self) -> Vec<(NodeMut, i32, Context)> {
44        self.stmts
45            .iter_mut()
46            .filter_map(|s|
47            // RawStmt  ->  Node   ->    NodeEnum           ->              NodeMut
48            s.stmt.as_mut().and_then(|s| s.node.as_mut()).map(|n| n.nodes_mut()))
49            .flatten()
50            .collect()
51    }
52}
53
54/// Result from calling [parse]
55#[derive(Debug)]
56pub struct ParseResult {
57    pub protobuf: protobuf::ParseResult,
58    pub warnings: Vec<String>,
59    pub tables: Vec<(String, Context)>,
60    pub aliases: HashMap<String, String>,
61    pub cte_names: Vec<String>,
62    pub functions: Vec<(String, Context)>,
63    pub filter_columns: Vec<(Option<String>, String)>,
64}
65
66impl ParseResult {
67    pub fn new(protobuf: protobuf::ParseResult, stderr: String) -> Self {
68        let warnings = stderr
69            .lines()
70            .filter_map(|l| {
71                if l.starts_with("WARNING") {
72                    Some(l.trim().into())
73                } else {
74                    None
75                }
76            })
77            .collect();
78        let mut tables: HashSet<(String, Context)> = HashSet::new();
79        let mut aliases: HashMap<String, String> = HashMap::new();
80        let mut cte_names: HashSet<String> = HashSet::new();
81        let mut functions: HashSet<(String, Context)> = HashSet::new();
82        let mut filter_columns: HashSet<(Option<String>, String)> = HashSet::new();
83
84        for (node, _depth, context, has_filter_columns) in protobuf.nodes().into_iter() {
85            match node {
86                NodeRef::CommonTableExpr(s) => {
87                    cte_names.insert(s.ctename.to_owned());
88                }
89                NodeRef::RangeVar(v) => {
90                    // TODO: this incorrectly returns no tables: parse('with f as (select * from f limit 1) select * from f')
91                    let table = if !v.schemaname.is_empty() {
92                        format!("{}.{}", v.schemaname, v.relname)
93                    } else {
94                        v.relname.to_owned()
95                    };
96                    if cte_names.contains(&table) {
97                        continue;
98                    }
99                    tables.insert((table.to_owned(), context));
100                    v.alias
101                        .as_ref()
102                        .and_then(|alias| aliases.insert(alias.aliasname.to_owned(), table));
103                }
104                NodeRef::FuncCall(c) => {
105                    let funcname = join(
106                        c.funcname.iter().filter_map(|n| {
107                            n.node.as_ref().map(|n| &cast!(n, NodeEnum::String).sval)
108                        }),
109                        ".",
110                    );
111                    functions.insert((funcname, Context::Call));
112                }
113                NodeRef::DropStmt(s) => {
114                    match protobuf::ObjectType::try_from(s.remove_type) {
115                        Ok(protobuf::ObjectType::ObjectTable) => {
116                            for o in &s.objects {
117                                if let Some(NodeEnum::List(list)) = &o.node {
118                                    let table = join(
119                                        list.items.iter().filter_map(|i| {
120                                            i.node
121                                                .as_ref()
122                                                .map(|n| &cast!(n, NodeEnum::String).sval)
123                                        }),
124                                        ".",
125                                    );
126                                    tables.insert((table, Context::DDL));
127                                };
128                            }
129                        }
130                        Ok(protobuf::ObjectType::ObjectRule)
131                        | Ok(protobuf::ObjectType::ObjectTrigger) => {
132                            for o in &s.objects {
133                                if let Some(NodeEnum::List(list)) = &o.node {
134                                    // Unlike ObjectTable, this ignores the last string (the rule/trigger name)
135                                    let table = join(
136                                        list.items[0..list.items.len() - 1].iter().filter_map(
137                                            |i| {
138                                                i.node
139                                                    .as_ref()
140                                                    .map(|n| &cast!(n, NodeEnum::String).sval)
141                                            },
142                                        ),
143                                        ".",
144                                    );
145                                    tables.insert((table, Context::DDL));
146                                };
147                            }
148                        }
149                        Ok(protobuf::ObjectType::ObjectFunction) => {
150                            // Only one function can be dropped in a statement
151                            if let Some(NodeEnum::ObjectWithArgs(object)) = &s.objects[0].node {
152                                if let Some(NodeEnum::String(string)) = &object.objname[0].node {
153                                    functions.insert((string.sval.to_string(), Context::DDL));
154                                }
155                            }
156                        }
157                        _ => (),
158                    }
159                }
160                NodeRef::CreateFunctionStmt(s) => {
161                    if let Some(NodeEnum::String(string)) = &s.funcname[0].node {
162                        functions.insert((string.sval.to_string(), Context::DDL));
163                    }
164                }
165                NodeRef::RenameStmt(s) => {
166                    if let Ok(protobuf::ObjectType::ObjectFunction) =
167                        protobuf::ObjectType::try_from(s.rename_type)
168                    {
169                        if let Some(object) = &s.object {
170                            if let Some(NodeEnum::ObjectWithArgs(object)) = &object.node {
171                                if let Some(NodeEnum::String(string)) = &object.objname[0].node {
172                                    functions.insert((string.sval.to_string(), Context::DDL));
173                                    functions.insert((s.newname.to_string(), Context::DDL));
174                                }
175                            }
176                        }
177                    }
178                }
179                NodeRef::ColumnRef(c) => {
180                    if !has_filter_columns {
181                        continue;
182                    }
183                    let f: Vec<String> = c
184                        .fields
185                        .iter()
186                        .filter_map(|n| match n.node.as_ref() {
187                            Some(NodeEnum::String(s)) => Some(s.sval.to_string()),
188                            _ => None,
189                        })
190                        .rev()
191                        .collect();
192                    if f.len() > 0 {
193                        filter_columns.insert((f.get(1).cloned(), f[0].to_string()));
194                    }
195                }
196                _ => (),
197            }
198        }
199
200        Self {
201            protobuf,
202            warnings,
203            tables: Vec::from_iter(tables),
204            aliases,
205            cte_names: Vec::from_iter(cte_names),
206            functions: Vec::from_iter(functions),
207            filter_columns: Vec::from_iter(filter_columns),
208        }
209    }
210
211    /// Returns all referenced tables in the query
212    pub fn tables(&self) -> Vec<String> {
213        let mut tables = HashSet::new();
214        self.tables.iter().for_each(|(t, _c)| {
215            tables.insert(t.to_string());
216        });
217        Vec::from_iter(tables)
218    }
219
220    /// Returns only tables that were selected from
221    pub fn select_tables(&self) -> Vec<String> {
222        self.tables
223            .iter()
224            .filter_map(|(table, context)| match context {
225                Context::Select => Some(table.to_string()),
226                _ => None,
227            })
228            .collect()
229    }
230
231    /// Returns only tables that were modified by the query
232    pub fn dml_tables(&self) -> Vec<String> {
233        self.tables
234            .iter()
235            .filter_map(|(table, context)| match context {
236                Context::DML => Some(table.to_string()),
237                _ => None,
238            })
239            .collect()
240    }
241
242    /// Returns only tables that were modified by DDL statements
243    pub fn ddl_tables(&self) -> Vec<String> {
244        self.tables
245            .iter()
246            .filter_map(|(table, context)| match context {
247                Context::DDL => Some(table.to_string()),
248                _ => None,
249            })
250            .collect()
251    }
252
253    /// Returns all function references
254    pub fn functions(&self) -> Vec<String> {
255        let mut functions = HashSet::new();
256        self.functions.iter().for_each(|(f, _c)| {
257            functions.insert(f.to_string());
258        });
259        Vec::from_iter(functions)
260    }
261
262    /// Returns DDL functions
263    pub fn ddl_functions(&self) -> Vec<String> {
264        self.functions
265            .iter()
266            .filter_map(|(function, context)| match context {
267                Context::DDL => Some(function.to_string()),
268                _ => None,
269            })
270            .collect()
271    }
272
273    /// Returns functions that were called
274    pub fn call_functions(&self) -> Vec<String> {
275        self.functions
276            .iter()
277            .filter_map(|(function, context)| match context {
278                Context::Call => Some(function.to_string()),
279                _ => None,
280            })
281            .collect()
282    }
283
284    /// Converts the parsed query back into a SQL string
285    pub fn deparse(&self) -> Result<String> {
286        crate::deparse(&self.protobuf)
287    }
288
289    /// Intelligently truncates queries to a max length.
290    ///
291    /// # Example
292    ///
293    /// ```rust
294    /// let query = "INSERT INTO \"x\" (a, b, c, d, e, f) VALUES ($1)";
295    /// let result = pg_query::parse(query).unwrap();
296    /// assert_eq!(result.truncate(32).unwrap(), "INSERT INTO x (...) VALUES (...)")
297    /// ```
298    pub fn truncate(&self, max_length: usize) -> Result<String> {
299        crate::truncate(&self.protobuf, max_length)
300    }
301
302    /// Returns all statement types in the query
303    pub fn statement_types(&self) -> Vec<&str> {
304        self.protobuf
305            .stmts
306            .iter()
307            .filter_map(|s| match s.stmt.as_ref().and_then(|s| s.node.as_ref()) {
308                Some(NodeEnum::InsertStmt(..)) => Some("InsertStmt"),
309                Some(NodeEnum::DeleteStmt(..)) => Some("DeleteStmt"),
310                Some(NodeEnum::UpdateStmt(..)) => Some("UpdateStmt"),
311                Some(NodeEnum::SelectStmt(..)) => Some("SelectStmt"),
312                Some(NodeEnum::MergeStmt(..)) => Some("MergeStmt"),
313                Some(NodeEnum::AlterTableStmt(..)) => Some("AlterTableStmt"),
314                Some(NodeEnum::AlterTableCmd(..)) => Some("AlterTableCmd"),
315                Some(NodeEnum::AlterDomainStmt(..)) => Some("AlterDomainStmt"),
316                Some(NodeEnum::SetOperationStmt(..)) => Some("SetOperationStmt"),
317                Some(NodeEnum::GrantStmt(..)) => Some("GrantStmt"),
318                Some(NodeEnum::GrantRoleStmt(..)) => Some("GrantRoleStmt"),
319                Some(NodeEnum::AlterDefaultPrivilegesStmt(..)) => {
320                    Some("AlterDefaultPrivilegesStmt")
321                }
322                Some(NodeEnum::ClosePortalStmt(..)) => Some("ClosePortalStmt"),
323                Some(NodeEnum::ClusterStmt(..)) => Some("ClusterStmt"),
324                Some(NodeEnum::CopyStmt(..)) => Some("CopyStmt"),
325                Some(NodeEnum::CreateStmt(..)) => Some("CreateStmt"),
326                Some(NodeEnum::DefineStmt(..)) => Some("DefineStmt"),
327                Some(NodeEnum::DropStmt(..)) => Some("DropStmt"),
328                Some(NodeEnum::TruncateStmt(..)) => Some("TruncateStmt"),
329                Some(NodeEnum::CommentStmt(..)) => Some("CommentStmt"),
330                Some(NodeEnum::FetchStmt(..)) => Some("FetchStmt"),
331                Some(NodeEnum::IndexStmt(..)) => Some("IndexStmt"),
332                Some(NodeEnum::CreateFunctionStmt(..)) => Some("CreateFunctionStmt"),
333                Some(NodeEnum::AlterFunctionStmt(..)) => Some("AlterFunctionStmt"),
334                Some(NodeEnum::DoStmt(..)) => Some("DoStmt"),
335                Some(NodeEnum::RenameStmt(..)) => Some("RenameStmt"),
336                Some(NodeEnum::RuleStmt(..)) => Some("RuleStmt"),
337                Some(NodeEnum::NotifyStmt(..)) => Some("NotifyStmt"),
338                Some(NodeEnum::ListenStmt(..)) => Some("ListenStmt"),
339                Some(NodeEnum::UnlistenStmt(..)) => Some("UnlistenStmt"),
340                Some(NodeEnum::TransactionStmt(..)) => Some("TransactionStmt"),
341                Some(NodeEnum::ViewStmt(..)) => Some("ViewStmt"),
342                Some(NodeEnum::LoadStmt(..)) => Some("LoadStmt"),
343                Some(NodeEnum::CreateDomainStmt(..)) => Some("CreateDomainStmt"),
344                Some(NodeEnum::CreatedbStmt(..)) => Some("CreatedbStmt"),
345                Some(NodeEnum::DropdbStmt(..)) => Some("DropdbStmt"),
346                Some(NodeEnum::VacuumStmt(..)) => Some("VacuumStmt"),
347                Some(NodeEnum::ExplainStmt(..)) => Some("ExplainStmt"),
348                Some(NodeEnum::CreateTableAsStmt(..)) => Some("CreateTableAsStmt"),
349                Some(NodeEnum::CreateSeqStmt(..)) => Some("CreateSeqStmt"),
350                Some(NodeEnum::AlterSeqStmt(..)) => Some("AlterSeqStmt"),
351                Some(NodeEnum::VariableSetStmt(..)) => Some("VariableSetStmt"),
352                Some(NodeEnum::VariableShowStmt(..)) => Some("VariableShowStmt"),
353                Some(NodeEnum::DiscardStmt(..)) => Some("DiscardStmt"),
354                Some(NodeEnum::CreateTrigStmt(..)) => Some("CreateTrigStmt"),
355                // CreatePLangStmt is capitalized differently to match C implementation.
356                Some(NodeEnum::CreatePlangStmt(..)) => Some("CreatePLangStmt"),
357                Some(NodeEnum::CreateRoleStmt(..)) => Some("CreateRoleStmt"),
358                Some(NodeEnum::AlterRoleStmt(..)) => Some("AlterRoleStmt"),
359                Some(NodeEnum::DropRoleStmt(..)) => Some("DropRoleStmt"),
360                Some(NodeEnum::LockStmt(..)) => Some("LockStmt"),
361                Some(NodeEnum::ConstraintsSetStmt(..)) => Some("ConstraintsSetStmt"),
362                Some(NodeEnum::ReindexStmt(..)) => Some("ReindexStmt"),
363                Some(NodeEnum::CheckPointStmt(..)) => Some("CheckPointStmt"),
364                Some(NodeEnum::CreateSchemaStmt(..)) => Some("CreateSchemaStmt"),
365                Some(NodeEnum::AlterDatabaseStmt(..)) => Some("AlterDatabaseStmt"),
366                Some(NodeEnum::AlterDatabaseSetStmt(..)) => Some("AlterDatabaseSetStmt"),
367                Some(NodeEnum::AlterRoleSetStmt(..)) => Some("AlterRoleSetStmt"),
368                Some(NodeEnum::CreateConversionStmt(..)) => Some("CreateConversionStmt"),
369                Some(NodeEnum::CreateCastStmt(..)) => Some("CreateCastStmt"),
370                Some(NodeEnum::CreateOpClassStmt(..)) => Some("CreateOpClassStmt"),
371                Some(NodeEnum::CreateOpFamilyStmt(..)) => Some("CreateOpFamilyStmt"),
372                Some(NodeEnum::AlterOpFamilyStmt(..)) => Some("AlterOpFamilyStmt"),
373                Some(NodeEnum::PrepareStmt(..)) => Some("PrepareStmt"),
374                Some(NodeEnum::ExecuteStmt(..)) => Some("ExecuteStmt"),
375                Some(NodeEnum::DeallocateStmt(..)) => Some("DeallocateStmt"),
376                Some(NodeEnum::DeclareCursorStmt(..)) => Some("DeclareCursorStmt"),
377                Some(NodeEnum::CreateTableSpaceStmt(..)) => Some("CreateTableSpaceStmt"),
378                Some(NodeEnum::DropTableSpaceStmt(..)) => Some("DropTableSpaceStmt"),
379                Some(NodeEnum::AlterObjectDependsStmt(..)) => Some("AlterObjectDependsStmt"),
380                Some(NodeEnum::AlterObjectSchemaStmt(..)) => Some("AlterObjectSchemaStmt"),
381                Some(NodeEnum::AlterOwnerStmt(..)) => Some("AlterOwnerStmt"),
382                Some(NodeEnum::AlterOperatorStmt(..)) => Some("AlterOperatorStmt"),
383                Some(NodeEnum::AlterTypeStmt(..)) => Some("AlterTypeStmt"),
384                Some(NodeEnum::DropOwnedStmt(..)) => Some("DropOwnedStmt"),
385                Some(NodeEnum::ReassignOwnedStmt(..)) => Some("ReassignOwnedStmt"),
386                Some(NodeEnum::CompositeTypeStmt(..)) => Some("CompositeTypeStmt"),
387                Some(NodeEnum::CreateEnumStmt(..)) => Some("CreateEnumStmt"),
388                Some(NodeEnum::CreateRangeStmt(..)) => Some("CreateRangeStmt"),
389                Some(NodeEnum::AlterEnumStmt(..)) => Some("AlterEnumStmt"),
390                // AlterTSDictionaryStmt is capitalized differently to match C implementation.
391                Some(NodeEnum::AlterTsdictionaryStmt(..)) => Some("AlterTSDictionaryStmt"),
392                // AlterTSConfigurationStmt is capitalized differently to match C implementation.
393                Some(NodeEnum::AlterTsconfigurationStmt(..)) => Some("AlterTSConfigurationStmt"),
394                Some(NodeEnum::CreateFdwStmt(..)) => Some("CreateFdwStmt"),
395                Some(NodeEnum::AlterFdwStmt(..)) => Some("AlterFdwStmt"),
396                Some(NodeEnum::CreateForeignServerStmt(..)) => Some("CreateForeignServerStmt"),
397                Some(NodeEnum::AlterForeignServerStmt(..)) => Some("AlterForeignServerStmt"),
398                Some(NodeEnum::CreateUserMappingStmt(..)) => Some("CreateUserMappingStmt"),
399                Some(NodeEnum::AlterUserMappingStmt(..)) => Some("AlterUserMappingStmt"),
400                Some(NodeEnum::DropUserMappingStmt(..)) => Some("DropUserMappingStmt"),
401                Some(NodeEnum::AlterTableSpaceOptionsStmt(..)) => {
402                    Some("AlterTableSpaceOptionsStmt")
403                }
404                Some(NodeEnum::AlterTableMoveAllStmt(..)) => Some("AlterTableMoveAllStmt"),
405                Some(NodeEnum::SecLabelStmt(..)) => Some("SecLabelStmt"),
406                Some(NodeEnum::CreateForeignTableStmt(..)) => Some("CreateForeignTableStmt"),
407                Some(NodeEnum::ImportForeignSchemaStmt(..)) => Some("ImportForeignSchemaStmt"),
408                Some(NodeEnum::CreateExtensionStmt(..)) => Some("CreateExtensionStmt"),
409                Some(NodeEnum::AlterExtensionStmt(..)) => Some("AlterExtensionStmt"),
410                Some(NodeEnum::AlterExtensionContentsStmt(..)) => {
411                    Some("AlterExtensionContentsStmt")
412                }
413                Some(NodeEnum::CreateEventTrigStmt(..)) => Some("CreateEventTrigStmt"),
414                Some(NodeEnum::AlterEventTrigStmt(..)) => Some("AlterEventTrigStmt"),
415                Some(NodeEnum::RefreshMatViewStmt(..)) => Some("RefreshMatViewStmt"),
416                Some(NodeEnum::ReplicaIdentityStmt(..)) => Some("ReplicaIdentityStmt"),
417                Some(NodeEnum::AlterSystemStmt(..)) => Some("AlterSystemStmt"),
418                Some(NodeEnum::CreatePolicyStmt(..)) => Some("CreatePolicyStmt"),
419                Some(NodeEnum::AlterPolicyStmt(..)) => Some("AlterPolicyStmt"),
420                Some(NodeEnum::CreateTransformStmt(..)) => Some("CreateTransformStmt"),
421                Some(NodeEnum::CreateAmStmt(..)) => Some("CreateAmStmt"),
422                Some(NodeEnum::CreatePublicationStmt(..)) => Some("CreatePublicationStmt"),
423                Some(NodeEnum::AlterPublicationStmt(..)) => Some("AlterPublicationStmt"),
424                Some(NodeEnum::CreateSubscriptionStmt(..)) => Some("CreateSubscriptionStmt"),
425                Some(NodeEnum::AlterSubscriptionStmt(..)) => Some("AlterSubscriptionStmt"),
426                Some(NodeEnum::DropSubscriptionStmt(..)) => Some("DropSubscriptionStmt"),
427                Some(NodeEnum::CreateStatsStmt(..)) => Some("CreateStatsStmt"),
428                Some(NodeEnum::AlterCollationStmt(..)) => Some("AlterCollationStmt"),
429                Some(NodeEnum::CallStmt(..)) => Some("CallStmt"),
430                Some(NodeEnum::AlterStatsStmt(..)) => Some("AlterStatsStmt"),
431                _ => None,
432            })
433            .collect()
434    }
435}