Skip to main content

grpc_macro/
lib.rs

1//! Procedural macros for `rust-grpc-lib`.
2//!
3//! This crate exposes four public macros:
4//!
5//! - [`GrpcClient`] — derive macro that generates
6//!   `from_endpoint_with_provider` on a tonic client struct, wiring it into
7//!   the process-wide channel pool with JWT auth via `ClientJwtInterceptor`.
8//! - [`GrpcNoAuthClient`] — derive macro that generates `from_endpoint` (no
9//!   auth) on a tonic client struct. For test harnesses only; requires the
10//!   `unauthenticated` feature on `rust-grpc-lib`.
11//! - [`keycloak_authenticated_service`] — attribute macro applied to an `impl Trait for Type`
12//!   block that injects Keycloak role-checking guards into methods annotated
13//!   with `#[roles(...)]`.
14//! - [`roles`] — marker attribute consumed by [`keycloak_authenticated_service`]. A pass-through
15//!   no-op when used without `#[keycloak_authenticated_service]` on the enclosing `impl` block.
16//!
17//! Internal helpers (`RolesSpec`, `RolesArgs`, `extract_roles_attr`,
18//! `first_param_ident`, `build_guard`) are private to this crate.
19
20use std::{iter, mem::take};
21
22use proc_macro::TokenStream;
23use proc_macro2::{Span, TokenStream as TokenStream2};
24
25use quote::quote;
26use syn::{
27    Attribute, DeriveInput, FnArg, Ident, ImplItem, ImplItemFn, ItemImpl, Lit, Meta, MetaList, Pat,
28    Path, Token,
29    parse::{Parse, ParseStream},
30    parse_macro_input,
31};
32
33#[proc_macro_derive(GrpcClient)]
34pub fn grpc_client_derive(input: TokenStream) -> TokenStream {
35    let parsed = parse_macro_input!(input as DeriveInput);
36    let name = &parsed.ident;
37
38    // Use `::rust_grpc_lib` so the impl resolves against the `core` crate when this
39    // derive is expanded, not against any re-export in a consuming crate.
40    let as_grpc_client = quote! {
41        impl #name<::tonic::transport::Channel>
42        {
43            pub fn from_endpoint_with_provider<P: ::rust_grpc_lib::auth::TokenProvider>(
44                endpoint: &str,
45                provider: P,
46            ) -> Result<#name<::tonic::service::interceptor::InterceptedService<::tonic::transport::Channel, ::rust_grpc_lib::auth::ClientJwtInterceptor<P>>>, ::tonic::transport::Error>
47            {
48                let channel = ::rust_grpc_lib::pool::get_channel(endpoint)?;
49                Ok(#name::new(::tonic::service::interceptor::InterceptedService::new(channel, ::rust_grpc_lib::auth::ClientJwtInterceptor::new(provider))))
50            }
51        }
52    };
53    TokenStream::from(as_grpc_client)
54}
55
56#[proc_macro_derive(GrpcNoAuthClient)]
57pub fn grpc_no_auth_client_derive(input: TokenStream) -> TokenStream {
58    let parsed = parse_macro_input!(input as DeriveInput);
59    let name = &parsed.ident;
60
61    // Use `::rust_grpc_lib` so the impl resolves against the `core` crate when this
62    // derive is expanded, not against any re-export in a consuming crate.
63    let as_grpc_client = quote! {
64        impl #name<::tonic::transport::Channel>
65        {
66            pub fn from_endpoint(
67                endpoint: &str,
68            ) -> Result<Self, ::tonic::transport::Error>
69            {
70                let channel = ::rust_grpc_lib::pool::get_channel(endpoint)?;
71                Ok(#name::new(channel))
72            }
73        }
74    };
75    TokenStream::from(as_grpc_client)
76}
77
78/// Marker attribute consumed by [`keycloak_authenticated_service`].
79///
80/// When used standalone (without `#[keycloak_authenticated_service]` on the enclosing `impl`
81/// block) this attribute is a pass-through no-op so that the code still
82/// compiles. `#[keycloak_authenticated_service]` strips and processes it before emitting the
83/// final `impl` block.
84///
85/// # Usage
86///
87/// ```rust,ignore
88/// #[keycloak_authenticated_service]
89/// impl MyService for MyServer {
90///     #[roles(any("operator", "admin"))]
91///     async fn set_data(&self, request: Request<SetDataRequest>)
92///         -> Result<Response<SetDataResponse>, Status>
93///     {
94///         // ...
95///     }
96/// }
97/// ```
98#[proc_macro_attribute]
99pub fn roles(_attr: TokenStream, item: TokenStream) -> TokenStream {
100    // Pass-through; keycloak_authenticated_service strips and processes this attribute.
101    item
102}
103
104/// Attribute macro applied to an `impl Trait for Type` block that injects
105/// Keycloak role-checking guards into methods annotated with `#[roles(...)]`.
106///
107/// # Role check variants
108///
109/// - `#[roles(any("r1", "r2"))]` — at least one of the listed roles must be
110///   present in the JWT claims.
111/// - `#[roles(all("r1", "r2"))]` — every listed role must be present.
112/// - No `#[roles(...)]` attribute — the method is left untouched (a valid JWT
113///   is still required by the `JwtValidationLayer`, but no role check is
114///   injected by this macro).
115///
116/// # Generated code shape
117///
118/// ```rust,ignore
119/// async fn set_data(&self, request: Request<SetDataRequest>)
120///     -> Result<Response<SetDataResponse>, Status>
121/// {
122///     {
123///         let __claims = request
124///             .extensions()
125///             .get::<::rust_grpc_lib::auth::KeycloakClaims>()
126///             .ok_or_else(|| ::tonic::Status::internal(
127///                 "JWT claims not populated; ensure JwtValidationLayer is installed",
128///             ))?;
129///         if !["operator", "admin"].iter().any(|r| __claims.has_role(r)) {
130///             return Err(::tonic::Status::permission_denied("required role not present"));
131///         }
132///     }
133///     // original body …
134/// }
135/// ```
136#[proc_macro_attribute]
137pub fn keycloak_authenticated_service(_attr: TokenStream, item: TokenStream) -> TokenStream {
138    let mut impl_block = parse_macro_input!(item as ItemImpl);
139
140    for impl_item in &mut impl_block.items {
141        let ImplItem::Fn(method) = impl_item else {
142            continue;
143        };
144
145        // Find and remove the #[roles(...)] attribute, capturing its content.
146        let roles_attr = match extract_roles_attr(&mut method.attrs) {
147            Ok(roles_attr) => roles_attr,
148            Err(e) => return e,
149        };
150
151        let Some(roles_spec) = roles_attr else {
152            // No roles attribute — leave method untouched.
153            continue;
154        };
155
156        // Determine the name of the first non-self parameter so we can call
157        // `.extensions()` on it.
158        let request_ident = first_param_ident(method);
159
160        // Build the guard block.
161        let guard = build_guard(request_ident, &roles_spec);
162
163        // Prepend the guard to the existing method body.
164        let original_stmts = take(&mut method.block.stmts);
165        match syn::parse2(quote! { #guard }) {
166            Ok(guard_stmt) => {
167                method.block.stmts = iter::once(guard_stmt).chain(original_stmts).collect()
168            }
169            Err(err) => return TokenStream::from(err.to_compile_error()),
170        }
171    }
172
173    TokenStream::from(quote! { #impl_block })
174}
175
176// Internal helpers
177
178/// Describes the parsed content of a `#[roles(...)]` attribute.
179enum RolesSpec {
180    Any(Vec<String>),
181    All(Vec<String>),
182}
183
184/// Parses `any("r1", "r2")` or `all("r1", "r2")` from a [`MetaList`] token
185/// stream.
186struct RolesArgs {
187    spec: RolesSpec,
188}
189
190impl Parse for RolesArgs {
191    fn parse(input: ParseStream<'_>) -> syn::Result<Self> {
192        // Expect an identifier: `any` or `all`.
193        let variant: syn::Ident = input.parse()?;
194        let variant_str = variant.to_string();
195
196        // Expect parenthesised list of string literals.
197        let content;
198        syn::parenthesized!(content in input);
199
200        let mut roles = Vec::new();
201        loop {
202            if content.is_empty() {
203                break;
204            }
205            let Lit::Str(s) = content.parse()? else {
206                return Err(syn::Error::new_spanned(
207                    &variant,
208                    "roles must be string literals",
209                ));
210            };
211            roles.push(s.value());
212            if content.is_empty() {
213                break;
214            }
215            let _comma: Token![,] = content.parse()?;
216        }
217
218        let spec = match variant_str.as_str() {
219            "any" => RolesSpec::Any(roles),
220            "all" => RolesSpec::All(roles),
221            other => {
222                return Err(syn::Error::new_spanned(
223                    &variant,
224                    format!("expected `any` or `all`, found `{other}`"),
225                ));
226            }
227        };
228
229        Ok(RolesArgs { spec })
230    }
231}
232
233/// Returns `true` if the [`Path`] is a single-segment path equal to `"roles"`.
234fn path_is_roles(path: &Path) -> bool {
235    path.get_ident().map(|i| i == "roles").unwrap_or(false)
236}
237
238/// Finds the first `#[roles(...)]` attribute in `attrs`, removes it, and
239/// returns the parsed [`RolesSpec`]. Returns `None` if no such attribute
240/// exists.
241fn extract_roles_attr(attrs: &mut Vec<Attribute>) -> Result<Option<RolesSpec>, TokenStream> {
242    let Some(tokens) = attrs
243        .iter()
244        .position(|a| matches!(&a.meta, Meta::List(ml) if path_is_roles(&ml.path)))
245        .and_then(|pos| {
246            let attr = attrs.remove(pos);
247            match attr.meta {
248                Meta::List(MetaList { tokens, .. }) => Some(tokens),
249                _ => None,
250            }
251        })
252    else {
253        return Ok(None);
254    };
255
256    syn::parse2::<RolesArgs>(tokens)
257        .map(|args| Some(args.spec))
258        .map_err(|e| TokenStream::from(e.to_compile_error()))
259}
260
261/// Returns the [`Ident`] of the first non-`self` parameter of `method`,
262/// or falls back to the identifier `request` if none can be found.
263fn first_param_ident(method: &ImplItemFn) -> Ident {
264    for input in &method.sig.inputs {
265        if let FnArg::Typed(pat_type) = input
266            && let Pat::Ident(pat_ident) = pat_type.pat.as_ref()
267        {
268            return pat_ident.ident.clone();
269        }
270    }
271    // Fallback — should not happen for well-formed gRPC service methods.
272    Ident::new("request", Span::call_site())
273}
274
275/// Builds the role-check guard block as a [`TokenStream2`].
276fn build_guard(request_ident: Ident, spec: &RolesSpec) -> TokenStream2 {
277    let (roles, iter_method) = match spec {
278        RolesSpec::Any(r) => (r, quote! { any }),
279        RolesSpec::All(r) => (r, quote! { all }),
280    };
281
282    // Build the array literal: ["r1", "r2", ...]
283    let role_literals: Vec<TokenStream2> = roles.iter().map(|r| quote! { #r }).collect();
284
285    quote! {
286        {
287            let __claims = #request_ident
288                .extensions()
289                .get::<::rust_grpc_lib::auth::KeycloakClaims>()
290                .ok_or_else(|| ::tonic::Status::internal(
291                    "JWT claims not populated; ensure JwtValidationLayer is installed"
292                ))?;
293            if ![#(#role_literals),*].iter().#iter_method(|r| __claims.has_role(r)) {
294                return Err(::tonic::Status::permission_denied("required role not present"));
295            }
296        }
297    }
298}