use proc_macro::TokenStream;
use proc_macro2::TokenStream as TokenStream2;
use quote::{format_ident, quote};
use syn::{
parse_macro_input, spanned::Spanned, Error, FnArg, Ident, ImplItem, ItemImpl, Pat, Type,
};
#[proc_macro_attribute]
pub fn handlers(_attr: TokenStream, item: TokenStream) -> TokenStream {
let mut input = parse_macro_input!(item as ItemImpl);
let mut req_structs: Vec<TokenStream2> = Vec::new();
let mut reg_fns: Vec<TokenStream2> = Vec::new();
for impl_item in &mut input.items {
let ImplItem::Fn(method) = impl_item else {
continue;
};
let Some(attr_idx) = method
.attrs
.iter()
.position(|a| a.path().is_ident("handler"))
else {
continue;
};
let attr = method.attrs.remove(attr_idx);
let args = match parse_handler_args(&attr) {
Ok(v) => v,
Err(e) => return e.to_compile_error().into(),
};
match build_handler(method, args) {
Ok(Built { req_struct, reg_fn }) => {
if let Some(req) = req_struct {
req_structs.push(req);
}
reg_fns.push(reg_fn);
}
Err(e) => return e.to_compile_error().into(),
}
}
let self_ty = &input.self_ty;
quote! {
#input
#(#req_structs)*
impl #self_ty {
#(#reg_fns)*
}
}
.into()
}
struct Built {
req_struct: Option<TokenStream2>,
reg_fn: TokenStream2,
}
#[derive(Clone, Copy, PartialEq)]
enum ReturnMode {
Default,
Graph,
GraphNode,
Bytes,
OkError,
}
struct HandlerArgs {
returns: ReturnMode,
mode: Option<bool>, }
fn parse_handler_args(attr: &syn::Attribute) -> syn::Result<HandlerArgs> {
let mut out = HandlerArgs {
returns: ReturnMode::Default,
mode: None,
};
match &attr.meta {
syn::Meta::Path(_) => Ok(out),
syn::Meta::List(_) => {
attr.parse_nested_meta(|meta| {
if meta.path.is_ident("returns") {
let value = meta.value()?;
let ident: Ident = value.parse()?;
out.returns =
match ident.to_string().as_str() {
"graph" => ReturnMode::Graph,
"graph_node" => ReturnMode::GraphNode,
"bytes" => ReturnMode::Bytes,
"ok_error" => ReturnMode::OkError,
_ => return Err(meta.error(
"unknown `returns` mode (expected `graph`, `graph_node`, `bytes`, or `ok_error`)",
)),
};
Ok(())
} else if meta.path.is_ident("send") {
out.mode = Some(true);
Ok(())
} else if meta.path.is_ident("post") {
out.mode = Some(false);
Ok(())
} else {
Err(meta.error("unknown `#[handler]` option"))
}
})?;
Ok(out)
}
syn::Meta::NameValue(_) => Err(Error::new(
attr.span(),
"`#[handler]` takes no `= value` form",
)),
}
}
fn result_ok_type(ty: &Type) -> Option<&Type> {
let Type::Path(p) = ty else { return None };
let seg = p.path.segments.last()?;
if seg.ident != "Result" {
return None;
}
let syn::PathArguments::AngleBracketed(args) = &seg.arguments else {
return None;
};
args.args.iter().find_map(|a| match a {
syn::GenericArgument::Type(t) => Some(t),
_ => None,
})
}
fn returns_unit(ret: &syn::ReturnType) -> bool {
match ret {
syn::ReturnType::Default => true,
syn::ReturnType::Type(_, ty) => matches!(&**ty, Type::Tuple(t) if t.elems.is_empty()),
}
}
struct Param {
call_expr: TokenStream2,
field: Option<(Ident, TokenStream2)>,
}
fn build_handler(method: &syn::ImplItemFn, args: HandlerArgs) -> syn::Result<Built> {
let returns = args.returns;
let name = &method.sig.ident;
let kind = name.to_string();
let mut params: Vec<Param> = Vec::new();
for input in &method.sig.inputs {
let FnArg::Typed(pat_ty) = input else {
continue; };
let Pat::Ident(pat_ident) = &*pat_ty.pat else {
return Err(Error::new(
pat_ty.pat.span(),
"#[handler] parameters must be plain identifiers",
));
};
let ident = pat_ident.ident.clone();
params.push(parse_param(ident, &pat_ty.ty)?);
}
let fields: Vec<&(Ident, TokenStream2)> =
params.iter().filter_map(|p| p.field.as_ref()).collect();
let req_ident = format_ident!("{}Req", pascal_case(&kind));
let req_struct = if fields.is_empty() {
None
} else {
let field_defs = fields.iter().map(|(id, ty)| quote!(#id: #ty));
Some(quote! {
#[derive(serde::Deserialize)]
#[cfg_attr(feature = "ts-export", derive(ts_rs::TS))]
pub(crate) struct #req_ident {
#(#field_defs,)*
}
})
};
let payload_ident = if fields.is_empty() {
format_ident!("_payload")
} else {
format_ident!("payload")
};
let uses_bytes = params.iter().any(|p| p.field.is_none());
let bytes_ident = if uses_bytes {
format_ident!("bytes")
} else {
format_ident!("_bytes")
};
let decode = if fields.is_empty() {
quote!()
} else {
quote!(let req: #req_ident = crate::engine::protocol::decode(payload)?;)
};
let call_args = params.iter().map(|p| &p.call_expr);
let call = quote!(let __result = engine.#name(#(#call_args),*););
let convert = match returns {
ReturnMode::Graph => quote!(crate::engine::protocol::graph_result(__result)),
ReturnMode::GraphNode => quote!(crate::engine::protocol::graph_node_result(__result)),
ReturnMode::Bytes => quote!(crate::engine::protocol::bytes_result(__result)),
ReturnMode::OkError => quote!(Ok(crate::engine::protocol::ok_or_error(__result))),
ReturnMode::Default => quote! {{
use crate::engine::protocol::{JsonResponseKind as _, ResultResponseKind as _};
let __tag = (&__result).response_kind();
__tag.into_response(__result)
}},
};
let req_call = if fields.is_empty() {
quote!()
} else {
quote!(.req::<#req_ident>())
};
let bytes_in_call = if uses_bytes {
quote!(.bytes_in())
} else {
quote!()
};
let ret_ty = &method.sig.output;
let is_unit = returns_unit(ret_ty);
let resp_call = match returns {
ReturnMode::Graph => quote!(.resp_literal("{ graph: JsonValue } | { error: string }")),
ReturnMode::GraphNode => quote!(.resp_literal(
"{ graph: JsonValue, added_node_id: string } | { error: string }"
)),
ReturnMode::OkError => quote!(.resp_literal("null | { error: string }")),
ReturnMode::Bytes => quote!(.bytes_out()),
ReturnMode::Default => {
if is_unit {
quote!()
} else if let syn::ReturnType::Type(_, ty) = ret_ty {
let resp_ty = result_ok_type(ty).unwrap_or(ty);
quote!(.resp::<#resp_ty>())
} else {
quote!()
}
}
};
let default_send = !(is_unit && returns == ReturnMode::Default);
let send = args.mode.unwrap_or(default_send);
let mode_call = if send {
quote!(.send())
} else {
quote!(.post())
};
let reg_fn_ident = format_ident!("__darkly_handler_{}", name);
let reg_fn = quote! {
#[doc(hidden)]
pub(crate) fn #reg_fn_ident() -> crate::engine::protocol::RequestRegistration {
crate::engine::protocol::RequestRegistration::new(
#kind,
|engine, #payload_ident, #bytes_ident| {
#decode
#call
#convert
},
)
#mode_call
#bytes_in_call
#req_call
#resp_call
}
};
Ok(Built { req_struct, reg_fn })
}
fn parse_param(ident: Ident, ty: &Type) -> syn::Result<Param> {
if ident == "bytes" {
if let Type::Reference(r) = ty {
if let Type::Slice(slice) = &*r.elem {
if is_u8(&slice.elem) {
return Ok(Param {
call_expr: quote!(bytes),
field: None,
});
}
}
}
return Err(Error::new(
ty.span(),
"a `bytes` handler parameter must be `&[u8]` (the protocol side-channel)",
));
}
match ty {
Type::Reference(r) if is_str(&r.elem) => Ok(Param {
call_expr: quote!(&req.#ident),
field: Some((ident.clone(), quote!(String))),
}),
Type::Reference(r) => {
if let Type::Slice(slice) = &*r.elem {
let elem = &slice.elem;
Ok(Param {
call_expr: quote!(&req.#ident),
field: Some((ident.clone(), quote!(Vec<#elem>))),
})
} else {
Err(Error::new(
ty.span(),
"unsupported reference parameter (only `&str`, `&[T]`, and `bytes: &[u8]` are handled)",
))
}
}
_ => Ok(Param {
call_expr: quote!(req.#ident),
field: Some((ident.clone(), quote!(#ty))),
}),
}
}
fn is_u8(ty: &Type) -> bool {
matches!(ty, Type::Path(p) if p.path.is_ident("u8"))
}
fn is_str(ty: &Type) -> bool {
matches!(ty, Type::Path(p) if p.path.is_ident("str"))
}
fn pascal_case(snake: &str) -> String {
snake
.split('_')
.filter(|s| !s.is_empty())
.map(|word| {
let mut chars = word.chars();
match chars.next() {
Some(c) => c.to_uppercase().chain(chars).collect::<String>(),
None => String::new(),
}
})
.collect()
}