Skip to main content

rorpc_parse/codegen/
router.rs

1//! Code generation for the `router!(...)` proc macro.
2//!
3//! Parses 0–2 arguments (state expression and/or module path pattern) in any
4//! order, then emits an Axum `Router` that merges all matching registered
5//! handlers from the `inventory`.
6
7use proc_macro2::TokenStream;
8use quote::quote;
9use syn::{
10    Expr, ExprArray, LitStr, Token,
11    parse::{Parse, ParseStream},
12};
13
14// ---------------------------------------------------------------------------
15// RouterArgs — parsed from router!(state?, pattern?)
16// ---------------------------------------------------------------------------
17
18/// Parsed arguments for the `router!(...)` macro.
19///
20/// Both arguments are optional and may appear in any order.
21pub struct RouterArgs {
22    pub state: Option<Expr>,
23    pub pattern: Option<Pattern>,
24}
25
26impl std::fmt::Debug for RouterArgs {
27    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
28        f.debug_struct("RouterArgs")
29            .field("has_state", &self.state.is_some())
30            .field("pattern", &self.pattern)
31            .finish()
32    }
33}
34
35/// Module path filter pattern.
36#[derive(Clone, Debug)]
37pub enum Pattern {
38    /// `"handlers::planet"` — single pattern
39    Single(String),
40    /// `["handlers::planet", "handlers::user"]` — multiple patterns
41    Multiple(Vec<String>),
42}
43
44impl Parse for RouterArgs {
45    fn parse(input: ParseStream) -> syn::Result<Self> {
46        let mut state: Option<Expr> = None;
47        let mut pattern: Option<Pattern> = None;
48
49        if !input.is_empty() {
50            match parse_one_arg(input)? {
51                Arg::Pattern(p) => pattern = Some(p),
52                Arg::State(s) => state = Some(s),
53            }
54
55            if input.peek(Token![,]) {
56                input.parse::<Token![,]>()?;
57                if !input.is_empty() {
58                    match parse_one_arg(input)? {
59                        Arg::Pattern(p) => {
60                            if pattern.is_some() {
61                                return Err(syn::Error::new(
62                                    input.span(),
63                                    "router!(): pattern specified twice",
64                                ));
65                            }
66                            pattern = Some(p);
67                        }
68                        Arg::State(s) => {
69                            if state.is_some() {
70                                return Err(syn::Error::new(
71                                    input.span(),
72                                    "router!(): state specified twice",
73                                ));
74                            }
75                            state = Some(s);
76                        }
77                    }
78                }
79            }
80        }
81
82        Ok(RouterArgs { state, pattern })
83    }
84}
85
86enum Arg {
87    Pattern(Pattern),
88    State(Expr),
89}
90
91fn parse_one_arg(input: ParseStream) -> syn::Result<Arg> {
92    if let Ok(lit) = input.parse::<LitStr>() {
93        return Ok(Arg::Pattern(Pattern::Single(lit.value())));
94    }
95    if let Ok(array) = input.parse::<ExprArray>() {
96        let patterns: syn::Result<Vec<String>> = array
97            .elems
98            .iter()
99            .map(|elem| match elem {
100                Expr::Lit(expr_lit) => {
101                    if let syn::Lit::Str(s) = &expr_lit.lit {
102                        Ok(s.value())
103                    } else {
104                        Err(syn::Error::new_spanned(
105                            elem,
106                            "router!(): array elements must be string literals",
107                        ))
108                    }
109                }
110                _ => Err(syn::Error::new_spanned(
111                    elem,
112                    "router!(): array elements must be string literals",
113                )),
114            })
115            .collect();
116        return Ok(Arg::Pattern(Pattern::Multiple(patterns?)));
117    }
118    Ok(Arg::State(input.parse::<Expr>()?))
119}
120
121// ---------------------------------------------------------------------------
122// expand_router
123// ---------------------------------------------------------------------------
124
125/// Generate the `router!(...)` expansion.
126pub fn expand_router(args: RouterArgs) -> TokenStream {
127    let state_expr = match &args.state {
128        Some(expr) => quote! { ::std::sync::Arc::new(#expr) },
129        None => quote! { ::std::sync::Arc::new(()) },
130    };
131
132    let filter = build_filter(&args.pattern);
133
134    quote! {
135        {
136            let state: ::std::sync::Arc<dyn ::std::any::Any + Send + Sync> = #state_expr;
137            let mut app: ::axum::Router = ::axum::Router::new();
138            
139            // Build namespace lookup map: module_path -> prefix
140            let mut namespace_map: ::std::collections::HashMap<&'static str, &'static str> = 
141                ::std::collections::HashMap::new();
142            for ns in ::rorpc::inventory::iter::<::rorpc::NamespaceMetadata> {
143                namespace_map.insert(ns.module_path, ns.prefix);
144            }
145            
146            for reg in ::rorpc::inventory::iter::<::rorpc::HandlerRegistration> {
147                let matches = #filter;
148                if matches {
149                    // Look up namespace for this handler and compose final path
150                    let mut final_path: String = reg.path.to_string();
151                    for metadata in ::rorpc::inventory::iter::<::rorpc::HandlerMetadata> {
152                        if metadata.path == reg.path && metadata.method == reg.method {
153                            // Check if handler's module or any parent has a namespace
154                            let module_parts: Vec<&str> = metadata.module_path.split("::").collect();
155                            for i in (0..=module_parts.len()).rev() {
156                                let parent_path = module_parts[..i].join("::");
157                                if let Some(prefix) = namespace_map.get(parent_path.as_str()) {
158                                    // Compose namespace + handler path
159                                    final_path = format!("{}{}", prefix, reg.path);
160                                    break;
161                                }
162                            }
163                            break;
164                        }
165                    }
166                    
167                    // Call factory with composed path
168                    let route = (reg.factory)(::std::sync::Arc::clone(&state), &final_path);
169                    app = app.merge(route);
170                }
171            }
172            app
173        }
174    }
175}
176
177// ---------------------------------------------------------------------------
178// Filter predicate generation
179// ---------------------------------------------------------------------------
180
181fn build_filter(pattern: &Option<Pattern>) -> TokenStream {
182    let Some(p) = pattern else {
183        return quote! { true };
184    };
185
186    let expanded: Vec<String> = match p {
187        Pattern::Single(s) => expand_pattern(s),
188        Pattern::Multiple(vec) => vec.iter().flat_map(|s| expand_pattern(s)).collect(),
189    };
190
191    if expanded.is_empty() {
192        return quote! { true };
193    }
194
195    let conditions: Vec<TokenStream> = expanded
196        .iter()
197        .map(|pat| {
198            let with_sep = format!("{}::", pat);
199            quote! {
200                (metadata.module_path == #pat || metadata.module_path.starts_with(#with_sep))
201            }
202        })
203        .collect();
204
205    quote! {
206        {
207            let mut matches = false;
208            for metadata in ::rorpc::inventory::iter::<::rorpc::HandlerMetadata> {
209                if metadata.path == reg.path && metadata.method == reg.method {
210                    matches = #(#conditions)||*;
211                    break;
212                }
213            }
214            matches
215        }
216    }
217}
218
219/// Expand a pattern string, handling brace groups and wildcards.
220///
221/// - `"handlers::{planet,user}"` → `["handlers::planet", "handlers::user"]`
222/// - `"handlers::*"` → `["handlers::"]` (prefix match)
223fn expand_pattern(pattern: &str) -> Vec<String> {
224    if let (Some(start), Some(end)) = (pattern.find('{'), pattern.find('}')) {
225        let prefix = &pattern[..start];
226        let group = &pattern[start + 1..end];
227        let suffix = &pattern[end + 1..];
228        return group
229            .split(',')
230            .map(|seg| normalise(format!("{}{}{}", prefix, seg.trim(), suffix)))
231            .collect();
232    }
233    vec![normalise(pattern.to_string())]
234}
235
236/// Normalise a single pattern: strip trailing `*` after `::`.
237fn normalise(pattern: String) -> String {
238    if pattern.ends_with("::*") {
239        pattern[..pattern.len() - 1].to_string()
240    } else {
241        pattern
242    }
243}
244
245// ---------------------------------------------------------------------------
246// Tests
247// ---------------------------------------------------------------------------
248
249#[cfg(test)]
250mod tests {
251    use super::*;
252
253    #[test]
254    fn expand_simple() {
255        assert_eq!(expand_pattern("handlers::planet"), vec!["handlers::planet"]);
256    }
257
258    #[test]
259    fn expand_wildcard() {
260        assert_eq!(expand_pattern("handlers::*"), vec!["handlers::"]);
261    }
262
263    #[test]
264    fn expand_brace_group() {
265        let mut result = expand_pattern("handlers::{planet,user}");
266        result.sort();
267        assert_eq!(result, vec!["handlers::planet", "handlers::user"]);
268    }
269
270    #[test]
271    fn expand_brace_with_whitespace() {
272        let mut result = expand_pattern("handlers::{ planet , user }");
273        result.sort();
274        assert_eq!(result, vec!["handlers::planet", "handlers::user"]);
275    }
276
277    #[test]
278    fn parse_no_args() {
279        let args: RouterArgs = syn::parse_str("").unwrap();
280        assert!(args.state.is_none());
281        assert!(args.pattern.is_none());
282    }
283
284    #[test]
285    fn parse_pattern_only() {
286        let args: RouterArgs = syn::parse_str("\"handlers::planet\"").unwrap();
287        assert!(args.state.is_none());
288        assert!(matches!(args.pattern, Some(Pattern::Single(_))));
289    }
290
291    #[test]
292    fn parse_array_pattern() {
293        let args: RouterArgs =
294            syn::parse_str("[\"handlers::planet\", \"handlers::user\"]").unwrap();
295        assert!(matches!(args.pattern, Some(Pattern::Multiple(_))));
296    }
297
298    #[test]
299    fn parse_duplicate_pattern_errors() {
300        let result: syn::Result<RouterArgs> =
301            syn::parse_str("\"handlers::planet\", \"handlers::user\"");
302        assert!(result.is_err());
303        assert!(
304            result
305                .unwrap_err()
306                .to_string()
307                .contains("pattern specified twice")
308        );
309    }
310}