sark-gen 0.9.0

Sark proc-macro generators
Documentation
use proc_macro2::TokenStream;
use quote::{format_ident, quote};
use syn::{FnArg, ItemFn, Result};

use crate::codegen::route_spec;
use crate::parse::{RouteConfig, take_route_config};
use crate::util::TypeExt;

pub(super) fn attr_fn(mut fun: ItemFn) -> Result<TokenStream> {
    fun.modifiers.require_empty()?;
    if let syn::Safety::Unsafe(unsafe_token) = &fun.sig.safety {
        return Err(syn::Error::new_spanned(
            unsafe_token,
            "#[sark_gen::handler] does not support unsafe functions",
        ));
    }
    if fun.sig.inputs.len() == 1 {
        let request_ident =
            format_ident!("__Sark{}Request", upper_camel(&fun.sig.ident.to_string()));
        let request = crate::request::Mode::empty().expand(syn::parse_quote! {
            struct #request_ident {}
        })?;
        fun.sig
            .inputs
            .insert(0, syn::parse_quote!(_request: #request_ident));
        let handler = attr_fn(fun)?;
        return Ok(quote! {
            #request
            #handler
        });
    }
    let name = fun.sig.ident.clone();
    let vis = fun.vis.clone();
    let hidden_fn = format_ident!("__{}_fn", name);
    let output_ty = match &fun.sig.output {
        syn::ReturnType::Type(_, ty) => (**ty).clone(),
        syn::ReturnType::Default => syn::parse_quote!(()),
    };
    fun.sig.ident = hidden_fn.clone();

    let RouteConfig {
        static_response,
        max_body,
        head_skip,
    } = take_route_config(&mut fun.attrs)?;
    let is_async = fun.sig.asyncness.is_some();
    let wants_timer = fun.sig.inputs.len() == 3;

    if is_async && has_non_static_lifetime(&output_ty) {
        return Err(syn::Error::new_spanned(
            &output_ty,
            "async handler responses must own request-derived data",
        ));
    }
    if fun.sig.inputs.len() != 2 && !(is_async && wants_timer) {
        return Err(syn::Error::new_spanned(
            &fun.sig.inputs,
            "#[sark_gen::handler] requires `(state)`, `(request, state)`, or async `(request, state, timer)`",
        ));
    }
    if wants_timer && !is_async {
        return Err(syn::Error::new_spanned(
            &fun.sig.inputs,
            "#[sark_gen::handler] timer argument is only valid on `async` handlers",
        ));
    }

    let request_arg_ty = match fun.sig.inputs.first() {
        Some(FnArg::Typed(pat)) => (*pat.ty).clone(),
        other => {
            return Err(syn::Error::new_spanned(
                other,
                "#[sark_gen::handler] request argument must be typed",
            ));
        }
    };
    let state_arg_ty = match fun.sig.inputs.iter().nth(1) {
        Some(FnArg::Typed(pat)) => (*pat.ty).clone(),
        other => {
            return Err(syn::Error::new_spanned(
                other,
                "#[sark_gen::handler] state argument must be typed",
            ));
        }
    };
    let state_arg_inner = match &state_arg_ty {
        syn::Type::Reference(reference) => (*reference.elem).clone(),
        other => {
            return Err(syn::Error::new_spanned(
                other,
                "#[sark_gen::handler] state argument must be a reference (`&T`)",
            ));
        }
    };

    let request_ty = request_arg_ty;
    let state_ty = state_arg_inner;

    let request_ident = request_ty.type_ident()?;
    let request_inner_ident = format_ident!("{}View", request_ident);
    let request_raw_headers_ident = format_ident!("{}HeadersRaw", request_ident);
    let request_raw_params_ident = format_ident!("{}ParamsRaw", request_ident);
    let request_params_inner_ident = format_ident!("{}Params", request_ident);
    let request_headers_inner_ident = format_ident!("{}Headers", request_ident);
    let request_header_slot_ident = format_ident!("{}HeaderSlot", request_ident);
    let header_slot_ty = quote!(#request_header_slot_ident);

    let state_has_lifetime = StateLifetimes::any(&state_ty);
    let state_ty_state = StateLifetimes::normalize_to(state_ty.clone(), "state");
    let state_ty_d = StateLifetimes::normalize_to(state_ty.clone(), "d");
    let state_lt_use = state_has_lifetime.then(|| quote!(<'state>));
    let state_outlives = state_has_lifetime.then(|| quote!('state: 'a,));
    let hidden_state_ty = if is_async {
        &state_ty_d
    } else {
        &state_ty_state
    };

    if !fun
        .sig
        .generics
        .params
        .iter()
        .any(|param| matches!(param, syn::GenericParam::Lifetime(lt) if lt.lifetime.ident == "req"))
    {
        fun.sig.generics.params.insert(0, syn::parse_quote!('req));
    }
    if is_async
        && !fun.sig.generics.params.iter().any(
            |param| matches!(param, syn::GenericParam::Lifetime(lt) if lt.lifetime.ident == "d"),
        )
    {
        fun.sig.generics.params.insert(0, syn::parse_quote!('d));
    }
    if state_has_lifetime && !is_async {
        fun.sig.generics.params.insert(0, syn::parse_quote!('state));
    }
    if let Some(FnArg::Typed(pat)) = fun.sig.inputs.first_mut() {
        *pat.ty = syn::parse_quote!(#request_inner_ident<'req>);
    }
    if let Some(FnArg::Typed(pat)) = fun.sig.inputs.iter_mut().nth(1) {
        *pat.ty = if is_async {
            syn::parse_quote!(&'req #hidden_state_ty)
        } else {
            syn::parse_quote!(&#hidden_state_ty)
        };
    }
    if wants_timer && let Some(FnArg::Typed(pat)) = fun.sig.inputs.iter_mut().nth(2) {
        *pat.ty = syn::parse_quote!(&'req ::sark::Timer<'d>);
    }
    if is_async {
        let where_clause = fun.sig.generics.make_where_clause();
        where_clause
            .predicates
            .push(syn::parse_quote!(#hidden_state_ty: 'req));
        where_clause.predicates.push(syn::parse_quote!('d: 'req));
        fun.attrs.push(syn::parse_quote!(#[::sark::fiber_fn('d)]));
    }

    let parsed_body_ty = quote! {
        <#request_ident as sark::service::RouteRequestImpl>::ParsedBody<'__req>
    };
    let parse_body = quote! {
        <#request_ident as sark::service::RouteRequestImpl>::parse_body(raw)
    };

    let output_ty_req = StateLifetimes::normalize_to(output_ty.clone(), "__req");
    let output_ty_static = StateLifetimes::normalize_to(output_ty.clone(), "static");
    let kind_ty = if is_async {
        quote!(sark::service::manifold::NativeFiber)
    } else {
        quote! {
            <#output_ty_static as sark::service::manifold::NativeResponse<'static>>::Kind
        }
    };
    let native_response_ty = (!is_async).then_some(&output_ty_req);
    let native_response_ty_static = (!is_async).then_some(&output_ty_static);

    let route_spec_impl = route_spec::Config {
        name: &name,
        request_ident: &request_ident,
        raw_params_ident: &request_raw_params_ident,
        raw_headers_ident: &request_raw_headers_ident,
        params_inner_ident: &request_params_inner_ident,
        headers_inner_ident: &request_headers_inner_ident,
        header_slot_ty: &header_slot_ty,
        static_response,
        kind_ty: &kind_ty,
        native_response_ty,
        native_response_ty_static,
        async_response_ty: is_async.then_some(&output_ty),
        parsed_body_ty: Some(&parsed_body_ty),
        parse_body_body: Some(&parse_body),
        max_body: max_body.as_ref(),
        head_skip,
    }
    .build();

    let native_impl = if is_async {
        TokenStream::new()
    } else {
        quote! {
            impl #state_lt_use ::sark::service::manifold::Route<#state_ty_state> for #name {
                fn invoke<'req, 'a>(
                    &'a self,
                    params: <Self as ::sark::service::RouteSpec>::Params<'req>,
                    req: &::sark::request::Ref<'req>,
                    headers: <Self as ::sark::service::RouteSpec>::Headers<'req>,
                    parsed_body: <Self as ::sark::service::RouteSpec>::ParsedBody<'req>,
                    state: &'a #state_ty_state,
                ) -> <Self as ::sark::service::RouteSpec>::Response<'req>
                where
                    'req: 'a,
                    #state_outlives
                {
                    let request = #request_inner_ident::<'req>::from_parts(
                        params,
                        headers,
                        parsed_body,
                        req,
                    );
                    let response = #hidden_fn(request, state);
                    ::sark::service::manifold::NativeResponse::into_route_response(response)
                }
            }
        }
    };

    let timer_call = wants_timer.then(|| quote!(, timer));
    let borrowed_request = quote! {
        #request_inner_ident::<'req>::from_parts(
            params,
            headers,
            parsed_body,
            &req,
        )
    };
    let task_impl = if is_async {
        quote! {
            impl<'d> ::sark::service::manifold::TaskRoute<'d, #state_ty_d> for #name {
                fn invoke_task<'req>(
                    &'req self,
                    params: <Self as ::sark::service::RouteSpec>::Params<'req>,
                    req: ::sark::request::Ref<'req>,
                    headers: <Self as ::sark::service::RouteSpec>::Headers<'req>,
                    parsed_body: <Self as ::sark::service::RouteSpec>::ParsedBody<'req>,
                    state: &'req #state_ty_d,
                    timer: &'req ::sark::Timer<'d>,
                ) -> impl ::sark::fiber::Fiber<'d, Output = #output_ty> + 'req
                where
                    #state_ty_d: 'req,
                    'd: 'req,
                {
                    let _ = self;
                    let request = #borrowed_request;
                    #hidden_fn(request, state #timer_call)
                }
            }
        }
    } else {
        quote! {
            impl<'d> ::sark::service::manifold::TaskRoute<'d, #state_ty_d> for #name {
                fn invoke_task<'req>(
                    &'req self,
                    _params: <Self as ::sark::service::RouteSpec>::Params<'req>,
                    _req: ::sark::request::Ref<'req>,
                    _headers: <Self as ::sark::service::RouteSpec>::Headers<'req>,
                    _parsed_body: <Self as ::sark::service::RouteSpec>::ParsedBody<'req>,
                    _state: &'req #state_ty_d,
                    _timer: &'req ::sark::Timer<'d>,
                ) -> impl ::sark::fiber::Fiber<'d, Output = ()> + 'req
                where
                    #state_ty_d: 'req,
                    'd: 'req,
                {
                    ::sark::service::manifold::ready()
                }
            }
        }
    };

    Ok(quote! {
        #[allow(unreachable_code)]
        #fun

        #[allow(non_camel_case_types)]
        #vis struct #name;

        #route_spec_impl
        #native_impl
        #task_impl
    })
}

fn upper_camel(value: &str) -> String {
    let mut output = String::with_capacity(value.len());
    for part in value.split('_').filter(|part| !part.is_empty()) {
        let mut chars = part.chars();
        if let Some(first) = chars.next() {
            output.extend(first.to_uppercase());
            output.extend(chars);
        }
    }
    output
}

fn has_non_static_lifetime(ty: &syn::Type) -> bool {
    struct Detect(bool);

    impl<'ast> syn::visit::Visit<'ast> for Detect {
        fn visit_lifetime(&mut self, lifetime: &'ast syn::Lifetime) {
            self.0 |= lifetime.ident != "static";
        }

        fn visit_type_reference(&mut self, reference: &'ast syn::TypeReference) {
            self.0 |= reference
                .lifetime
                .as_ref()
                .is_none_or(|lifetime| lifetime.ident != "static");
            syn::visit::visit_type_reference(self, reference);
        }
    }

    let mut detect = Detect(false);
    syn::visit::Visit::visit_type(&mut detect, ty);
    detect.0
}

struct StateLifetimes;

impl StateLifetimes {
    fn any(ty: &syn::Type) -> bool {
        struct Detect(bool);
        impl<'ast> syn::visit::Visit<'ast> for Detect {
            fn visit_lifetime(&mut self, _lt: &'ast syn::Lifetime) {
                self.0 = true;
            }
        }
        let mut detect = Detect(false);
        syn::visit::Visit::visit_type(&mut detect, ty);
        detect.0
    }

    fn normalize_to(mut ty: syn::Type, ident: &str) -> syn::Type {
        struct Rewrite<'a>(&'a str);
        impl syn::visit_mut::VisitMut for Rewrite<'_> {
            fn visit_lifetime_mut(&mut self, lifetime: &mut syn::Lifetime) {
                lifetime.ident = format_ident!("{}", self.0);
            }
        }
        syn::visit_mut::VisitMut::visit_type_mut(&mut Rewrite(ident), &mut ty);
        ty
    }
}