1use 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
23const ATTR_METHOD: &str = "method";
28const ATTR_PATH: &str = "path";
29const ATTR_DATA: &str = "data";
30
31pub struct OrpcArgs {
37 pub method: String,
38 pub path: String,
39 pub stream_event: Option<syn::Path>,
40}
41
42pub 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 let path_lit: syn::LitStr = input.parse()?;
58 let path = path_lit.value();
59
60 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 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
127impl 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 Expr::Lit(ExprLit {
208 lit: Lit::Str(s), ..
209 }) => {
210 let path_str = s.value();
211 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 stream_event = Some(parsed_path);
224 }
225 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
281pub 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 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 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 let stream_event_witness = match &args.stream_event {
339 Some(type_path) => {
340 quote! {
341 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 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
407fn 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
468fn emit_schema_registrations(func: &ItemFn) -> TokenStream {
473 let mut seen = std::collections::HashSet::new();
474 let mut registrations = Vec::new();
475
476 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 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 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 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
536fn emit_stream_event_schema_registration(stream_event: &Option<syn::Path>) -> TokenStream {
540 match stream_event {
541 Some(type_path) => {
542 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#[cfg(test)]
569mod tests {
570 use super::*;
571 use syn::parse_quote;
572
573 #[test]
574 fn parse_data_type_string() {
575 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 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 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 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 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 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 let tokens = result.unwrap().to_string();
658
659 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 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#[test]
679fn parse_shorthand_path_only() {
680 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 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 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 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"); assert_eq!(args.method, "GET"); }