Skip to main content

rorpc_parse/codegen/
orpc.rs

1//! Code generation for the `#[rorpc(method, path)]` attribute macro.
2//!
3//! Parses the attribute arguments, analyses the handler signature, and emits:
4//! - The original function unchanged
5//! - An `inventory::submit!` for `HandlerMetadata`
6//! - An `inventory::submit!` for `HandlerRegistration` (Axum router factory)
7//! - `inventory::submit!` blocks for `SchemaRegistration` fallback schemas
8
9use proc_macro2::TokenStream;
10use quote::quote;
11use syn::{
12    Expr, ExprLit, ItemFn, Lit, MetaNameValue, Token,
13    parse::{Parse, ParseStream},
14    punctuated::Punctuated,
15};
16
17use crate::{
18    errors::{Error, Result, type_display},
19    functions::extract_handler_signature,
20    types::{JSON, QUERY, RESULT, innermost_custom_type, is_primitive, try_extract_wrapper},
21};
22
23// ---------------------------------------------------------------------------
24// Attribute name constants — centralized for easy renaming
25// ---------------------------------------------------------------------------
26
27const ATTR_METHOD: &str = "method";
28const ATTR_PATH: &str = "path";
29const ATTR_DATA: &str = "data";
30
31// ---------------------------------------------------------------------------
32// OrpcArgs — parsed from #[orpc(method = "...", path = "...", data = TypePath)]
33// ---------------------------------------------------------------------------
34
35/// Parsed arguments for the `#[orpc(...)]` attribute.
36pub struct OrpcArgs {
37    pub method: String,
38    pub path: String,
39    pub stream_event: Option<syn::Path>,
40}
41
42// ---------------------------------------------------------------------------
43// MethodShorthandArgs — parsed from #[orpc::get("/path")] or #[orpc::post("/path", data = "Type")]
44// ---------------------------------------------------------------------------
45
46/// Parsed arguments for method-specific shorthand macros like `#[orpc::get("/path")]`.
47///
48/// Syntax: `#[orpc::get("/path")]` or `#[orpc::post("/path", data = "StreamEvent")]`
49pub struct MethodShorthandArgs {
50    pub path: String,
51    pub data: Option<syn::Path>,
52}
53
54impl Parse for MethodShorthandArgs {
55    fn parse(input: ParseStream) -> syn::Result<Self> {
56        // First token must be a string literal (the path)
57        let path_lit: syn::LitStr = input.parse()?;
58        let path = path_lit.value();
59
60        // Optional: comma + data = "Type"
61        let mut data = None;
62
63        if input.peek(Token![,]) {
64            input.parse::<Token![,]>()?;
65
66            let pairs = Punctuated::<MetaNameValue, Token![,]>::parse_terminated(input)?;
67
68            for pair in &pairs {
69                let key = pair
70                    .path
71                    .get_ident()
72                    .map(|i| i.to_string())
73                    .unwrap_or_default();
74
75                let span = pair
76                    .path
77                    .get_ident()
78                    .map(|i| i.span())
79                    .unwrap_or_else(proc_macro2::Span::call_site);
80
81                match key.as_str() {
82                    ATTR_DATA => match &pair.value {
83                        Expr::Lit(ExprLit {
84                            lit: Lit::Str(s), ..
85                        }) => {
86                            let path_str = s.value();
87                            // Parse string content as a Rust path
88                            let parsed_path: syn::Path =
89                                syn::parse_str(&path_str).map_err(|e| {
90                                    syn::Error::new(
91                                        span,
92                                        format!(
93                                            "{} string \"{}\" is not a valid Rust type path: {}",
94                                            ATTR_DATA, path_str, e
95                                        ),
96                                    )
97                                })?;
98                            data = Some(parsed_path);
99                        }
100                        Expr::Path(expr_path) => {
101                            data = Some(expr_path.path.clone());
102                        }
103                        _ => {
104                            return Err(syn::Error::new(
105                                span,
106                                format!(
107                                    "{} must be a string literal (\"StreamEvent\") or type path (StreamEvent)",
108                                    ATTR_DATA
109                                ),
110                            ));
111                        }
112                    },
113                    _ => {
114                        return Err(syn::Error::new(
115                            span,
116                            Error::unknown_key(span, &key, &[ATTR_DATA]).to_string(),
117                        ));
118                    }
119                }
120            }
121        }
122
123        Ok(MethodShorthandArgs { path, data })
124    }
125}
126
127/// Convert method shorthand args to standard OrpcArgs.
128///
129/// This allows method-specific macros like `#[orpc::get("/path")]` to reuse
130/// all the existing codegen logic without duplication.
131impl MethodShorthandArgs {
132    pub fn into_orpc_args(self, method: &str) -> OrpcArgs {
133        OrpcArgs {
134            method: method.to_uppercase(),
135            path: self.path,
136            stream_event: self.data,
137        }
138    }
139}
140
141const VALID_KEYS: &[&str] = &[ATTR_METHOD, ATTR_PATH, ATTR_DATA];
142
143impl Parse for OrpcArgs {
144    fn parse(input: ParseStream) -> syn::Result<Self> {
145        let pairs = Punctuated::<MetaNameValue, Token![,]>::parse_terminated(input)?;
146
147        let mut method = None;
148        let mut path = None;
149        let mut stream_event = None;
150
151        for pair in &pairs {
152            let key = pair
153                .path
154                .get_ident()
155                .map(|i| i.to_string())
156                .unwrap_or_default();
157
158            let span = pair
159                .path
160                .get_ident()
161                .map(|i| i.span())
162                .unwrap_or_else(proc_macro2::Span::call_site);
163
164            match key.as_str() {
165                ATTR_METHOD => {
166                    if let Expr::Lit(ExprLit {
167                        lit: Lit::Str(s), ..
168                    }) = &pair.value
169                    {
170                        method = Some(s.value().to_uppercase());
171                    } else {
172                        return Err(syn::Error::new(
173                            span,
174                            Error::invalid_attr_value(
175                                span,
176                                &key,
177                                "a string literal",
178                                "non-string expression",
179                            )
180                            .to_string(),
181                        ));
182                    }
183                }
184                ATTR_PATH => {
185                    if let Expr::Lit(ExprLit {
186                        lit: Lit::Str(s), ..
187                    }) = &pair.value
188                    {
189                        path = Some(s.value());
190                    } else {
191                        return Err(syn::Error::new(
192                            span,
193                            Error::invalid_attr_value(
194                                span,
195                                &key,
196                                "a string literal",
197                                "non-string expression",
198                            )
199                            .to_string(),
200                        ));
201                    }
202                }
203                ATTR_DATA => {
204                    match &pair.value {
205                        // String literal containing a type path: data = "crate::SseEvent"
206                        // Parse the string as a Rust path for compile-time validation
207                        Expr::Lit(ExprLit {
208                            lit: Lit::Str(s), ..
209                        }) => {
210                            let path_str = s.value();
211                            // Parse string content as a Rust path to validate syntax
212                            let parsed_path: syn::Path =
213                                syn::parse_str(&path_str).map_err(|e| {
214                                    syn::Error::new(
215                                        span,
216                                        format!(
217                                            "{} string \"{}\" is not a valid Rust type path: {}",
218                                            ATTR_DATA, path_str, e
219                                        ),
220                                    )
221                                })?;
222                            // Store the parsed path for witness generation
223                            stream_event = Some(parsed_path);
224                        }
225                        // Type path: data = StreamEvent (deprecated but still supported)
226                        Expr::Path(expr_path) => {
227                            stream_event = Some(expr_path.path.clone());
228                        }
229                        _ => {
230                            return Err(syn::Error::new(
231                                span,
232                                format!(
233                                    "{} must be a string literal (\"StreamEvent\") or type path (StreamEvent)",
234                                    ATTR_DATA
235                                ),
236                            ));
237                        }
238                    }
239                }
240                _ => {
241                    return Err(syn::Error::new(
242                        span,
243                        Error::unknown_key(span, &key, VALID_KEYS).to_string(),
244                    ));
245                }
246            }
247        }
248
249        let method = method.ok_or_else(|| {
250            syn::Error::new(
251                proc_macro2::Span::call_site(),
252                Error::missing_required_attr(
253                    proc_macro2::Span::call_site(),
254                    ATTR_METHOD,
255                    "add `method = \"GET\"` to #[orpc]",
256                )
257                .to_string(),
258            )
259        })?;
260
261        let path = path.ok_or_else(|| {
262            syn::Error::new(
263                proc_macro2::Span::call_site(),
264                Error::missing_required_attr(
265                    proc_macro2::Span::call_site(),
266                    ATTR_PATH,
267                    "add `path = \"/your/route\"` to #[orpc]",
268                )
269                .to_string(),
270            )
271        })?;
272
273        Ok(OrpcArgs {
274            method,
275            path,
276            stream_event,
277        })
278    }
279}
280
281// ---------------------------------------------------------------------------
282// expand_orpc
283// ---------------------------------------------------------------------------
284
285/// Generate the full expansion for `#[rorpc::route(method, path)]` or shorthand macros.
286///
287/// Returns the original function unchanged plus all inventory registrations.
288pub fn expand_orpc(args: OrpcArgs, func: ItemFn) -> TokenStream {
289    match try_expand_orpc(args, func) {
290        Ok(ts) => ts,
291        Err(e) => e.to_compile_error(),
292    }
293}
294
295fn try_expand_orpc(args: OrpcArgs, func: ItemFn) -> Result<TokenStream> {
296    let sig = extract_handler_signature(&func)?;
297
298    let fn_name = &func.sig.ident;
299    let fn_name_str = sig.fn_name.as_str();
300    let method = &args.method;
301    let path = &args.path;
302
303    // Validate route path
304    if !path.starts_with('/') {
305        return Err(syn::Error::new_spanned(path, "route path must start with '/'").into());
306    }
307
308    if path.contains("..") {
309        return Err(
310            syn::Error::new_spanned(path, "route path cannot contain '..' path traversal").into(),
311        );
312    }
313
314    let output_type_str = type_display(&sig.output_type);
315
316    let error_type_token = match &sig.error_type {
317        Some(ty) => {
318            let s = type_display(ty);
319            quote! { Some(#s) }
320        }
321        None => quote! { None },
322    };
323
324    let stream_event_token = match &args.stream_event {
325        Some(type_path) => {
326            // Extract bare name (last segment) for metadata storage
327            let bare_name = type_path
328                .segments
329                .last()
330                .map(|seg| seg.ident.to_string())
331                .unwrap_or_else(|| "Unknown".to_string());
332            quote! { Some(#bare_name) }
333        }
334        None => quote! { None },
335    };
336
337    // Generate compile-time witness for SSE event type validation
338    let stream_event_witness = match &args.stream_event {
339        Some(type_path) => {
340            quote! {
341                // Compile-time witness: validates type is Serialize
342                // ZodTs derive is validated by inventory schema registration
343                const _: () = {
344                    fn assert_serialize<T: ::serde::Serialize>() {}
345                    fn check() {
346                        assert_serialize::<#type_path>();
347                    }
348                };
349            }
350        }
351        None => quote! {},
352    };
353
354    let input_type_str = match &sig.input_type {
355        Some(ty) => type_display(ty),
356        None => "()".to_string(),
357    };
358
359    let query_type_token = match &sig.query_type {
360        Some(ty) => {
361            let s = type_display(ty);
362            quote! { Some(#s) }
363        }
364        None => quote! { None },
365    };
366
367    // Encode path param types as comma-separated string: "i32,String"
368    // Order matches the path template param order (declaration order in signature)
369    let path_param_types_str = sig
370        .path_params
371        .iter()
372        .map(|(_, ty)| type_display(ty))
373        .collect::<Vec<_>>()
374        .join(",");
375
376    let registration = emit_handler_registration(fn_name, method, path, &sig.state_type);
377    let schema_registrations = emit_schema_registrations(&func);
378    let stream_event_schema_reg = emit_stream_event_schema_registration(&args.stream_event);
379
380    Ok(quote! {
381        #func
382
383        #stream_event_witness
384
385        ::rorpc::inventory::submit! {
386            ::rorpc::HandlerMetadata {
387                name: #fn_name_str,
388                method: #method,
389                path: #path,
390                input_type_name: #input_type_str,
391                query_type_name: #query_type_token,
392                output_type_name: #output_type_str,
393                module_path: ::std::module_path!(),
394                namespace: None,
395                error_type_name: #error_type_token,
396                stream_event_type_name: #stream_event_token,
397                path_param_types: #path_param_types_str,
398            }
399        }
400
401        #registration
402        #schema_registrations
403        #stream_event_schema_reg
404    })
405}
406
407// ---------------------------------------------------------------------------
408// Handler registration factory
409// ---------------------------------------------------------------------------
410
411fn emit_handler_registration(
412    fn_name: &syn::Ident,
413    method: &str,
414    path: &str,
415    state_type: &Option<syn::Type>,
416) -> TokenStream {
417    if let Some(state_ty) = state_type {
418        quote! {
419            ::rorpc::inventory::submit! {
420                ::rorpc::HandlerRegistration {
421                    path: #path,
422                    method: #method,
423                    factory: |state: ::std::sync::Arc<dyn ::std::any::Any + Send + Sync>, final_path: &str| {
424                        use ::axum::routing::{delete, get, patch, post, put};
425                        let method_router = match #method {
426                            "GET"    => get(#fn_name),
427                            "POST"   => post(#fn_name),
428                            "PUT"    => put(#fn_name),
429                            "PATCH"  => patch(#fn_name),
430                            "DELETE" => delete(#fn_name),
431                            _        => post(#fn_name),
432                        };
433                        if let Some(typed_state) = state.downcast_ref::<#state_ty>() {
434                            ::axum::Router::new()
435                                .route(final_path, method_router)
436                                .with_state(typed_state.clone())
437                        } else {
438                            ::axum::Router::new()
439                        }
440                    },
441                }
442            }
443        }
444    } else {
445        quote! {
446            ::rorpc::inventory::submit! {
447                ::rorpc::HandlerRegistration {
448                    path: #path,
449                    method: #method,
450                    factory: |_state: ::std::sync::Arc<dyn ::std::any::Any + Send + Sync>, final_path: &str| {
451                        use ::axum::routing::{delete, get, patch, post, put};
452                        let method_router = match #method {
453                            "GET"    => get(#fn_name),
454                            "POST"   => post(#fn_name),
455                            "PUT"    => put(#fn_name),
456                            "PATCH"  => patch(#fn_name),
457                            "DELETE" => delete(#fn_name),
458                            _        => post(#fn_name),
459                        };
460                        ::axum::Router::new().route(final_path, method_router)
461                    },
462                }
463            }
464        }
465    }
466}
467
468// ---------------------------------------------------------------------------
469// Schema registrations — z.unknown() fallback for types without #[derive(ZodTs)]
470// ---------------------------------------------------------------------------
471
472fn emit_schema_registrations(func: &ItemFn) -> TokenStream {
473    let mut seen = std::collections::HashSet::new();
474    let mut registrations = Vec::new();
475
476    // Collect candidate types from Json<T> and Query<T> params and return type
477    let mut candidates: Vec<&syn::Type> = Vec::new();
478
479    for arg in &func.sig.inputs {
480        if let syn::FnArg::Typed(pat_type) = arg {
481            // Check for Json<T>
482            if let Some(m) = try_extract_wrapper(&pat_type.ty, JSON)
483                && let Some(inner) = m.first_type()
484            {
485                candidates.push(inner);
486            }
487            // Check for Query<T>
488            if let Some(m) = try_extract_wrapper(&pat_type.ty, QUERY)
489                && let Some(inner) = m.first_type()
490            {
491                candidates.push(inner);
492            }
493        }
494    }
495
496    if let syn::ReturnType::Type(_, ty) = &func.sig.output {
497        // Handle both Json<T> and Result<Json<T>, E>
498        if let Some(m) = try_extract_wrapper(ty, JSON) {
499            if let Some(inner) = m.first_type() {
500                candidates.push(inner);
501            }
502        } else if let Some(result_m) = try_extract_wrapper(ty, RESULT)
503            && let Some(first) = result_m.first_type()
504            && let Some(json_m) = try_extract_wrapper(first, JSON)
505            && let Some(inner) = json_m.first_type()
506        {
507            candidates.push(inner);
508        }
509    }
510
511    for ty in candidates {
512        if let Some(custom_ty) = innermost_custom_type(ty) {
513            if is_primitive(custom_ty) {
514                continue;
515            }
516            let name = type_display(custom_ty);
517            if !seen.insert(name.clone()) {
518                continue;
519            }
520            registrations.push(quote! {
521                ::rorpc::inventory::submit! {
522                    ::rorpc::SchemaRegistration {
523                        type_name: #name,
524                        module_path: "",
525                        schema_def: ::rorpc::SchemaDef::Unknown,
526                        dependent_types: || vec![],
527                    }
528                }
529            });
530        }
531    }
532
533    quote! { #(#registrations)* }
534}
535
536/// Generate schema registration for SSE event type specified via `data` attribute.
537///
538/// Similar to `emit_schema_registrations`, but for the stream event type.
539fn emit_stream_event_schema_registration(stream_event: &Option<syn::Path>) -> TokenStream {
540    match stream_event {
541        Some(type_path) => {
542            // Extract bare name for registration
543            let bare_name = type_path
544                .segments
545                .last()
546                .map(|seg| seg.ident.to_string())
547                .unwrap_or_else(|| "Unknown".to_string());
548
549            quote! {
550                ::rorpc::inventory::submit! {
551                    ::rorpc::SchemaRegistration {
552                        type_name: #bare_name,
553                        module_path: "",
554                        schema_def: ::rorpc::SchemaDef::Unknown,
555                        dependent_types: || vec![],
556                    }
557                }
558            }
559        }
560        None => quote! {},
561    }
562}
563
564// ---------------------------------------------------------------------------
565// Tests
566// ---------------------------------------------------------------------------
567
568#[cfg(test)]
569mod tests {
570    use super::*;
571    use syn::parse_quote;
572
573    #[test]
574    fn parse_data_type_string() {
575        // Test that data = "StreamEvent" (string literal) parses correctly
576        let args: OrpcArgs = syn::parse_quote! {
577            method = "GET", path = "/stream", data = "StreamEvent"
578        };
579
580        assert_eq!(args.method, "GET");
581        assert_eq!(args.path, "/stream");
582        assert!(args.stream_event.is_some());
583        let path = args.stream_event.unwrap();
584        assert_eq!(path.segments.len(), 1);
585        assert_eq!(
586            path.segments.first().unwrap().ident.to_string(),
587            "StreamEvent"
588        );
589    }
590
591    #[test]
592    fn parse_data_qualified_path_string() {
593        // Test that data = "crate::models::StreamEvent" works
594        let args: OrpcArgs = syn::parse_quote! {
595            method = "GET", path = "/stream", data = "crate::models::StreamEvent"
596        };
597
598        assert_eq!(args.method, "GET");
599        assert_eq!(args.path, "/stream");
600        assert!(args.stream_event.is_some());
601        // Verify the path was parsed correctly
602        let path = args.stream_event.unwrap();
603        assert_eq!(path.segments.len(), 3);
604        assert_eq!(
605            path.segments.last().unwrap().ident.to_string(),
606            "StreamEvent"
607        );
608    }
609
610    #[test]
611    fn parse_data_type_path_backward_compat() {
612        // Test backward compatibility: data = StreamEvent (bare path)
613        let args: OrpcArgs = syn::parse_quote! {
614            method = "GET", path = "/stream", data = StreamEvent
615        };
616
617        assert_eq!(args.method, "GET");
618        assert_eq!(args.path, "/stream");
619        assert!(args.stream_event.is_some());
620        let path = args.stream_event.unwrap();
621        assert_eq!(path.segments.len(), 1);
622        assert_eq!(
623            path.segments.first().unwrap().ident.to_string(),
624            "StreamEvent"
625        );
626    }
627
628    #[test]
629    fn parse_without_data() {
630        // Test that data is optional
631        let args: OrpcArgs = syn::parse_quote! {
632            method = "POST", path = "/create"
633        };
634
635        assert_eq!(args.method, "POST");
636        assert_eq!(args.path, "/create");
637        assert_eq!(args.stream_event, None);
638    }
639
640    #[test]
641    fn data_type_converts_to_string_literal() {
642        // Verify that when we generate the metadata, data becomes a string literal
643        let args: OrpcArgs = syn::parse_quote! {
644            method = "GET", path = "/stream", data = "StreamEvent"
645        };
646
647        let func: syn::ItemFn = parse_quote! {
648            async fn stream_test() -> Sse<impl Stream<Item = Event>> {
649                todo!()
650            }
651        };
652
653        let result = try_expand_orpc(args, func);
654        assert!(result.is_ok(), "expand_orpc should succeed");
655
656        // Check that the generated code contains Some("StreamEvent") as a string literal
657        let tokens = result.unwrap().to_string();
658
659        // quote! serialises with spaces between tokens, so `Some("StreamEvent")` becomes
660        // `Some ("StreamEvent")`. Check the field name and the quoted value separately.
661        assert!(
662            tokens.contains("stream_event_type_name") && tokens.contains(r#""StreamEvent""#),
663            "Generated code should contain stream_event_type_name: Some(\"StreamEvent\"), got: {}",
664            tokens
665        );
666        // Also assert it is NOT a bare identifier (which would be a type error at compile time)
667        assert!(
668            !tokens.contains("Some (StreamEvent)") && !tokens.contains("Some(StreamEvent)"),
669            "stream_event_type_name must be a string literal, not a bare identifier"
670        );
671    }
672}
673
674// ---------------------------------------------------------------------------
675// MethodShorthandArgs tests
676// ---------------------------------------------------------------------------
677
678#[test]
679fn parse_shorthand_path_only() {
680    // Test: #[rorpc::get("/planet/list")]
681    let args: MethodShorthandArgs = syn::parse_quote! { "/planet/list" };
682
683    assert_eq!(args.path, "/planet/list");
684    assert_eq!(args.data, None);
685}
686
687#[test]
688fn parse_shorthand_with_data_string() {
689    // Test: #[rorpc::get("/stream", data = "EventData")]
690    let args: MethodShorthandArgs = syn::parse_quote! { "/stream", data = "EventData" };
691
692    assert_eq!(args.path, "/stream");
693    assert!(args.data.is_some());
694    let path = args.data.unwrap();
695    assert_eq!(
696        path.segments.first().unwrap().ident.to_string(),
697        "EventData"
698    );
699}
700
701#[test]
702fn parse_shorthand_with_qualified_data() {
703    // Test: #[rorpc::get("/stream", data = "models::EventData")]
704    let args: MethodShorthandArgs = syn::parse_quote! { "/stream", data = "models::EventData" };
705
706    assert_eq!(args.path, "/stream");
707    assert!(args.data.is_some());
708    let path = args.data.unwrap();
709    assert_eq!(path.segments.len(), 2);
710    assert_eq!(path.segments.last().unwrap().ident.to_string(), "EventData");
711}
712
713#[test]
714fn parse_shorthand_with_data_type_path() {
715    // Test backward compat: #[rorpc::get("/stream", data = EventData)]
716    let args: MethodShorthandArgs = syn::parse_quote! { "/stream", data = EventData };
717
718    assert_eq!(args.path, "/stream");
719    assert!(args.data.is_some());
720    let path = args.data.unwrap();
721    assert_eq!(
722        path.segments.first().unwrap().ident.to_string(),
723        "EventData"
724    );
725}
726
727#[test]
728fn shorthand_converts_to_orpc_args() {
729    let shorthand: MethodShorthandArgs = syn::parse_quote! { "/planet/list" };
730    let args = shorthand.into_orpc_args("GET");
731
732    assert_eq!(args.method, "GET");
733    assert_eq!(args.path, "/planet/list");
734    assert_eq!(args.stream_event, None);
735}
736
737#[test]
738fn shorthand_with_data_converts_to_orpc_args() {
739    let shorthand: MethodShorthandArgs = syn::parse_quote! { "/stream", data = "EventData" };
740    let args = shorthand.into_orpc_args("GET");
741
742    assert_eq!(args.method, "GET");
743    assert_eq!(args.path, "/stream");
744    assert!(args.stream_event.is_some());
745    let path = args.stream_event.unwrap();
746    assert_eq!(
747        path.segments.first().unwrap().ident.to_string(),
748        "EventData"
749    );
750}
751
752#[test]
753fn shorthand_method_normalized_to_uppercase() {
754    let shorthand: MethodShorthandArgs = syn::parse_quote! { "/test" };
755    let args = shorthand.into_orpc_args("get"); // lowercase input
756
757    assert_eq!(args.method, "GET"); // should be uppercase
758}