use heck::ToUpperCamelCase as _;
use proc_macro2::TokenStream;
use quote::{format_ident, quote};
use syn::{spanned::Spanned as _, Error, FnArg, ItemFn, ReturnType, Type};
struct ParsedLocatorFn {
name: syn::Ident,
key_ty: Type,
asset_ty: Type,
}
pub fn generate_asset_locator(input_fn: ItemFn) -> Result<TokenStream, Error> {
let parsed = parse_locator_function(&input_fn)?;
let struct_name = format_ident!("{}", parsed.name.to_string().to_upper_camel_case());
let vis = &input_fn.vis;
let key_ty = &parsed.key_ty;
let asset_ty = &parsed.asset_ty;
let fn_name = &parsed.name;
Ok(quote! {
#input_fn
#vis struct #struct_name;
impl ::query_flow::AssetLocator<#key_ty> for #struct_name {
fn locate(
&self,
db: &impl ::query_flow::Db,
key: &#key_ty,
) -> ::std::result::Result<::query_flow::LocateResult<#asset_ty>, ::query_flow::QueryError> {
#fn_name(db, key)
}
}
})
}
fn parse_locator_function(input_fn: &ItemFn) -> Result<ParsedLocatorFn, Error> {
let name = input_fn.sig.ident.clone();
let mut params = input_fn.sig.inputs.iter();
let first_param = params.next().ok_or_else(|| {
Error::new(
input_fn.sig.span(),
"asset locator function must have `db: &impl Db` as first parameter",
)
})?;
validate_db_param(first_param)?;
let second_param = params.next().ok_or_else(|| {
Error::new(
input_fn.sig.span(),
"asset locator function must have `key: &KeyType` as second parameter",
)
})?;
let key_ty = parse_key_param(second_param)?;
if params.next().is_some() {
return Err(Error::new(
input_fn.sig.span(),
"asset locator function should have exactly 2 parameters: (db, key)",
));
}
let asset_ty = parse_locator_return_type(&input_fn.sig.output)?;
Ok(ParsedLocatorFn {
name,
key_ty,
asset_ty,
})
}
fn validate_db_param(arg: &FnArg) -> Result<(), Error> {
match arg {
FnArg::Typed(_) => {
Ok(())
}
FnArg::Receiver(_) => Err(Error::new(
arg.span(),
"first parameter must be `db: &impl Db`, not `self`",
)),
}
}
fn parse_key_param(arg: &FnArg) -> Result<Type, Error> {
match arg {
FnArg::Typed(pat_type) => {
let ty = &*pat_type.ty;
if let Type::Reference(ref_ty) = ty {
Ok((*ref_ty.elem).clone())
} else {
Err(Error::new(
ty.span(),
"key parameter should be a reference type `&KeyType`",
))
}
}
FnArg::Receiver(_) => Err(Error::new(
arg.span(),
"second parameter must be `key: &KeyType`",
)),
}
}
fn parse_locator_return_type(ret: &ReturnType) -> Result<Type, Error> {
match ret {
ReturnType::Default => Err(Error::new(
ret.span(),
"asset locator must return `Result<LocateResult<AssetType>, QueryError>`",
)),
ReturnType::Type(_, ty) => extract_locate_result_asset_type(ty),
}
}
fn extract_locate_result_asset_type(ty: &Type) -> Result<Type, Error> {
if let Type::Path(type_path) = ty {
if let Some(segment) = type_path.path.segments.last() {
if segment.ident == "Result" {
if let syn::PathArguments::AngleBracketed(args) = &segment.arguments {
if let Some(syn::GenericArgument::Type(ok_ty)) = args.args.first() {
return extract_locate_result_inner(ok_ty);
}
}
}
}
}
Err(Error::new(
ty.span(),
"expected `Result<LocateResult<AssetType>, QueryError>` return type",
))
}
fn extract_locate_result_inner(ty: &Type) -> Result<Type, Error> {
if let Type::Path(type_path) = ty {
if let Some(segment) = type_path.path.segments.last() {
if segment.ident == "LocateResult" {
if let syn::PathArguments::AngleBracketed(args) = &segment.arguments {
if let Some(syn::GenericArgument::Type(asset_ty)) = args.args.first() {
return Ok(asset_ty.clone());
}
}
}
}
}
Err(Error::new(
ty.span(),
"expected `LocateResult<AssetType>` in return type",
))
}
#[cfg(test)]
mod tests {
use super::*;
use syn::ItemFn;
fn normalize_tokens(tokens: TokenStream) -> String {
tokens
.to_string()
.split_whitespace()
.collect::<Vec<_>>()
.join(" ")
.replace("> >", ">>")
}
#[test]
fn test_asset_locator_macro_basic() {
let input_fn: ItemFn = syn::parse_quote! {
fn pending_locator(_db: &impl Db, _key: &ConfigFile) -> Result<LocateResult<String>, QueryError> {
Ok(LocateResult::Pending)
}
};
let output = generate_asset_locator(input_fn).unwrap();
let expected = quote! {
fn pending_locator(_db: &impl Db, _key: &ConfigFile) -> Result<LocateResult<String>, QueryError> {
Ok(LocateResult::Pending)
}
struct PendingLocator;
impl ::query_flow::AssetLocator<ConfigFile> for PendingLocator {
fn locate(
&self,
db: &impl ::query_flow::Db,
key: &ConfigFile,
) -> ::std::result::Result<::query_flow::LocateResult<String>, ::query_flow::QueryError> {
pending_locator(db, key)
}
}
};
assert_eq!(normalize_tokens(output), normalize_tokens(expected));
}
#[test]
fn test_asset_locator_with_generic_asset() {
let input_fn: ItemFn = syn::parse_quote! {
fn my_locator(db: &impl Db, key: &MyKey) -> Result<LocateResult<Vec<u8>>, QueryError> {
Ok(LocateResult::Pending)
}
};
let output = generate_asset_locator(input_fn).unwrap();
let expected = quote! {
fn my_locator(db: &impl Db, key: &MyKey) -> Result<LocateResult<Vec<u8>>, QueryError> {
Ok(LocateResult::Pending)
}
struct MyLocator;
impl ::query_flow::AssetLocator<MyKey> for MyLocator {
fn locate(
&self,
db: &impl ::query_flow::Db,
key: &MyKey,
) -> ::std::result::Result<::query_flow::LocateResult<Vec<u8>>, ::query_flow::QueryError> {
my_locator(db, key)
}
}
};
assert_eq!(normalize_tokens(output), normalize_tokens(expected));
}
#[test]
fn test_asset_locator_with_vis() {
let input_fn: ItemFn = syn::parse_quote! {
pub fn my_locator(db: &impl Db, key: &MyKey) -> Result<LocateResult<String>, QueryError> {
Ok(LocateResult::Pending)
}
};
let output = generate_asset_locator(input_fn).unwrap();
let expected = quote! {
pub fn my_locator(db: &impl Db, key: &MyKey) -> Result<LocateResult<String>, QueryError> {
Ok(LocateResult::Pending)
}
pub struct MyLocator;
impl ::query_flow::AssetLocator<MyKey> for MyLocator {
fn locate(
&self,
db: &impl ::query_flow::Db,
key: &MyKey,
) -> ::std::result::Result<::query_flow::LocateResult<String>, ::query_flow::QueryError> {
my_locator(db, key)
}
}
};
assert_eq!(normalize_tokens(output), normalize_tokens(expected));
}
}