Skip to main content

protovalidate_buffa_macros/
lib.rs

1//! `#[connect_impl]` — validates the request view at the top of every Connect
2//! service handler method in an `impl` block whose request parameter is an
3//! `OwnedView<_>` or `connectrpc::ServiceRequest<'_, _>`. Single-site safety
4//! net: add it once to the service impl and every present-and-future handler is
5//! validated on entry.
6//!
7//! Validation borrows `OwnedView::reborrow()` or `ServiceRequest::view()` and
8//! requires `Validate` on the generated view type. It does not convert the
9//! request to an owned message.
10//!
11//! Non-handler `async fn`s inside the same `impl` block are left alone
12//! (they lack a recognized request parameter, so the macro skips them).
13
14use proc_macro::TokenStream;
15use proc_macro2::TokenStream as TokenStream2;
16use quote::{quote, quote_spanned};
17use syn::{Error, FnArg, ImplItem, ItemImpl, PatType, Type, TypePath, parse_macro_input};
18
19#[proc_macro_attribute]
20pub fn connect_impl(attr: TokenStream, input: TokenStream) -> TokenStream {
21    if !attr.is_empty() {
22        return Error::new_spanned(
23            TokenStream2::from(attr),
24            "protovalidate_buffa::connect_impl takes no arguments",
25        )
26        .to_compile_error()
27        .into();
28    }
29
30    let mut item = parse_macro_input!(input as ItemImpl);
31
32    for impl_item in &mut item.items {
33        if let ImplItem::Fn(f) = impl_item
34            && let Some((arg_ident, request_type)) = find_request_arg(&f.sig)
35        {
36            // Point missing view-validator errors at the request parameter.
37            let span = arg_ident.span();
38            let view = match request_type {
39                RequestType::OwnedView => quote_spanned! {span=> #arg_ident.reborrow() },
40                RequestType::ServiceRequest => quote_spanned! {span=> #arg_ident.view() },
41            };
42            let validate: syn::Stmt = syn::parse_quote_spanned! {span=>
43                <_ as ::protovalidate_buffa::Validate>::validate(#view)
44                    .map_err(::protovalidate_buffa::ValidationError::into_connect_error)?;
45            };
46
47            f.block.stmts.insert(0, validate);
48        }
49    }
50
51    TokenStream::from(quote! { #item })
52}
53
54enum RequestType {
55    OwnedView,
56    ServiceRequest,
57}
58
59/// Returns the first recognized request parameter and its view accessor kind.
60/// Non-handler methods that lack such a parameter return `None`.
61fn find_request_arg(sig: &syn::Signature) -> Option<(syn::Ident, RequestType)> {
62    for arg in &sig.inputs {
63        if let FnArg::Typed(PatType { pat, ty, .. }) = arg
64            && let Some(request_type) = request_type(ty)
65            && let syn::Pat::Ident(pat_ident) = pat.as_ref()
66        {
67            return Some((pat_ident.ident.clone(), request_type));
68        }
69    }
70    None
71}
72
73fn request_type(ty: &Type) -> Option<RequestType> {
74    if let Type::Path(TypePath { path, .. }) = ty
75        && let Some(last) = path.segments.last()
76    {
77        if last.ident == "OwnedView" {
78            return Some(RequestType::OwnedView);
79        }
80        if last.ident == "ServiceRequest" {
81            return Some(RequestType::ServiceRequest);
82        }
83    }
84    None
85}