feather_macros/
lib.rs

1use proc_macro::TokenStream;
2use quote::quote;
3#[cfg(feature = "jwt")]
4use syn::{Data, DeriveInput, Fields};
5use syn::{ItemFn, parse_macro_input};
6
7/// This macro derives the `Claim` trait for a struct, allowing it to be used as JWT claims.
8/// It checks for fields annotated with `#[required]` and `#[exp]` to
9/// validate the claims when decoding a JWT token.
10/// The `#[required]` attribute ensures that the field is not empty, and the `#[exp]` attribute checks if the field is a valid expiration time.
11#[cfg(feature = "jwt")]
12#[proc_macro_derive(Claim, attributes(required, exp))]
13pub fn derive_claim(input: TokenStream) -> TokenStream {
14    let input = parse_macro_input!(input as DeriveInput);
15    let name = &input.ident;
16    let mut checks = Vec::new();
17
18    if let Data::Struct(data_struct) = &input.data {
19        if let Fields::Named(fields) = &data_struct.fields {
20            for field in &fields.named {
21                let field_name = &field.ident;
22                for attr in &field.attrs {
23                    if attr.path().is_ident("required") {
24                        checks.push(quote! {
25                            if self.#field_name.is_empty() {
26                                return Err(feather::jwt::Error::from(feather::jwt::ErrorKind::InvalidToken));
27                            }
28                        });
29                    }
30                    if attr.path().is_ident("exp") {
31                        checks.push(quote! {
32                            if self.#field_name < ::std::time::SystemTime::now().duration_since(::std::time::UNIX_EPOCH).unwrap().as_secs() as usize {
33                                return Err(feather::jwt::Error::from(feather::jwt::ErrorKind::ExpiredSignature));
34                            }
35                        });
36                    }
37                }
38            }
39        }
40    }
41
42    let expanded = quote! {
43        impl feather::jwt::Claim for #name {
44            fn validate(&self) -> Result<(), feather::jwt::Error> {
45                #(#checks)*
46                Ok(())
47            }
48        }
49    };
50    TokenStream::from(expanded)
51}
52
53/// This macro defines a middleware function that can be used in Feather applications.  
54/// It allows you to write middleware functions without repeating the type signatures for request, response, and context.
55/// Example:
56/// ```rust,ignore
57/// use feather::{middleware_fn, Outcome, next};
58/// #[middleware_fn]
59/// fn my_middleware() -> Outcome {
60///     res.send_text("Hello from middleware!");
61///     next!()
62/// }
63///     // Your middleware logic here
64#[proc_macro_attribute]
65pub fn middleware_fn(_attr: TokenStream, item: TokenStream) -> TokenStream {
66    let input = parse_macro_input!(item as ItemFn);
67    let vis = &input.vis;
68    let sig: &syn::Signature = &input.sig;
69    let block = &input.block;
70    let fn_name = &sig.ident;
71
72    let expanded = quote! {
73        #vis fn #fn_name(
74            req: &mut feather::Request,
75            res: &mut feather::Response,
76            ctx: &feather::AppContext
77        ) -> feather::Outcome {
78            #block
79        }
80    };
81    TokenStream::from(expanded)
82}
83
84/// This macro is used to define a JWT-required middleware function.
85/// It expects a function with a specific signature that includes a claims argument.
86/// The claims argument must implement the `feather::jwt::Claim` trait.
87/// Example:
88/// ```rust,ignore
89/// use feather::{jwt_required, middleware_fn, Outcome, next};
90/// use feather::jwt::{JwtManager, SimpleClaims};
91/// #[jwt_required]
92/// #[middleware_fn]
93/// fn protected_route(claims: SimpleClaims) -> Outcome {
94///   // Your Logic Here
95///   next!()
96/// }
97#[cfg(feature = "jwt")]
98#[proc_macro_attribute]
99pub fn jwt_required(_attr: TokenStream, item: TokenStream) -> TokenStream {
100    let input = parse_macro_input!(item as ItemFn);
101    let fn_name = &input.sig.ident;
102    let vis = &input.vis;
103    let block = &input.block;
104    let inputs = &input.sig.inputs;
105
106    let claims_ident = inputs.iter().find_map(|arg| {
107        if let syn::FnArg::Typed(pat_type) = arg {
108            if let syn::Pat::Ident(ident) = &*pat_type.pat {
109                Some((&ident.ident, &*pat_type.ty))
110            } else {
111                None
112            }
113        } else {
114            None
115        }
116    });
117
118    let (claims_name, claims_type) = match claims_ident {
119        Some(x) => x,
120        None => {
121            return syn::Error::new_spanned(&input.sig, "expected a `claims: T` argument for #[jwt_required]").to_compile_error().into();
122        }
123    };
124
125    let expanded = quote! {
126        #vis fn #fn_name(req: &mut feather::Request, res: &mut feather::Response, ctx: &feather::AppContext) -> feather::Outcome {
127            let manager = ctx.jwt();
128            let token = match req
129                .headers
130                .get("Authorization")
131                .and_then(|h| h.to_str().ok())
132                .and_then(|h| h.strip_prefix("Bearer ")) {
133                    Some(t) => t,
134                    None => {
135                        res.set_status(401);
136                        res.send_text("Missing or invalid Authorization header");
137                        return feather::next!();
138                    }
139                };
140
141            let #claims_name: #claims_type = match manager.decode(token) {
142                Ok(c) => c,
143                Err(_) => {
144                    res.set_status(401);
145                    res.send_text("Invalid or expired token");
146                    return feather::next!();
147                }
148            };
149
150            if let Err(_) = #claims_name.validate() {
151                res.set_status(401);
152                res.send_text("Invalid or expired token");
153                return feather::next!();
154            }
155
156            #block
157        }
158    };
159
160    TokenStream::from(expanded)
161}