Skip to main content

rustra_macros/
lib.rs

1//! # rustra-macros — rustra용 proc macro
2//!
3//! `#[command]` 속성 매크로와 `register!` / `build!` 매크로를 제공합니다.
4//!
5//! 직접 이 crate를 사용하지 말고 `rustra` crate를 통해 사용하세요:
6//!
7//! ```rust
8//! use rustra::prelude::*;
9//!
10//! #[bridge_type]
11//! struct AddInput { a: i64, b: i64 }
12//! #[bridge_type]
13//! struct AddOutput { sum: i64 }
14//!
15//! #[command]
16//! fn add_numbers(input: AddInput) -> Result<AddOutput> {
17//!     Ok(AddOutput { sum: input.a + input.b })
18//! }
19//! ```
20
21use proc_macro::TokenStream;
22use proc_macro2::TokenStream as TokenStream2;
23use quote::quote;
24use syn::{
25    DeriveInput, GenericArgument, Ident, ItemFn, LitStr, PathArguments, ReturnType, Token, Type,
26    parse::Parse, parse::ParseStream, parse_macro_input,
27};
28
29/// `#[command]` 속성의 파싱 결과입니다.
30///
31/// `#[command]` 또는 `#[command(name = "customName")]` 형태를 지원합니다.
32struct CommandAttr {
33    /// 명시적으로 지정한 명령 이름. 없으면 함수 이름에서 자동 추론합니다.
34    name: Option<String>,
35}
36
37/// `#[command]` 속성의 입력을 파싱합니다.
38///
39/// 빈 입력(`#[command]`)이면 `name: None`, `#[command(name = "foo")]`면 `name: Some("foo")`입니다.
40impl Parse for CommandAttr {
41    fn parse(input: ParseStream) -> syn::Result<Self> {
42        if input.is_empty() {
43            return Ok(CommandAttr { name: None });
44        }
45
46        let key: Ident = input.parse()?;
47        if key != "name" {
48            return Err(syn::Error::new(key.span(), "expected `name`"));
49        }
50        let _: Token![=] = input.parse()?;
51        let name: LitStr = input.parse()?;
52
53        Ok(CommandAttr {
54            name: Some(name.value()),
55        })
56    }
57}
58
59/// 함수를 rustra 명령으로 등록하는 속성 매크로입니다.
60///
61/// ## 지원 기능
62///
63/// - 동기(`fn`) 및 비동기(`async fn`) 함수 지원
64/// - 0개 인자(`fn ping() -> Result<()>`) 및 1개 데이터 인자(`fn add(input: Input) -> Result<Output>`)
65/// - `State<T>` 의존성 자동 주입 (`fn get(input: In, db: State<Db>) -> Result<Out>`)
66/// - Rust doc comment(`///`)를 추출하여 메타데이터에 보존
67///
68/// ## 명령 이름 규칙
69///
70/// - `#[command]`: 함수 이름에서 `_command` 접미사를 제거한 뒤 lowerCamelCase로 변환
71///   (예: `add_numbers` → `addNumbers`)
72/// - `#[command(name = "customName")]`: 지정한 이름을 그대로 사용
73#[proc_macro_attribute]
74pub fn command(attr: TokenStream, item: TokenStream) -> TokenStream {
75    let attr = parse_macro_input!(attr as CommandAttr);
76    let func = parse_macro_input!(item as ItemFn);
77
78    // Extract doc comments
79    let docs: Vec<String> = func
80        .attrs
81        .iter()
82        .filter_map(|attr| {
83            if attr.path().is_ident("doc")
84                && let syn::Meta::NameValue(nv) = &attr.meta
85                && let syn::Expr::Lit(syn::ExprLit {
86                    lit: syn::Lit::Str(s),
87                    ..
88                }) = &nv.value
89            {
90                return Some(s.value().trim().to_string());
91            }
92            None
93        })
94        .collect();
95    let doc_comment = docs.join("\n");
96
97    let is_async = func.sig.asyncness.is_some();
98    let fn_name = &func.sig.ident;
99    let vis = &func.vis;
100    let output_type = match &func.sig.output {
101        ReturnType::Type(_, ty) => match extract_result_inner(ty) {
102            Some(inner) => inner,
103            None => {
104                return syn::Error::new_spanned(
105                    ty,
106                    "#[command] must return `Result<O>` where O: Serialize + JsonSchema",
107                )
108                .to_compile_error()
109                .into();
110            }
111        },
112        _ => {
113            return syn::Error::new_spanned(
114                &func.sig,
115                "#[command] must have an explicit return type `Result<O>`",
116            )
117            .to_compile_error()
118            .into();
119        }
120    };
121
122    // Analyze parameters
123    struct ParamInfo {
124        pat: syn::Pat,
125        ty: Type,
126        is_state: bool,
127        state_inner: Option<Type>,
128    }
129
130    let mut params = Vec::new();
131    for input in &func.sig.inputs {
132        match input {
133            syn::FnArg::Receiver(_) => {
134                return syn::Error::new_spanned(input, "#[command] functions cannot accept `self`")
135                    .to_compile_error()
136                    .into();
137            }
138            syn::FnArg::Typed(pat_type) => {
139                let is_state_inner = extract_state_inner(&pat_type.ty);
140                let is_state = is_state_inner.is_some();
141                params.push(ParamInfo {
142                    pat: (*pat_type.pat).clone(),
143                    ty: (*pat_type.ty).clone(),
144                    is_state,
145                    state_inner: is_state_inner,
146                });
147            }
148        }
149    }
150
151    let data_params: Vec<&ParamInfo> = params.iter().filter(|p| !p.is_state).collect();
152    if data_params.len() > 1 {
153        return syn::Error::new_spanned(
154            &func.sig.inputs,
155            "#[command] supports at most one input data parameter (plus optional State<T> parameters)",
156        )
157        .to_compile_error()
158        .into();
159    }
160
161    let input_type = if let Some(data) = data_params.first() {
162        let ty = &data.ty;
163        quote! { #ty }
164    } else {
165        quote! { () }
166    };
167
168    let inner_fn_name = Ident::new(
169        &format!("__rustra_inner_{}", fn_name),
170        proc_macro2::Span::call_site(),
171    );
172    let mut inner_func = func.clone();
173    inner_func.sig.ident = inner_fn_name.clone();
174
175    let command_name = attr.name.unwrap_or_else(|| {
176        let raw = fn_name.to_string();
177        snake_to_lower_camel(raw.trim_end_matches("_command"))
178    });
179    let meta_ident = Ident::new(
180        &format!("__RUstra_meta_{}", fn_name),
181        proc_macro2::Span::call_site(),
182    );
183    let doc_ident = Ident::new(
184        &format!("__RUstra_doc_{}", fn_name),
185        proc_macro2::Span::call_site(),
186    );
187
188    // Prepare state bindings and call args
189    let mut state_bindings = Vec::new();
190    let mut call_args = Vec::new();
191
192    for param in &params {
193        if param.is_state {
194            let pat = &param.pat;
195            let ty = &param.ty;
196            let inner_ty = param.state_inner.as_ref().unwrap();
197            state_bindings.push(quote! {
198                let #pat: #ty = rustra::get_state::<#inner_ty>()
199                    .ok_or_else(|| rustra::RustraError::internal(concat!("State<", stringify!(#inner_ty), "> not managed in package")))?;
200            });
201            call_args.push(quote! { #pat });
202        } else {
203            call_args.push(quote! { __rustra_input });
204        }
205    }
206
207    let outer_input_arg = if data_params.is_empty() {
208        quote! { _: () }
209    } else {
210        quote! { __rustra_input: #input_type }
211    };
212
213    let inner_invocation = if is_async {
214        quote! {
215            rustra::__private::block_on(async move {
216                #inner_fn_name(#(#call_args),*).await
217            })
218        }
219    } else {
220        quote! {
221            #inner_fn_name(#(#call_args),*)
222        }
223    };
224
225    let expanded = quote! {
226        #inner_func
227
228        #vis fn #fn_name(#outer_input_arg) -> rustra::Result<#output_type> {
229            #(#state_bindings)*
230            #inner_invocation
231        }
232
233        #[allow(non_upper_case_globals, dead_code)]
234        const #meta_ident: &str = #command_name;
235
236        #[allow(non_upper_case_globals, dead_code)]
237        const #doc_ident: &str = #doc_comment;
238
239        #[allow(dead_code)]
240        const _: () = {
241            fn _assert_command_bounds<
242                __I: rustra::__private::CommandInput,
243                __O: rustra::__private::CommandOutput,
244            >() {
245            }
246            fn _check_command_bounds() {
247                _assert_command_bounds::<#input_type, #output_type>();
248            }
249        };
250    };
251
252    expanded.into()
253}
254
255/// Type이 `State<T>` 형태인지 검사하고 내부 `T`를 반환합니다.
256fn extract_state_inner(ty: &Type) -> Option<Type> {
257    let Type::Path(type_path) = ty else {
258        return None;
259    };
260    let segment = type_path.path.segments.last()?;
261    if segment.ident == "State"
262        && let PathArguments::AngleBracketed(args) = &segment.arguments
263        && let Some(GenericArgument::Type(inner_ty)) = args.args.first()
264    {
265        return Some(inner_ty.clone());
266    }
267    None
268}
269
270/// `register!` 매크로의 파싱 결과입니다.
271///
272/// 형태: `<builder_expr>, <fn_ident>, <fn_ident>, ...`
273struct RegisterInput {
274    /// 패키지 빌더 표현식 (예: `Package::builder("my.pkg")`)
275    builder: syn::Expr,
276    /// 등록할 `#[command]` 함수 식별자 목록입니다.
277    commands: Vec<Ident>,
278}
279
280/// `register!` 매크로 입력을 파싱합니다.
281///
282/// 형태: `<builder>, <fn1>, <fn2>, ...` (쉼표로 구분, 마지막 쉼표는 선택)
283impl Parse for RegisterInput {
284    fn parse(input: ParseStream) -> syn::Result<Self> {
285        let builder: syn::Expr = input.parse()?;
286        let _: Token![,] = input.parse()?;
287
288        let mut commands = Vec::new();
289        loop {
290            let name: Ident = input.parse()?;
291            commands.push(name);
292            if input.parse::<Token![,]>().is_err() {
293                break;
294            }
295        }
296
297        Ok(RegisterInput { builder, commands })
298    }
299}
300
301/// Registers `#[command]` functions with a package builder.
302///
303/// ```rust
304/// use rustra::prelude::*;
305///
306/// #[bridge_type]
307/// struct AddInput { a: i64, b: i64 }
308/// #[bridge_type]
309/// struct AddOutput { sum: i64 }
310///
311/// #[command]
312/// fn add_numbers(input: AddInput) -> Result<AddOutput> {
313///     Ok(AddOutput { sum: input.a + input.b })
314/// }
315///
316/// let pkg = rustra::register!(Package::builder("my.pkg"), add_numbers).build();
317/// ```
318#[proc_macro]
319pub fn register(input: TokenStream) -> TokenStream {
320    let input = parse_macro_input!(input as RegisterInput);
321
322    if input.commands.is_empty() {
323        return syn::Error::new(
324            proc_macro2::Span::call_site(),
325            "register! requires at least one command function after the builder expression",
326        )
327        .to_compile_error()
328        .into();
329    }
330
331    let builder = &input.builder;
332    let chain: TokenStream2 = input
333        .commands
334        .iter()
335        .map(|fn_name| {
336            let meta_ident = Ident::new(
337                &format!("__RUstra_meta_{}", fn_name),
338                proc_macro2::Span::call_site(),
339            );
340            quote! { .command(#meta_ident, #fn_name) }
341        })
342        .collect();
343
344    let expanded = quote! {
345        #builder #chain
346    };
347
348    expanded.into()
349}
350
351/// `Result<O>` 타입에서 내부 `O` 타입을 추출합니다.
352///
353/// `Result<O>`가 아니면 `None`을 반환합니다.
354fn extract_result_inner(ty: &Type) -> Option<TokenStream2> {
355    let Type::Path(type_path) = ty else {
356        return None;
357    };
358    let segment = type_path.path.segments.last()?;
359    if segment.ident != "Result" {
360        return None;
361    }
362    let PathArguments::AngleBracketed(args) = &segment.arguments else {
363        return None;
364    };
365    let GenericArgument::Type(inner_ty) = args.args.first()? else {
366        return None;
367    };
368    Some(quote! { #inner_ty })
369}
370
371/// struct/enum에 bridge에 필요한 derive와 serde 설정을 자동 추가하는 속성 매크로입니다.
372///
373/// 다음을 자동으로 추가합니다:
374/// - `#[derive(Debug, serde::Serialize, serde::Deserialize, schemars::JsonSchema)]`
375/// - `#[serde(rename_all = "camelCase")]` (기존 serde rename 속성이 없을 시)
376///
377/// ## 예제
378///
379/// ```rust
380/// use rustra::prelude::*;
381///
382/// #[bridge_type]
383/// #[derive(Clone)]
384/// struct AddNumbersInput { a: i64, b: i64 }
385/// ```
386#[proc_macro_attribute]
387pub fn bridge_type(_attr: TokenStream, item: TokenStream) -> TokenStream {
388    let mut input = parse_macro_input!(item as DeriveInput);
389
390    // Add derive attributes for bridge-required traits
391    input.attrs.push(syn::parse_quote! {
392        #[derive(Debug, serde::Serialize, serde::Deserialize, schemars::JsonSchema)]
393    });
394
395    // Add serde rename_all = "camelCase" only if no serde(rename_all = ...) exists
396    let has_serde_rename = input.attrs.iter().any(|attr| {
397        if !attr.path().is_ident("serde") {
398            return false;
399        }
400        let Ok(nested) = attr.parse_args_with(
401            syn::punctuated::Punctuated::<syn::MetaNameValue, syn::Token![,]>::parse_terminated,
402        ) else {
403            return false;
404        };
405        nested.iter().any(|nv| nv.path.is_ident("rename_all"))
406    });
407
408    if !has_serde_rename {
409        input.attrs.push(syn::parse_quote! {
410            #[serde(rename_all = "camelCase")]
411        });
412    }
413
414    quote! { #input }.into()
415}
416
417/// `build!` 매크로의 파싱 결과입니다.
418///
419/// 형태: `"package.name", <fn_ident>, <fn_ident>, ...`
420struct BuildInput {
421    package_name: LitStr,
422    commands: Vec<Ident>,
423}
424
425impl Parse for BuildInput {
426    fn parse(input: ParseStream) -> syn::Result<Self> {
427        let package_name: LitStr = input.parse()?;
428        let _: Token![,] = input.parse()?;
429
430        let mut commands = Vec::new();
431        loop {
432            if input.is_empty() {
433                break;
434            }
435            let name: Ident = input.parse()?;
436            commands.push(name);
437            if input.parse::<Token![,]>().is_err() {
438                break;
439            }
440        }
441
442        if commands.is_empty() {
443            return Err(syn::Error::new(
444                package_name.span(),
445                "build! requires at least one command function after the package name",
446            ));
447        }
448
449        Ok(BuildInput {
450            package_name,
451            commands,
452        })
453    }
454}
455
456/// `#[command]` 함수들을 간결하게 등록하는 매크로입니다.
457///
458/// `register!(Package::builder("name"), fn1, fn2).build()` 대신
459/// `rustra::build!("name", fn1, fn2).done()`을 사용할 수 있습니다.
460///
461/// ## 예제
462///
463/// ```rust
464/// use rustra::prelude::*;
465///
466/// #[bridge_type]
467/// struct AddInput { a: i64, b: i64 }
468/// #[bridge_type]
469/// struct AddOutput { sum: i64 }
470///
471/// #[command]
472/// fn add_numbers(input: AddInput) -> Result<AddOutput> {
473///     Ok(AddOutput { sum: input.a + input.b })
474/// }
475///
476/// let pkg = rustra::build!("com.example.my", add_numbers).done();
477/// ```
478#[proc_macro]
479pub fn build(input: TokenStream) -> TokenStream {
480    let input = parse_macro_input!(input as BuildInput);
481
482    let package_name = &input.package_name;
483    let chain: TokenStream2 = input
484        .commands
485        .iter()
486        .map(|fn_name| {
487            let meta_ident = Ident::new(
488                &format!("__RUstra_meta_{}", fn_name),
489                proc_macro2::Span::call_site(),
490            );
491            quote! { .command(#meta_ident, #fn_name) }
492        })
493        .collect();
494
495    let expanded = quote! {
496        rustra::Package::builder(#package_name) #chain
497    };
498
499    expanded.into()
500}
501
502/// snake_case, kebab-case를 lowerCamelCase로 변환합니다.
503///
504/// 예: `add_numbers` → `addNumbers`, `my_func_command` → `myFunc`
505fn snake_to_lower_camel(name: &str) -> String {
506    let mut output = String::new();
507    let mut uppercase_next = false;
508
509    for character in name.chars() {
510        if character == '_' || character == '-' || character == '.' {
511            uppercase_next = true;
512            continue;
513        }
514
515        if output.is_empty() {
516            output.push(character.to_ascii_lowercase());
517        } else if uppercase_next {
518            output.push(character.to_ascii_uppercase());
519            uppercase_next = false;
520        } else {
521            output.push(character);
522        }
523    }
524
525    output
526}