1use proc_macro::TokenStream;
2use quote::quote;
3#[cfg(feature = "jwt")]
4use syn::{Data, DeriveInput, Fields};
5use syn::{ItemFn, parse_macro_input};
6
7#[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#[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#[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}