1use std::collections::BTreeMap;
4use std::ffi::{c_void, CStr, CString};
5
6use crate::bindings::*;
7use crate::{Error, Result};
8
9#[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#[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 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
144pub 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#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
160pub enum PlpgsqlCompileMode {
161 #[default]
163 Validate,
164 Runtime,
167}
168
169pub 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
179pub 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;