Skip to main content

elfo_macros_impl/
message.rs

1use proc_macro2::TokenStream;
2use quote::{quote, ToTokens};
3use syn::{
4    parenthesized,
5    parse::{Error as ParseError, Parse, ParseStream},
6    parse_macro_input,
7    spanned::Spanned,
8    Data, DeriveInput, Ident, LitStr, Path, Token, Type,
9};
10
11use crate::errors::emit_error;
12
13#[derive(Debug)]
14struct MessageArgs {
15    name: Option<LitStr>,
16    protocol: Option<LitStr>,
17    ret: Option<Type>,
18    part: bool,
19    transparent: bool,
20    dumping_allowed: Option<bool>,
21    crate_: Option<Path>,
22    not: Vec<String>,
23}
24
25impl Parse for MessageArgs {
26    fn parse(input: ParseStream<'_>) -> Result<Self, ParseError> {
27        let mut args = MessageArgs {
28            ret: None,
29            name: None,
30            protocol: None,
31            part: false,
32            transparent: false,
33            dumping_allowed: None,
34            crate_: None,
35            not: Vec::new(),
36        };
37
38        // `#[message]`
39        // `#[message(name = "N")]`
40        // `#[message(protocol = "P")]`
41        // `#[message(ret = A)]`
42        // `#[message(part)]`
43        // `#[message(part, transparent)]`
44        // `#[message(elfo = some)]`
45        // `#[message(not(Debug))]`
46        // `#[message(dumping = "disabled")]`
47        while !input.is_empty() {
48            let ident: Ident = input.parse()?;
49
50            match ident.to_string().as_str() {
51                "name" => {
52                    let _: Token![=] = input.parse()?;
53                    args.name = Some(input.parse()?);
54                }
55                "protocol" => {
56                    let _: Token![=] = input.parse()?;
57                    args.protocol = Some(input.parse()?);
58                }
59                "ret" => {
60                    let _: Token![=] = input.parse()?;
61                    args.ret = Some(input.parse()?);
62                }
63                "part" => args.part = true,
64                "transparent" => args.transparent = true,
65                "dumping" => {
66                    // TODO: introduce `DumpingMode`.
67                    let _: Token![=] = input.parse()?;
68                    let s: LitStr = input.parse()?;
69
70                    if s.value() == "disabled" {
71                        args.dumping_allowed = Some(false);
72                    } else {
73                        return Err(input.error("only `dumping = \"disabled\"` is supported"));
74                    }
75                }
76                // TODO: call it `crate` like in linkme?
77                "elfo" => {
78                    let _: Token![=] = input.parse()?;
79                    args.crate_ = Some(input.parse()?);
80                }
81                "not" => {
82                    let content;
83                    parenthesized!(content in input);
84                    args.not = content
85                        .parse_terminated(Ident::parse, Token![,])?
86                        .iter()
87                        .map(|ident| ident.to_string())
88                        .collect();
89                }
90                _ => return Err(input.error("unknown attribute")),
91            }
92
93            if !input.is_empty() {
94                let _: Token![,] = input.parse()?;
95            }
96        }
97
98        Ok(args)
99    }
100}
101
102impl MessageArgs {
103    fn validate(&self) {
104        if self.part {
105            fn incompatible(spanned: &Option<impl Spanned>, name: &str) {
106                if let Some(span) = spanned.as_ref().map(|s| s.span()) {
107                    emit_error!(span, "`part` and `{name}` attributes are incompatible");
108                }
109            }
110
111            incompatible(&self.ret, "ret");
112            incompatible(&self.name, "name");
113            incompatible(&self.protocol, "protocol");
114            incompatible(&self.dumping_allowed, "dumping_allowed");
115        }
116    }
117}
118
119fn gen_derive_attr(blacklist: &[String], name: &str, path: TokenStream) -> TokenStream {
120    blacklist
121        .iter()
122        .all(|x| x != name)
123        .then(|| quote! { #[derive(#path)] })
124        .into_token_stream()
125}
126
127fn gen_impl_debug(input: &DeriveInput) -> TokenStream {
128    let name = &input.ident;
129    let field = match &input.data {
130        Data::Struct(data) if data.fields.len() == 1 => Some(data.fields.iter().next().unwrap()),
131        _ => None,
132    };
133
134    let Some(field) = field else {
135        emit_error!(
136            name.span(),
137            "`transparent` is applicable only for structs with one field"
138        );
139        return TokenStream::new();
140    };
141
142    let propagate_fmt = if let Some(ident) = field.ident.as_ref() {
143        quote! { ::std::fmt::Debug::fmt(&self.#ident, f) }
144    } else {
145        quote! { ::std::fmt::Debug::fmt(&self.0, f) }
146    };
147
148    quote! {
149        #[automatically_derived]
150        impl ::std::fmt::Debug for #name {
151            #[inline]
152            fn fmt(&self, f: &mut ::std::fmt::Formatter<'_>) -> ::std::fmt::Result {
153                #propagate_fmt
154            }
155        }
156    }
157}
158
159/// Implementation of the `#[message]` macro.
160pub fn message_impl(
161    args: proc_macro::TokenStream,
162    input: proc_macro::TokenStream,
163    default_path_to_elfo: Path,
164) -> proc_macro::TokenStream {
165    let args = parse_macro_input!(args as MessageArgs);
166    args.validate();
167
168    let crate_ = args.crate_.unwrap_or(default_path_to_elfo);
169
170    // TODO: what about parsing into something cheaper?
171    let input = parse_macro_input!(input as DeriveInput);
172    let name = &input.ident;
173    let serde_crate = format!("{}::_priv::serde", crate_.to_token_stream());
174    let internal = quote![#crate_::_priv];
175
176    let name_str = args
177        .name
178        .as_ref()
179        .map(LitStr::value)
180        .unwrap_or_else(|| input.ident.to_string());
181
182    let derive_debug =
183        (!args.transparent).then(|| gen_derive_attr(&args.not, "Debug", quote![Debug]));
184    let derive_clone = gen_derive_attr(&args.not, "Clone", quote![Clone]);
185    let derive_serialize =
186        gen_derive_attr(&args.not, "Serialize", quote![#internal::serde::Serialize]);
187    let derive_deserialize = gen_derive_attr(
188        &args.not,
189        "Deserialize",
190        quote![#internal::serde::Deserialize],
191    );
192
193    let serde_crate_attr = (!derive_serialize.is_empty() || !derive_deserialize.is_empty())
194        .then(|| quote! { #[serde(crate = #serde_crate)] });
195
196    let serde_transparent_attr = args.transparent.then(|| quote! { #[serde(transparent)] });
197
198    // TODO: pass to `ElfoResponseWrapper`.
199    let dumping_allowed = args.dumping_allowed.unwrap_or(true);
200
201    let protocol = if let Some(protocol) = &args.protocol {
202        quote! { #protocol }
203    } else {
204        quote! { #crate_::get_protocol!() }
205    };
206
207    let impl_message = (!args.part).then(|| {
208        quote! {
209            #[automatically_derived]
210            impl #crate_::Message for #name {
211                #[inline(always)]
212                fn _type_id() -> #internal::MessageTypeId {
213                    #internal::MessageTypeId::new(VTABLE)
214                }
215
216                #[inline(always)]
217                fn _vtable(&self) -> &'static #internal::MessageVTable {
218                    VTABLE
219                }
220            }
221
222            #[#internal::linkme::distributed_slice(#internal::MESSAGE_VTABLES_LIST)]
223            #[linkme(crate = #internal::linkme)]
224            static VTABLE: &#internal::MessageVTable = &#internal::MessageVTable::new::<#name>(
225                #name_str,
226                #protocol,
227                #dumping_allowed
228            );
229        }
230    });
231
232    let impl_request = args.ret.as_ref().map(|ret| {
233        let wrapper_name_str = format!("{name_str}::Response");
234        let protocol = args.protocol.as_ref().map(|p| quote! { protocol = #p, });
235
236        quote! {
237            #[automatically_derived]
238            impl #crate_::Request for #name {
239                type Response = #ret;
240                type Wrapper = ElfoResponseWrapper;
241            }
242
243            #[message(not(Debug), #protocol name = #wrapper_name_str, elfo = #crate_)]
244            pub struct ElfoResponseWrapper(#ret);
245
246            #[automatically_derived]
247            impl ::std::fmt::Debug for ElfoResponseWrapper {
248                #[inline]
249                fn fmt(&self, f: &mut ::std::fmt::Formatter<'_>) -> ::std::fmt::Result {
250                    ::std::fmt::Debug::fmt(&self.0, f)
251                }
252            }
253
254            #[automatically_derived]
255            impl From<#ret> for ElfoResponseWrapper {
256                #[inline]
257                fn from(inner: #ret) -> Self {
258                    ElfoResponseWrapper(inner)
259                }
260            }
261
262            #[automatically_derived]
263            impl From<ElfoResponseWrapper> for #ret {
264                #[inline]
265                fn from(wrapper: ElfoResponseWrapper) -> Self {
266                    wrapper.0
267                }
268            }
269        }
270    });
271
272    let impl_debug =
273        (args.transparent && args.not.iter().all(|x| x != "Debug")).then(|| gen_impl_debug(&input));
274
275    // Don't add `use` statements here to avoid possible collisions with user code.
276    let expanded = quote! {
277        #derive_debug
278        #derive_clone
279        #derive_serialize
280        #derive_deserialize
281        #serde_crate_attr
282        #serde_transparent_attr
283        #input
284
285        #[doc(hidden)]
286        #[allow(unreachable_code)] // for `enum Impossible {}`
287        const _: () = {
288            #impl_message
289            #impl_request
290            #impl_debug
291        };
292    };
293
294    // Errors must be checked after expansion, otherwise some errors can be lost.
295    if let Some(errors) = crate::errors::into_tokens() {
296        quote! { #expanded #errors }.into()
297    } else {
298        expanded.into()
299    }
300}