Skip to main content

rustpython_literal/
escape.rs

1use alloc::string::String;
2use rustpython_wtf8::{CodePoint, Wtf8};
3
4#[derive(Debug, PartialEq, Eq, Copy, Clone, Hash, is_macro::Is)]
5pub enum Quote {
6    Single,
7    Double,
8}
9
10impl Quote {
11    #[inline]
12    #[must_use]
13    pub const fn swap(self) -> Self {
14        match self {
15            Self::Single => Self::Double,
16            Self::Double => Self::Single,
17        }
18    }
19
20    #[inline]
21    #[must_use]
22    pub const fn to_byte(&self) -> u8 {
23        match self {
24            Self::Single => b'\'',
25            Self::Double => b'"',
26        }
27    }
28
29    #[inline]
30    #[must_use]
31    pub const fn to_char(&self) -> char {
32        match self {
33            Self::Single => '\'',
34            Self::Double => '"',
35        }
36    }
37}
38
39pub struct EscapeLayout {
40    pub quote: Quote,
41    pub len: Option<usize>,
42}
43
44/// Represents string types that can be escape-printed.
45///
46/// # Safety
47///
48/// `source_len` and `layout` must be accurate, and `layout.len` must not be equal
49/// to `Some(source_len)` if the string contains non-printable characters.
50pub unsafe trait Escape {
51    fn source_len(&self) -> usize;
52    fn layout(&self) -> &EscapeLayout;
53    fn changed(&self) -> bool {
54        self.layout().len != Some(self.source_len())
55    }
56
57    /// Write the body of the string directly to the formatter.
58    ///
59    /// # Safety
60    ///
61    /// This string must only contain printable characters.
62    unsafe fn write_source(&self, formatter: &mut impl core::fmt::Write) -> core::fmt::Result;
63    fn write_body_slow(&self, formatter: &mut impl core::fmt::Write) -> core::fmt::Result;
64    fn write_body(&self, formatter: &mut impl core::fmt::Write) -> core::fmt::Result {
65        if self.changed() {
66            self.write_body_slow(formatter)
67        } else {
68            // SAFETY: verified the string contains only printable characters.
69            unsafe { self.write_source(formatter) }
70        }
71    }
72}
73
74/// Returns the outer quotes to use and the number of quotes that need to be
75/// escaped.
76pub(crate) const fn choose_quote(
77    single_count: usize,
78    double_count: usize,
79    preferred_quote: Quote,
80) -> (Quote, usize) {
81    let (primary_count, secondary_count) = match preferred_quote {
82        Quote::Single => (single_count, double_count),
83        Quote::Double => (double_count, single_count),
84    };
85
86    // always use primary unless we have primary but no secondary
87    let use_secondary = primary_count > 0 && secondary_count == 0;
88    if use_secondary {
89        (preferred_quote.swap(), secondary_count)
90    } else {
91        (preferred_quote, primary_count)
92    }
93}
94
95pub struct UnicodeEscape<'a> {
96    source: &'a Wtf8,
97    layout: EscapeLayout,
98}
99
100impl<'a> UnicodeEscape<'a> {
101    #[inline]
102    #[must_use]
103    pub const fn with_forced_quote(source: &'a Wtf8, quote: Quote) -> Self {
104        let layout = EscapeLayout { quote, len: None };
105        Self { source, layout }
106    }
107    #[inline]
108    #[must_use]
109    pub fn with_preferred_quote(source: &'a Wtf8, quote: Quote) -> Self {
110        let layout = Self::repr_layout(source, quote);
111        Self { source, layout }
112    }
113    #[inline]
114    #[must_use]
115    pub fn new_repr(source: &'a Wtf8) -> Self {
116        Self::with_preferred_quote(source, Quote::Single)
117    }
118    #[inline]
119    #[must_use]
120    pub const fn str_repr<'r>(&'a self) -> StrRepr<'r, 'a> {
121        StrRepr(self)
122    }
123}
124
125pub struct StrRepr<'r, 'a>(&'r UnicodeEscape<'a>);
126
127impl StrRepr<'_, '_> {
128    pub fn write(&self, formatter: &mut impl core::fmt::Write) -> core::fmt::Result {
129        let quote = self.0.layout().quote.to_char();
130        formatter.write_char(quote)?;
131        self.0.write_body(formatter)?;
132        formatter.write_char(quote)
133    }
134
135    #[must_use]
136    pub fn to_string(&self) -> Option<String> {
137        let mut s = String::with_capacity(self.0.layout().len?);
138        self.write(&mut s).unwrap();
139        Some(s)
140    }
141}
142
143impl core::fmt::Display for StrRepr<'_, '_> {
144    fn fmt(&self, formatter: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
145        self.write(formatter)
146    }
147}
148
149impl UnicodeEscape<'_> {
150    const REPR_RESERVED_LEN: usize = 2; // for quotes
151
152    #[must_use]
153    pub fn repr_layout(source: &Wtf8, preferred_quote: Quote) -> EscapeLayout {
154        Self::output_layout_with_checker(source, preferred_quote, |a, b| {
155            Some((a as isize).checked_add(b as isize)? as usize)
156        })
157    }
158
159    fn output_layout_with_checker(
160        source: &Wtf8,
161        preferred_quote: Quote,
162        length_add: impl Fn(usize, usize) -> Option<usize>,
163    ) -> EscapeLayout {
164        let mut out_len = Self::REPR_RESERVED_LEN;
165        let mut single_count = 0;
166        let mut double_count = 0;
167
168        for ch in source.code_points() {
169            let incr = match ch.to_char() {
170                Some('\'') => {
171                    single_count += 1;
172                    1
173                }
174                Some('"') => {
175                    double_count += 1;
176                    1
177                }
178                _ => Self::escaped_char_len(ch),
179            };
180            let Some(new_len) = length_add(out_len, incr) else {
181                #[cold]
182                const fn stop(
183                    single_count: usize,
184                    double_count: usize,
185                    preferred_quote: Quote,
186                ) -> EscapeLayout {
187                    EscapeLayout {
188                        quote: choose_quote(single_count, double_count, preferred_quote).0,
189                        len: None,
190                    }
191                }
192                return stop(single_count, double_count, preferred_quote);
193            };
194            out_len = new_len;
195        }
196
197        let (quote, num_escaped_quotes) = choose_quote(single_count, double_count, preferred_quote);
198        // we'll be adding backslashes in front of the existing inner quotes
199        let Some(out_len) = length_add(out_len, num_escaped_quotes) else {
200            return EscapeLayout { quote, len: None };
201        };
202
203        EscapeLayout {
204            quote,
205            len: Some(out_len - Self::REPR_RESERVED_LEN),
206        }
207    }
208
209    fn escaped_char_len(ch: CodePoint) -> usize {
210        // surrogates are \uHHHH
211        let Some(ch) = ch.to_char() else { return 6 };
212        match ch {
213            '\\' | '\t' | '\r' | '\n' => 2,
214            ch if ch < ' ' || ch as u32 == 0x7f => 4, // \xHH
215            ch if ch.is_ascii() => 1,
216            ch if rustpython_unicode::classify::is_repr_printable(ch) => {
217                // max = std::cmp::max(ch, max);
218                ch.len_utf8()
219            }
220            ch if (ch as u32) < 0x100 => 4,   // \xHH
221            ch if (ch as u32) < 0x10000 => 6, // \uHHHH
222            _ => 10,                          // \uHHHHHHHH
223        }
224    }
225
226    fn write_char(
227        ch: CodePoint,
228        quote: Quote,
229        formatter: &mut impl core::fmt::Write,
230    ) -> core::fmt::Result {
231        let Some(ch) = ch.to_char() else {
232            return write!(formatter, "\\u{:04x}", ch.to_u32());
233        };
234        match ch {
235            '\n' => formatter.write_str("\\n"),
236            '\t' => formatter.write_str("\\t"),
237            '\r' => formatter.write_str("\\r"),
238            // these 2 branches *would* be handled below, but we shouldn't have to do a
239            // unicodedata lookup just for ascii characters
240            '\x20'..='\x7e' => {
241                // printable ascii range
242                if ch == quote.to_char() || ch == '\\' {
243                    formatter.write_char('\\')?;
244                }
245                formatter.write_char(ch)
246            }
247            ch if ch.is_ascii() => {
248                write!(formatter, "\\x{:02x}", ch as u8)
249            }
250            ch if rustpython_unicode::classify::is_repr_printable(ch) => formatter.write_char(ch),
251            '\0'..='\u{ff}' => {
252                write!(formatter, "\\x{:02x}", ch as u32)
253            }
254            '\0'..='\u{ffff}' => {
255                write!(formatter, "\\u{:04x}", ch as u32)
256            }
257            _ => {
258                write!(formatter, "\\U{:08x}", ch as u32)
259            }
260        }
261    }
262}
263
264unsafe impl Escape for UnicodeEscape<'_> {
265    fn source_len(&self) -> usize {
266        self.source.len()
267    }
268
269    fn layout(&self) -> &EscapeLayout {
270        &self.layout
271    }
272
273    unsafe fn write_source(&self, formatter: &mut impl core::fmt::Write) -> core::fmt::Result {
274        formatter.write_str(unsafe {
275            // SAFETY: this function must be called only when source is printable characters (i.e. no surrogates)
276            core::str::from_utf8_unchecked(self.source.as_bytes())
277        })
278    }
279
280    #[cold]
281    fn write_body_slow(&self, formatter: &mut impl core::fmt::Write) -> core::fmt::Result {
282        for ch in self.source.code_points() {
283            Self::write_char(ch, self.layout().quote, formatter)?;
284        }
285        Ok(())
286    }
287}
288
289pub struct AsciiEscape<'a> {
290    source: &'a [u8],
291    layout: EscapeLayout,
292}
293
294impl<'a> AsciiEscape<'a> {
295    #[inline]
296    #[must_use]
297    pub const fn new(source: &'a [u8], layout: EscapeLayout) -> Self {
298        Self { source, layout }
299    }
300    #[inline]
301    #[must_use]
302    pub const fn with_forced_quote(source: &'a [u8], quote: Quote) -> Self {
303        let layout = EscapeLayout { quote, len: None };
304        Self { source, layout }
305    }
306    #[inline]
307    #[must_use]
308    pub fn with_preferred_quote(source: &'a [u8], quote: Quote) -> Self {
309        let layout = Self::repr_layout(source, quote);
310        Self { source, layout }
311    }
312    #[inline]
313    #[must_use]
314    pub fn new_repr(source: &'a [u8]) -> Self {
315        Self::with_preferred_quote(source, Quote::Single)
316    }
317    #[inline]
318    #[must_use]
319    pub const fn bytes_repr<'r>(&'a self) -> BytesRepr<'r, 'a> {
320        BytesRepr(self)
321    }
322}
323
324impl AsciiEscape<'_> {
325    #[must_use]
326    pub fn repr_layout(source: &[u8], preferred_quote: Quote) -> EscapeLayout {
327        Self::output_layout_with_checker(source, preferred_quote, 3, |a, b| {
328            Some((a as isize).checked_add(b as isize)? as usize)
329        })
330    }
331
332    #[must_use]
333    pub fn named_repr_layout(source: &[u8], name: &str) -> EscapeLayout {
334        Self::output_layout_with_checker(source, Quote::Single, name.len() + 2 + 3, |a, b| {
335            Some((a as isize).checked_add(b as isize)? as usize)
336        })
337    }
338
339    fn output_layout_with_checker(
340        source: &[u8],
341        preferred_quote: Quote,
342        reserved_len: usize,
343        length_add: impl Fn(usize, usize) -> Option<usize>,
344    ) -> EscapeLayout {
345        let mut out_len = reserved_len;
346        let mut single_count = 0;
347        let mut double_count = 0;
348
349        for ch in source {
350            let incr = match ch {
351                b'\'' => {
352                    single_count += 1;
353                    1
354                }
355                b'"' => {
356                    double_count += 1;
357                    1
358                }
359                c => Self::escaped_char_len(*c),
360            };
361            let Some(new_len) = length_add(out_len, incr) else {
362                #[cold]
363                const fn stop(
364                    single_count: usize,
365                    double_count: usize,
366                    preferred_quote: Quote,
367                ) -> EscapeLayout {
368                    EscapeLayout {
369                        quote: choose_quote(single_count, double_count, preferred_quote).0,
370                        len: None,
371                    }
372                }
373                return stop(single_count, double_count, preferred_quote);
374            };
375            out_len = new_len;
376        }
377
378        let (quote, num_escaped_quotes) = choose_quote(single_count, double_count, preferred_quote);
379        // we'll be adding backslashes in front of the existing inner quotes
380        let Some(out_len) = length_add(out_len, num_escaped_quotes) else {
381            return EscapeLayout { quote, len: None };
382        };
383
384        EscapeLayout {
385            quote,
386            len: Some(out_len - reserved_len),
387        }
388    }
389
390    const fn escaped_char_len(ch: u8) -> usize {
391        match ch {
392            b'\\' | b'\t' | b'\r' | b'\n' => 2,
393            0x20..=0x7e => 1,
394            _ => 4, // \xHH
395        }
396    }
397
398    fn write_char(
399        ch: u8,
400        quote: Quote,
401        formatter: &mut impl core::fmt::Write,
402    ) -> core::fmt::Result {
403        match ch {
404            b'\t' => formatter.write_str("\\t"),
405            b'\n' => formatter.write_str("\\n"),
406            b'\r' => formatter.write_str("\\r"),
407            0x20..=0x7e => {
408                // printable ascii range
409                if ch == quote.to_byte() || ch == b'\\' {
410                    formatter.write_char('\\')?;
411                }
412                formatter.write_char(ch as char)
413            }
414            ch => write!(formatter, "\\x{ch:02x}"),
415        }
416    }
417}
418
419unsafe impl Escape for AsciiEscape<'_> {
420    fn source_len(&self) -> usize {
421        self.source.len()
422    }
423
424    fn layout(&self) -> &EscapeLayout {
425        &self.layout
426    }
427
428    unsafe fn write_source(&self, formatter: &mut impl core::fmt::Write) -> core::fmt::Result {
429        formatter.write_str(unsafe {
430            // SAFETY: this function must be called only when source is printable ascii characters
431            core::str::from_utf8_unchecked(self.source)
432        })
433    }
434
435    #[cold]
436    fn write_body_slow(&self, formatter: &mut impl core::fmt::Write) -> core::fmt::Result {
437        for ch in self.source {
438            Self::write_char(*ch, self.layout().quote, formatter)?;
439        }
440        Ok(())
441    }
442}
443
444pub struct BytesRepr<'r, 'a>(&'r AsciiEscape<'a>);
445
446impl BytesRepr<'_, '_> {
447    pub fn write(&self, formatter: &mut impl core::fmt::Write) -> core::fmt::Result {
448        let quote = self.0.layout().quote.to_char();
449        formatter.write_char('b')?;
450        formatter.write_char(quote)?;
451        self.0.write_body(formatter)?;
452        formatter.write_char(quote)
453    }
454
455    #[must_use]
456    pub fn to_string(&self) -> Option<String> {
457        let mut s = String::with_capacity(self.0.layout().len?);
458        self.write(&mut s).unwrap();
459        Some(s)
460    }
461}
462
463impl core::fmt::Display for BytesRepr<'_, '_> {
464    fn fmt(&self, formatter: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
465        self.write(formatter)
466    }
467}
468
469#[cfg(test)]
470mod unicode_escape_tests {
471    use super::*;
472
473    #[test]
474    fn changed() {
475        fn test(s: &str) -> bool {
476            UnicodeEscape::new_repr(s.as_ref()).changed()
477        }
478        assert!(!test("hello"));
479        assert!(!test("'hello'"));
480        assert!(!test("\"hello\""));
481
482        assert!(test("'\"hello"));
483        assert!(test("hello\n"));
484    }
485}