Skip to main content

rustpython_vm/
codecs.rs

1use alloc::borrow::Cow;
2use core::ops::{Deref, Range};
3use std::collections::HashMap;
4
5use rustpython_common::{
6    ascii,
7    borrow::BorrowedValue,
8    encodings::{
9        CodecContext, DecodeContext, DecodeErrorHandler, EncodeContext, EncodeErrorHandler,
10        EncodeReplace, StrBuffer, StrSize, errors,
11    },
12    lock::{OnceCell, PyRwLock},
13    str::StrKind,
14    wtf8::{CodePoint, Wtf8, Wtf8Buf},
15};
16
17use crate::{
18    AsObject, Context, Py, PyObject, PyObjectRef, PyResult, TryFromBorrowedObject, TryFromObject,
19    VirtualMachine,
20    builtins::{
21        PyBaseExceptionRef, PyByteArray, PyBytes, PyBytesRef, PyStr, PyStrRef, PyTuple, PyTupleRef,
22        PyUtf8Str, PyUtf8StrRef,
23    },
24    convert::ToPyObject,
25    function::{ArgBytesLike, PyMethodDef},
26};
27
28pub struct CodecsRegistry {
29    inner: PyRwLock<RegistryInner>,
30}
31
32struct RegistryInner {
33    search_path: Vec<PyObjectRef>,
34    search_cache: HashMap<String, PyCodec>,
35    errors: HashMap<String, PyObjectRef>,
36}
37
38pub(crate) const DEFAULT_ENCODING: &str = "utf-8";
39
40#[derive(Clone)]
41#[repr(transparent)]
42pub struct PyCodec(PyTupleRef);
43
44impl PyCodec {
45    #[inline]
46    pub fn from_tuple(tuple: PyTupleRef) -> Result<Self, PyTupleRef> {
47        if tuple.as_slice().len() == 4 {
48            Ok(Self(tuple))
49        } else {
50            Err(tuple)
51        }
52    }
53
54    #[inline]
55    pub fn into_tuple(self) -> PyTupleRef {
56        self.0
57    }
58
59    #[inline]
60    pub fn as_tuple(&self) -> &Py<PyTuple> {
61        &self.0
62    }
63
64    #[inline]
65    pub fn get_encode_func(&self) -> &PyObject {
66        &self.0.as_slice()[0]
67    }
68
69    #[inline]
70    pub fn get_decode_func(&self) -> &PyObject {
71        &self.0.as_slice()[1]
72    }
73
74    pub fn is_text_codec(&self, vm: &VirtualMachine) -> PyResult<bool> {
75        let is_text = vm.get_attribute_opt(self.0.as_object(), "_is_text_encoding")?;
76        is_text.map_or(Ok(true), |is_text| is_text.try_to_bool(vm))
77    }
78
79    pub fn encode(
80        &self,
81        obj: PyObjectRef,
82        errors: Option<PyUtf8StrRef>,
83        vm: &VirtualMachine,
84    ) -> PyResult {
85        let args = match errors {
86            Some(errors) => vec![obj, errors.into_wtf8().into()],
87            None => vec![obj],
88        };
89        let res = self.get_encode_func().call(args, vm)?;
90        let res = res
91            .downcast::<PyTuple>()
92            .ok()
93            .filter(|tuple| tuple.as_slice().len() == 2)
94            .ok_or_else(|| vm.new_type_error("encoder must return a tuple (object, integer)"))?;
95        // we don't actually care about the integer
96        Ok(res.as_slice()[0].clone())
97    }
98
99    pub fn decode(
100        &self,
101        obj: PyObjectRef,
102        errors: Option<PyUtf8StrRef>,
103        vm: &VirtualMachine,
104    ) -> PyResult {
105        let args = match errors {
106            Some(errors) => vec![obj, errors.into_wtf8().into()],
107            None => vec![obj],
108        };
109        let res = self.get_decode_func().call(args, vm)?;
110        let res = res
111            .downcast::<PyTuple>()
112            .ok()
113            .filter(|tuple| tuple.as_slice().len() == 2)
114            .ok_or_else(|| vm.new_type_error("decoder must return a tuple (object,integer)"))?;
115        // we don't actually care about the integer
116        Ok(res.as_slice()[0].clone())
117    }
118
119    pub fn get_incremental_encoder(
120        &self,
121        errors: Option<PyStrRef>,
122        vm: &VirtualMachine,
123    ) -> PyResult {
124        let args = errors.map_or_else(Vec::new, |e| vec![e.into()]);
125        vm.call_method(self.0.as_object(), "incrementalencoder", args)
126    }
127
128    pub fn get_incremental_decoder(
129        &self,
130        errors: Option<PyStrRef>,
131        vm: &VirtualMachine,
132    ) -> PyResult {
133        let args = errors.map_or_else(Vec::new, |e| vec![e.into()]);
134        vm.call_method(self.0.as_object(), "incrementaldecoder", args)
135    }
136}
137
138impl TryFromObject for PyCodec {
139    fn try_from_object(vm: &VirtualMachine, obj: PyObjectRef) -> PyResult<Self> {
140        obj.downcast::<PyTuple>()
141            .ok()
142            .and_then(|tuple| Self::from_tuple(tuple).ok())
143            .ok_or_else(|| vm.new_type_error("codec search functions must return 4-tuples"))
144    }
145}
146
147impl ToPyObject for PyCodec {
148    #[inline]
149    fn to_pyobject(self, _vm: &VirtualMachine) -> PyObjectRef {
150        self.0.into()
151    }
152}
153
154impl CodecsRegistry {
155    /// Reset the inner RwLock to unlocked state after fork().
156    ///
157    /// # Safety
158    /// Must only be called after fork() in the child process when no other
159    /// threads exist.
160    #[cfg(all(unix, feature = "threading", feature = "host_env"))]
161    pub(crate) unsafe fn reinit_after_fork(&self) {
162        unsafe { crate::common::lock::reinit_rwlock_after_fork(&self.inner) };
163    }
164
165    pub(crate) fn new(ctx: &Context) -> Self {
166        ::rustpython_vm::common::static_cell! {
167            static METHODS: Box<[PyMethodDef]>;
168        }
169
170        let methods = METHODS.get_or_init(|| {
171            macro_rules! error_handler {
172                ($name:literal, $func:ident) => {{
173                    #[cfg(feature = "doc")]
174                    const DOC: crate::function::ItemDoc =
175                        crate::function::ItemDoc::db(concat!("codecs.", $name));
176                    crate::function::PyMethodDef {
177                        name: $name,
178                        func: crate::function::static_func($func),
179                        flags: crate::function::PyMethodFlags::O,
180                        #[cfg(feature = "doc")]
181                        doc_off: DOC.offset,
182                        #[cfg(feature = "doc")]
183                        doc_len: DOC.len,
184                        #[cfg(feature = "doc")]
185                        doc_body_pending: false,
186                        doc: Some(concat!($name, "($self, object, /)\n--\n\n")),
187                    }
188                }};
189            }
190            vec![
191                error_handler!("strict_errors", strict_errors),
192                error_handler!("ignore_errors", ignore_errors),
193                error_handler!("replace_errors", replace_errors),
194                error_handler!("xmlcharrefreplace_errors", xmlcharrefreplace_errors),
195                error_handler!("backslashreplace_errors", backslashreplace_errors),
196                error_handler!("namereplace_errors", namereplace_errors),
197                error_handler!("surrogatepass_errors", surrogatepass_errors),
198                error_handler!("surrogateescape_errors", surrogateescape_errors),
199            ]
200            .into_boxed_slice()
201        });
202
203        let errors = [
204            ("strict", methods[0].build_function(ctx)),
205            ("ignore", methods[1].build_function(ctx)),
206            ("replace", methods[2].build_function(ctx)),
207            ("xmlcharrefreplace", methods[3].build_function(ctx)),
208            ("backslashreplace", methods[4].build_function(ctx)),
209            ("namereplace", methods[5].build_function(ctx)),
210            ("surrogatepass", methods[6].build_function(ctx)),
211            ("surrogateescape", methods[7].build_function(ctx)),
212        ]
213        .into_iter()
214        .map(|(name, f)| (name.to_owned(), f.into()))
215        .collect();
216
217        let inner = RegistryInner {
218            search_path: Vec::new(),
219            search_cache: HashMap::new(),
220            errors,
221        };
222
223        Self {
224            inner: PyRwLock::new(inner),
225        }
226    }
227
228    pub fn register(&self, search_function: PyObjectRef, vm: &VirtualMachine) -> PyResult<()> {
229        if !search_function.is_callable() {
230            return Err(vm.new_type_error("argument must be callable"));
231        }
232
233        self.inner.write().search_path.push(search_function);
234        Ok(())
235    }
236
237    pub fn unregister(&self, search_function: &PyObject) {
238        let mut inner = self.inner.write();
239        // Do nothing if search_path is not created yet or was cleared.
240        if inner.search_path.is_empty() {
241            return;
242        }
243
244        for (i, item) in inner.search_path.iter().enumerate() {
245            if item.get_id() == search_function.get_id() {
246                if !inner.search_cache.is_empty() {
247                    inner.search_cache.clear();
248                }
249                inner.search_path.remove(i);
250                return;
251            }
252        }
253    }
254
255    pub(crate) fn register_manual(&self, name: &str, codec: PyCodec) {
256        let name = normalize_encoding_name(name);
257        self.inner
258            .write()
259            .search_cache
260            .insert(name.into_owned(), codec);
261    }
262
263    pub fn lookup(&self, encoding: &str, vm: &VirtualMachine) -> PyResult<PyCodec> {
264        let encoding = normalize_encoding_name(encoding);
265        let search_path = {
266            let inner = self.inner.read();
267            if let Some(codec) = inner.search_cache.get(encoding.as_ref()) {
268                // hit cache
269                return Ok(codec.clone());
270            }
271            inner.search_path.clone()
272        };
273
274        let encoding: PyUtf8StrRef = vm.ctx.new_utf8_str(encoding.as_ref());
275        for func in search_path {
276            let res = func.call((encoding.clone(),), vm)?;
277            let res: Option<PyCodec> = res.try_into_value(vm)?;
278            if let Some(codec) = res {
279                let mut inner = self.inner.write();
280                // someone might have raced us to this, so use theirs
281                let codec = inner
282                    .search_cache
283                    .entry(encoding.as_str().to_owned())
284                    .or_insert(codec);
285                return Ok(codec.clone());
286            }
287        }
288
289        Err(vm.new_lookup_error(format!("unknown encoding: {encoding}")))
290    }
291
292    fn _lookup_text_encoding(
293        &self,
294        encoding: &str,
295        generic_func: &str,
296        vm: &VirtualMachine,
297    ) -> PyResult<PyCodec> {
298        let codec = self.lookup(encoding, vm)?;
299
300        if codec.is_text_codec(vm)? {
301            Ok(codec)
302        } else {
303            Err(vm.new_lookup_error(format!(
304                "'{encoding}' is not a text encoding; use {generic_func} to handle arbitrary codecs"
305            )))
306        }
307    }
308
309    pub fn forget(&self, encoding: &str) -> Option<PyCodec> {
310        let encoding = normalize_encoding_name(encoding);
311        self.inner.write().search_cache.remove(encoding.as_ref())
312    }
313
314    pub fn encode(
315        &self,
316        obj: PyObjectRef,
317        encoding: &str,
318        errors: Option<PyUtf8StrRef>,
319        vm: &VirtualMachine,
320    ) -> PyResult {
321        let codec = self.lookup(encoding, vm)?;
322        codec.encode(obj, errors, vm).inspect_err(|exc| {
323            Self::add_codec_note(exc, "encoding", encoding, vm);
324        })
325    }
326
327    pub fn decode(
328        &self,
329        obj: PyObjectRef,
330        encoding: &str,
331        errors: Option<PyUtf8StrRef>,
332        vm: &VirtualMachine,
333    ) -> PyResult {
334        let codec = self.lookup(encoding, vm)?;
335        codec.decode(obj, errors, vm).inspect_err(|exc| {
336            Self::add_codec_note(exc, "decoding", encoding, vm);
337        })
338    }
339
340    pub fn encode_text(
341        &self,
342        obj: PyStrRef,
343        encoding: &str,
344        errors: Option<PyUtf8StrRef>,
345        vm: &VirtualMachine,
346    ) -> PyResult<PyBytesRef> {
347        if let Some(b) =
348            Self::encode_fast(&obj, encoding, errors.as_deref(), vm).inspect_err(|exc| {
349                Self::add_codec_note(exc, "encoding", encoding, vm);
350            })?
351        {
352            return Ok(b);
353        }
354        let codec = self._lookup_text_encoding(encoding, "codecs.encode()", vm)?;
355        codec
356            .encode(obj.into(), errors, vm)
357            .inspect_err(|exc| {
358                Self::add_codec_note(exc, "encoding", encoding, vm);
359            })?
360            .downcast()
361            .map_err(|obj| {
362                vm.new_type_error(format!(
363                    "'{}' encoder returned '{}' instead of 'bytes'; use codecs.encode() to \
364                     encode to arbitrary types",
365                    encoding,
366                    obj.class().name(),
367                ))
368            })
369    }
370
371    /// The text an encoded object holds. An object holding nothing decodes to
372    /// the empty string without the encoding ever being looked up.
373    /// = PyUnicode_FromEncodedObject
374    pub fn decode_text_object(
375        &self,
376        obj: PyObjectRef,
377        encoding: &str,
378        errors: Option<PyUtf8StrRef>,
379        vm: &VirtualMachine,
380    ) -> PyResult<PyStrRef> {
381        let empty = if let Some(bytes) = obj.downcast_ref::<PyBytes>() {
382            bytes.as_bytes().is_empty()
383        } else if let Some(bytes) = obj.downcast_ref::<PyByteArray>() {
384            bytes.borrow_buf().is_empty()
385        } else {
386            false
387        };
388        if empty {
389            return Ok(vm.ctx.empty_str.to_owned());
390        }
391        self.decode_text(obj, encoding, errors, vm)
392    }
393
394    pub fn decode_text(
395        &self,
396        obj: PyObjectRef,
397        encoding: &str,
398        errors: Option<PyUtf8StrRef>,
399        vm: &VirtualMachine,
400    ) -> PyResult<PyStrRef> {
401        if let Some(s) =
402            Self::decode_fast(&obj, encoding, errors.as_deref(), vm).inspect_err(|exc| {
403                Self::add_codec_note(exc, "decoding", encoding, vm);
404            })?
405        {
406            return Ok(s);
407        }
408        let codec = self._lookup_text_encoding(encoding, "codecs.decode()", vm)?;
409        codec
410            .decode(obj, errors, vm)
411            .inspect_err(|exc| {
412                Self::add_codec_note(exc, "decoding", encoding, vm);
413            })?
414            .downcast()
415            .map_err(|obj| {
416                vm.new_type_error(format!(
417                    "'{}' decoder returned '{}' instead of 'str'; use codecs.decode() to \
418                 decode to arbitrary types",
419                    encoding,
420                    obj.class().name(),
421                ))
422            })
423    }
424
425    /// Fast path for decoding with the "utf-8", "ascii" or "latin-1"
426    /// encodings (and their common aliases), used by `bytes.decode()` /
427    /// `str(bytes, encoding)`.
428    ///
429    /// CPython's `PyUnicode_Decode` special-cases a handful of built-in
430    /// encoding names -- these among them -- and decodes them directly in C
431    /// without ever consulting the codec registry, so a `codecs.register()`
432    /// override does not affect `str(b"...", "utf-8")` there either. This
433    /// mirrors that behavior, skipping the registry lookup, the
434    /// `_is_text_encoding` attribute probe, and the round trip through the
435    /// pure-Python `encodings.*.decode` wrapper and its 2-tuple result.
436    fn decode_fast(
437        obj: &PyObject,
438        encoding: &str,
439        errors: Option<&Py<PyUtf8Str>>,
440        vm: &VirtualMachine,
441    ) -> PyResult<Option<PyStrRef>> {
442        let Some(fast) = FastCodec::classify(encoding) else {
443            return Ok(None);
444        };
445        let Ok(data) = ArgBytesLike::try_from_object(vm, obj.to_owned()) else {
446            return Ok(None);
447        };
448        let errors_handler = ErrorsHandler::new(errors, vm);
449        let (decoded, _consumed) = match fast {
450            FastCodec::Utf8 => {
451                let ctx = PyDecodeContext::new(DEFAULT_ENCODING, &data, vm);
452                crate::common::encodings::utf8::decode(ctx, &errors_handler, true)?
453            }
454            FastCodec::Ascii => {
455                let ctx =
456                    PyDecodeContext::new(crate::common::encodings::ascii::ENCODING_NAME, &data, vm);
457                crate::common::encodings::ascii::decode(ctx, &errors_handler)?
458            }
459            FastCodec::Latin1 => {
460                let ctx = PyDecodeContext::new(
461                    crate::common::encodings::latin_1::ENCODING_NAME,
462                    &data,
463                    vm,
464                );
465                crate::common::encodings::latin_1::decode(ctx, &errors_handler)?
466            }
467        };
468        Ok(Some(vm.ctx.new_str(decoded)))
469    }
470
471    /// Fast path for encoding with the "utf-8", "ascii" or "latin-1"
472    /// encodings (and their common aliases), used by `str.encode()` /
473    /// `bytes(s, encoding)`. See [`Self::decode_fast`] for the rationale.
474    fn encode_fast(
475        obj: &Py<PyStr>,
476        encoding: &str,
477        errors: Option<&Py<PyUtf8Str>>,
478        vm: &VirtualMachine,
479    ) -> PyResult<Option<PyBytesRef>> {
480        let Some(fast) = FastCodec::classify(encoding) else {
481            return Ok(None);
482        };
483        if fast == FastCodec::Utf8 && obj.is_utf8() {
484            // No surrogates in the string, so encoding is just its existing
485            // (already UTF-8) bytes, whatever the error handler is: it can
486            // never be reached since there's nothing to fail to encode.
487            return Ok(Some(vm.ctx.new_bytes(obj.as_bytes().to_vec())));
488        }
489        let errors_handler = ErrorsHandler::new(errors, vm);
490        let encoded = match fast {
491            FastCodec::Utf8 => {
492                let ctx = PyEncodeContext::new(DEFAULT_ENCODING, obj, vm);
493                crate::common::encodings::utf8::encode(ctx, &errors_handler)?
494            }
495            FastCodec::Ascii => {
496                let ctx =
497                    PyEncodeContext::new(crate::common::encodings::ascii::ENCODING_NAME, obj, vm);
498                crate::common::encodings::ascii::encode(ctx, &errors_handler)?
499            }
500            FastCodec::Latin1 => {
501                let ctx =
502                    PyEncodeContext::new(crate::common::encodings::latin_1::ENCODING_NAME, obj, vm);
503                crate::common::encodings::latin_1::encode(ctx, &errors_handler)?
504            }
505        };
506        Ok(Some(vm.ctx.new_bytes(encoded)))
507    }
508
509    fn add_codec_note(
510        exc: &crate::builtins::PyBaseExceptionRef,
511        operation: &str,
512        encoding: &str,
513        vm: &VirtualMachine,
514    ) {
515        let note = format!("{operation} with '{encoding}' codec failed");
516        let _ = vm.call_method(exc.as_object(), "add_note", (vm.ctx.new_str(note),));
517    }
518
519    pub fn register_error(&self, name: String, handler: PyObjectRef) -> Option<PyObjectRef> {
520        self.inner.write().errors.insert(name, handler)
521    }
522
523    pub fn unregister_error(&self, name: &str, vm: &VirtualMachine) -> PyResult<bool> {
524        const BUILTIN_ERROR_HANDLERS: &[&str] = &[
525            "strict",
526            "ignore",
527            "replace",
528            "xmlcharrefreplace",
529            "backslashreplace",
530            "namereplace",
531            "surrogatepass",
532            "surrogateescape",
533        ];
534        if BUILTIN_ERROR_HANDLERS.contains(&name) {
535            return Err(vm.new_value_error(format!(
536                "cannot un-register built-in error handler '{name}'"
537            )));
538        }
539        Ok(self.inner.write().errors.remove(name).is_some())
540    }
541
542    pub fn lookup_error_opt(&self, name: &str) -> Option<PyObjectRef> {
543        self.inner.read().errors.get(name).cloned()
544    }
545
546    pub fn lookup_error(&self, name: &str, vm: &VirtualMachine) -> PyResult<PyObjectRef> {
547        self.lookup_error_opt(name)
548            .ok_or_else(|| vm.new_lookup_error(format!("unknown error handler name '{name}'")))
549    }
550}
551
552/// Encodings recognized by [`CodecsRegistry::decode_fast`] and
553/// [`CodecsRegistry::encode_fast`], covering the built-in encodings whose
554/// native (de)coders are cheap to call directly, bypassing the codec
555/// registry round trip. Matches the alias sets in `Lib/encodings/aliases.py`.
556#[derive(Clone, Copy, Eq, PartialEq)]
557enum FastCodec {
558    Utf8,
559    Ascii,
560    Latin1,
561}
562
563impl FastCodec {
564    fn classify(encoding: &str) -> Option<Self> {
565        const UTF8_ALIASES: &[&str] = &[
566            "utf_8",
567            "utf8",
568            "u8",
569            "utf",
570            "utf8_ucs2",
571            "utf8_ucs4",
572            "cp65001",
573        ];
574        const ASCII_ALIASES: &[&str] = &[
575            "ascii",
576            "646",
577            "ansi_x3.4_1968",
578            "ansi_x3_4_1968",
579            "ansi_x3.4_1986",
580            "cp367",
581            "csascii",
582            "ibm367",
583            "iso646_us",
584            "iso_646.irv_1991",
585            "iso_ir_6",
586            "us",
587            "us_ascii",
588        ];
589        const LATIN1_ALIASES: &[&str] = &[
590            "latin_1",
591            "8859",
592            "cp819",
593            "csisolatin1",
594            "ibm819",
595            "iso8859",
596            "iso8859_1",
597            "iso_8859_1",
598            "iso_8859_1_1987",
599            "iso_ir_100",
600            "l1",
601            "latin",
602            "latin1",
603        ];
604        // The spellings used most often, matched before normalizing allocates.
605        match encoding {
606            "utf-8" | "utf8" | "UTF-8" => return Some(Self::Utf8),
607            "ascii" => return Some(Self::Ascii),
608            "latin-1" | "latin1" | "iso-8859-1" => return Some(Self::Latin1),
609            _ => {}
610        }
611        let normalized = normalize_encoding_name(encoding);
612        let normalized = normalized.as_ref();
613        if UTF8_ALIASES.contains(&normalized) {
614            Some(Self::Utf8)
615        } else if ASCII_ALIASES.contains(&normalized) {
616            Some(Self::Ascii)
617        } else if LATIN1_ALIASES.contains(&normalized) {
618            Some(Self::Latin1)
619        } else {
620            None
621        }
622    }
623}
624
625fn normalize_encoding_name(encoding: &str) -> Cow<'_, str> {
626    // _Py_normalize_encoding: collapse non-alphanumeric/non-dot chars into
627    // single underscore, strip non-ASCII, lowercase ASCII letters.
628    let needs_transform = encoding
629        .bytes()
630        .any(|b| b.is_ascii_uppercase() || !b.is_ascii_alphanumeric() && b != b'.');
631    if !needs_transform {
632        return encoding.into();
633    }
634    let mut out = String::with_capacity(encoding.len());
635    let mut punct = false;
636    for c in encoding.chars() {
637        if c.is_ascii_alphanumeric() || c == '.' {
638            if punct && !out.is_empty() {
639                out.push('_');
640            }
641            out.push(c.to_ascii_lowercase());
642            punct = false;
643        } else {
644            punct = true;
645        }
646    }
647    out.into()
648}
649
650#[derive(Clone, Copy, Eq, PartialEq)]
651enum StandardEncoding {
652    Utf8,
653    Utf16Be,
654    Utf16Le,
655    Utf32Be,
656    Utf32Le,
657}
658
659impl StandardEncoding {
660    const UTF_16_NE: Self = cfg_select! {
661        target_endian = "little" => Self::Utf16Le,
662        target_endian = "big" => Self::Utf16Be,
663    };
664
665    const UTF_32_NE: Self = cfg_select! {
666        target_endian = "little" => Self::Utf32Le,
667        target_endian = "big" => Self::Utf32Be,
668    };
669
670    fn parse(encoding: &str) -> Option<Self> {
671        if let Some(encoding) = encoding.to_lowercase().strip_prefix("utf") {
672            let encoding = encoding
673                .strip_prefix(|c| ['-', '_'].contains(&c))
674                .unwrap_or(encoding);
675
676            if encoding == "8" {
677                Some(Self::Utf8)
678            } else if let Some(encoding) = encoding.strip_prefix("16") {
679                if encoding.is_empty() {
680                    return Some(Self::UTF_16_NE);
681                }
682
683                let encoding = encoding.strip_prefix(['-', '_']).unwrap_or(encoding);
684                match encoding {
685                    "be" => Some(Self::Utf16Be),
686                    "le" => Some(Self::Utf16Le),
687                    _ => None,
688                }
689            } else if let Some(encoding) = encoding.strip_prefix("32") {
690                if encoding.is_empty() {
691                    return Some(Self::UTF_32_NE);
692                }
693
694                let encoding = encoding.strip_prefix(['-', '_']).unwrap_or(encoding);
695                match encoding {
696                    "be" => Some(Self::Utf32Be),
697                    "le" => Some(Self::Utf32Le),
698                    _ => None,
699                }
700            } else {
701                None
702            }
703        } else if encoding == "cp65001" {
704            Some(Self::Utf8)
705        } else {
706            None
707        }
708    }
709}
710
711struct SurrogatePass;
712
713impl<'a> EncodeErrorHandler<PyEncodeContext<'a>> for SurrogatePass {
714    fn handle_encode_error(
715        &self,
716        ctx: &mut PyEncodeContext<'a>,
717        range: Range<StrSize>,
718        reason: Option<&str>,
719    ) -> PyResult<(EncodeReplace<PyEncodeContext<'a>>, StrSize)> {
720        let standard_encoding = StandardEncoding::parse(ctx.encoding)
721            .ok_or_else(|| ctx.error_encoding(range.clone(), reason))?;
722        let err_str = &ctx.full_data()[range.start.bytes..range.end.bytes];
723        let num_chars = range.end.chars - range.start.chars;
724        let mut out: Vec<u8> = Vec::with_capacity(num_chars * 4);
725        for ch in err_str.code_points() {
726            let c = ch.to_u32();
727
728            if !(0xd800..=0xdfff).contains(&c) {
729                // Not a surrogate, fail with original exception
730                return Err(ctx.error_encoding(range, reason));
731            }
732
733            match standard_encoding {
734                StandardEncoding::Utf8 => out.extend(ch.encode_wtf8(&mut [0; 4]).as_bytes()),
735                StandardEncoding::Utf16Le => out.extend((c as u16).to_le_bytes()),
736                StandardEncoding::Utf16Be => out.extend((c as u16).to_be_bytes()),
737                StandardEncoding::Utf32Le => out.extend(c.to_le_bytes()),
738                StandardEncoding::Utf32Be => out.extend(c.to_be_bytes()),
739            }
740        }
741        Ok((EncodeReplace::Bytes(ctx.bytes(out)), range.end))
742    }
743}
744
745impl<'a> DecodeErrorHandler<PyDecodeContext<'a>> for SurrogatePass {
746    fn handle_decode_error(
747        &self,
748        ctx: &mut PyDecodeContext<'a>,
749        byte_range: Range<usize>,
750        reason: Option<&str>,
751    ) -> PyResult<(PyStrRef, usize)> {
752        let standard_encoding = StandardEncoding::parse(ctx.encoding)
753            .ok_or_else(|| ctx.error_decoding(byte_range.clone(), reason))?;
754
755        let s = ctx.full_data();
756        debug_assert!(byte_range.start <= s.len().saturating_sub(1));
757        debug_assert!(byte_range.end >= 1.min(s.len()));
758        debug_assert!(byte_range.end <= s.len());
759
760        // Try decoding a single surrogate character. If there are more,
761        // let the codec call us again.
762        let p = &s[byte_range.start..];
763
764        fn slice<const N: usize>(p: &[u8]) -> Option<[u8; N]> {
765            p.first_chunk().copied()
766        }
767
768        let c = match standard_encoding {
769            StandardEncoding::Utf8 => {
770                // it's a three-byte code
771                slice::<3>(p)
772                    .filter(|&[a, b, c]| {
773                        (u32::from(a) & 0xf0) == 0xe0
774                            && (u32::from(b) & 0xc0) == 0x80
775                            && (u32::from(c) & 0xc0) == 0x80
776                    })
777                    .map(|[a, b, c]| {
778                        ((u32::from(a) & 0x0f) << 12)
779                            + ((u32::from(b) & 0x3f) << 6)
780                            + (u32::from(c) & 0x3f)
781                    })
782            }
783            StandardEncoding::Utf16Le => slice(p).map(u16::from_le_bytes).map(u32::from),
784            StandardEncoding::Utf16Be => slice(p).map(u16::from_be_bytes).map(u32::from),
785            StandardEncoding::Utf32Le => slice(p).map(u32::from_le_bytes),
786            StandardEncoding::Utf32Be => slice(p).map(u32::from_be_bytes),
787        };
788        let byte_length = match standard_encoding {
789            StandardEncoding::Utf8 => 3,
790            StandardEncoding::Utf16Be | StandardEncoding::Utf16Le => 2,
791            StandardEncoding::Utf32Be | StandardEncoding::Utf32Le => 4,
792        };
793
794        // !Py_UNICODE_IS_SURROGATE
795        let c = c
796            .and_then(CodePoint::from_u32)
797            .filter(|c| matches!(c.to_u32(), 0xd800..=0xdfff))
798            .ok_or_else(|| ctx.error_decoding(byte_range.clone(), reason))?;
799
800        Ok((ctx.string(c.into()), byte_range.start + byte_length))
801    }
802}
803
804pub(crate) struct PyEncodeContext<'a> {
805    vm: &'a VirtualMachine,
806    encoding: &'a str,
807    data: &'a Py<PyStr>,
808    pos: StrSize,
809    exception: OnceCell<PyBaseExceptionRef>,
810}
811
812impl<'a> PyEncodeContext<'a> {
813    pub(crate) fn new(encoding: &'a str, data: &'a Py<PyStr>, vm: &'a VirtualMachine) -> Self {
814        Self {
815            vm,
816            encoding,
817            data,
818            pos: StrSize::default(),
819            exception: OnceCell::new(),
820        }
821    }
822}
823
824impl CodecContext for PyEncodeContext<'_> {
825    type Error = PyBaseExceptionRef;
826
827    type StrBuf = PyStrRef;
828
829    type BytesBuf = PyBytesRef;
830
831    fn string(&self, s: Wtf8Buf) -> Self::StrBuf {
832        self.vm.ctx.new_str(s)
833    }
834
835    fn bytes(&self, b: Vec<u8>) -> Self::BytesBuf {
836        self.vm.ctx.new_bytes(b)
837    }
838}
839
840impl EncodeContext for PyEncodeContext<'_> {
841    fn full_data(&self) -> &Wtf8 {
842        self.data.as_wtf8()
843    }
844
845    fn data_len(&self) -> StrSize {
846        StrSize {
847            bytes: self.data.byte_len(),
848            chars: self.data.char_len(),
849        }
850    }
851
852    fn remaining_data(&self) -> &Wtf8 {
853        &self.full_data()[self.pos.bytes..]
854    }
855
856    fn position(&self) -> StrSize {
857        self.pos
858    }
859
860    fn restart_from(&mut self, pos: StrSize) -> Result<(), Self::Error> {
861        if pos.chars > self.data.char_len() {
862            return Err(self.vm.new_index_error(format!(
863                "position {} from error handler out of bounds",
864                pos.chars
865            )));
866        }
867        assert!(
868            self.data.as_wtf8().is_code_point_boundary(pos.bytes),
869            "invalid pos {pos:?} for {:?}",
870            self.data.as_wtf8()
871        );
872        self.pos = pos;
873        Ok(())
874    }
875
876    fn error_encoding(&self, range: Range<StrSize>, reason: Option<&str>) -> Self::Error {
877        let vm = self.vm;
878        match self.exception.get() {
879            Some(exc) => {
880                match update_unicode_error_attrs(
881                    exc.as_object(),
882                    range.start.chars,
883                    range.end.chars,
884                    reason,
885                    vm,
886                ) {
887                    Ok(()) => exc.clone(),
888                    Err(e) => e,
889                }
890            }
891            None => self
892                .exception
893                .get_or_init(|| {
894                    let reason = reason.expect(
895                        "should only ever pass reason: None if an exception is already set",
896                    );
897                    vm.new_unicode_encode_error(
898                        vm.ctx.new_str(self.encoding),
899                        self.data.to_owned(),
900                        range.start.chars,
901                        range.end.chars,
902                        vm.ctx.new_str(reason),
903                    )
904                })
905                .clone(),
906        }
907    }
908}
909
910pub(crate) struct PyDecodeContext<'a> {
911    vm: &'a VirtualMachine,
912    encoding: &'a str,
913    data: PyDecodeData<'a>,
914    orig_bytes: Option<&'a Py<PyBytes>>,
915    pos: usize,
916    exception: OnceCell<PyBaseExceptionRef>,
917}
918
919enum PyDecodeData<'a> {
920    Original(BorrowedValue<'a, [u8]>),
921    Modified(PyBytesRef),
922}
923
924impl Deref for PyDecodeData<'_> {
925    type Target = [u8];
926
927    fn deref(&self) -> &Self::Target {
928        match self {
929            PyDecodeData::Original(data) => data,
930            PyDecodeData::Modified(data) => data.as_bytes(),
931        }
932    }
933}
934
935impl<'a> PyDecodeContext<'a> {
936    pub(crate) fn new(encoding: &'a str, data: &'a ArgBytesLike, vm: &'a VirtualMachine) -> Self {
937        Self {
938            vm,
939            encoding,
940            data: PyDecodeData::Original(data.borrow_buf()),
941            orig_bytes: data.as_object().downcast_ref(),
942            pos: 0,
943            exception: OnceCell::new(),
944        }
945    }
946}
947
948impl CodecContext for PyDecodeContext<'_> {
949    type Error = PyBaseExceptionRef;
950
951    type StrBuf = PyStrRef;
952
953    type BytesBuf = PyBytesRef;
954
955    fn string(&self, s: Wtf8Buf) -> Self::StrBuf {
956        self.vm.ctx.new_str(s)
957    }
958
959    fn bytes(&self, b: Vec<u8>) -> Self::BytesBuf {
960        self.vm.ctx.new_bytes(b)
961    }
962}
963
964impl DecodeContext for PyDecodeContext<'_> {
965    fn full_data(&self) -> &[u8] {
966        &self.data
967    }
968
969    fn remaining_data(&self) -> &[u8] {
970        &self.data[self.pos..]
971    }
972
973    fn position(&self) -> usize {
974        self.pos
975    }
976
977    fn advance(&mut self, by: usize) {
978        self.pos += by;
979    }
980
981    fn restart_from(&mut self, pos: usize) -> Result<(), Self::Error> {
982        if pos > self.data.len() {
983            return Err(self
984                .vm
985                .new_index_error(format!("position {pos} from error handler out of bounds",)));
986        }
987        self.pos = pos;
988        Ok(())
989    }
990
991    fn error_decoding(&self, byte_range: Range<usize>, reason: Option<&str>) -> Self::Error {
992        let vm = self.vm;
993
994        match self.exception.get() {
995            Some(exc) => {
996                match update_unicode_error_attrs(
997                    exc.as_object(),
998                    byte_range.start,
999                    byte_range.end,
1000                    reason,
1001                    vm,
1002                ) {
1003                    Ok(()) => exc.clone(),
1004                    Err(e) => e,
1005                }
1006            }
1007            None => self
1008                .exception
1009                .get_or_init(|| {
1010                    let reason = reason.expect(
1011                        "should only ever pass reason: None if an exception is already set",
1012                    );
1013                    let data = if let Some(bytes) = self.orig_bytes {
1014                        bytes.to_owned()
1015                    } else {
1016                        vm.ctx.new_bytes(self.data.to_vec())
1017                    };
1018                    vm.new_unicode_decode_error(
1019                        vm.ctx.new_str(self.encoding),
1020                        data,
1021                        byte_range.start,
1022                        byte_range.end,
1023                        vm.ctx.new_str(reason),
1024                    )
1025                })
1026                .clone(),
1027        }
1028    }
1029}
1030
1031#[derive(strum_macros::EnumString)]
1032#[strum(serialize_all = "lowercase")]
1033enum StandardError {
1034    Strict,
1035    Ignore,
1036    Replace,
1037    XmlCharRefReplace,
1038    BackslashReplace,
1039    SurrogatePass,
1040    SurrogateEscape,
1041}
1042
1043impl<'a> EncodeErrorHandler<PyEncodeContext<'a>> for StandardError {
1044    fn handle_encode_error(
1045        &self,
1046        ctx: &mut PyEncodeContext<'a>,
1047        range: Range<StrSize>,
1048        reason: Option<&str>,
1049    ) -> PyResult<(EncodeReplace<PyEncodeContext<'a>>, StrSize)> {
1050        match self {
1051            Self::Strict => errors::Strict.handle_encode_error(ctx, range, reason),
1052            Self::Ignore => errors::Ignore.handle_encode_error(ctx, range, reason),
1053            Self::Replace => errors::Replace.handle_encode_error(ctx, range, reason),
1054            Self::XmlCharRefReplace => {
1055                errors::XmlCharRefReplace.handle_encode_error(ctx, range, reason)
1056            }
1057            Self::BackslashReplace => {
1058                errors::BackslashReplace.handle_encode_error(ctx, range, reason)
1059            }
1060            Self::SurrogatePass => SurrogatePass.handle_encode_error(ctx, range, reason),
1061            Self::SurrogateEscape => {
1062                errors::SurrogateEscape.handle_encode_error(ctx, range, reason)
1063            }
1064        }
1065    }
1066}
1067
1068impl<'a> DecodeErrorHandler<PyDecodeContext<'a>> for StandardError {
1069    fn handle_decode_error(
1070        &self,
1071        ctx: &mut PyDecodeContext<'a>,
1072        byte_range: Range<usize>,
1073        reason: Option<&str>,
1074    ) -> PyResult<(PyStrRef, usize)> {
1075        match self {
1076            Self::Strict => errors::Strict.handle_decode_error(ctx, byte_range, reason),
1077            Self::Ignore => errors::Ignore.handle_decode_error(ctx, byte_range, reason),
1078            Self::Replace => errors::Replace.handle_decode_error(ctx, byte_range, reason),
1079            Self::XmlCharRefReplace => Err(ctx
1080                .vm
1081                .new_type_error("don't know how to handle UnicodeDecodeError in error callback")),
1082            Self::BackslashReplace => {
1083                errors::BackslashReplace.handle_decode_error(ctx, byte_range, reason)
1084            }
1085            Self::SurrogatePass => self::SurrogatePass.handle_decode_error(ctx, byte_range, reason),
1086            Self::SurrogateEscape => {
1087                errors::SurrogateEscape.handle_decode_error(ctx, byte_range, reason)
1088            }
1089        }
1090    }
1091}
1092
1093pub(crate) struct ErrorsHandler<'a> {
1094    errors: &'a Py<PyUtf8Str>,
1095    resolved: OnceCell<ResolvedError>,
1096}
1097
1098enum ResolvedError {
1099    Standard(StandardError),
1100    Handler(PyObjectRef),
1101}
1102
1103impl<'a> ErrorsHandler<'a> {
1104    #[inline]
1105    pub(crate) fn new(errors: Option<&'a Py<PyUtf8Str>>, vm: &VirtualMachine) -> Self {
1106        if let Some(errors) = errors {
1107            Self {
1108                errors,
1109                resolved: OnceCell::new(),
1110            }
1111        } else {
1112            Self {
1113                errors: identifier_utf8!(vm, strict),
1114                resolved: OnceCell::from(ResolvedError::Standard(StandardError::Strict)),
1115            }
1116        }
1117    }
1118
1119    #[inline]
1120    fn resolve(&self, vm: &VirtualMachine) -> PyResult<&ResolvedError> {
1121        if let Some(val) = self.resolved.get() {
1122            return Ok(val);
1123        }
1124        let errors_str = self.errors.as_str();
1125        let val = if let Ok(standard) = errors_str.parse() {
1126            ResolvedError::Standard(standard)
1127        } else {
1128            vm.state
1129                .codec_registry
1130                .lookup_error(errors_str, vm)
1131                .map(ResolvedError::Handler)?
1132        };
1133        let _ = self.resolved.set(val);
1134        Ok(self.resolved.get().unwrap())
1135    }
1136}
1137
1138impl StrBuffer for PyStrRef {
1139    fn is_compatible_with(&self, kind: StrKind) -> bool {
1140        self.kind() <= kind
1141    }
1142}
1143
1144impl<'a> EncodeErrorHandler<PyEncodeContext<'a>> for ErrorsHandler<'_> {
1145    fn handle_encode_error(
1146        &self,
1147        ctx: &mut PyEncodeContext<'a>,
1148        range: Range<StrSize>,
1149        reason: Option<&str>,
1150    ) -> PyResult<(EncodeReplace<PyEncodeContext<'a>>, StrSize)> {
1151        let vm = ctx.vm;
1152        let handler = match self.resolve(vm)? {
1153            ResolvedError::Standard(standard) => {
1154                return standard.handle_encode_error(ctx, range, reason);
1155            }
1156            ResolvedError::Handler(handler) => handler,
1157        };
1158        let encode_exc = ctx.error_encoding(range.clone(), reason);
1159        let res = handler.call((encode_exc,), vm)?;
1160        let tuple_err =
1161            || vm.new_type_error("encoding error handler must return (str/bytes, int) tuple");
1162        let (replace, restart) = match res.downcast_ref::<PyTuple>().map(|tup| tup.as_slice()) {
1163            Some([replace, restart]) => (replace.clone(), restart),
1164            _ => return Err(tuple_err()),
1165        };
1166        let replace = match_class!(match replace {
1167            s @ PyStr => EncodeReplace::Str(s),
1168            b @ PyBytes => EncodeReplace::Bytes(b),
1169            _ => return Err(tuple_err()),
1170        });
1171        let restart = isize::try_from_borrowed_object(vm, restart).map_err(|_| tuple_err())?;
1172        let restart = if restart < 0 {
1173            // will still be out of bounds if it underflows ¯\_(ツ)_/¯
1174            ctx.data.char_len().wrapping_sub(restart.unsigned_abs())
1175        } else {
1176            restart as usize
1177        };
1178        let restart = if restart == range.end.chars {
1179            range.end
1180        } else {
1181            StrSize {
1182                chars: restart,
1183                bytes: ctx
1184                    .data
1185                    .as_wtf8()
1186                    .code_point_indices()
1187                    .nth(restart)
1188                    .map_or_else(|| ctx.data.byte_len(), |(i, _)| i),
1189            }
1190        };
1191        Ok((replace, restart))
1192    }
1193}
1194
1195impl<'a> DecodeErrorHandler<PyDecodeContext<'a>> for ErrorsHandler<'_> {
1196    fn handle_decode_error(
1197        &self,
1198        ctx: &mut PyDecodeContext<'a>,
1199        byte_range: Range<usize>,
1200        reason: Option<&str>,
1201    ) -> PyResult<(PyStrRef, usize)> {
1202        let vm = ctx.vm;
1203        let handler = match self.resolve(vm)? {
1204            ResolvedError::Standard(standard) => {
1205                return standard.handle_decode_error(ctx, byte_range, reason);
1206            }
1207            ResolvedError::Handler(handler) => handler,
1208        };
1209        let decode_exc = ctx.error_decoding(byte_range, reason);
1210        let data_bytes: PyObjectRef = decode_exc.as_object().get_attr("object", vm)?;
1211        let res = handler.call((decode_exc.clone(),), vm)?;
1212        let new_data = decode_exc.as_object().get_attr("object", vm)?;
1213        if !new_data.is(&data_bytes) {
1214            let new_data: PyBytesRef = new_data
1215                .downcast()
1216                .map_err(|_| vm.new_type_error("object attribute must be bytes"))?;
1217            ctx.data = PyDecodeData::Modified(new_data);
1218        }
1219        let data = &*ctx.data;
1220        let tuple_err = || vm.new_type_error("decoding error handler must return (str, int) tuple");
1221        match res.downcast_ref::<PyTuple>().map(|tup| tup.as_slice()) {
1222            Some([replace, restart]) => {
1223                let replace = replace
1224                    .downcast_ref::<PyStr>()
1225                    .ok_or_else(tuple_err)?
1226                    .to_owned();
1227                let restart =
1228                    isize::try_from_borrowed_object(vm, restart).map_err(|_| tuple_err())?;
1229                let restart = if restart < 0 {
1230                    // will still be out of bounds if it underflows ¯\_(ツ)_/¯
1231                    data.len().wrapping_sub(restart.unsigned_abs())
1232                } else {
1233                    restart as usize
1234                };
1235                Ok((replace, restart))
1236            }
1237            _ => Err(tuple_err()),
1238        }
1239    }
1240}
1241
1242fn call_native_encode_error<E>(
1243    handler: E,
1244    err: PyObjectRef,
1245    vm: &VirtualMachine,
1246) -> PyResult<(PyObjectRef, usize)>
1247where
1248    for<'a> E: EncodeErrorHandler<PyEncodeContext<'a>>,
1249{
1250    // let err = err.
1251    let range = extract_unicode_error_range(&err, vm)?;
1252    let s = PyStrRef::try_from_object(vm, err.get_attr("object", vm)?)?;
1253    let s_encoding = PyUtf8StrRef::try_from_object(vm, err.get_attr("encoding", vm)?)?;
1254    let mut ctx = PyEncodeContext {
1255        vm,
1256        encoding: s_encoding.as_str(),
1257        data: &s,
1258        pos: StrSize::default(),
1259        exception: OnceCell::from(err.downcast().unwrap()),
1260    };
1261    let mut iter = s.as_wtf8().code_point_indices();
1262    let start = StrSize {
1263        chars: range.start,
1264        bytes: iter.nth(range.start).unwrap().0,
1265    };
1266    let end = StrSize {
1267        chars: range.end,
1268        bytes: if let Some(n) = range.len().checked_sub(1) {
1269            iter.nth(n).map_or_else(|| s.byte_len(), |(i, _)| i)
1270        } else {
1271            start.bytes
1272        },
1273    };
1274    let (replace, restart) = handler.handle_encode_error(&mut ctx, start..end, None)?;
1275    let replace = match replace {
1276        EncodeReplace::Str(s) => s.into(),
1277        EncodeReplace::Bytes(b) => b.into(),
1278    };
1279    Ok((replace, restart.chars))
1280}
1281
1282fn call_native_decode_error<E>(
1283    handler: E,
1284    err: PyObjectRef,
1285    vm: &VirtualMachine,
1286) -> PyResult<(PyObjectRef, usize)>
1287where
1288    for<'a> E: DecodeErrorHandler<PyDecodeContext<'a>>,
1289{
1290    let range = extract_unicode_error_range(&err, vm)?;
1291    let s = ArgBytesLike::try_from_object(vm, err.get_attr("object", vm)?)?;
1292    let s_encoding = PyUtf8StrRef::try_from_object(vm, err.get_attr("encoding", vm)?)?;
1293    let mut ctx = PyDecodeContext {
1294        vm,
1295        encoding: s_encoding.as_str(),
1296        data: PyDecodeData::Original(s.borrow_buf()),
1297        orig_bytes: s.as_object().downcast_ref(),
1298        pos: 0,
1299        exception: OnceCell::from(err.downcast().unwrap()),
1300    };
1301    let (replace, restart) = handler.handle_decode_error(&mut ctx, range, None)?;
1302    Ok((replace.into(), restart))
1303}
1304
1305// this is a hack, for now
1306fn call_native_translate_error<E>(
1307    handler: E,
1308    err: PyObjectRef,
1309    vm: &VirtualMachine,
1310) -> PyResult<(PyObjectRef, usize)>
1311where
1312    for<'a> E: EncodeErrorHandler<PyEncodeContext<'a>>,
1313{
1314    // let err = err.
1315    let range = extract_unicode_error_range(&err, vm)?;
1316    let s = PyStrRef::try_from_object(vm, err.get_attr("object", vm)?)?;
1317    let mut ctx = PyEncodeContext {
1318        vm,
1319        encoding: "",
1320        data: &s,
1321        pos: StrSize::default(),
1322        exception: OnceCell::from(err.downcast().unwrap()),
1323    };
1324    let mut iter = s.as_wtf8().code_point_indices();
1325    let start = StrSize {
1326        chars: range.start,
1327        bytes: iter.nth(range.start).unwrap().0,
1328    };
1329    let end = StrSize {
1330        chars: range.end,
1331        bytes: if let Some(n) = range.len().checked_sub(1) {
1332            iter.nth(n).map_or_else(|| s.byte_len(), |(i, _)| i)
1333        } else {
1334            start.bytes
1335        },
1336    };
1337    let (replace, restart) = handler.handle_encode_error(&mut ctx, start..end, None)?;
1338    let replace = match replace {
1339        EncodeReplace::Str(s) => s.into(),
1340        EncodeReplace::Bytes(b) => b.into(),
1341    };
1342    Ok((replace, restart.chars))
1343}
1344
1345// TODO: exceptions with custom payloads
1346fn extract_unicode_error_range(err: &PyObject, vm: &VirtualMachine) -> PyResult<Range<usize>> {
1347    let start = err.get_attr("start", vm)?;
1348    let start = start.try_into_value(vm)?;
1349
1350    let end = err.get_attr("end", vm)?;
1351    let end = end.try_into_value(vm)?;
1352
1353    Ok(Range { start, end })
1354}
1355
1356fn update_unicode_error_attrs(
1357    err: &PyObject,
1358    start: usize,
1359    end: usize,
1360    reason: Option<&str>,
1361    vm: &VirtualMachine,
1362) -> PyResult<()> {
1363    err.set_attr("start", start.to_pyobject(vm), vm)?;
1364    err.set_attr("end", end.to_pyobject(vm), vm)?;
1365    if let Some(reason) = reason {
1366        err.set_attr("reason", reason.to_pyobject(vm), vm)?;
1367    }
1368    Ok(())
1369}
1370
1371#[inline]
1372fn is_encode_err(err: &PyObject, vm: &VirtualMachine) -> bool {
1373    err.fast_isinstance(vm.ctx.exceptions.unicode_encode_error)
1374}
1375
1376#[inline]
1377fn is_decode_err(err: &PyObject, vm: &VirtualMachine) -> bool {
1378    err.fast_isinstance(vm.ctx.exceptions.unicode_decode_error)
1379}
1380
1381#[inline]
1382fn is_translate_err(err: &PyObject, vm: &VirtualMachine) -> bool {
1383    err.fast_isinstance(vm.ctx.exceptions.unicode_translate_error)
1384}
1385
1386fn bad_err_type(err: &PyObject, vm: &VirtualMachine) -> PyBaseExceptionRef {
1387    vm.new_type_error(format!(
1388        "don't know how to handle {} in error callback",
1389        err.class().name()
1390    ))
1391}
1392
1393fn strict_errors(err: PyObjectRef, vm: &VirtualMachine) -> PyResult {
1394    Err(err
1395        .downcast()
1396        .unwrap_or_else(|_| vm.new_type_error("codec must pass exception instance")))
1397}
1398
1399fn ignore_errors(err: PyObjectRef, vm: &VirtualMachine) -> PyResult<(PyObjectRef, usize)> {
1400    if is_encode_err(&err, vm) || is_decode_err(&err, vm) || is_translate_err(&err, vm) {
1401        let range = extract_unicode_error_range(&err, vm)?;
1402        Ok((vm.ctx.new_str(ascii!("")).into(), range.end))
1403    } else {
1404        Err(bad_err_type(&err, vm))
1405    }
1406}
1407
1408fn replace_errors(err: PyObjectRef, vm: &VirtualMachine) -> PyResult<(PyObjectRef, usize)> {
1409    if is_encode_err(&err, vm) {
1410        call_native_encode_error(errors::Replace, err, vm)
1411    } else if is_decode_err(&err, vm) {
1412        call_native_decode_error(errors::Replace, err, vm)
1413    } else if is_translate_err(&err, vm) {
1414        // char::REPLACEMENT_CHARACTER as a str
1415        let replacement_char = "\u{FFFD}";
1416        let range = extract_unicode_error_range(&err, vm)?;
1417        let replace = replacement_char.repeat(range.end - range.start);
1418        Ok((replace.to_pyobject(vm), range.end))
1419    } else {
1420        Err(bad_err_type(&err, vm))
1421    }
1422}
1423
1424fn xmlcharrefreplace_errors(
1425    err: PyObjectRef,
1426    vm: &VirtualMachine,
1427) -> PyResult<(PyObjectRef, usize)> {
1428    if is_encode_err(&err, vm) {
1429        call_native_encode_error(errors::XmlCharRefReplace, err, vm)
1430    } else {
1431        Err(bad_err_type(&err, vm))
1432    }
1433}
1434
1435fn backslashreplace_errors(
1436    err: PyObjectRef,
1437    vm: &VirtualMachine,
1438) -> PyResult<(PyObjectRef, usize)> {
1439    if is_decode_err(&err, vm) {
1440        call_native_decode_error(errors::BackslashReplace, err, vm)
1441    } else if is_encode_err(&err, vm) {
1442        call_native_encode_error(errors::BackslashReplace, err, vm)
1443    } else if is_translate_err(&err, vm) {
1444        call_native_translate_error(errors::BackslashReplace, err, vm)
1445    } else {
1446        Err(bad_err_type(&err, vm))
1447    }
1448}
1449
1450fn namereplace_errors(err: PyObjectRef, vm: &VirtualMachine) -> PyResult<(PyObjectRef, usize)> {
1451    if is_encode_err(&err, vm) {
1452        call_native_encode_error(errors::NameReplace, err, vm)
1453    } else {
1454        Err(bad_err_type(&err, vm))
1455    }
1456}
1457
1458fn surrogatepass_errors(err: PyObjectRef, vm: &VirtualMachine) -> PyResult<(PyObjectRef, usize)> {
1459    if is_encode_err(&err, vm) {
1460        call_native_encode_error(SurrogatePass, err, vm)
1461    } else if is_decode_err(&err, vm) {
1462        call_native_decode_error(SurrogatePass, err, vm)
1463    } else {
1464        Err(bad_err_type(&err, vm))
1465    }
1466}
1467
1468fn surrogateescape_errors(err: PyObjectRef, vm: &VirtualMachine) -> PyResult<(PyObjectRef, usize)> {
1469    if is_encode_err(&err, vm) {
1470        call_native_encode_error(errors::SurrogateEscape, err, vm)
1471    } else if is_decode_err(&err, vm) {
1472        call_native_decode_error(errors::SurrogateEscape, err, vm)
1473    } else {
1474        Err(bad_err_type(&err, vm))
1475    }
1476}