Skip to main content

teloxide_plugins_macros/
lib.rs

1#![allow(non_snake_case)]
2
3use proc_macro::TokenStream;
4use proc_macro2;
5use quote::quote;
6use syn::{
7    parse::Parse, parse::ParseStream, parse_macro_input, punctuated::Punctuated, Expr, ExprArray,
8    ExprLit, ItemFn, Lit, LitStr, Meta, MetaNameValue, Token,
9};
10
11const COMMANDS_IDENT: &str = "commands";
12const PREFIXES_IDENT: &str = "prefixes";
13const REGEX_IDENT: &str = "regex";
14const CALLBACK_IDENT: &str = "callback";
15
16struct PluginArgs {
17    metas: Punctuated<Meta, Token![,]>,
18}
19
20impl Parse for PluginArgs {
21    fn parse(input: ParseStream) -> syn::Result<Self> {
22        Ok(PluginArgs {
23            metas: Punctuated::parse_terminated(input)?,
24        })
25    }
26}
27
28fn extract_strings_from_array(expr: &Expr) -> syn::Result<Vec<String>> {
29    match expr {
30        Expr::Array(ExprArray { elems, .. }) => {
31            let mut strings = Vec::new();
32            for elem in elems {
33                match elem {
34                    Expr::Lit(ExprLit {
35                        lit: Lit::Str(lit_str),
36                        ..
37                    }) => {
38                        strings.push(lit_str.value());
39                    }
40                    _ => {
41                        return Err(syn::Error::new_spanned(
42                            elem,
43                            "expected string literal in array",
44                        ));
45                    }
46                }
47            }
48            Ok(strings)
49        }
50        _ => Err(syn::Error::new_spanned(
51            expr,
52            "expected array of string literals",
53        )),
54    }
55}
56
57fn create_optional_string_literal(value: Option<&String>) -> proc_macro2::TokenStream {
58    value
59        .map(|s| {
60            let lit_str = LitStr::new(s, proc_macro2::Span::call_site());
61            quote! { Some(#lit_str) }
62        })
63        .unwrap_or_else(|| quote! { None })
64}
65
66fn parse_plugin_args(
67    args: TokenStream,
68) -> syn::Result<(Vec<String>, Vec<String>, Option<String>, Option<String>)> {
69    let mut commands = Vec::new();
70    let mut prefixes = Vec::new();
71    let mut regex = None;
72    let mut callback_filter = None;
73
74    let plugin_args: PluginArgs = syn::parse(args)?;
75
76    for meta in plugin_args.metas {
77        if let Meta::NameValue(MetaNameValue { path, value, .. }) = meta {
78            if let Some(ident) = path.get_ident() {
79                match ident.to_string().as_str() {
80                    COMMANDS_IDENT => {
81                        commands = extract_strings_from_array(&value)?;
82                    }
83                    PREFIXES_IDENT => {
84                        prefixes = extract_strings_from_array(&value)?;
85                    }
86                    REGEX_IDENT => {
87                        let patterns = extract_strings_from_array(&value)?;
88                        if !patterns.is_empty() {
89                            if patterns.len() == 1 {
90                                regex = Some(patterns[0].clone());
91                            } else {
92                                let combined_pattern = patterns.join("|");
93                                regex = Some(combined_pattern);
94                            }
95                        }
96                    }
97                    CALLBACK_IDENT => {
98                        let patterns = extract_strings_from_array(&value)?;
99                        if !patterns.is_empty() {
100                            if patterns.len() == 1 {
101                                callback_filter = Some(patterns[0].clone());
102                            } else {
103                                let combined_pattern = patterns.join("|");
104                                callback_filter = Some(combined_pattern);
105                            }
106                        }
107                    }
108                    _ => {}
109                }
110            }
111        }
112    }
113
114    Ok((commands, prefixes, regex, callback_filter))
115}
116
117fn determine_handler_type(
118    commands: &[String],
119    prefixes: &[String],
120    regex: &Option<String>,
121    callback_filter: &Option<String>,
122) -> syn::Result<bool> {
123    let has_message_triggers = !commands.is_empty() || !prefixes.is_empty() || regex.is_some();
124    let has_callback_triggers = callback_filter.is_some();
125
126    match (has_message_triggers, has_callback_triggers) {
127        (true, true) => {
128            Err(syn::Error::new(
129                proc_macro2::Span::call_site(),
130                "plugin cannot handle both message triggers (commands/prefixes/regex) and callback triggers simultaneously"
131            ))
132        }
133        (true, false) => Ok(false),
134        (false, true) => Ok(true),
135        (false, false) => {
136            Err(syn::Error::new(
137                proc_macro2::Span::call_site(),
138                "plugin must specify at least one trigger: commands, prefixes, regex, or callback"
139            ))
140        }
141    }
142}
143
144fn create_callback_handler(fn_name: &syn::Ident, is_callback: bool) -> proc_macro2::TokenStream {
145    if is_callback {
146        quote! {
147            |ctx| Box::pin(async move {
148                if let Some(cq) = ctx.callback_query {
149                    #fn_name(ctx.bot, cq).await;
150                }
151            })
152        }
153    } else {
154        quote! {
155            |ctx| Box::pin(async move {
156                if let Some(msg) = ctx.message {
157                    #fn_name(ctx.bot, msg).await;
158                }
159            })
160        }
161    }
162}
163
164#[proc_macro_attribute]
165pub fn TeloxidePlugin(args: TokenStream, input: TokenStream) -> TokenStream {
166    let input_fn = parse_macro_input!(input as ItemFn);
167
168    let fn_name = &input_fn.sig.ident;
169    let fn_name_str = fn_name.to_string();
170    let ctor_fn_name = syn::Ident::new(&format!("{}_ctor", fn_name_str), fn_name.span());
171    let static_name = syn::Ident::new(&format!("{}_meta", fn_name_str), fn_name.span());
172    let vis = &input_fn.vis;
173    let sig = &input_fn.sig;
174    let block = &input_fn.block;
175
176    let (commands, prefixes, regex, callback_filter) = match parse_plugin_args(args) {
177        Ok(result) => result,
178        Err(err) => return err.to_compile_error().into(),
179    };
180
181    let is_callback_handler =
182        match determine_handler_type(&commands, &prefixes, &regex, &callback_filter) {
183            Ok(is_callback) => is_callback,
184            Err(err) => return err.to_compile_error().into(),
185        };
186
187    let commands_lit = commands
188        .iter()
189        .map(|c| LitStr::new(c, proc_macro2::Span::call_site()));
190    let prefixes_lit = prefixes
191        .iter()
192        .map(|p| LitStr::new(p, proc_macro2::Span::call_site()));
193    let regex_lit = create_optional_string_literal(regex.as_ref());
194    let callback_filter_lit = create_optional_string_literal(callback_filter.as_ref());
195
196    let callback_handler = create_callback_handler(fn_name, is_callback_handler);
197
198    let expanded = quote! {
199        #vis #sig #block
200
201        #[allow(non_upper_case_globals)]
202        #[doc(hidden)]
203        static #static_name: &teloxide_plugins::registry::PluginMeta = &teloxide_plugins::registry::PluginMeta {
204            name: #fn_name_str,
205            commands: &[#(#commands_lit),*],
206            prefixes: &[#(#prefixes_lit),*],
207            regex: #regex_lit,
208            callback_filter: #callback_filter_lit,
209            callback: #callback_handler,
210        };
211
212        #[ctor::ctor]
213        fn #ctor_fn_name() {
214            teloxide_plugins::registry::register_plugin(#static_name);
215        }
216    };
217
218    TokenStream::from(expanded)
219}