Skip to main content

formatjs_cli/
rust_extractor.rs

1use crate::extractor::{
2    MessageDescriptor, MessageExtraction, flatten_message_descriptors, normalize_whitespace,
3};
4use crate::id_generator::IdGenerator;
5use anyhow::{Result, bail};
6use proc_macro2::{LineColumn, Span};
7use serde_json::Value;
8use std::path::Path;
9use syn::parse::{Parse, ParseStream};
10use syn::spanned::Spanned;
11use syn::visit::{self, Visit};
12use syn::{Expr, Ident, LitStr, Macro, Token, braced};
13
14const RUST_ID_INTERPOLATION_PATTERN: &str = "[sha512:contenthash:base64:10]";
15
16pub fn extract_messages_from_rust_source(
17    source_text: &str,
18    file_path: &Path,
19    extract_source_location: bool,
20    preserve_whitespace: bool,
21    flatten: bool,
22    throws: bool,
23) -> Result<Vec<MessageDescriptor>> {
24    let extraction = extract_messages_from_rust_source_with_diagnostics(
25        source_text,
26        file_path,
27        extract_source_location,
28        preserve_whitespace,
29        flatten,
30        throws,
31    )?;
32    for error in extraction.errors {
33        eprintln!("{error}");
34    }
35    Ok(extraction.messages)
36}
37
38pub fn extract_messages_from_rust_source_with_diagnostics(
39    source_text: &str,
40    file_path: &Path,
41    extract_source_location: bool,
42    preserve_whitespace: bool,
43    flatten: bool,
44    throws: bool,
45) -> Result<MessageExtraction> {
46    let file = syn::parse_file(source_text)?;
47    let id_generator = IdGenerator::new(RUST_ID_INTERPOLATION_PATTERN)?;
48    let mut extractor = RustMessageExtractor {
49        source_text,
50        file_path,
51        extract_source_location,
52        preserve_whitespace,
53        id_generator,
54        messages: Vec::new(),
55        errors: Vec::new(),
56    };
57    extractor.visit_file(&file);
58    if throws && let Some(error) = extractor.errors.first() {
59        bail!("{error}");
60    }
61    let messages =
62        flatten_message_descriptors(extractor.messages, source_text, file_path, flatten)?;
63    Ok(MessageExtraction {
64        messages,
65        errors: extractor.errors,
66    })
67}
68
69struct RustMessageExtractor<'a> {
70    source_text: &'a str,
71    file_path: &'a Path,
72    extract_source_location: bool,
73    preserve_whitespace: bool,
74    id_generator: IdGenerator,
75    messages: Vec<MessageDescriptor>,
76    errors: Vec<String>,
77}
78
79impl RustMessageExtractor<'_> {
80    fn descriptor(&self, arguments: MessageArgs, span: Span) -> Result<MessageDescriptor> {
81        let (start, end) = if self.extract_source_location {
82            (
83                Some(line_column_to_offset(self.source_text, span.start())),
84                Some(line_column_to_offset(self.source_text, span.end())),
85            )
86        } else {
87            (None, None)
88        };
89        let normalized_default_message = normalize_whitespace(&arguments.default_message);
90        let description = arguments.description.map(Value::String);
91        let id = match arguments.id {
92            Some(id) => id,
93            None => {
94                self.id_generator
95                    .generate(Some(&normalized_default_message), &description, None)?
96            }
97        };
98        let default_message = if self.preserve_whitespace {
99            arguments.default_message
100        } else {
101            normalized_default_message
102        };
103        Ok(MessageDescriptor {
104            id: Some(id),
105            default_message: Some(default_message),
106            description,
107            file: self
108                .extract_source_location
109                .then(|| self.file_path.to_string_lossy().into_owned()),
110            start,
111            end,
112        })
113    }
114
115    fn error(&mut self, span: Span, message: impl AsRef<str>) {
116        let start = span.start();
117        self.errors.push(format!(
118            "{}:{}:{}: {}",
119            self.file_path.display(),
120            start.line,
121            start.column + 1,
122            message.as_ref()
123        ));
124    }
125}
126
127impl<'ast> Visit<'ast> for RustMessageExtractor<'_> {
128    fn visit_macro(&mut self, node: &'ast Macro) {
129        let name = node.path.segments.last().map(|segment| &segment.ident);
130        if name.is_some_and(|name| name == "message_descriptor" || name == "format_message") {
131            let arguments = if name.is_some_and(|name| name == "format_message") {
132                syn::parse2::<FormatMessageArgs>(node.tokens.clone()).map(|args| args.message)
133            } else {
134                syn::parse2::<MessageArgs>(node.tokens.clone())
135            };
136            match arguments {
137                Ok(arguments) => match self.descriptor(arguments, node.span()) {
138                    Ok(descriptor) => self.messages.push(descriptor),
139                    Err(error) => self.error(node.span(), error.to_string()),
140                },
141                Err(error) => self.error(node.span(), error.to_string()),
142            }
143        }
144        visit::visit_macro(self, node);
145    }
146}
147
148struct FormatMessageArgs {
149    message: MessageArgs,
150}
151
152impl Parse for FormatMessageArgs {
153    fn parse(input: ParseStream<'_>) -> syn::Result<Self> {
154        input.parse::<Expr>()?;
155        input.parse::<Token![,]>()?;
156        let message = parse_message_args(input, true)?;
157        Ok(Self { message })
158    }
159}
160
161fn line_column_to_offset(source: &str, location: LineColumn) -> u32 {
162    let line_offset: usize = source
163        .split_inclusive('\n')
164        .take(location.line.saturating_sub(1))
165        .map(str::len)
166        .sum();
167    line_offset
168        .saturating_add(location.column)
169        .min(source.len()) as u32
170}
171
172struct MessageArgs {
173    id: Option<String>,
174    default_message: String,
175    description: Option<String>,
176}
177
178impl Parse for MessageArgs {
179    fn parse(input: ParseStream<'_>) -> syn::Result<Self> {
180        parse_message_args(input, false)
181    }
182}
183
184fn parse_message_args(input: ParseStream<'_>, allow_values: bool) -> syn::Result<MessageArgs> {
185    let mut id = None;
186    let mut default_message = None;
187    let mut description = None;
188    let mut values = false;
189    while !input.is_empty() {
190        let key: Ident = input.parse()?;
191        if input.peek(Token![:]) {
192            input.parse::<Token![:]>()?;
193        } else {
194            input.parse::<Token![=]>()?;
195        }
196        match key.to_string().as_str() {
197            "values" if allow_values && !values => {
198                if input.peek(syn::token::Brace) {
199                    let content;
200                    braced!(content in input);
201                    while !content.is_empty() {
202                        content.parse::<Ident>()?;
203                        content.parse::<Token![:]>()?;
204                        content.parse::<Expr>()?;
205                        if content.peek(Token![,]) {
206                            content.parse::<Token![,]>()?;
207                        } else if !content.is_empty() {
208                            return Err(content.error("expected comma"));
209                        }
210                    }
211                } else {
212                    input.parse::<Expr>()?;
213                }
214                values = true;
215            }
216            "values" if allow_values => {
217                return Err(syn::Error::new(key.span(), "duplicate message field"));
218            }
219            "values" => return Err(syn::Error::new(key.span(), "unknown message field")),
220            field => {
221                let value: LitStr = input.parse()?;
222                match field {
223                    "id" if id.is_none() => id = Some(value.value()),
224                    "default_message" if default_message.is_none() => {
225                        default_message = Some(value.value())
226                    }
227                    "description" if description.is_none() => description = Some(value.value()),
228                    "id" | "default_message" | "description" => {
229                        return Err(syn::Error::new(key.span(), "duplicate message field"));
230                    }
231                    _ => return Err(syn::Error::new(key.span(), "unknown message field")),
232                }
233            }
234        }
235        if input.peek(Token![,]) {
236            input.parse::<Token![,]>()?;
237        } else if !input.is_empty() {
238            return Err(input.error("expected comma"));
239        }
240    }
241
242    Ok(MessageArgs {
243        id,
244        default_message: default_message
245            .ok_or_else(|| syn::Error::new(Span::call_site(), "default_message is required"))?,
246        description,
247    })
248}
249
250#[cfg(test)]
251mod tests {
252    use super::*;
253
254    #[test]
255    fn extracts_message_descriptors() {
256        let source = r#"fn main() {
257            let descriptor = message_descriptor!(
258                default_message: "Hello, {name}!",
259                description: "Greeting"
260            );
261        }"#;
262        let messages = extract_messages_from_rust_source(
263            source,
264            Path::new("src/main.rs"),
265            false,
266            false,
267            false,
268            true,
269        )
270        .unwrap();
271        assert_eq!(messages.len(), 1);
272        assert_eq!(messages[0].id.as_deref(), Some("EG1xJTTqQy"));
273        assert_eq!(
274            messages[0].default_message.as_deref(),
275            Some("Hello, {name}!")
276        );
277        assert_eq!(
278            messages[0].description,
279            Some(Value::String("Greeting".to_owned()))
280        );
281    }
282
283    #[test]
284    fn preserves_explicit_id() {
285        let messages = extract_messages_from_rust_source(
286            r#"fn main() { message_descriptor!(id: "hello", default_message: "Hello"); }"#,
287            Path::new("src/main.rs"),
288            false,
289            false,
290            false,
291            true,
292        )
293        .unwrap();
294        assert_eq!(messages[0].id.as_deref(), Some("hello"));
295    }
296
297    #[test]
298    fn extracts_format_message_macros() {
299        let messages = extract_messages_from_rust_source(
300            r#"fn render(intl: &Intl, values: &Values<String>) {
301                format_message!(
302                    &intl,
303                    default_message: "Hello, {name}!",
304                    description: "Greeting",
305                    values: values,
306                );
307                format_message!(
308                    &intl,
309                    default_message: "{count, plural, one {# task} other {{name} has # tasks}}",
310                    values: { count: count, name: user.name() },
311                );
312                formatjs_intl::format_message!(
313                    &intl,
314                    id: "approval.title",
315                    default_message: "Approve to continue",
316                );
317            }"#,
318            Path::new("src/main.rs"),
319            false,
320            false,
321            false,
322            true,
323        )
324        .unwrap();
325
326        assert_eq!(messages.len(), 3);
327        assert_eq!(messages[0].id.as_deref(), Some("EG1xJTTqQy"));
328        assert_eq!(
329            messages[1].default_message.as_deref(),
330            Some("{count, plural, one {# task} other {{name} has # tasks}}")
331        );
332        assert_eq!(messages[2].id.as_deref(), Some("approval.title"));
333    }
334
335    #[test]
336    fn reports_source_offsets() {
337        let source = "fn main() {\n    message_descriptor!(default_message: \"Hello\");\n}\n";
338        let messages = extract_messages_from_rust_source(
339            source,
340            Path::new("src/main.rs"),
341            true,
342            false,
343            false,
344            true,
345        )
346        .unwrap();
347        assert_eq!(messages[0].file.as_deref(), Some("src/main.rs"));
348        assert_eq!(messages[0].start, Some(16));
349        assert!(messages[0].end.unwrap() > messages[0].start.unwrap());
350    }
351}