Skip to main content

pg_query/
plpgsql_catalog.rs

1//! Catalog snapshots for PostgreSQL's PL/pgSQL type lookup callbacks.
2
3use std::collections::BTreeMap;
4use std::ffi::{c_void, CStr, CString};
5
6use crate::bindings::*;
7use crate::{Error, Result};
8
9/// The pg_type attributes needed to distinguish scalar, domain, array, and composite declarations.
10#[derive(Debug, Clone)]
11pub struct PlpgsqlType {
12    pub oid: u32,
13    pub namespace_oid: u32,
14    pub name: String,
15    pub length: i16,
16    pub by_value: bool,
17    pub type_kind: u8,
18    pub category: u8,
19    pub alignment: u8,
20    pub storage: u8,
21    pub array_oid: u32,
22    pub element_oid: u32,
23    pub base_type_oid: u32,
24    pub collation_oid: u32,
25    pub subscript_handler_oid: u32,
26}
27
28/// An immutable catalog snapshot for one parse. `search_path` is the caller's effective namespace order, including pg_catalog where appropriate. Built-in types remain available through PostgreSQL's own fallback catalog.
29#[derive(Debug, Clone, Default)]
30pub struct PlpgsqlCatalog {
31    pub namespaces: BTreeMap<String, u32>,
32    pub search_path: Vec<String>,
33    pub types: Vec<PlpgsqlType>,
34}
35
36struct CatalogContext<'a> {
37    catalog: &'a PlpgsqlCatalog,
38    names: Vec<CString>,
39}
40
41impl CatalogContext<'_> {
42    fn write_type(&self, index: usize, output: *mut PgQueryPlpgsqlTypeMetadata) {
43        let ty = &self.catalog.types[index];
44        let value = PgQueryPlpgsqlTypeMetadata {
45            oid: ty.oid,
46            namespace_oid: ty.namespace_oid,
47            name: self.names[index].as_ptr(),
48            length: ty.length,
49            by_value: ty.by_value,
50            type_kind: ty.type_kind as _,
51            category: ty.category as _,
52            alignment: ty.alignment as _,
53            storage: ty.storage as _,
54            array_oid: ty.array_oid,
55            element_oid: ty.element_oid,
56            base_type_oid: ty.base_type_oid,
57            collation_oid: ty.collation_oid,
58            subscript_handler_oid: ty.subscript_handler_oid,
59        };
60        // The C parser supplies a valid output pointer and copies this metadata before the next callback.
61        unsafe { output.write(value) };
62    }
63}
64
65unsafe extern "C" fn lookup_namespace(
66    context: *mut c_void,
67    name: *const std::os::raw::c_char,
68    output: *mut u32,
69) -> PgQueryCatalogLookupResult {
70    let context = &*(context as *const CatalogContext<'_>);
71    let name = CStr::from_ptr(name).to_bytes();
72    match context
73        .catalog
74        .namespaces
75        .iter()
76        .find(|(candidate, _)| candidate.as_bytes() == name)
77    {
78        Some((_, oid)) => {
79            output.write(*oid);
80            PgQueryCatalogLookupResult_PG_QUERY_CATALOG_LOOKUP_FOUND
81        }
82        None => PgQueryCatalogLookupResult_PG_QUERY_CATALOG_LOOKUP_NOT_FOUND,
83    }
84}
85
86unsafe extern "C" fn lookup_type_by_name(
87    context: *mut c_void,
88    schema: *const std::os::raw::c_char,
89    name: *const std::os::raw::c_char,
90    output: *mut PgQueryPlpgsqlTypeMetadata,
91) -> PgQueryCatalogLookupResult {
92    let context = &*(context as *const CatalogContext<'_>);
93    let name = CStr::from_ptr(name).to_bytes();
94    let in_schema = |schema: &[u8]| {
95        let namespace = context
96            .catalog
97            .namespaces
98            .iter()
99            .find(|(candidate, _)| candidate.as_bytes() == schema)
100            .map(|(_, oid)| *oid)?;
101        context
102            .catalog
103            .types
104            .iter()
105            .position(|ty| ty.namespace_oid == namespace && ty.name.as_bytes() == name)
106    };
107    let found = if schema.is_null() {
108        context
109            .catalog
110            .search_path
111            .iter()
112            .find_map(|schema| in_schema(schema.as_bytes()))
113    } else {
114        in_schema(CStr::from_ptr(schema).to_bytes())
115    };
116    match found {
117        Some(index) => {
118            context.write_type(index, output);
119            PgQueryCatalogLookupResult_PG_QUERY_CATALOG_LOOKUP_FOUND
120        }
121        None => PgQueryCatalogLookupResult_PG_QUERY_CATALOG_LOOKUP_NOT_FOUND,
122    }
123}
124
125unsafe extern "C" fn lookup_type_by_oid(
126    context: *mut c_void,
127    oid: u32,
128    output: *mut PgQueryPlpgsqlTypeMetadata,
129) -> PgQueryCatalogLookupResult {
130    let context = &*(context as *const CatalogContext<'_>);
131    match context.catalog.types.iter().position(|ty| ty.oid == oid) {
132        Some(index) => {
133            context.write_type(index, output);
134            PgQueryCatalogLookupResult_PG_QUERY_CATALOG_LOOKUP_FOUND
135        }
136        None => PgQueryCatalogLookupResult_PG_QUERY_CATALOG_LOOKUP_NOT_FOUND,
137    }
138}
139
140unsafe extern "C" fn catalog_error(_context: *mut c_void) -> *const std::os::raw::c_char {
141    std::ptr::null()
142}
143
144/// Parse PL/pgSQL with the caller's catalog types instead of treating every unknown type as a record. The snapshot and callback storage remain local to this synchronous parse; no pointers or callback state escape it.
145pub fn parse_plpgsql_with_catalog(
146    stmt: &str,
147    catalog: &PlpgsqlCatalog,
148) -> Result<serde_json::Value> {
149    parse_plpgsql_with_options(stmt, Some(catalog), crate::ParseOptions::default())
150        .result
151        .map_err(|error| match error {
152            Error::ParseDiagnostic(diagnostic) => Error::Parse(diagnostic.message),
153            other => other,
154        })
155}
156
157/// Whether a PL/pgSQL structure is checked at definition or compiled for execution.
158/// Runtime compilation keeps embedded SQL text for its first reached preparation.
159#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
160pub enum PlpgsqlCompileMode {
161    /// Validate a declaration, using representative types for polymorphic arguments.
162    #[default]
163    Validate,
164    /// Compile a caller-specialized declaration without checking unreached SQL.
165    /// Argument and result types must already be concrete.
166    Runtime,
167}
168
169/// Parse PL/pgSQL with a synchronous catalog snapshot, per-call scanner settings
170/// and structured diagnostics. A missing catalog uses PostgreSQL's builtin types.
171pub fn parse_plpgsql_with_options(
172    stmt: &str,
173    catalog: Option<&PlpgsqlCatalog>,
174    options: crate::ParseOptions,
175) -> crate::ParseOutcome<serde_json::Value> {
176    parse_plpgsql_with_mode(stmt, catalog, options, PlpgsqlCompileMode::Validate)
177}
178
179/// Compile a PL/pgSQL structure with the caller-selected validator/runtime boundary.
180/// This changes only PostgreSQL's own validator checks, not warning delivery.
181pub fn parse_plpgsql_with_mode(
182    stmt: &str,
183    catalog: Option<&PlpgsqlCatalog>,
184    options: crate::ParseOptions,
185    mode: PlpgsqlCompileMode,
186) -> crate::ParseOutcome<serde_json::Value> {
187    use crate::parse_options::{capture_diagnostic, parse_error};
188    use crate::{Diagnostic, ParseOutcome};
189    let input = match CString::new(stmt) {
190        Ok(input) => input,
191        Err(error) => return ParseOutcome::error(error.into()),
192    };
193    let mut context = match catalog
194        .map(|catalog| {
195            catalog
196                .types
197                .iter()
198                .map(|ty| CString::new(ty.name.as_str()))
199                .collect::<std::result::Result<Vec<_>, _>>()
200                .map(|names| CatalogContext { catalog, names })
201        })
202        .transpose()
203    {
204        Ok(context) => context,
205        Err(error) => return ParseOutcome::error(error.into()),
206    };
207    let callbacks = context.as_mut().map(|context| PgQueryPlpgsqlCatalog {
208        context: (context as *mut CatalogContext<'_>).cast(),
209        lookup_namespace: Some(lookup_namespace),
210        lookup_type_by_name: Some(lookup_type_by_name),
211        lookup_type_by_oid: Some(lookup_type_by_oid),
212        get_error: Some(catalog_error),
213    });
214    let mut diagnostics = Vec::<Diagnostic>::new();
215    let result = unsafe {
216        pg_query_parse_plpgsql_with_options(
217            input.as_ptr(),
218            callbacks
219                .as_ref()
220                .map_or(std::ptr::null(), |callbacks| callbacks),
221            options.bits()
222                | match mode {
223                    PlpgsqlCompileMode::Validate => 0,
224                    PlpgsqlCompileMode::Runtime => PG_QUERY_PLPGSQL_RUNTIME as i32,
225                },
226            Some(capture_diagnostic),
227            (&mut diagnostics as *mut Vec<Diagnostic>).cast(),
228        )
229    };
230    let structure = if !result.error.is_null() {
231        Err(unsafe { parse_error(result.error, &diagnostics) })
232    } else if result.plpgsql_funcs.is_null() {
233        Err(Error::InvalidPointer)
234    } else {
235        let raw = unsafe { CStr::from_ptr(result.plpgsql_funcs) };
236        serde_json::from_str(&raw.to_string_lossy())
237            .map_err(|error| Error::InvalidJson(error.to_string()))
238    };
239    unsafe { pg_query_free_plpgsql_parse_result(result) };
240    ParseOutcome {
241        result: structure,
242        diagnostics,
243    }
244}
245
246#[cfg(test)]
247mod first_use_tests;