extern crate proc_macro;
mod path;
mod query;
mod schema;
mod validation;
use darling::FromDeriveInput;
use darling::FromField;
use darling::FromVariant;
use darling::util::SpannedValue;
use proc_macro::TokenStream;
use proc_macro2::TokenStream as TokenStream2;
use quote::quote;
use syn::DeriveInput;
use syn::parse_macro_input;
#[derive(Debug, FromDeriveInput)]
#[darling(attributes(response), supports(struct_named, enum_newtype))]
struct ResponseInput {
ident: syn::Ident,
generics: syn::Generics,
data: darling::ast::Data<ResponseVariant, ResponseField>,
#[darling(default)]
schema: Option<String>,
#[darling(default)]
root_type: Option<SpannedValue<String>>,
}
#[derive(Debug, FromField)]
#[darling(attributes(field))]
struct ResponseField {
ident: Option<syn::Ident>,
ty: syn::Type,
path: SpannedValue<String>,
#[darling(default)]
skip_schema_validation: bool,
}
#[derive(Debug, FromField)]
struct VariantInner {
ty: syn::Type,
}
#[derive(Debug, FromVariant)]
#[darling(attributes(response))]
struct ResponseVariant {
ident: syn::Ident,
fields: darling::ast::Fields<VariantInner>,
#[darling(default)]
on: Option<SpannedValue<String>>,
}
#[proc_macro_derive(Response, attributes(response, field))]
pub fn derive_query_response(input: TokenStream) -> TokenStream {
let input = parse_macro_input!(input as DeriveInput);
match derive_query_response_impl(input) {
Ok(tokens) => tokens.into(),
Err(err) => err.to_compile_error().into(),
}
}
#[proc_macro]
pub fn graphql_query(input: TokenStream) -> TokenStream {
query::expand(input)
}
fn derive_query_response_impl(input: DeriveInput) -> Result<TokenStream2, syn::Error> {
let parsed = ResponseInput::from_derive_input(&input)?;
let loaded_schema = if let Some(path) = &parsed.schema {
let base_dir = std::env::var("SUI_GRAPHQL_SCHEMA_DIR")
.or_else(|_| std::env::var("CARGO_MANIFEST_DIR"))
.unwrap();
let full_path = std::path::Path::new(&base_dir).join(path);
let sdl = std::fs::read_to_string(&full_path).map_err(|e| {
syn::Error::new(
proc_macro2::Span::call_site(),
format!(
"Failed to read schema from '{}': {}",
full_path.display(),
e
),
)
})?;
Some(schema::Schema::from_sdl(&sdl)?)
} else {
None
};
let schema = if let Some(schema) = &loaded_schema {
schema
} else {
schema::Schema::load()?
};
let root_type = parsed
.root_type
.as_ref()
.map(|s| s.as_str())
.unwrap_or("Query");
if !schema.has_type(root_type) {
use std::fmt::Write;
let type_names = schema.type_names();
let suggestion = validation::find_similar(&type_names, root_type);
let mut msg = format!("Type '{}' not found in GraphQL schema", root_type);
if let Some(suggested) = suggestion {
write!(msg, ". Did you mean '{}'?", suggested).unwrap();
}
let span = parsed.root_type.as_ref().unwrap().span();
return Err(syn::Error::new(span, msg));
}
match parsed.data {
darling::ast::Data::Struct(ref fields) => {
generate_struct_impl(&parsed, &fields.fields, schema, root_type)
}
darling::ast::Data::Enum(ref variants) => {
generate_enum_impl(&parsed, variants, schema, root_type)
}
}
}
fn generate_struct_impl(
input: &ResponseInput,
fields: &[ResponseField],
schema: &schema::Schema,
root_type: &str,
) -> Result<TokenStream2, syn::Error> {
let ident = &input.ident;
let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl();
let mut field_extractions = Vec::new();
let mut field_names = Vec::new();
for field in fields {
let field_ident = field
.ident
.as_ref()
.expect("darling ensures named fields only");
let spanned_path = &field.path;
let parsed_path = path::ParsedPath::parse(spanned_path.as_str())
.map_err(|e| syn::Error::new(spanned_path.span(), e.to_string()))?;
let terminal_type = if !field.skip_schema_validation {
Some(validation::validate_path_against_schema(
schema,
root_type,
&parsed_path,
spanned_path.span(),
)?)
} else {
None
};
let skip_vec_excess_check = field.skip_schema_validation
|| terminal_type.is_some_and(validation::is_object_like_scalar);
validation::validate_type_matches_path(&parsed_path, &field.ty, skip_vec_excess_check)?;
let type_structure = validation::analyze_type(&field.ty);
let extraction = generate_field_extraction(&parsed_path, &type_structure, field_ident);
field_extractions.push(extraction);
field_names.push(field_ident);
}
Ok(quote! {
impl #impl_generics #ident #ty_generics #where_clause {
pub fn from_value(value: serde_json::Value) -> Result<Self, String> {
#(#field_extractions)*
Ok(Self {
#(#field_names),*
})
}
}
impl<'de> serde::Deserialize<'de> for #ident #ty_generics #where_clause {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
let value = serde_json::Value::deserialize(deserializer)?;
Self::from_value(value).map_err(serde::de::Error::custom)
}
}
})
}
fn generate_enum_impl(
input: &ResponseInput,
variants: &[ResponseVariant],
schema: &schema::Schema,
root_type: &str,
) -> Result<TokenStream2, syn::Error> {
let ident = &input.ident;
let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl();
let root_type_span = input
.root_type
.as_ref()
.map(|s| s.span())
.unwrap_or_else(|| ident.span());
if !schema.is_union(root_type) {
return Err(syn::Error::new(
root_type_span,
format!(
"'{}' is not a union type. \
Enum Response requires root_type to be a GraphQL union",
root_type
),
));
}
let mut match_arms = Vec::new();
for variant in variants {
let variant_ident = &variant.ident;
let graphql_typename = variant
.on
.as_ref()
.map(|s| s.as_str().to_string())
.unwrap_or_else(|| variant_ident.to_string());
let span = variant
.on
.as_ref()
.map(|s| s.span())
.unwrap_or_else(|| variant_ident.span());
if let Err(mut err) =
validation::validate_union_member(schema, root_type, &graphql_typename, span)
{
if variant.on.is_none() {
err.combine(syn::Error::new(
span,
"hint: use #[response(on = \"...\")] to specify a GraphQL type name different from the variant name",
));
}
return Err(err);
}
let inner_ty = &variant.fields.fields[0].ty;
match_arms.push(quote! {
#graphql_typename => {
Ok(Self::#variant_ident(
<#inner_ty>::from_value(value)?
))
}
});
}
let root_type_str = root_type;
let enum_name_str = ident.to_string();
Ok(quote! {
impl #impl_generics #ident #ty_generics #where_clause {
pub fn from_value(value: serde_json::Value) -> Result<Self, String> {
let typename = value.get("__typename")
.and_then(|v| v.as_str())
.ok_or_else(|| format!(
"union '{}' requires '__typename' in the response to distinguish variants. \
Make sure your query requests '__typename' on this field ({})",
#root_type_str, #enum_name_str
))?;
match typename {
#(#match_arms)*
other => Err(format!(
"unknown __typename '{}' for union '{}' ({})",
other, #root_type_str, #enum_name_str
)),
}
}
}
impl<'de> serde::Deserialize<'de> for #ident #ty_generics #where_clause {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
let value = serde_json::Value::deserialize(deserializer)?;
Self::from_value(value).map_err(serde::de::Error::custom)
}
}
})
}
fn generate_field_extraction(
path: &path::ParsedPath,
type_structure: &validation::TypeStructure,
field_ident: &syn::Ident,
) -> TokenStream2 {
let full_path = &path.raw;
let inner = generate_from_segments(full_path, &path.segments, type_structure);
quote! {
let #field_ident = {
let current = &value;
#inner?
};
}
}
fn generate_from_segments(
full_path: &str,
segments: &[path::PathSegment],
type_structure: &validation::TypeStructure,
) -> TokenStream2 {
let (is_optional, inner_type) = match type_structure {
validation::TypeStructure::Optional(inner) => (true, inner.as_ref()),
other => (false, other),
};
let core = generate_from_segments_core(full_path, segments, inner_type);
if is_optional {
quote! {
(|| {
if current.is_null() { return Ok(None) }
#core.map(Some)
})()
}
} else {
core
}
}
fn generate_from_segments_core(
full_path: &str,
segments: &[path::PathSegment],
type_structure: &validation::TypeStructure,
) -> TokenStream2 {
let Some((segment, rest)) = segments.split_first() else {
return quote! {
serde_json::from_value(current.clone())
.map_err(|e| format!("failed to deserialize '{}': {}", #full_path, e))
};
};
let name = segment.field;
let json_key = segment.json_key();
let on_null = if segment.is_nullable {
quote! { return Ok(None) }
} else {
quote! {
return Err(format!("null value at '{}' in path '{}'", #name, #full_path))
}
};
if segment.is_list() {
let element_type = match type_structure {
validation::TypeStructure::Vector(inner) => inner.as_ref(),
_ => unreachable!("validated: list segment requires Vec type"),
};
let rest_code = generate_from_segments(full_path, rest, element_type);
quote! {
let field_value = current.get(#json_key).unwrap_or(&serde_json::Value::Null);
if field_value.is_null() {
#on_null
}
let array = field_value.as_array()
.ok_or_else(|| format!("expected array at '{}' in path '{}'", #json_key, #full_path))?;
array.iter()
.map(|current| { #rest_code })
.collect::<Result<Vec<_>, String>>()
}
} else {
let rest_code = generate_from_segments_core(full_path, rest, type_structure);
quote! {
let current = current.get(#json_key).unwrap_or(&serde_json::Value::Null);
if current.is_null() {
#on_null
}
#rest_code
}
}
}