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    let output_type_str = type_display(&sig.output_type);
304
305    let error_type_token = match &sig.error_type {
306        Some(ty) => {
307            let s = type_display(ty);
308            quote! { Some(#s) }
309        }
310        None => quote! { None },
311    };
312
313    let stream_event_token = match &args.stream_event {
314        Some(type_path) => {
315            // Extract bare name (last segment) for metadata storage
316            let bare_name = type_path
317                .segments
318                .last()
319                .map(|seg| seg.ident.to_string())
320                .unwrap_or_else(|| "Unknown".to_string());
321            quote! { Some(#bare_name) }
322        }
323        None => quote! { None },
324    };
325
326    // Generate compile-time witness for SSE event type validation
327    let stream_event_witness = match &args.stream_event {
328        Some(type_path) => {
329            quote! {
330                // Compile-time witness: validates type is Serialize
331                // ZodTs derive is validated by inventory schema registration
332                const _: () = {
333                    fn assert_serialize<T: ::serde::Serialize>() {}
334                    fn check() {
335                        assert_serialize::<#type_path>();
336                    }
337                };
338            }
339        }
340        None => quote! {},
341    };
342
343    let input_type_str = match &sig.input_type {
344        Some(ty) => type_display(ty),
345        None => "()".to_string(),
346    };
347
348    let query_type_token = match &sig.query_type {
349        Some(ty) => {
350            let s = type_display(ty);
351            quote! { Some(#s) }
352        }
353        None => quote! { None },
354    };
355
356    // Encode path param types as comma-separated string: "i32,String"
357    // Order matches the path template param order (declaration order in signature)
358    let path_param_types_str = sig
359        .path_params
360        .iter()
361        .map(|(_, ty)| type_display(ty))
362        .collect::<Vec<_>>()
363        .join(",");
364
365    let registration = emit_handler_registration(fn_name, method, path, &sig.state_type);
366    let schema_registrations = emit_schema_registrations(&func);
367    let stream_event_schema_reg = emit_stream_event_schema_registration(&args.stream_event);
368
369    Ok(quote! {
370        #func
371
372        #stream_event_witness
373
374        ::rorpc::inventory::submit! {
375            ::rorpc::HandlerMetadata {
376                name: #fn_name_str,
377                method: #method,
378                path: #path,
379                input_type_name: #input_type_str,
380                query_type_name: #query_type_token,
381                output_type_name: #output_type_str,
382                module_path: ::std::module_path!(),
383                namespace: None,
384                error_type_name: #error_type_token,
385                stream_event_type_name: #stream_event_token,
386                path_param_types: #path_param_types_str,
387            }
388        }
389
390        #registration
391        #schema_registrations
392        #stream_event_schema_reg
393    })
394}
395
396// ---------------------------------------------------------------------------
397// Handler registration factory
398// ---------------------------------------------------------------------------
399
400fn emit_handler_registration(
401    fn_name: &syn::Ident,
402    method: &str,
403    path: &str,
404    state_type: &Option<syn::Type>,
405) -> TokenStream {
406    if let Some(state_ty) = state_type {
407        quote! {
408            ::rorpc::inventory::submit! {
409                ::rorpc::HandlerRegistration {
410                    path: #path,
411                    method: #method,
412                    factory: |state: ::std::sync::Arc<dyn ::std::any::Any + Send + Sync>, final_path: &str| {
413                        use ::axum::routing::{delete, get, patch, post, put};
414                        let method_router = match #method {
415                            "GET"    => get(#fn_name),
416                            "POST"   => post(#fn_name),
417                            "PUT"    => put(#fn_name),
418                            "PATCH"  => patch(#fn_name),
419                            "DELETE" => delete(#fn_name),
420                            _        => post(#fn_name),
421                        };
422                        if let Some(typed_state) = state.downcast_ref::<#state_ty>() {
423                            ::axum::Router::new()
424                                .route(final_path, method_router)
425                                .with_state(typed_state.clone())
426                        } else {
427                            ::axum::Router::new()
428                        }
429                    },
430                }
431            }
432        }
433    } else {
434        quote! {
435            ::rorpc::inventory::submit! {
436                ::rorpc::HandlerRegistration {
437                    path: #path,
438                    method: #method,
439                    factory: |_state: ::std::sync::Arc<dyn ::std::any::Any + Send + Sync>, final_path: &str| {
440                        use ::axum::routing::{delete, get, patch, post, put};
441                        let method_router = match #method {
442                            "GET"    => get(#fn_name),
443                            "POST"   => post(#fn_name),
444                            "PUT"    => put(#fn_name),
445                            "PATCH"  => patch(#fn_name),
446                            "DELETE" => delete(#fn_name),
447                            _        => post(#fn_name),
448                        };
449                        ::axum::Router::new().route(final_path, method_router)
450                    },
451                }
452            }
453        }
454    }
455}
456
457// ---------------------------------------------------------------------------
458// Schema registrations — z.unknown() fallback for types without #[derive(ZodTs)]
459// ---------------------------------------------------------------------------
460
461fn emit_schema_registrations(func: &ItemFn) -> TokenStream {
462    let mut seen = std::collections::HashSet::new();
463    let mut registrations = Vec::new();
464
465    // Collect candidate types from Json<T> and Query<T> params and return type
466    let mut candidates: Vec<&syn::Type> = Vec::new();
467
468    for arg in &func.sig.inputs {
469        if let syn::FnArg::Typed(pat_type) = arg {
470            // Check for Json<T>
471            if let Some(m) = try_extract_wrapper(&pat_type.ty, JSON)
472                && let Some(inner) = m.first_type()
473            {
474                candidates.push(inner);
475            }
476            // Check for Query<T>
477            if let Some(m) = try_extract_wrapper(&pat_type.ty, QUERY)
478                && let Some(inner) = m.first_type()
479            {
480                candidates.push(inner);
481            }
482        }
483    }
484
485    if let syn::ReturnType::Type(_, ty) = &func.sig.output {
486        // Handle both Json<T> and Result<Json<T>, E>
487        if let Some(m) = try_extract_wrapper(ty, JSON) {
488            if let Some(inner) = m.first_type() {
489                candidates.push(inner);
490            }
491        } else if let Some(result_m) = try_extract_wrapper(ty, RESULT)
492            && let Some(first) = result_m.first_type()
493            && let Some(json_m) = try_extract_wrapper(first, JSON)
494            && let Some(inner) = json_m.first_type()
495        {
496            candidates.push(inner);
497        }
498    }
499
500    for ty in candidates {
501        if let Some(custom_ty) = innermost_custom_type(ty) {
502            if is_primitive(custom_ty) {
503                continue;
504            }
505            let name = type_display(custom_ty);
506            if !seen.insert(name.clone()) {
507                continue;
508            }
509            registrations.push(quote! {
510                ::rorpc::inventory::submit! {
511                    ::rorpc::SchemaRegistration {
512                        type_name: #name,
513                        module_path: "",
514                        schema_def: ::rorpc::SchemaDef::Unknown,
515                        dependent_types: || vec![],
516                    }
517                }
518            });
519        }
520    }
521
522    quote! { #(#registrations)* }
523}
524
525/// Generate schema registration for SSE event type specified via `data` attribute.
526///
527/// Similar to `emit_schema_registrations`, but for the stream event type.
528fn emit_stream_event_schema_registration(stream_event: &Option<syn::Path>) -> TokenStream {
529    match stream_event {
530        Some(type_path) => {
531            // Extract bare name for registration
532            let bare_name = type_path
533                .segments
534                .last()
535                .map(|seg| seg.ident.to_string())
536                .unwrap_or_else(|| "Unknown".to_string());
537
538            quote! {
539                ::rorpc::inventory::submit! {
540                    ::rorpc::SchemaRegistration {
541                        type_name: #bare_name,
542                        module_path: "",
543                        schema_def: ::rorpc::SchemaDef::Unknown,
544                        dependent_types: || vec![],
545                    }
546                }
547            }
548        }
549        None => quote! {},
550    }
551}
552
553// ---------------------------------------------------------------------------
554// Tests
555// ---------------------------------------------------------------------------
556
557#[cfg(test)]
558mod tests {
559    use super::*;
560    use syn::parse_quote;
561
562    #[test]
563    fn parse_data_type_string() {
564        // Test that data = "StreamEvent" (string literal) parses correctly
565        let args: OrpcArgs = syn::parse_quote! {
566            method = "GET", path = "/stream", data = "StreamEvent"
567        };
568
569        assert_eq!(args.method, "GET");
570        assert_eq!(args.path, "/stream");
571        assert!(args.stream_event.is_some());
572        let path = args.stream_event.unwrap();
573        assert_eq!(path.segments.len(), 1);
574        assert_eq!(
575            path.segments.first().unwrap().ident.to_string(),
576            "StreamEvent"
577        );
578    }
579
580    #[test]
581    fn parse_data_qualified_path_string() {
582        // Test that data = "crate::models::StreamEvent" works
583        let args: OrpcArgs = syn::parse_quote! {
584            method = "GET", path = "/stream", data = "crate::models::StreamEvent"
585        };
586
587        assert_eq!(args.method, "GET");
588        assert_eq!(args.path, "/stream");
589        assert!(args.stream_event.is_some());
590        // Verify the path was parsed correctly
591        let path = args.stream_event.unwrap();
592        assert_eq!(path.segments.len(), 3);
593        assert_eq!(
594            path.segments.last().unwrap().ident.to_string(),
595            "StreamEvent"
596        );
597    }
598
599    #[test]
600    fn parse_data_type_path_backward_compat() {
601        // Test backward compatibility: data = StreamEvent (bare path)
602        let args: OrpcArgs = syn::parse_quote! {
603            method = "GET", path = "/stream", data = StreamEvent
604        };
605
606        assert_eq!(args.method, "GET");
607        assert_eq!(args.path, "/stream");
608        assert!(args.stream_event.is_some());
609        let path = args.stream_event.unwrap();
610        assert_eq!(path.segments.len(), 1);
611        assert_eq!(
612            path.segments.first().unwrap().ident.to_string(),
613            "StreamEvent"
614        );
615    }
616
617    #[test]
618    fn parse_without_data() {
619        // Test that data is optional
620        let args: OrpcArgs = syn::parse_quote! {
621            method = "POST", path = "/create"
622        };
623
624        assert_eq!(args.method, "POST");
625        assert_eq!(args.path, "/create");
626        assert_eq!(args.stream_event, None);
627    }
628
629    #[test]
630    fn data_type_converts_to_string_literal() {
631        // Verify that when we generate the metadata, data becomes a string literal
632        let args: OrpcArgs = syn::parse_quote! {
633            method = "GET", path = "/stream", data = "StreamEvent"
634        };
635
636        let func: syn::ItemFn = parse_quote! {
637            async fn stream_test() -> Sse<impl Stream<Item = Event>> {
638                todo!()
639            }
640        };
641
642        let result = try_expand_orpc(args, func);
643        assert!(result.is_ok(), "expand_orpc should succeed");
644
645        // Check that the generated code contains Some("StreamEvent") as a string literal
646        let tokens = result.unwrap().to_string();
647
648        // quote! serialises with spaces between tokens, so `Some("StreamEvent")` becomes
649        // `Some ("StreamEvent")`. Check the field name and the quoted value separately.
650        assert!(
651            tokens.contains("stream_event_type_name") && tokens.contains(r#""StreamEvent""#),
652            "Generated code should contain stream_event_type_name: Some(\"StreamEvent\"), got: {}",
653            tokens
654        );
655        // Also assert it is NOT a bare identifier (which would be a type error at compile time)
656        assert!(
657            !tokens.contains("Some (StreamEvent)") && !tokens.contains("Some(StreamEvent)"),
658            "stream_event_type_name must be a string literal, not a bare identifier"
659        );
660    }
661}
662
663// ---------------------------------------------------------------------------
664// MethodShorthandArgs tests
665// ---------------------------------------------------------------------------
666
667#[test]
668fn parse_shorthand_path_only() {
669    // Test: #[rorpc::get("/planet/list")]
670    let args: MethodShorthandArgs = syn::parse_quote! { "/planet/list" };
671
672    assert_eq!(args.path, "/planet/list");
673    assert_eq!(args.data, None);
674}
675
676#[test]
677fn parse_shorthand_with_data_string() {
678    // Test: #[rorpc::get("/stream", data = "EventData")]
679    let args: MethodShorthandArgs = syn::parse_quote! { "/stream", data = "EventData" };
680
681    assert_eq!(args.path, "/stream");
682    assert!(args.data.is_some());
683    let path = args.data.unwrap();
684    assert_eq!(
685        path.segments.first().unwrap().ident.to_string(),
686        "EventData"
687    );
688}
689
690#[test]
691fn parse_shorthand_with_qualified_data() {
692    // Test: #[rorpc::get("/stream", data = "models::EventData")]
693    let args: MethodShorthandArgs = syn::parse_quote! { "/stream", data = "models::EventData" };
694
695    assert_eq!(args.path, "/stream");
696    assert!(args.data.is_some());
697    let path = args.data.unwrap();
698    assert_eq!(path.segments.len(), 2);
699    assert_eq!(path.segments.last().unwrap().ident.to_string(), "EventData");
700}
701
702#[test]
703fn parse_shorthand_with_data_type_path() {
704    // Test backward compat: #[rorpc::get("/stream", data = EventData)]
705    let args: MethodShorthandArgs = syn::parse_quote! { "/stream", data = EventData };
706
707    assert_eq!(args.path, "/stream");
708    assert!(args.data.is_some());
709    let path = args.data.unwrap();
710    assert_eq!(
711        path.segments.first().unwrap().ident.to_string(),
712        "EventData"
713    );
714}
715
716#[test]
717fn shorthand_converts_to_orpc_args() {
718    let shorthand: MethodShorthandArgs = syn::parse_quote! { "/planet/list" };
719    let args = shorthand.into_orpc_args("GET");
720
721    assert_eq!(args.method, "GET");
722    assert_eq!(args.path, "/planet/list");
723    assert_eq!(args.stream_event, None);
724}
725
726#[test]
727fn shorthand_with_data_converts_to_orpc_args() {
728    let shorthand: MethodShorthandArgs = syn::parse_quote! { "/stream", data = "EventData" };
729    let args = shorthand.into_orpc_args("GET");
730
731    assert_eq!(args.method, "GET");
732    assert_eq!(args.path, "/stream");
733    assert!(args.stream_event.is_some());
734    let path = args.stream_event.unwrap();
735    assert_eq!(
736        path.segments.first().unwrap().ident.to_string(),
737        "EventData"
738    );
739}
740
741#[test]
742fn shorthand_method_normalized_to_uppercase() {
743    let shorthand: MethodShorthandArgs = syn::parse_quote! { "/test" };
744    let args = shorthand.into_orpc_args("get"); // lowercase input
745
746    assert_eq!(args.method, "GET"); // should be uppercase
747}