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 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 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 let stream_event_witness = match &args.stream_event {
328 Some(type_path) => {
329 quote! {
330 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 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
396fn 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
457fn emit_schema_registrations(func: &ItemFn) -> TokenStream {
462 let mut seen = std::collections::HashSet::new();
463 let mut registrations = Vec::new();
464
465 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 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 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 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
525fn emit_stream_event_schema_registration(stream_event: &Option<syn::Path>) -> TokenStream {
529 match stream_event {
530 Some(type_path) => {
531 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#[cfg(test)]
558mod tests {
559 use super::*;
560 use syn::parse_quote;
561
562 #[test]
563 fn parse_data_type_string() {
564 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 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 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 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 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 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 let tokens = result.unwrap().to_string();
647
648 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 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#[test]
668fn parse_shorthand_path_only() {
669 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 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 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 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"); assert_eq!(args.method, "GET"); }