extern crate proc_macro;
use proc_macro::TokenStream;
use std::collections::HashSet;
use proc_macro2::Ident;
use quote::{quote, ToTokens};
use syn::spanned::Spanned;
use syn::{parse_macro_input, Attribute, Data, DeriveInput, Item, ItemImpl};
use operators::{impl_postgres_eq, impl_postgres_hash, impl_postgres_ord};
use pgx_sql_entity_graph::{
parse_extern_attributes, CodeEnrichment, ExtensionSql, ExtensionSqlFile, ExternArgs,
PgAggregate, PgExtern, PostgresEnum, PostgresType, Schema,
};
use crate::rewriter::PgGuardRewriter;
mod operators;
mod rewriter;
#[proc_macro_attribute]
pub fn pg_guard(_attr: TokenStream, item: TokenStream) -> TokenStream {
let ast = parse_macro_input!(item as syn::Item);
let rewriter = PgGuardRewriter::new();
let res = match ast {
Item::ForeignMod(block) => Ok(rewriter.extern_block(block)),
Item::Fn(func) => rewriter.item_fn_without_rewrite(func),
unknown => Err(syn::Error::new(
unknown.span(),
"#[pg_guard] can only be applied to extern \"C\" blocks and top-level functions",
)),
};
res.unwrap_or_else(|e| e.into_compile_error()).into()
}
#[proc_macro_attribute]
pub fn pg_test(attr: TokenStream, item: TokenStream) -> TokenStream {
let mut stream = proc_macro2::TokenStream::new();
let args = parse_extern_attributes(proc_macro2::TokenStream::from(attr.clone()));
let mut expected_error = None;
args.into_iter().for_each(|v| {
if let ExternArgs::Error(message) = v {
expected_error = Some(message)
}
});
let ast = parse_macro_input!(item as syn::Item);
match ast {
Item::Fn(mut func) => {
let mut test_attributes = Vec::new();
let mut non_test_attributes = Vec::new();
for attribute in func.attrs.iter() {
if let Some(ident) = attribute.path.get_ident() {
let ident_str = ident.to_string();
if ident_str == "ignore" || ident_str == "should_panic" {
test_attributes.push(attribute.clone());
} else {
non_test_attributes.push(attribute.clone());
}
} else {
non_test_attributes.push(attribute.clone());
}
}
func.attrs = non_test_attributes;
stream.extend(proc_macro2::TokenStream::from(pg_extern(
attr,
Item::Fn(func.clone()).to_token_stream().into(),
)));
let expected_error = match expected_error {
Some(msg) => quote! {Some(#msg)},
None => quote! {None},
};
let sql_funcname = func.sig.ident.to_string();
let test_func_name =
Ident::new(&format!("pg_{}", func.sig.ident.to_string()), func.span());
let attributes = func.attrs;
let mut att_stream = proc_macro2::TokenStream::new();
for a in attributes.iter() {
let as_str = a.tokens.to_string();
att_stream.extend(quote! {
options.push(#as_str);
});
}
stream.extend(quote! {
#[test]
#(#test_attributes)*
fn #test_func_name() {
let mut options = Vec::new();
#att_stream
crate::pg_test::setup(options);
let res = pgx_tests::run_test(#sql_funcname, #expected_error, crate::pg_test::postgresql_conf_options());
match res {
Ok(()) => (),
Err(e) => panic!("{:?}", e)
}
}
});
}
thing => {
return syn::Error::new(
thing.span(),
"#[pg_test] can only be applied to top-level functions",
)
.to_compile_error()
.into()
}
}
stream.into()
}
#[proc_macro_attribute]
pub fn initialize(_attr: TokenStream, item: TokenStream) -> TokenStream {
item
}
#[proc_macro_attribute]
pub fn pg_operator(attr: TokenStream, item: TokenStream) -> TokenStream {
pg_extern(attr, item)
}
#[proc_macro_attribute]
pub fn opname(_attr: TokenStream, item: TokenStream) -> TokenStream {
item
}
#[proc_macro_attribute]
pub fn commutator(_attr: TokenStream, item: TokenStream) -> TokenStream {
item
}
#[proc_macro_attribute]
pub fn negator(_attr: TokenStream, item: TokenStream) -> TokenStream {
item
}
#[proc_macro_attribute]
pub fn restrict(_attr: TokenStream, item: TokenStream) -> TokenStream {
item
}
#[proc_macro_attribute]
pub fn join(_attr: TokenStream, item: TokenStream) -> TokenStream {
item
}
#[proc_macro_attribute]
pub fn hashes(_attr: TokenStream, item: TokenStream) -> TokenStream {
item
}
#[proc_macro_attribute]
pub fn merges(_attr: TokenStream, item: TokenStream) -> TokenStream {
item
}
#[proc_macro_attribute]
pub fn pg_schema(_attr: TokenStream, input: TokenStream) -> TokenStream {
fn wrapped(input: TokenStream) -> Result<TokenStream, syn::Error> {
let pgx_schema: Schema = syn::parse(input)?;
Ok(pgx_schema.to_token_stream().into())
}
match wrapped(input) {
Ok(tokens) => tokens,
Err(e) => {
let msg = e.to_string();
TokenStream::from(quote! {
compile_error!(#msg);
})
}
}
}
#[proc_macro]
pub fn extension_sql(input: TokenStream) -> TokenStream {
fn wrapped(input: TokenStream) -> Result<TokenStream, syn::Error> {
let ext_sql: CodeEnrichment<ExtensionSql> = syn::parse(input)?;
Ok(ext_sql.to_token_stream().into())
}
match wrapped(input) {
Ok(tokens) => tokens,
Err(e) => {
let msg = e.to_string();
TokenStream::from(quote! {
compile_error!(#msg);
})
}
}
}
#[proc_macro]
pub fn extension_sql_file(input: TokenStream) -> TokenStream {
fn wrapped(input: TokenStream) -> Result<TokenStream, syn::Error> {
let ext_sql: CodeEnrichment<ExtensionSqlFile> = syn::parse(input)?;
Ok(ext_sql.to_token_stream().into())
}
match wrapped(input) {
Ok(tokens) => tokens,
Err(e) => {
let msg = e.to_string();
TokenStream::from(quote! {
compile_error!(#msg);
})
}
}
}
#[proc_macro_attribute]
pub fn search_path(_attr: TokenStream, item: TokenStream) -> TokenStream {
item
}
#[proc_macro_attribute]
pub fn pg_extern(attr: TokenStream, item: TokenStream) -> TokenStream {
fn wrapped(attr: TokenStream, item: TokenStream) -> Result<TokenStream, syn::Error> {
let pg_extern_item = PgExtern::new(attr.clone().into(), item.clone().into())?;
Ok(pg_extern_item.to_token_stream().into())
}
match wrapped(attr, item) {
Ok(tokens) => tokens,
Err(e) => {
let msg = e.to_string();
TokenStream::from(quote! {
compile_error!(#msg);
})
}
}
}
#[proc_macro_derive(PostgresEnum, attributes(requires, pgx))]
pub fn postgres_enum(input: TokenStream) -> TokenStream {
let ast = parse_macro_input!(input as syn::DeriveInput);
impl_postgres_enum(ast).unwrap_or_else(|e| e.to_compile_error()).into()
}
fn impl_postgres_enum(ast: DeriveInput) -> syn::Result<proc_macro2::TokenStream> {
let mut stream = proc_macro2::TokenStream::new();
let sql_graph_entity_ast = ast.clone();
let enum_ident = &ast.ident;
let enum_name = enum_ident.to_string();
let enum_data = match ast.data {
Data::Enum(e) => e,
_ => {
return Err(syn::Error::new(
ast.span(),
"#[derive(PostgresEnum)] can only be applied to enums",
))
}
};
let mut from_datum = proc_macro2::TokenStream::new();
let mut into_datum = proc_macro2::TokenStream::new();
for d in enum_data.variants.clone() {
let label_ident = &d.ident;
let label_string = label_ident.to_string();
from_datum.extend(quote! { #label_string => Some(#enum_ident::#label_ident), });
into_datum.extend(quote! { #enum_ident::#label_ident => Some(::pgx::enum_helper::lookup_enum_by_label(#enum_name, #label_string)), });
}
stream.extend(quote! {
impl ::pgx::datum::FromDatum for #enum_ident {
#[inline]
unsafe fn from_polymorphic_datum(datum: ::pgx::pg_sys::Datum, is_null: bool, typeoid: ::pgx::pg_sys::Oid) -> Option<#enum_ident> {
if is_null {
None
} else {
let (name, _, _) = ::pgx::enum_helper::lookup_enum_by_oid(unsafe { ::pgx::pg_sys::Oid::from_datum(datum, is_null)? } );
match name.as_str() {
#from_datum
_ => panic!("invalid enum value: {}", name)
}
}
}
}
impl ::pgx::datum::IntoDatum for #enum_ident {
#[inline]
fn into_datum(self) -> Option<::pgx::pg_sys::Datum> {
match self {
#into_datum
}
}
fn type_oid() -> ::pgx::pg_sys::Oid {
::pgx::wrappers::regtypein(#enum_name)
}
}
});
let sql_graph_entity_item = PostgresEnum::from_derive_input(sql_graph_entity_ast)?;
sql_graph_entity_item.to_tokens(&mut stream);
Ok(stream)
}
#[proc_macro_derive(PostgresType, attributes(inoutfuncs, pgvarlena_inoutfuncs, requires, pgx))]
pub fn postgres_type(input: TokenStream) -> TokenStream {
let ast = parse_macro_input!(input as syn::DeriveInput);
impl_postgres_type(ast).unwrap_or_else(|e| e.to_compile_error()).into()
}
fn impl_postgres_type(ast: DeriveInput) -> syn::Result<proc_macro2::TokenStream> {
let name = &ast.ident;
let generics = &ast.generics;
let has_lifetimes = generics.lifetimes().next();
let funcname_in = Ident::new(&format!("{}_in", name).to_lowercase(), name.span());
let funcname_out = Ident::new(&format!("{}_out", name).to_lowercase(), name.span());
let mut args = parse_postgres_type_args(&ast.attrs);
let mut stream = proc_macro2::TokenStream::new();
match ast.data {
Data::Struct(_) => { }
Data::Enum(_) => {
}
_ => {
return Err(syn::Error::new(
ast.span(),
"#[derive(PostgresType)] can only be applied to structs or enums",
))
}
}
if args.is_empty() {
args.insert(PostgresTypeAttribute::Default);
}
let lifetime = match has_lifetimes {
Some(lifetime) => quote! {#lifetime},
None => quote! {'static},
};
stream.extend(quote! {
impl #generics ::pgx::PostgresType for #name #generics { }
});
if args.contains(&PostgresTypeAttribute::Default) {
let inout_generics = if has_lifetimes.is_some() {
quote! {#generics}
} else {
quote! {<'_>}
};
stream.extend(quote! {
impl #generics ::pgx::inoutfuncs::JsonInOutFuncs #inout_generics for #name #generics {}
#[doc(hidden)]
#[::pgx::pgx_macros::pg_extern(immutable,parallel_safe)]
pub fn #funcname_in #generics(input: Option<&#lifetime ::core::ffi::CStr>) -> Option<#name #generics> {
input.map_or_else(|| {
for m in <#name as ::pgx::inoutfuncs::JsonInOutFuncs>::NULL_ERROR_MESSAGE {
::pgx::pg_sys::error!("{}", m);
}
None
}, |i| Some(<#name as ::pgx::inoutfuncs::JsonInOutFuncs>::input(i)))
}
#[doc(hidden)]
#[::pgx::pgx_macros::pg_extern(immutable,parallel_safe)]
pub fn #funcname_out #generics(input: #name #generics) -> &#lifetime ::core::ffi::CStr {
let mut buffer = ::pgx::stringinfo::StringInfo::new();
::pgx::inoutfuncs::JsonInOutFuncs::output(&input, &mut buffer);
buffer.into()
}
});
} else if args.contains(&PostgresTypeAttribute::InOutFuncs) {
stream.extend(quote! {
#[doc(hidden)]
#[::pgx::pgx_macros::pg_extern(immutable,parallel_safe)]
pub fn #funcname_in #generics(input: Option<&#lifetime ::core::ffi::CStr>) -> Option<#name #generics> {
input.map_or_else(|| {
for m in <#name as ::pgx::inoutfuncs::InOutFuncs>::NULL_ERROR_MESSAGE {
::pgx::pg_sys::error!("{}", m);
}
None
}, |i| Some(<#name as ::pgx::inoutfuncs::InOutFuncs>::input(i)))
}
#[doc(hidden)]
#[::pgx::pgx_macros::pg_extern(immutable,parallel_safe)]
pub fn #funcname_out #generics(input: #name #generics) -> &#lifetime ::core::ffi::CStr {
let mut buffer = ::pgx::stringinfo::StringInfo::new();
::pgx::inoutfuncs::InOutFuncs::output(&input, &mut buffer);
buffer.into()
}
});
} else if args.contains(&PostgresTypeAttribute::PgVarlenaInOutFuncs) {
stream.extend(quote! {
#[doc(hidden)]
#[::pgx::pgx_macros::pg_extern(immutable,parallel_safe)]
pub fn #funcname_in #generics(input: Option<&#lifetime ::core::ffi::CStr>) -> Option<::pgx::datum::PgVarlena<#name #generics>> {
input.map_or_else(|| {
for m in <#name as ::pgx::inoutfuncs::PgVarlenaInOutFuncs>::NULL_ERROR_MESSAGE {
::pgx::pg_sys::error!("{}", m);
}
None
}, |i| Some(<#name as ::pgx::inoutfuncs::PgVarlenaInOutFuncs>::input(i)))
}
#[doc(hidden)]
#[::pgx::pgx_macros::pg_extern(immutable,parallel_safe)]
pub fn #funcname_out #generics(input: ::pgx::datum::PgVarlena<#name #generics>) -> &#lifetime ::core::ffi::CStr {
let mut buffer = ::pgx::stringinfo::StringInfo::new();
::pgx::inoutfuncs::PgVarlenaInOutFuncs::output(&*input, &mut buffer);
buffer.into()
}
});
}
let sql_graph_entity_item = PostgresType::from_derive_input(ast)?;
sql_graph_entity_item.to_tokens(&mut stream);
Ok(stream)
}
#[proc_macro_derive(PostgresGucEnum, attributes(hidden))]
pub fn postgres_guc_enum(input: TokenStream) -> TokenStream {
let ast = parse_macro_input!(input as syn::DeriveInput);
impl_guc_enum(ast).unwrap_or_else(|e| e.to_compile_error()).into()
}
fn impl_guc_enum(ast: DeriveInput) -> syn::Result<proc_macro2::TokenStream> {
let mut stream = proc_macro2::TokenStream::new();
let enum_data = match ast.data {
Data::Enum(e) => e,
_ => {
return Err(syn::Error::new(
ast.span(),
"#[derive(PostgresGucEnum)] can only be applied to enums",
))
}
};
let enum_name = ast.ident;
let enum_len = enum_data.variants.len();
let mut from_match_arms = proc_macro2::TokenStream::new();
for (idx, e) in enum_data.variants.iter().enumerate() {
let label = &e.ident;
let idx = idx as i32;
from_match_arms.extend(quote! { #idx => #enum_name::#label, })
}
from_match_arms.extend(quote! { _ => panic!("Unrecognized ordinal ")});
let mut ordinal_match_arms = proc_macro2::TokenStream::new();
for (idx, e) in enum_data.variants.iter().enumerate() {
let label = &e.ident;
let idx = idx as i32;
ordinal_match_arms.extend(quote! { #enum_name::#label => #idx, });
}
let mut build_array_body = proc_macro2::TokenStream::new();
for (idx, e) in enum_data.variants.iter().enumerate() {
let label = e.ident.to_string();
let mut hidden = false;
for att in e.attrs.iter() {
let att = quote! {#att}.to_string();
if att == "# [hidden]" {
hidden = true;
break;
}
}
build_array_body.extend(quote! {
::pgx::pgbox::PgBox::<_, ::pgx::pgbox::AllocatedByPostgres>::with(&mut slice[#idx], |v| {
v.name = ::pgx::memcxt::PgMemoryContexts::TopMemoryContext.pstrdup(#label);
v.val = #idx as i32;
v.hidden = #hidden;
});
});
}
stream.extend(quote! {
impl ::pgx::guc::GucEnum<#enum_name> for #enum_name {
fn from_ordinal(ordinal: i32) -> #enum_name {
match ordinal {
#from_match_arms
}
}
fn to_ordinal(&self) -> i32 {
match *self {
#ordinal_match_arms
}
}
unsafe fn config_matrix(&self) -> *const ::pgx::pg_sys::config_enum_entry {
let slice = ::pgx::memcxt::PgMemoryContexts::TopMemoryContext.palloc0_slice::<::pgx::pg_sys::config_enum_entry>(#enum_len + 1usize);
#build_array_body
slice.as_ptr()
}
}
});
Ok(stream)
}
#[derive(Debug, Hash, Ord, PartialOrd, Eq, PartialEq)]
enum PostgresTypeAttribute {
InOutFuncs,
PgVarlenaInOutFuncs,
Default,
}
fn parse_postgres_type_args(attributes: &[Attribute]) -> HashSet<PostgresTypeAttribute> {
let mut categorized_attributes = HashSet::new();
for a in attributes {
let path = &a.path;
let path = quote! {#path}.to_string();
match path.as_str() {
"inoutfuncs" => {
categorized_attributes.insert(PostgresTypeAttribute::InOutFuncs);
}
"pgvarlena_inoutfuncs" => {
categorized_attributes.insert(PostgresTypeAttribute::PgVarlenaInOutFuncs);
}
_ => {
}
};
}
categorized_attributes
}
#[proc_macro_derive(PostgresEq, attributes(pgx))]
pub fn postgres_eq(input: TokenStream) -> TokenStream {
let ast = parse_macro_input!(input as syn::DeriveInput);
impl_postgres_eq(ast).unwrap_or_else(syn::Error::into_compile_error).into()
}
#[proc_macro_derive(PostgresOrd, attributes(pgx))]
pub fn postgres_ord(input: TokenStream) -> TokenStream {
let ast = parse_macro_input!(input as syn::DeriveInput);
impl_postgres_ord(ast).unwrap_or_else(syn::Error::into_compile_error).into()
}
#[proc_macro_derive(PostgresHash, attributes(pgx))]
pub fn postgres_hash(input: TokenStream) -> TokenStream {
let ast = parse_macro_input!(input as syn::DeriveInput);
impl_postgres_hash(ast).unwrap_or_else(syn::Error::into_compile_error).into()
}
#[proc_macro_attribute]
pub fn pg_aggregate(_attr: TokenStream, item: TokenStream) -> TokenStream {
fn wrapped(item_impl: ItemImpl) -> Result<TokenStream, syn::Error> {
let sql_graph_entity_item = PgAggregate::new(item_impl.into())?;
Ok(sql_graph_entity_item.to_token_stream().into())
}
let parsed_base = parse_macro_input!(item as syn::ItemImpl);
match wrapped(parsed_base) {
Ok(tokens) => tokens,
Err(e) => {
let msg = e.to_string();
TokenStream::from(quote! {
compile_error!(#msg);
})
}
}
}
#[proc_macro_attribute]
pub fn pgx(_attr: TokenStream, item: TokenStream) -> TokenStream {
item
}
#[proc_macro_attribute]
pub fn pg_trigger(attrs: TokenStream, input: TokenStream) -> TokenStream {
fn wrapped(attrs: TokenStream, input: TokenStream) -> Result<TokenStream, syn::Error> {
use pgx_sql_entity_graph::{PgTrigger, PgTriggerAttribute};
use syn::parse::Parser;
use syn::punctuated::Punctuated;
use syn::Token;
let attributes =
Punctuated::<PgTriggerAttribute, Token![,]>::parse_terminated.parse(attrs)?;
let item_fn: syn::ItemFn = syn::parse(input)?;
let trigger_item = PgTrigger::new(item_fn, attributes)?;
let trigger_tokens = trigger_item.to_token_stream();
Ok(trigger_tokens.into())
}
match wrapped(attrs, input) {
Ok(tokens) => tokens,
Err(e) => {
let msg = e.to_string();
TokenStream::from(quote! {
compile_error!(#msg);
})
}
}
}