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