Skip to main content

rustra_macros/
lib.rs

1//! # rustra-macros — rustra용 proc macro
2//!
3//! `#[command]` 속성 매크로와 `register!` 매크로를 제공합니다.
4//!
5//! 직접 이 crate를 사용하지 말고 `rustra` crate를 통해 사용하세요:
6//!
7//! ```rust,ignore
8//! use rustra::prelude::*;
9//!
10//! #[command]
11//! fn add_numbers(input: AddInput) -> Result<AddOutput> { ... }
12//! ```
13
14use proc_macro::TokenStream;
15use proc_macro2::TokenStream as TokenStream2;
16use quote::quote;
17use syn::{
18    DeriveInput, GenericArgument, Ident, ItemFn, LitStr, PathArguments, ReturnType, Token, Type,
19    parse::Parse, parse::ParseStream, parse_macro_input,
20};
21
22/// `#[command]` 속성의 파싱 결과입니다.
23///
24/// `#[command]` 또는 `#[command(name = "customName")]` 형태를 지원합니다.
25struct CommandAttr {
26    /// 명시적으로 지정한 명령 이름. 없으면 함수 이름에서 자동 추론합니다.
27    name: Option<String>,
28}
29
30/// `#[command]` 속성의 입력을 파싱합니다.
31///
32/// 빈 입력(`#[command]`)이면 `name: None`, `#[command(name = "foo")]`면 `name: Some("foo")`입니다.
33impl Parse for CommandAttr {
34    fn parse(input: ParseStream) -> syn::Result<Self> {
35        if input.is_empty() {
36            return Ok(CommandAttr { name: None });
37        }
38
39        let key: Ident = input.parse()?;
40        if key != "name" {
41            return Err(syn::Error::new(key.span(), "expected `name`"));
42        }
43        let _: Token![=] = input.parse()?;
44        let name: LitStr = input.parse()?;
45
46        Ok(CommandAttr {
47            name: Some(name.value()),
48        })
49    }
50}
51
52/// 함수를 rustra 명령으로 등록하는 속성 매크로입니다.
53///
54/// ## 제약 사항
55///
56/// - 정확히 하나의 입력 파라미터를 가져야 합니다.
57/// - 반환 타입은 `Result<O>` 형태여야 합니다.
58/// - 입력 타입 `I`는 `DeserializeOwned + JsonSchema`를 충족해야 합니다.
59/// - 출력 타입 `O`는 `Serialize + JsonSchema`를 충족해야 합니다.
60///
61/// ## 명령 이름 규칙
62///
63/// - `#[command]`: 함수 이름에서 `_command` 접미사를 제거한 뒤 lowerCamelCase로 변환
64///   (예: `add_numbers` → `addNumbers`)
65/// - `#[command(name = "customName")]`: 지정한 이름을 그대로 사용
66///
67/// ## 컴파일 타임 검증
68///
69/// 매크로는 다음을 자동으로 검증합니다:
70/// - 입력 파라미터가 정확히 하나인지
71/// - 반환 타입이 `Result<O>` 형태인지
72/// - 입출력 타입이 필요한 trait bound를 충족하는지
73///
74/// ## 예제
75///
76/// ```rust,ignore
77/// #[command]
78/// fn add_numbers(input: AddNumbersInput) -> Result<AddNumbersOutput> {
79///     Ok(AddNumbersOutput { value: input.a + input.b })
80/// }
81///
82/// #[command(name = "multiply")]
83/// fn mul(input: MulInput) -> Result<MulOutput> { ... }
84/// ```
85#[proc_macro_attribute]
86pub fn command(attr: TokenStream, item: TokenStream) -> TokenStream {
87    let attr = parse_macro_input!(attr as CommandAttr);
88    let func = parse_macro_input!(item as ItemFn);
89
90    if func.sig.inputs.is_empty() {
91        return syn::Error::new_spanned(
92            &func.sig,
93            "#[command] requires exactly one input parameter, but none were provided",
94        )
95        .to_compile_error()
96        .into();
97    }
98
99    if func.sig.inputs.len() > 1 {
100        return syn::Error::new_spanned(
101            &func.sig.inputs[1],
102            "#[command] requires exactly one input parameter; remove additional parameters",
103        )
104        .to_compile_error()
105        .into();
106    }
107
108    if func.sig.asyncness.is_some() {
109        return syn::Error::new_spanned(
110            func.sig.fn_token,
111            "#[command] functions must be synchronous — use `fn` not `async fn`",
112        )
113        .to_compile_error()
114        .into();
115    }
116
117    let input_type = match &func.sig.inputs[0] {
118        syn::FnArg::Typed(pat_type) => &*pat_type.ty,
119        _ => {
120            return syn::Error::new_spanned(
121                &func.sig.inputs[0],
122                "#[command] parameter must be a typed value (e.g., `input: MyInput`), not `self`",
123            )
124            .to_compile_error()
125            .into();
126        }
127    };
128
129    let output_type = match &func.sig.output {
130        ReturnType::Type(_, ty) => match extract_result_inner(ty) {
131            Some(inner) => inner,
132            None => {
133                return syn::Error::new_spanned(
134                    ty,
135                    "#[command] must return `Result<O>` where O: Serialize + JsonSchema",
136                )
137                .to_compile_error()
138                .into();
139            }
140        },
141        _ => {
142            return syn::Error::new_spanned(
143                &func.sig,
144                "#[command] must have an explicit return type `Result<O>`",
145            )
146            .to_compile_error()
147            .into();
148        }
149    };
150
151    let fn_name = &func.sig.ident;
152    let command_name = attr.name.unwrap_or_else(|| {
153        let raw = fn_name.to_string();
154        snake_to_lower_camel(raw.trim_end_matches("_command"))
155    });
156    let meta_ident = Ident::new(
157        &format!("__RUstra_meta_{}", fn_name),
158        proc_macro2::Span::call_site(),
159    );
160
161    let expanded = quote! {
162        #func
163
164        #[allow(non_upper_case_globals, dead_code)]
165        const #meta_ident: &str = #command_name;
166
167        #[allow(dead_code)]
168        const _: () = {
169            fn _assert_command_bounds<
170                __I: rustra::__private::CommandInput,
171                __O: rustra::__private::CommandOutput,
172            >() {
173            }
174            fn _check_command_bounds() {
175                _assert_command_bounds::<#input_type, #output_type>();
176            }
177        };
178    };
179
180    expanded.into()
181}
182
183/// `register!` 매크로의 파싱 결과입니다.
184///
185/// 형태: `<builder_expr>, <fn_ident>, <fn_ident>, ...`
186struct RegisterInput {
187    /// 패키지 빌더 표현식 (예: `Package::builder("my.pkg")`)
188    builder: syn::Expr,
189    /// 등록할 `#[command]` 함수 식별자 목록입니다.
190    commands: Vec<Ident>,
191}
192
193/// `register!` 매크로 입력을 파싱합니다.
194///
195/// 형태: `<builder>, <fn1>, <fn2>, ...` (쉼표로 구분, 마지막 쉼표는 선택)
196impl Parse for RegisterInput {
197    fn parse(input: ParseStream) -> syn::Result<Self> {
198        let builder: syn::Expr = input.parse()?;
199        let _: Token![,] = input.parse()?;
200
201        let mut commands = Vec::new();
202        loop {
203            let name: Ident = input.parse()?;
204            commands.push(name);
205            if input.parse::<Token![,]>().is_err() {
206                break;
207            }
208        }
209
210        Ok(RegisterInput { builder, commands })
211    }
212}
213
214/// Registers `#[command]` functions with a package builder.
215///
216/// ```ignore
217/// rustra::register!(Package::builder("my.pkg"), add_numbers, multiply)
218///     .build()
219/// ```
220#[proc_macro]
221pub fn register(input: TokenStream) -> TokenStream {
222    let input = parse_macro_input!(input as RegisterInput);
223
224    if input.commands.is_empty() {
225        return syn::Error::new(
226            proc_macro2::Span::call_site(),
227            "register! requires at least one command function after the builder expression",
228        )
229        .to_compile_error()
230        .into();
231    }
232
233    let builder = &input.builder;
234    let chain: TokenStream2 = input
235        .commands
236        .iter()
237        .map(|fn_name| {
238            let meta_ident = Ident::new(
239                &format!("__RUstra_meta_{}", fn_name),
240                proc_macro2::Span::call_site(),
241            );
242            quote! { .command(#meta_ident, #fn_name) }
243        })
244        .collect();
245
246    let expanded = quote! {
247        #builder #chain
248    };
249
250    expanded.into()
251}
252
253/// `Result<O>` 타입에서 내부 `O` 타입을 추출합니다.
254///
255/// `Result<O>`가 아니면 `None`을 반환합니다.
256fn extract_result_inner(ty: &Type) -> Option<TokenStream2> {
257    let Type::Path(type_path) = ty else {
258        return None;
259    };
260    let segment = type_path.path.segments.last()?;
261    if segment.ident != "Result" {
262        return None;
263    }
264    let PathArguments::AngleBracketed(args) = &segment.arguments else {
265        return None;
266    };
267    let GenericArgument::Type(inner_ty) = args.args.first()? else {
268        return None;
269    };
270    Some(quote! { #inner_ty })
271}
272
273/// struct/enum에 bridge에 필요한 derive와 serde 설정을 자동 추가하는 속성 매크로입니다.
274///
275/// 다음을 자동으로 추가합니다:
276/// - `#[derive(Debug, serde::Serialize, serde::Deserialize, schemars::JsonSchema)]`
277/// - `#[serde(rename_all = "camelCase")]` (기존 serde rename 속성이 없을 시)
278///
279/// ## 예제
280///
281/// ```rust,ignore
282/// #[bridge_type]
283/// #[derive(Clone)]
284/// struct AddNumbersInput { a: i64, b: i64 }
285/// ```
286#[proc_macro_attribute]
287pub fn bridge_type(_attr: TokenStream, item: TokenStream) -> TokenStream {
288    let mut input = parse_macro_input!(item as DeriveInput);
289
290    // Add derive attributes for bridge-required traits
291    input.attrs.push(syn::parse_quote! {
292        #[derive(Debug, serde::Serialize, serde::Deserialize, schemars::JsonSchema)]
293    });
294
295    // Add serde rename_all = "camelCase" only if no serde(rename_all = ...) exists
296    let has_serde_rename = input.attrs.iter().any(|attr| {
297        if !attr.path().is_ident("serde") {
298            return false;
299        }
300        let Ok(nested) = attr.parse_args_with(
301            syn::punctuated::Punctuated::<syn::MetaNameValue, syn::Token![,]>::parse_terminated,
302        ) else {
303            return false;
304        };
305        nested.iter().any(|nv| nv.path.is_ident("rename_all"))
306    });
307
308    if !has_serde_rename {
309        input.attrs.push(syn::parse_quote! {
310            #[serde(rename_all = "camelCase")]
311        });
312    }
313
314    quote! { #input }.into()
315}
316
317/// `build!` 매크로의 파싱 결과입니다.
318///
319/// 형태: `"package.name", <fn_ident>, <fn_ident>, ...`
320struct BuildInput {
321    package_name: LitStr,
322    commands: Vec<Ident>,
323}
324
325impl Parse for BuildInput {
326    fn parse(input: ParseStream) -> syn::Result<Self> {
327        let package_name: LitStr = input.parse()?;
328        let _: Token![,] = input.parse()?;
329
330        let mut commands = Vec::new();
331        loop {
332            if input.is_empty() {
333                break;
334            }
335            let name: Ident = input.parse()?;
336            commands.push(name);
337            if input.parse::<Token![,]>().is_err() {
338                break;
339            }
340        }
341
342        if commands.is_empty() {
343            return Err(syn::Error::new(
344                package_name.span(),
345                "build! requires at least one command function after the package name",
346            ));
347        }
348
349        Ok(BuildInput {
350            package_name,
351            commands,
352        })
353    }
354}
355
356/// `#[command]` 함수들을 간결하게 등록하는 매크로입니다.
357///
358/// `register!(Package::builder("name"), fn1, fn2).build()` 대신
359/// `rustra::build!("name", fn1, fn2).done()`을 사용할 수 있습니다.
360///
361/// ## 예제
362///
363/// ```rust,ignore
364/// pub fn my_package() -> Package {
365///     rustra::build!("com.example.my", add_numbers, multiply).done()
366/// }
367/// ```
368#[proc_macro]
369pub fn build(input: TokenStream) -> TokenStream {
370    let input = parse_macro_input!(input as BuildInput);
371
372    let package_name = &input.package_name;
373    let chain: TokenStream2 = input
374        .commands
375        .iter()
376        .map(|fn_name| {
377            let meta_ident = Ident::new(
378                &format!("__RUstra_meta_{}", fn_name),
379                proc_macro2::Span::call_site(),
380            );
381            quote! { .command(#meta_ident, #fn_name) }
382        })
383        .collect();
384
385    let expanded = quote! {
386        rustra::Package::builder(#package_name) #chain
387    };
388
389    expanded.into()
390}
391
392/// snake_case, kebab-case를 lowerCamelCase로 변환합니다.
393///
394/// 예: `add_numbers` → `addNumbers`, `my_func_command` → `myFunc`
395fn snake_to_lower_camel(name: &str) -> String {
396    let mut output = String::new();
397    let mut uppercase_next = false;
398
399    for character in name.chars() {
400        if character == '_' || character == '-' || character == '.' {
401            uppercase_next = true;
402            continue;
403        }
404
405        if output.is_empty() {
406            output.push(character.to_ascii_lowercase());
407        } else if uppercase_next {
408            output.push(character.to_ascii_uppercase());
409            uppercase_next = false;
410        } else {
411            output.push(character);
412        }
413    }
414
415    output
416}