Skip to main content

rustpython_common/
encodings.rs

1use core::ops::{self, Range};
2
3use num_traits::ToPrimitive;
4
5use crate::str::StrKind;
6use crate::wtf8::{CodePoint, Wtf8, Wtf8Buf};
7
8#[cfg(feature = "cjk-codecs")]
9pub mod cjk;
10mod wide;
11pub use wide::ByteOrder;
12pub mod escape;
13pub mod raw_unicode_escape;
14pub mod unicode_escape;
15pub mod utf16;
16pub mod utf32;
17pub mod utf7;
18
19pub trait StrBuffer: AsRef<Wtf8> {
20    fn is_compatible_with(&self, kind: StrKind) -> bool {
21        let s = self.as_ref();
22        match kind {
23            StrKind::Ascii => s.is_ascii(),
24            StrKind::Utf8 => s.is_utf8(),
25            StrKind::Wtf8 => true,
26        }
27    }
28}
29
30pub trait CodecContext: Sized {
31    type Error;
32    type StrBuf: StrBuffer;
33    type BytesBuf: AsRef<[u8]>;
34
35    fn string(&self, s: Wtf8Buf) -> Self::StrBuf;
36    fn bytes(&self, b: Vec<u8>) -> Self::BytesBuf;
37}
38
39pub trait EncodeContext: CodecContext {
40    fn full_data(&self) -> &Wtf8;
41    fn data_len(&self) -> StrSize;
42
43    fn remaining_data(&self) -> &Wtf8;
44    fn position(&self) -> StrSize;
45
46    fn restart_from(&mut self, pos: StrSize) -> Result<(), Self::Error>;
47
48    fn error_encoding(&self, range: Range<StrSize>, reason: Option<&str>) -> Self::Error;
49
50    fn handle_error<E>(
51        &mut self,
52        errors: &E,
53        range: Range<StrSize>,
54        reason: Option<&str>,
55    ) -> Result<EncodeReplace<Self>, Self::Error>
56    where
57        E: EncodeErrorHandler<Self>,
58    {
59        let (replace, restart) = errors.handle_encode_error(self, range, reason)?;
60        self.restart_from(restart)?;
61        Ok(replace)
62    }
63}
64
65pub trait DecodeContext: CodecContext {
66    fn full_data(&self) -> &[u8];
67
68    fn remaining_data(&self) -> &[u8];
69    fn position(&self) -> usize;
70
71    fn advance(&mut self, by: usize);
72
73    fn restart_from(&mut self, pos: usize) -> Result<(), Self::Error>;
74
75    fn error_decoding(&self, byte_range: Range<usize>, reason: Option<&str>) -> Self::Error;
76
77    fn handle_error<E>(
78        &mut self,
79        errors: &E,
80        byte_range: Range<usize>,
81        reason: Option<&str>,
82    ) -> Result<Self::StrBuf, Self::Error>
83    where
84        E: DecodeErrorHandler<Self>,
85    {
86        let (replace, restart) = errors.handle_decode_error(self, byte_range, reason)?;
87        self.restart_from(restart)?;
88        Ok(replace)
89    }
90}
91
92pub trait EncodeErrorHandler<Ctx: EncodeContext> {
93    fn handle_encode_error(
94        &self,
95        ctx: &mut Ctx,
96        range: Range<StrSize>,
97        reason: Option<&str>,
98    ) -> Result<(EncodeReplace<Ctx>, StrSize), Ctx::Error>;
99}
100pub trait DecodeErrorHandler<Ctx: DecodeContext> {
101    fn handle_decode_error(
102        &self,
103        ctx: &mut Ctx,
104        byte_range: Range<usize>,
105        reason: Option<&str>,
106    ) -> Result<(Ctx::StrBuf, usize), Ctx::Error>;
107}
108
109pub enum EncodeReplace<Ctx: CodecContext> {
110    Str(Ctx::StrBuf),
111    Bytes(Ctx::BytesBuf),
112}
113
114#[derive(Copy, Clone, Default, Debug)]
115pub struct StrSize {
116    pub bytes: usize,
117    pub chars: usize,
118}
119
120fn iter_code_points(w: &Wtf8) -> impl Iterator<Item = (StrSize, CodePoint)> {
121    w.code_point_indices()
122        .enumerate()
123        .map(|(chars, (bytes, c))| (StrSize { bytes, chars }, c))
124}
125
126impl ops::Add for StrSize {
127    type Output = Self;
128    fn add(self, rhs: Self) -> Self::Output {
129        Self {
130            bytes: self.bytes + rhs.bytes,
131            chars: self.chars + rhs.chars,
132        }
133    }
134}
135
136impl ops::AddAssign for StrSize {
137    fn add_assign(&mut self, rhs: Self) {
138        self.bytes += rhs.bytes;
139        self.chars += rhs.chars;
140    }
141}
142
143struct DecodeError<'a> {
144    valid_prefix: &'a str,
145    rest: &'a [u8],
146    err_len: Option<usize>,
147}
148
149/// # Safety
150/// `v[..valid_up_to]` must be valid utf8
151const unsafe fn make_decode_err(
152    v: &[u8],
153    valid_up_to: usize,
154    err_len: Option<usize>,
155) -> DecodeError<'_> {
156    let (valid_prefix, rest) = unsafe { v.split_at_unchecked(valid_up_to) };
157    let valid_prefix = unsafe { core::str::from_utf8_unchecked(valid_prefix) };
158    DecodeError {
159        valid_prefix,
160        rest,
161        err_len,
162    }
163}
164
165enum HandleResult<'a> {
166    Done,
167    Error {
168        err_len: Option<usize>,
169        reason: &'a str,
170    },
171}
172
173fn decode_utf8_compatible<Ctx, E, DecodeF, ErrF>(
174    mut ctx: Ctx,
175    errors: &E,
176    decode: DecodeF,
177    handle_error: ErrF,
178) -> Result<(Wtf8Buf, usize), Ctx::Error>
179where
180    Ctx: DecodeContext,
181    E: DecodeErrorHandler<Ctx>,
182    DecodeF: Fn(&[u8]) -> Result<&str, DecodeError<'_>>,
183    ErrF: Fn(&[u8], Option<usize>) -> HandleResult<'static>,
184{
185    if ctx.remaining_data().is_empty() {
186        return Ok((Wtf8Buf::new(), 0));
187    }
188    let mut out = Wtf8Buf::with_capacity(ctx.remaining_data().len());
189    loop {
190        match decode(ctx.remaining_data()) {
191            Ok(decoded) => {
192                out.push_str(decoded);
193                ctx.advance(decoded.len());
194                break;
195            }
196            Err(e) => {
197                out.push_str(e.valid_prefix);
198                match handle_error(e.rest, e.err_len) {
199                    HandleResult::Done => {
200                        ctx.advance(e.valid_prefix.len());
201                        break;
202                    }
203                    HandleResult::Error { err_len, reason } => {
204                        let err_start = ctx.position() + e.valid_prefix.len();
205                        let err_end = match err_len {
206                            Some(len) => err_start + len,
207                            None => ctx.full_data().len(),
208                        };
209                        let err_range = err_start..err_end;
210                        let replace = ctx.handle_error(errors, err_range, Some(reason))?;
211                        out.push_wtf8(replace.as_ref());
212                        continue;
213                    }
214                }
215            }
216        }
217    }
218    Ok((out, ctx.position()))
219}
220
221#[inline]
222fn encode_utf8_compatible<Ctx, E>(
223    mut ctx: Ctx,
224    errors: &E,
225    err_reason: &str,
226    target_kind: StrKind,
227) -> Result<Vec<u8>, Ctx::Error>
228where
229    Ctx: EncodeContext,
230    E: EncodeErrorHandler<Ctx>,
231{
232    // let mut data = s.as_ref();
233    // let mut char_data_index = 0;
234    let mut out = Vec::<u8>::with_capacity(ctx.remaining_data().len());
235    loop {
236        let data = ctx.remaining_data();
237        let mut iter = iter_code_points(data);
238        let Some((i, _)) = iter.find(|(_, c)| !target_kind.can_encode(*c)) else {
239            break;
240        };
241
242        out.extend_from_slice(&ctx.remaining_data().as_bytes()[..i.bytes]);
243
244        let err_start = ctx.position() + i;
245        // number of non-compatible chars between the first non-compatible char and the next compatible char
246        let err_end = match { iter }.find(|(_, c)| target_kind.can_encode(*c)) {
247            Some((i, _)) => ctx.position() + i,
248            None => ctx.data_len(),
249        };
250
251        let range = err_start..err_end;
252        let replace = ctx.handle_error(errors, range.clone(), Some(err_reason))?;
253        match replace {
254            EncodeReplace::Str(s) => {
255                if s.is_compatible_with(target_kind) {
256                    out.extend_from_slice(s.as_ref().as_bytes());
257                } else {
258                    return Err(ctx.error_encoding(range, Some(err_reason)));
259                }
260            }
261            EncodeReplace::Bytes(b) => {
262                out.extend_from_slice(b.as_ref());
263            }
264        }
265    }
266    out.extend_from_slice(ctx.remaining_data().as_bytes());
267    Ok(out)
268}
269
270pub mod errors {
271    use crate::str::UnicodeEscapeCodepoint;
272
273    use super::*;
274    use core::fmt::Write;
275
276    #[derive(Clone, Copy)]
277    pub struct Strict;
278
279    impl<Ctx: EncodeContext> EncodeErrorHandler<Ctx> for Strict {
280        fn handle_encode_error(
281            &self,
282            ctx: &mut Ctx,
283            range: Range<StrSize>,
284            reason: Option<&str>,
285        ) -> Result<(EncodeReplace<Ctx>, StrSize), Ctx::Error> {
286            Err(ctx.error_encoding(range, reason))
287        }
288    }
289
290    impl<Ctx: DecodeContext> DecodeErrorHandler<Ctx> for Strict {
291        fn handle_decode_error(
292            &self,
293            ctx: &mut Ctx,
294            byte_range: Range<usize>,
295            reason: Option<&str>,
296        ) -> Result<(Ctx::StrBuf, usize), Ctx::Error> {
297            Err(ctx.error_decoding(byte_range, reason))
298        }
299    }
300
301    #[derive(Clone, Copy)]
302    pub struct Ignore;
303
304    impl<Ctx: EncodeContext> EncodeErrorHandler<Ctx> for Ignore {
305        fn handle_encode_error(
306            &self,
307            ctx: &mut Ctx,
308            range: Range<StrSize>,
309            _reason: Option<&str>,
310        ) -> Result<(EncodeReplace<Ctx>, StrSize), Ctx::Error> {
311            Ok((EncodeReplace::Bytes(ctx.bytes(b"".into())), range.end))
312        }
313    }
314
315    impl<Ctx: DecodeContext> DecodeErrorHandler<Ctx> for Ignore {
316        fn handle_decode_error(
317            &self,
318            ctx: &mut Ctx,
319            byte_range: Range<usize>,
320            _reason: Option<&str>,
321        ) -> Result<(Ctx::StrBuf, usize), Ctx::Error> {
322            Ok((ctx.string("".into()), byte_range.end))
323        }
324    }
325
326    #[derive(Clone, Copy)]
327    pub struct Replace;
328
329    impl<Ctx: EncodeContext> EncodeErrorHandler<Ctx> for Replace {
330        fn handle_encode_error(
331            &self,
332            ctx: &mut Ctx,
333            range: Range<StrSize>,
334            _reason: Option<&str>,
335        ) -> Result<(EncodeReplace<Ctx>, StrSize), Ctx::Error> {
336            let replace = "?".repeat(range.end.chars - range.start.chars);
337            Ok((EncodeReplace::Str(ctx.string(replace.into())), range.end))
338        }
339    }
340
341    impl<Ctx: DecodeContext> DecodeErrorHandler<Ctx> for Replace {
342        fn handle_decode_error(
343            &self,
344            ctx: &mut Ctx,
345            byte_range: Range<usize>,
346            _reason: Option<&str>,
347        ) -> Result<(Ctx::StrBuf, usize), Ctx::Error> {
348            Ok((
349                ctx.string(char::REPLACEMENT_CHARACTER.to_string().into()),
350                byte_range.end,
351            ))
352        }
353    }
354
355    #[derive(Clone, Copy)]
356    pub struct XmlCharRefReplace;
357
358    impl<Ctx: EncodeContext> EncodeErrorHandler<Ctx> for XmlCharRefReplace {
359        fn handle_encode_error(
360            &self,
361            ctx: &mut Ctx,
362            range: Range<StrSize>,
363            _reason: Option<&str>,
364        ) -> Result<(EncodeReplace<Ctx>, StrSize), Ctx::Error> {
365            let err_str = &ctx.full_data()[range.start.bytes..range.end.bytes];
366            let num_chars = range.end.chars - range.start.chars;
367            // capacity rough guess; assuming that the codepoints are 3 digits in decimal + the &#;
368            let mut out = String::with_capacity(num_chars * 6);
369            for c in err_str.code_points() {
370                write!(out, "&#{};", c.to_u32()).unwrap()
371            }
372            Ok((EncodeReplace::Str(ctx.string(out.into())), range.end))
373        }
374    }
375
376    #[derive(Clone, Copy)]
377    pub struct BackslashReplace;
378
379    impl<Ctx: EncodeContext> EncodeErrorHandler<Ctx> for BackslashReplace {
380        fn handle_encode_error(
381            &self,
382            ctx: &mut Ctx,
383            range: Range<StrSize>,
384            _reason: Option<&str>,
385        ) -> Result<(EncodeReplace<Ctx>, StrSize), Ctx::Error> {
386            let err_str = &ctx.full_data()[range.start.bytes..range.end.bytes];
387            let num_chars = range.end.chars - range.start.chars;
388            // minimum 4 output bytes per char: \xNN
389            let mut out = String::with_capacity(num_chars * 4);
390            for c in err_str.code_points() {
391                write!(out, "{}", UnicodeEscapeCodepoint(c)).unwrap();
392            }
393            Ok((EncodeReplace::Str(ctx.string(out.into())), range.end))
394        }
395    }
396
397    impl<Ctx: DecodeContext> DecodeErrorHandler<Ctx> for BackslashReplace {
398        fn handle_decode_error(
399            &self,
400            ctx: &mut Ctx,
401            byte_range: Range<usize>,
402            _reason: Option<&str>,
403        ) -> Result<(Ctx::StrBuf, usize), Ctx::Error> {
404            let err_bytes = &ctx.full_data()[byte_range.clone()];
405            let mut replace = String::with_capacity(4 * err_bytes.len());
406            for &c in err_bytes {
407                write!(replace, "\\x{c:02x}").unwrap();
408            }
409            Ok((ctx.string(replace.into()), byte_range.end))
410        }
411    }
412
413    #[derive(Clone, Copy)]
414    pub struct NameReplace;
415
416    impl<Ctx: EncodeContext> EncodeErrorHandler<Ctx> for NameReplace {
417        fn handle_encode_error(
418            &self,
419            ctx: &mut Ctx,
420            range: Range<StrSize>,
421            _reason: Option<&str>,
422        ) -> Result<(EncodeReplace<Ctx>, StrSize), Ctx::Error> {
423            let err_str = &ctx.full_data()[range.start.bytes..range.end.bytes];
424            let num_chars = range.end.chars - range.start.chars;
425            let mut out = String::with_capacity(num_chars * 4);
426            for c in err_str.code_points() {
427                let c_u32 = c.to_u32();
428                if let Some(c_name) = c.to_char().and_then(rustpython_unicode::character_name) {
429                    write!(out, "\\N{{{c_name}}}").unwrap();
430                } else if c_u32 >= 0x10000 {
431                    write!(out, "\\U{c_u32:08x}").unwrap();
432                } else if c_u32 >= 0x100 {
433                    write!(out, "\\u{c_u32:04x}").unwrap();
434                } else {
435                    write!(out, "\\x{c_u32:02x}").unwrap();
436                }
437            }
438            Ok((EncodeReplace::Str(ctx.string(out.into())), range.end))
439        }
440    }
441
442    #[derive(Clone, Copy)]
443    pub struct SurrogateEscape;
444
445    impl<Ctx: EncodeContext> EncodeErrorHandler<Ctx> for SurrogateEscape {
446        fn handle_encode_error(
447            &self,
448            ctx: &mut Ctx,
449            range: Range<StrSize>,
450            reason: Option<&str>,
451        ) -> Result<(EncodeReplace<Ctx>, StrSize), Ctx::Error> {
452            let err_str = &ctx.full_data()[range.start.bytes..range.end.bytes];
453            let num_chars = range.end.chars - range.start.chars;
454            let mut out = Vec::with_capacity(num_chars);
455            let mut pos = range.start;
456            for ch in err_str.code_points() {
457                let ch_u32 = ch.to_u32();
458                if !(0xdc80..=0xdcff).contains(&ch_u32) {
459                    if out.is_empty() {
460                        // Can't handle even the first character
461                        return Err(ctx.error_encoding(range, reason));
462                    }
463                    // Return partial result, restart from this character
464                    return Ok((EncodeReplace::Bytes(ctx.bytes(out)), pos));
465                }
466                out.push((ch_u32 - 0xdc00) as u8);
467                pos += StrSize {
468                    bytes: ch.len_wtf8(),
469                    chars: 1,
470                };
471            }
472            Ok((EncodeReplace::Bytes(ctx.bytes(out)), range.end))
473        }
474    }
475
476    impl<Ctx: DecodeContext> DecodeErrorHandler<Ctx> for SurrogateEscape {
477        fn handle_decode_error(
478            &self,
479            ctx: &mut Ctx,
480            byte_range: Range<usize>,
481            reason: Option<&str>,
482        ) -> Result<(Ctx::StrBuf, usize), Ctx::Error> {
483            let err_bytes = &ctx.full_data()[byte_range.clone()];
484            let mut consumed = 0;
485            let mut replace = Wtf8Buf::with_capacity(4 * byte_range.len());
486            while consumed < 4 && consumed < byte_range.len() {
487                let c = err_bytes[consumed] as u16;
488                // Refuse to escape ASCII bytes
489                if c < 128 {
490                    break;
491                }
492                replace.push(CodePoint::from(0xdc00 + c));
493                consumed += 1;
494            }
495            if consumed == 0 {
496                return Err(ctx.error_decoding(byte_range, reason));
497            }
498            Ok((ctx.string(replace), byte_range.start + consumed))
499        }
500    }
501}
502
503pub mod utf8 {
504    use super::*;
505
506    pub const ENCODING_NAME: &str = "utf-8";
507
508    #[inline]
509    pub fn encode<Ctx, E>(ctx: Ctx, errors: &E) -> Result<Vec<u8>, Ctx::Error>
510    where
511        Ctx: EncodeContext,
512        E: EncodeErrorHandler<Ctx>,
513    {
514        encode_utf8_compatible(ctx, errors, "surrogates not allowed", StrKind::Utf8)
515    }
516
517    pub fn decode<Ctx: DecodeContext, E: DecodeErrorHandler<Ctx>>(
518        ctx: Ctx,
519        errors: &E,
520        final_decode: bool,
521    ) -> Result<(Wtf8Buf, usize), Ctx::Error> {
522        decode_utf8_compatible(
523            ctx,
524            errors,
525            |v| {
526                core::str::from_utf8(v).map_err(|e| {
527                    // SAFETY: as specified in valid_up_to's documentation, input[..e.valid_up_to()]
528                    //         is valid utf8
529                    unsafe { make_decode_err(v, e.valid_up_to(), e.error_len()) }
530                })
531            },
532            |rest, err_len| {
533                let first_err = rest[0];
534                if matches!(first_err, 0x80..=0xc1 | 0xf5..=0xff) {
535                    HandleResult::Error {
536                        err_len: Some(1),
537                        reason: "invalid start byte",
538                    }
539                } else if err_len.is_none() {
540                    // error_len() == None means unexpected eof
541                    if final_decode {
542                        HandleResult::Error {
543                            err_len,
544                            reason: "unexpected end of data",
545                        }
546                    } else {
547                        HandleResult::Done
548                    }
549                } else if !final_decode && matches!(rest, [0xed, 0xa0..=0xbf]) {
550                    // truncated surrogate
551                    HandleResult::Done
552                } else {
553                    HandleResult::Error {
554                        err_len,
555                        reason: "invalid continuation byte",
556                    }
557                }
558            },
559        )
560    }
561}
562
563pub mod latin_1 {
564    use super::*;
565
566    pub const ENCODING_NAME: &str = "latin-1";
567
568    const ERR_REASON: &str = "ordinal not in range(256)";
569
570    #[inline]
571    pub fn encode<Ctx, E>(mut ctx: Ctx, errors: &E) -> Result<Vec<u8>, Ctx::Error>
572    where
573        Ctx: EncodeContext,
574        E: EncodeErrorHandler<Ctx>,
575    {
576        let mut out = Vec::<u8>::new();
577        loop {
578            let data = ctx.remaining_data();
579            let mut iter = iter_code_points(ctx.remaining_data());
580            let Some((i, ch)) = iter.find(|(_, c)| !c.is_ascii()) else {
581                break;
582            };
583            out.extend_from_slice(&data.as_bytes()[..i.bytes]);
584            let err_start = ctx.position() + i;
585            if let Some(byte) = ch.to_u32().to_u8() {
586                drop(iter);
587                out.push(byte);
588                // if the codepoint is between 128..=255, it's utf8-length is 2
589                ctx.restart_from(err_start + StrSize { bytes: 2, chars: 1 })?;
590            } else {
591                // number of non-latin_1 chars between the first non-latin_1 char and the next latin_1 char
592                let err_end = match { iter }.find(|(_, c)| c.to_u32() <= 255) {
593                    Some((i, _)) => ctx.position() + i,
594                    None => ctx.data_len(),
595                };
596                let err_range = err_start..err_end;
597                let replace = ctx.handle_error(errors, err_range.clone(), Some(ERR_REASON))?;
598                match replace {
599                    EncodeReplace::Str(s) => {
600                        if s.as_ref().code_points().any(|c| c.to_u32() > 255) {
601                            return Err(ctx.error_encoding(err_range, Some(ERR_REASON)));
602                        }
603                        out.extend(s.as_ref().code_points().map(|c| c.to_u32() as u8));
604                    }
605                    EncodeReplace::Bytes(b) => {
606                        out.extend_from_slice(b.as_ref());
607                    }
608                }
609            }
610        }
611        out.extend_from_slice(ctx.remaining_data().as_bytes());
612        Ok(out)
613    }
614
615    pub fn decode<Ctx: DecodeContext, E: DecodeErrorHandler<Ctx>>(
616        ctx: Ctx,
617        _errors: &E,
618    ) -> Result<(Wtf8Buf, usize), Ctx::Error> {
619        let out: String = ctx.remaining_data().iter().map(|c| *c as char).collect();
620        let out_len = out.len();
621        Ok((out.into(), out_len))
622    }
623}
624
625pub mod ascii {
626    use super::*;
627    use ::ascii::AsciiStr;
628
629    pub const ENCODING_NAME: &str = "ascii";
630
631    const ERR_REASON: &str = "ordinal not in range(128)";
632
633    #[inline]
634    pub fn encode<Ctx, E>(ctx: Ctx, errors: &E) -> Result<Vec<u8>, Ctx::Error>
635    where
636        Ctx: EncodeContext,
637        E: EncodeErrorHandler<Ctx>,
638    {
639        encode_utf8_compatible(ctx, errors, ERR_REASON, StrKind::Ascii)
640    }
641
642    pub fn decode<Ctx: DecodeContext, E: DecodeErrorHandler<Ctx>>(
643        ctx: Ctx,
644        errors: &E,
645    ) -> Result<(Wtf8Buf, usize), Ctx::Error> {
646        decode_utf8_compatible(
647            ctx,
648            errors,
649            |v| {
650                AsciiStr::from_ascii(v).map(|s| s.as_str()).map_err(|e| {
651                    // SAFETY: as specified in valid_up_to's documentation, input[..e.valid_up_to()]
652                    //         is valid ascii & therefore valid utf8
653                    unsafe { make_decode_err(v, e.valid_up_to(), Some(1)) }
654                })
655            },
656            |_rest, err_len| HandleResult::Error {
657                err_len,
658                reason: ERR_REASON,
659            },
660        )
661    }
662}