use std::collections::BTreeMap;
use std::ffi::{c_void, CStr, CString};
use crate::bindings::*;
use crate::{Error, Result};
#[derive(Debug, Clone)]
pub struct PlpgsqlType {
pub oid: u32,
pub namespace_oid: u32,
pub name: String,
pub length: i16,
pub by_value: bool,
pub type_kind: u8,
pub category: u8,
pub alignment: u8,
pub storage: u8,
pub array_oid: u32,
pub element_oid: u32,
pub base_type_oid: u32,
pub collation_oid: u32,
pub subscript_handler_oid: u32,
}
#[derive(Debug, Clone, Default)]
pub struct PlpgsqlCatalog {
pub namespaces: BTreeMap<String, u32>,
pub search_path: Vec<String>,
pub types: Vec<PlpgsqlType>,
}
struct CatalogContext<'a> {
catalog: &'a PlpgsqlCatalog,
names: Vec<CString>,
}
impl CatalogContext<'_> {
fn write_type(&self, index: usize, output: *mut PgQueryPlpgsqlTypeMetadata) {
let ty = &self.catalog.types[index];
let value = PgQueryPlpgsqlTypeMetadata {
oid: ty.oid,
namespace_oid: ty.namespace_oid,
name: self.names[index].as_ptr(),
length: ty.length,
by_value: ty.by_value,
type_kind: ty.type_kind as _,
category: ty.category as _,
alignment: ty.alignment as _,
storage: ty.storage as _,
array_oid: ty.array_oid,
element_oid: ty.element_oid,
base_type_oid: ty.base_type_oid,
collation_oid: ty.collation_oid,
subscript_handler_oid: ty.subscript_handler_oid,
};
unsafe { output.write(value) };
}
}
unsafe extern "C" fn lookup_namespace(
context: *mut c_void,
name: *const std::os::raw::c_char,
output: *mut u32,
) -> PgQueryCatalogLookupResult {
let context = &*(context as *const CatalogContext<'_>);
let name = CStr::from_ptr(name).to_bytes();
match context
.catalog
.namespaces
.iter()
.find(|(candidate, _)| candidate.as_bytes() == name)
{
Some((_, oid)) => {
output.write(*oid);
PgQueryCatalogLookupResult_PG_QUERY_CATALOG_LOOKUP_FOUND
}
None => PgQueryCatalogLookupResult_PG_QUERY_CATALOG_LOOKUP_NOT_FOUND,
}
}
unsafe extern "C" fn lookup_type_by_name(
context: *mut c_void,
schema: *const std::os::raw::c_char,
name: *const std::os::raw::c_char,
output: *mut PgQueryPlpgsqlTypeMetadata,
) -> PgQueryCatalogLookupResult {
let context = &*(context as *const CatalogContext<'_>);
let name = CStr::from_ptr(name).to_bytes();
let in_schema = |schema: &[u8]| {
let namespace = context
.catalog
.namespaces
.iter()
.find(|(candidate, _)| candidate.as_bytes() == schema)
.map(|(_, oid)| *oid)?;
context
.catalog
.types
.iter()
.position(|ty| ty.namespace_oid == namespace && ty.name.as_bytes() == name)
};
let found = if schema.is_null() {
context
.catalog
.search_path
.iter()
.find_map(|schema| in_schema(schema.as_bytes()))
} else {
in_schema(CStr::from_ptr(schema).to_bytes())
};
match found {
Some(index) => {
context.write_type(index, output);
PgQueryCatalogLookupResult_PG_QUERY_CATALOG_LOOKUP_FOUND
}
None => PgQueryCatalogLookupResult_PG_QUERY_CATALOG_LOOKUP_NOT_FOUND,
}
}
unsafe extern "C" fn lookup_type_by_oid(
context: *mut c_void,
oid: u32,
output: *mut PgQueryPlpgsqlTypeMetadata,
) -> PgQueryCatalogLookupResult {
let context = &*(context as *const CatalogContext<'_>);
match context.catalog.types.iter().position(|ty| ty.oid == oid) {
Some(index) => {
context.write_type(index, output);
PgQueryCatalogLookupResult_PG_QUERY_CATALOG_LOOKUP_FOUND
}
None => PgQueryCatalogLookupResult_PG_QUERY_CATALOG_LOOKUP_NOT_FOUND,
}
}
unsafe extern "C" fn catalog_error(_context: *mut c_void) -> *const std::os::raw::c_char {
std::ptr::null()
}
pub fn parse_plpgsql_with_catalog(
stmt: &str,
catalog: &PlpgsqlCatalog,
) -> Result<serde_json::Value> {
parse_plpgsql_with_options(stmt, Some(catalog), crate::ParseOptions::default())
.result
.map_err(|error| match error {
Error::ParseDiagnostic(diagnostic) => Error::Parse(diagnostic.message),
other => other,
})
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum PlpgsqlCompileMode {
#[default]
Validate,
Runtime,
}
pub fn parse_plpgsql_with_options(
stmt: &str,
catalog: Option<&PlpgsqlCatalog>,
options: crate::ParseOptions,
) -> crate::ParseOutcome<serde_json::Value> {
parse_plpgsql_with_mode(stmt, catalog, options, PlpgsqlCompileMode::Validate)
}
pub fn parse_plpgsql_with_mode(
stmt: &str,
catalog: Option<&PlpgsqlCatalog>,
options: crate::ParseOptions,
mode: PlpgsqlCompileMode,
) -> crate::ParseOutcome<serde_json::Value> {
use crate::parse_options::{capture_diagnostic, parse_error};
use crate::{Diagnostic, ParseOutcome};
let input = match CString::new(stmt) {
Ok(input) => input,
Err(error) => return ParseOutcome::error(error.into()),
};
let mut context = match catalog
.map(|catalog| {
catalog
.types
.iter()
.map(|ty| CString::new(ty.name.as_str()))
.collect::<std::result::Result<Vec<_>, _>>()
.map(|names| CatalogContext { catalog, names })
})
.transpose()
{
Ok(context) => context,
Err(error) => return ParseOutcome::error(error.into()),
};
let callbacks = context.as_mut().map(|context| PgQueryPlpgsqlCatalog {
context: (context as *mut CatalogContext<'_>).cast(),
lookup_namespace: Some(lookup_namespace),
lookup_type_by_name: Some(lookup_type_by_name),
lookup_type_by_oid: Some(lookup_type_by_oid),
get_error: Some(catalog_error),
});
let mut diagnostics = Vec::<Diagnostic>::new();
let result = unsafe {
pg_query_parse_plpgsql_with_options(
input.as_ptr(),
callbacks
.as_ref()
.map_or(std::ptr::null(), |callbacks| callbacks),
options.bits()
| match mode {
PlpgsqlCompileMode::Validate => 0,
PlpgsqlCompileMode::Runtime => PG_QUERY_PLPGSQL_RUNTIME as i32,
},
Some(capture_diagnostic),
(&mut diagnostics as *mut Vec<Diagnostic>).cast(),
)
};
let structure = if !result.error.is_null() {
Err(unsafe { parse_error(result.error, &diagnostics) })
} else if result.plpgsql_funcs.is_null() {
Err(Error::InvalidPointer)
} else {
let raw = unsafe { CStr::from_ptr(result.plpgsql_funcs) };
serde_json::from_str(&raw.to_string_lossy())
.map_err(|error| Error::InvalidJson(error.to_string()))
};
unsafe { pg_query_free_plpgsql_parse_result(result) };
ParseOutcome {
result: structure,
diagnostics,
}
}
#[cfg(test)]
mod first_use_tests;