Skip to main content

rustpython_vm/
buffer.rs

1use crate::{
2    AsObject, Py, PyObject, PyObjectRef, PyResult, TryFromObject, VirtualMachine,
3    builtins::{PyBaseExceptionRef, PyBytesRef, PyComplex, PyTuple, PyTupleRef, PyType, PyTypeRef},
4    common::{lock::PyRwLock, rc::PyRc, static_cell, str::wchar_t},
5    convert::ToPyObject,
6    exceptions,
7    function::{ArgBytesLike, ArgIntoBool, ArgIntoComplex, ArgIntoFloat},
8};
9
10use rustpython_common::wtf8::Wtf8Buf;
11
12use core::{fmt, iter::Peekable, mem};
13use half::f16;
14use itertools::Itertools;
15use malachite_bigint::BigInt;
16use num_complex::Complex64;
17use num_traits::{PrimInt, ToPrimitive};
18use std::{collections::HashMap, os::raw};
19
20type PackFunc = fn(&VirtualMachine, FormatType, PyObjectRef, &mut [u8]) -> Result<(), PackError>;
21type UnpackFunc = fn(&VirtualMachine, &[u8]) -> PyObjectRef;
22
23/// Why a value could not be packed.
24///
25/// `struct` reports both as `struct.error`, so the kind travels beside the
26/// exception rather than in it; `memoryview`, which reports them as TypeError
27/// and ValueError, is what needs to tell them apart.
28#[derive(Clone, Copy, Debug, Eq, PartialEq)]
29pub enum PackErrorKind {
30    /// The value was not the kind of thing the format takes.
31    Type,
32    /// The value was the right kind, and the format has no room for it.
33    Value,
34    /// The value's own code raised, and that error is the answer as it is.
35    Raised,
36}
37
38pub struct PackError {
39    pub kind: PackErrorKind,
40    pub exception: PyBaseExceptionRef,
41}
42
43impl PackError {
44    fn new<T: Into<Wtf8Buf>>(kind: PackErrorKind, vm: &VirtualMachine, msg: T) -> Self {
45        Self {
46            kind,
47            exception: new_struct_error(vm, msg),
48        }
49    }
50
51    /// An error raised by something other than the packing itself, such as a
52    /// conversion running the value's own code.
53    fn from_exception(exception: PyBaseExceptionRef, vm: &VirtualMachine) -> Self {
54        let kind = if exception.fast_isinstance(vm.ctx.exceptions.type_error) {
55            PackErrorKind::Type
56        } else if exception.fast_isinstance(vm.ctx.exceptions.overflow_error)
57            || exception.fast_isinstance(vm.ctx.exceptions.value_error)
58        {
59            PackErrorKind::Value
60        } else {
61            PackErrorKind::Raised
62        };
63        Self { kind, exception }
64    }
65
66    /// An error that is the answer exactly as it was raised. `pack_single()`
67    /// leaves `'?'` to `PyObject_IsTrue()` this way, with no message of its own.
68    fn raised(exception: PyBaseExceptionRef) -> Self {
69        Self {
70            kind: PackErrorKind::Raised,
71            exception,
72        }
73    }
74}
75
76static OVERFLOW_MSG: &str = "total struct size too long"; // not a const to reduce code size
77
78#[derive(Clone, Copy, Debug, Eq, PartialEq)]
79pub(crate) enum Endianness {
80    Native,
81    Little,
82    Big,
83    Host,
84}
85
86impl Endianness {
87    /// Parse endianness
88    /// See also: https://docs.python.org/3/library/struct.html?highlight=struct#byte-order-size-and-alignment
89    fn parse<I>(chars: &mut Peekable<I>) -> Self
90    where
91        I: Sized + Iterator<Item = u8>,
92    {
93        let e = match chars.peek() {
94            Some(b'@') => Self::Native,
95            Some(b'=') => Self::Host,
96            Some(b'<') => Self::Little,
97            Some(b'>' | b'!') => Self::Big,
98            _ => return Self::Native,
99        };
100
101        // SAFETY:
102        // We just ensured with `chars.peek()` that this is safe
103        unsafe {
104            let _ = chars.next().unwrap_unchecked();
105        }
106        e
107    }
108}
109
110trait ByteOrder {
111    fn convert<I: PrimInt>(i: I) -> I;
112}
113
114enum BigEndian {}
115
116impl ByteOrder for BigEndian {
117    fn convert<I: PrimInt>(i: I) -> I {
118        i.to_be()
119    }
120}
121
122enum LittleEndian {}
123
124impl ByteOrder for LittleEndian {
125    fn convert<I: PrimInt>(i: I) -> I {
126        i.to_le()
127    }
128}
129
130type NativeEndian = cfg_select! {
131    target_endian = "big" => BigEndian,
132    target_endian = "little" => LittleEndian,
133};
134
135#[derive(Copy, Clone, num_enum::TryFromPrimitive, Eq, PartialEq)]
136#[repr(u8)]
137pub(crate) enum FormatType {
138    Pad = b'x',
139    SByte = b'b',
140    UByte = b'B',
141    Char = b'c',
142    WideChar = b'u',
143    Ucs4Char = b'w',
144    Str = b's',
145    Pascal = b'p',
146    Short = b'h',
147    UShort = b'H',
148    Int = b'i',
149    UInt = b'I',
150    Long = b'l',
151    ULong = b'L',
152    SSizeT = b'n',
153    SizeT = b'N',
154    LongLong = b'q',
155    ULongLong = b'Q',
156    Bool = b'?',
157    Half = b'e',
158    Float = b'f',
159    Double = b'd',
160    LongDouble = b'g',
161    FloatComplex = b'F',
162    DoubleComplex = b'D',
163    VoidP = b'P',
164    PyObject = b'O',
165}
166
167impl fmt::Debug for FormatType {
168    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
169        fmt::Debug::fmt(&(*self as u8 as char), f)
170    }
171}
172
173impl FormatType {
174    fn info(self, e: Endianness) -> &'static FormatInfo {
175        use mem::{align_of, size_of};
176
177        macro_rules! native_info {
178            ($t:ty) => {{
179                &FormatInfo {
180                    size: size_of::<$t>(),
181                    align: align_of::<$t>(),
182                    pack: Some(<$t as Packable>::pack::<NativeEndian>),
183                    unpack: Some(<$t as Packable>::unpack::<NativeEndian>),
184                }
185            }};
186        }
187
188        macro_rules! nonnative_info {
189            ($t:ty, $end:ty) => {{
190                &FormatInfo {
191                    size: size_of::<$t>(),
192                    align: 0,
193                    pack: Some(<$t as Packable>::pack::<$end>),
194                    unpack: Some(<$t as Packable>::unpack::<$end>),
195                }
196            }};
197        }
198
199        macro_rules! match_nonnative {
200            ($zelf:expr, $end:ty) => {{
201                match $zelf {
202                    Self::Pad | Self::Str | Self::Pascal => &FormatInfo {
203                        size: size_of::<u8>(),
204                        align: 0,
205                        pack: None,
206                        unpack: None,
207                    },
208                    Self::SByte => nonnative_info!(i8, $end),
209                    Self::UByte => nonnative_info!(u8, $end),
210                    Self::Char => &FormatInfo {
211                        size: size_of::<u8>(),
212                        align: 0,
213                        pack: Some(pack_char),
214                        unpack: Some(unpack_char),
215                    },
216                    Self::Short => nonnative_info!(i16, $end),
217                    Self::UShort => nonnative_info!(u16, $end),
218                    Self::Int | Self::Long => nonnative_info!(i32, $end),
219                    Self::UInt | Self::ULong => nonnative_info!(u32, $end),
220                    Self::LongLong => nonnative_info!(i64, $end),
221                    Self::ULongLong => nonnative_info!(u64, $end),
222                    Self::Bool => nonnative_info!(bool, $end),
223                    Self::Half => nonnative_info!(f16, $end),
224                    Self::Float => nonnative_info!(f32, $end),
225                    Self::Double => nonnative_info!(f64, $end),
226                    Self::LongDouble => nonnative_info!(f64, $end), // long double same as double
227                    Self::FloatComplex => nonnative_info!(PackFloatComplex, $end),
228                    Self::DoubleComplex => nonnative_info!(PackDoubleComplex, $end),
229                    Self::PyObject => nonnative_info!(usize, $end), // pointer size
230                    _ => unreachable!(),                            // size_t or void*
231                }
232            }};
233        }
234
235        match e {
236            Endianness::Native => match self {
237                Self::Pad | Self::Str | Self::Pascal => &FormatInfo {
238                    size: size_of::<raw::c_char>(),
239                    align: 0,
240                    pack: None,
241                    unpack: None,
242                },
243                Self::SByte => native_info!(raw::c_schar),
244                Self::UByte => native_info!(raw::c_uchar),
245                Self::Char => &FormatInfo {
246                    size: size_of::<raw::c_char>(),
247                    align: 0,
248                    pack: Some(pack_char),
249                    unpack: Some(unpack_char),
250                },
251                Self::WideChar => native_info!(wchar_t),
252                Self::Ucs4Char => native_info!(u32),
253                Self::Short => native_info!(raw::c_short),
254                Self::UShort => native_info!(raw::c_ushort),
255                Self::Int => native_info!(raw::c_int),
256                Self::UInt => native_info!(raw::c_uint),
257                Self::Long => native_info!(raw::c_long),
258                Self::ULong => native_info!(raw::c_ulong),
259                Self::SSizeT => native_info!(isize), // ssize_t == isize
260                Self::SizeT => native_info!(usize),  //  size_t == usize
261                Self::LongLong => native_info!(raw::c_longlong),
262                Self::ULongLong => native_info!(raw::c_ulonglong),
263                Self::Bool => native_info!(bool),
264                Self::Half => native_info!(f16),
265                Self::Float => native_info!(raw::c_float),
266                Self::Double => native_info!(raw::c_double),
267                Self::LongDouble => native_info!(raw::c_double), // long double same as double for now
268                Self::FloatComplex => native_info!(PackFloatComplex),
269                Self::DoubleComplex => native_info!(PackDoubleComplex),
270                Self::VoidP => native_info!(*mut raw::c_void),
271                Self::PyObject => native_info!(*mut raw::c_void), // pointer to PyObject
272            },
273            Endianness::Big => match_nonnative!(self, BigEndian),
274            Endianness::Little => match_nonnative!(self, LittleEndian),
275            Endianness::Host => match_nonnative!(self, NativeEndian),
276        }
277    }
278}
279
280#[derive(Debug, Clone)]
281pub(crate) struct FormatCode {
282    pub repeat: usize,
283    pub code: FormatType,
284    pub info: &'static FormatInfo,
285    pub pre_padding: usize,
286}
287
288impl FormatCode {
289    pub(crate) const fn arg_count(&self) -> usize {
290        match self.code {
291            FormatType::Pad => 0,
292            FormatType::Str | FormatType::Pascal => 1,
293            _ => self.repeat,
294        }
295    }
296
297    pub(crate) fn parse<I>(
298        chars: &mut Peekable<I>,
299        endianness: Endianness,
300    ) -> Result<(Vec<Self>, usize, usize), String>
301    where
302        I: Sized + Iterator<Item = u8>,
303    {
304        let mut offset = 0isize;
305        let mut arg_count = 0usize;
306        let mut codes = vec![];
307        while chars.peek().is_some() {
308            // Skip whitespace before repeat count or format char
309            while let Some(b' ' | b'\t' | b'\n' | b'\r') = chars.peek() {
310                chars.next();
311            }
312
313            // determine repeat operator:
314            let repeat = match chars.peek() {
315                Some(b'0'..=b'9') => {
316                    let mut repeat = 0isize;
317                    while let Some(b'0'..=b'9') = chars.peek() {
318                        if let Some(c) = chars.next() {
319                            let current_digit = c - b'0';
320                            repeat = repeat
321                                .checked_mul(10)
322                                .and_then(|r| r.checked_add(current_digit as _))
323                                .ok_or_else(|| OVERFLOW_MSG.to_owned())?;
324                        }
325                    }
326                    repeat
327                }
328                _ => 1,
329            };
330
331            // determine format char:
332            let c = match chars.next() {
333                Some(c) => c,
334                None => {
335                    // If we have a repeat count but only whitespace follows, error
336                    if repeat != 1 {
337                        return Err("repeat count given without format specifier".to_owned());
338                    }
339                    // Otherwise, we're done parsing
340                    break;
341                }
342            };
343
344            // Check for embedded null character
345            if c == 0 {
346                return Err(exceptions::NulError.to_string());
347            }
348
349            // PEP3118: Handle extended format specifiers
350            // T{...} - struct, X{} - function pointer, (...) - array shape, :name: - field name
351            if c == b'T' || c == b'X' {
352                // Skip struct/function pointer: consume until matching '}'
353                if chars.peek() == Some(&b'{') {
354                    chars.next(); // consume '{'
355                    let mut depth = 1;
356                    while depth > 0 {
357                        match chars.next() {
358                            Some(b'{') => depth += 1,
359                            Some(b'}') => depth -= 1,
360                            None => return Err("unmatched '{' in format".to_owned()),
361                            _ => {}
362                        }
363                    }
364                    continue;
365                }
366            }
367
368            if c == b'(' {
369                // Skip array shape: consume until matching ')'
370                let mut depth = 1;
371                while depth > 0 {
372                    match chars.next() {
373                        Some(b'(') => depth += 1,
374                        Some(b')') => depth -= 1,
375                        None => return Err("unmatched '(' in format".to_owned()),
376                        _ => {}
377                    }
378                }
379                continue;
380            }
381
382            if c == b':' {
383                // Skip field name: consume until next ':'
384                loop {
385                    match chars.next() {
386                        Some(b':') => break,
387                        None => return Err("unmatched ':' in format".to_owned()),
388                        _ => {}
389                    }
390                }
391                continue;
392            }
393
394            if c == b'{'
395                || c == b'}'
396                || c == b'&'
397                || c == b'<'
398                || c == b'>'
399                || c == b'@'
400                || c == b'='
401                || c == b'!'
402            {
403                // Skip standalone braces (pointer targets, etc.), pointer prefix, and nested endianness markers
404                continue;
405            }
406
407            let code = FormatType::try_from(c)
408                .ok()
409                .filter(|c| match c {
410                    FormatType::SSizeT
411                    | FormatType::SizeT
412                    | FormatType::VoidP
413                    | FormatType::Ucs4Char => endianness == Endianness::Native,
414                    _ => true,
415                })
416                .ok_or_else(|| "bad char in struct format".to_owned())?;
417
418            let info = code.info(endianness);
419
420            let padding = compensate_alignment(offset as usize, info.align)
421                .ok_or_else(|| OVERFLOW_MSG.to_owned())?;
422            offset = padding
423                .to_isize()
424                .and_then(|extra| offset.checked_add(extra))
425                .ok_or_else(|| OVERFLOW_MSG.to_owned())?;
426
427            let code = Self {
428                repeat: repeat as usize,
429                code,
430                info,
431                pre_padding: padding,
432            };
433            arg_count += code.arg_count();
434            codes.push(code);
435
436            offset = (info.size as isize)
437                .checked_mul(repeat)
438                .and_then(|item_size| offset.checked_add(item_size))
439                .ok_or_else(|| OVERFLOW_MSG.to_owned())?;
440        }
441
442        Ok((codes, offset as usize, arg_count))
443    }
444}
445
446const fn compensate_alignment(offset: usize, align: usize) -> Option<usize> {
447    if align != 0 && offset != 0 {
448        // a % b == a & (b-1) if b is a power of 2
449        (align - 1).checked_sub((offset - 1) & (align - 1))
450    } else {
451        // alignment is already all good
452        Some(0)
453    }
454}
455
456pub(crate) struct FormatInfo {
457    pub size: usize,
458    pub align: usize,
459    pub pack: Option<PackFunc>,
460    pub unpack: Option<UnpackFunc>,
461}
462
463impl fmt::Debug for FormatInfo {
464    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
465        f.debug_struct("FormatInfo")
466            .field("size", &self.size)
467            .field("align", &self.align)
468            .finish()
469    }
470}
471
472#[derive(Debug, Clone)]
473pub struct FormatSpec {
474    #[allow(dead_code)]
475    pub(crate) endianness: Endianness,
476    pub(crate) codes: Vec<FormatCode>,
477    pub size: usize,
478    pub arg_count: usize,
479}
480
481/// Bounded, interpreter-local cache for the module-level `struct` functions.
482/// Keys and specifications contain no Python objects or user callbacks.
483#[derive(Default)]
484pub struct FormatSpecCache {
485    entries: PyRwLock<HashMap<Box<[u8]>, PyRc<FormatSpec>>>,
486}
487
488impl FormatSpecCache {
489    pub fn get_or_parse(&self, format: &[u8], vm: &VirtualMachine) -> PyResult<PyRc<FormatSpec>> {
490        if let Some(spec) = self.entries.read().get(format) {
491            return Ok(spec.clone());
492        }
493
494        // Parsing can raise; neither errors nor Python conversions run under the cache lock.
495        let spec = PyRc::new(FormatSpec::parse(format, vm)?);
496        let mut entries = self.entries.write();
497        if let Some(spec) = entries.get(format) {
498            return Ok(spec.clone());
499        }
500        if entries.len() >= 100 {
501            entries.clear();
502        }
503        entries.insert(format.into(), spec.clone());
504        Ok(spec)
505    }
506
507    pub fn clear(&self) {
508        *self.entries.write() = HashMap::new();
509    }
510}
511
512impl FormatSpec {
513    pub fn parse(fmt: &[u8], vm: &VirtualMachine) -> PyResult<Self> {
514        let mut chars = fmt.iter().copied().peekable();
515
516        // First determine "@", "<", ">","!" or "="
517        let endianness = Endianness::parse(&mut chars);
518
519        // Now, analyze struct string further:
520        let (codes, size, arg_count) =
521            FormatCode::parse(&mut chars, endianness).map_err(|err| new_struct_error(vm, err))?;
522
523        Ok(Self {
524            endianness,
525            codes,
526            size,
527            arg_count,
528        })
529    }
530
531    pub fn pack(&self, args: Vec<PyObjectRef>, vm: &VirtualMachine) -> PyResult<Vec<u8>> {
532        self.try_pack(args, vm).map_err(|e| e.exception)
533    }
534
535    /// [`Self::pack`], keeping why a value could not be packed.
536    pub fn try_pack(
537        &self,
538        args: Vec<PyObjectRef>,
539        vm: &VirtualMachine,
540    ) -> Result<Vec<u8>, PackError> {
541        // Create data vector:
542        let mut data = vm
543            .new_zeroed_bytes(self.size)
544            .map_err(|e| PackError::from_exception(e, vm))?;
545
546        self.try_pack_into(&mut data, args, vm)?;
547
548        Ok(data)
549    }
550
551    pub fn pack_into(
552        &self,
553        buffer: &mut [u8],
554        args: Vec<PyObjectRef>,
555        vm: &VirtualMachine,
556    ) -> PyResult<()> {
557        self.try_pack_into(buffer, args, vm)
558            .map_err(|e| e.exception)
559    }
560
561    /// [`Self::pack_into`], keeping why a value could not be packed.
562    pub fn try_pack_into(
563        &self,
564        mut buffer: &mut [u8],
565        args: Vec<PyObjectRef>,
566        vm: &VirtualMachine,
567    ) -> Result<(), PackError> {
568        if self.arg_count != args.len() {
569            return Err(PackError::new(
570                PackErrorKind::Type,
571                vm,
572                format!(
573                    "pack expected {} items for packing (got {})",
574                    self.codes.len(),
575                    args.len()
576                ),
577            ));
578        }
579
580        let mut args = args.into_iter();
581        // Loop over all opcodes:
582        for code in &self.codes {
583            buffer = &mut buffer[code.pre_padding..];
584            debug!("code: {code:?}");
585            match code.code {
586                FormatType::Str => {
587                    let (buf, rest) = buffer.split_at_mut(code.repeat);
588                    pack_string(vm, args.next().unwrap(), buf)
589                        .map_err(|e| PackError::from_exception(e, vm))?;
590                    buffer = rest;
591                }
592                FormatType::Pascal => {
593                    let (buf, rest) = buffer.split_at_mut(code.repeat);
594                    pack_pascal(vm, args.next().unwrap(), buf)
595                        .map_err(|e| PackError::from_exception(e, vm))?;
596                    buffer = rest;
597                }
598                FormatType::Pad => {
599                    let (pad_buf, rest) = buffer.split_at_mut(code.repeat);
600                    for el in pad_buf {
601                        *el = 0
602                    }
603                    buffer = rest;
604                }
605                _ => {
606                    let pack = code.info.pack.unwrap();
607                    for arg in args.by_ref().take(code.repeat) {
608                        let (item_buf, rest) = buffer.split_at_mut(code.info.size);
609                        pack(vm, code.code, arg, item_buf)?;
610                        buffer = rest;
611                    }
612                }
613            }
614        }
615
616        Ok(())
617    }
618
619    pub fn unpack(&self, mut data: &[u8], vm: &VirtualMachine) -> PyResult<PyTupleRef> {
620        if self.size != data.len() {
621            return Err(new_struct_error(
622                vm,
623                format!("unpack requires a buffer of {} bytes", self.size),
624            ));
625        }
626
627        let mut items = Vec::with_capacity(self.arg_count);
628        for code in &self.codes {
629            data = &data[code.pre_padding..];
630            debug!("unpack code: {code:?}");
631            match code.code {
632                FormatType::Pad => {
633                    data = &data[code.repeat..];
634                }
635                FormatType::Str => {
636                    let (str_data, rest) = data.split_at(code.repeat);
637                    // string is just stored inline
638                    items.push(vm.ctx.new_bytes(str_data.to_vec()).into());
639                    data = rest;
640                }
641                FormatType::Pascal => {
642                    let (str_data, rest) = data.split_at(code.repeat);
643                    items.push(unpack_pascal(vm, str_data));
644                    data = rest;
645                }
646                _ => {
647                    let unpack = code.info.unpack.unwrap();
648                    for _ in 0..code.repeat {
649                        let (item_data, rest) = data.split_at(code.info.size);
650                        items.push(unpack(vm, item_data));
651                        data = rest;
652                    }
653                }
654            };
655        }
656
657        Ok(PyTuple::new_ref(items, &vm.ctx))
658    }
659
660    #[inline]
661    #[must_use]
662    pub const fn size(&self) -> usize {
663        self.size
664    }
665
666    #[must_use]
667    pub fn codes_sizeof(&self) -> usize {
668        core::mem::size_of::<FormatCode>() * (self.codes.len() + 1)
669    }
670}
671
672trait Packable {
673    fn pack<E: ByteOrder>(
674        vm: &VirtualMachine,
675        code: FormatType,
676        arg: PyObjectRef,
677        data: &mut [u8],
678    ) -> Result<(), PackError>;
679    fn unpack<E: ByteOrder>(vm: &VirtualMachine, data: &[u8]) -> PyObjectRef;
680}
681
682trait PackInt: PrimInt {
683    fn pack_int<E: ByteOrder>(self, data: &mut [u8]);
684    fn unpack_int<E: ByteOrder>(data: &[u8]) -> Self;
685}
686
687macro_rules! make_pack_prim_int {
688    ($T:ty) => {
689        impl PackInt for $T {
690            fn pack_int<E: ByteOrder>(self, data: &mut [u8]) {
691                let i = E::convert(self);
692                data.copy_from_slice(&i.to_ne_bytes());
693            }
694            #[inline]
695            fn unpack_int<E: ByteOrder>(data: &[u8]) -> Self {
696                let mut x = [0; core::mem::size_of::<$T>()];
697                x.copy_from_slice(data);
698                E::convert(<$T>::from_ne_bytes(x))
699            }
700        }
701
702        impl Packable for $T {
703            fn pack<E: ByteOrder>(
704                vm: &VirtualMachine,
705                code: FormatType,
706                arg: PyObjectRef,
707                data: &mut [u8],
708            ) -> Result<(), PackError> {
709                let i: $T = get_int_or_index(vm, code, &arg)?;
710                i.pack_int::<E>(data);
711                Ok(())
712            }
713
714            fn unpack<E: ByteOrder>(vm: &VirtualMachine, rdr: &[u8]) -> PyObjectRef {
715                let i = <$T>::unpack_int::<E>(rdr);
716                vm.ctx.new_int(i).into()
717            }
718        }
719    };
720}
721
722fn get_int_or_index<T>(
723    vm: &VirtualMachine,
724    code: FormatType,
725    arg: &PyObject,
726) -> Result<T, PackError>
727where
728    T: PrimInt + fmt::Display + for<'a> TryFrom<&'a BigInt>,
729{
730    let index = match arg.try_index_opt(vm) {
731        None => {
732            return Err(PackError::new(
733                PackErrorKind::Type,
734                vm,
735                "required argument is not an integer",
736            ));
737        }
738        Some(Err(e)) => return Err(PackError::from_exception(e, vm)),
739        Some(Ok(index)) => index,
740    };
741    index.try_to_primitive(vm).map_err(|_| {
742        // A pointer is converted rather than checked against the range of a
743        // named format, so what it reports is the conversion failing.
744        let msg = if code == FormatType::VoidP {
745            "int too large to convert".to_owned()
746        } else {
747            format!(
748                "'{}' format requires {} <= number <= {}",
749                code as u8 as char,
750                T::min_value(),
751                T::max_value()
752            )
753        };
754        PackError::new(PackErrorKind::Value, vm, msg)
755    })
756}
757
758make_pack_prim_int!(i8);
759make_pack_prim_int!(u8);
760make_pack_prim_int!(i16);
761make_pack_prim_int!(u16);
762make_pack_prim_int!(i32);
763make_pack_prim_int!(u32);
764make_pack_prim_int!(i64);
765make_pack_prim_int!(u64);
766make_pack_prim_int!(usize);
767make_pack_prim_int!(isize);
768
769macro_rules! make_pack_float {
770    ($T:ty, $fmt:literal) => {
771        impl Packable for $T {
772            fn pack<E: ByteOrder>(
773                vm: &VirtualMachine,
774                _code: FormatType,
775                arg: PyObjectRef,
776                data: &mut [u8],
777            ) -> Result<(), PackError> {
778                let f_64 = ArgIntoFloat::try_from_object(vm, arg)
779                    .map_err(|e| PackError::from_exception(e, vm))?
780                    .into_float();
781                let f = f_64 as $T;
782                if f.is_infinite() != f_64.is_infinite() {
783                    return Err(PackError {
784                        kind: PackErrorKind::Value,
785                        exception: vm.new_overflow_error(concat!(
786                            "float too large to pack with ",
787                            $fmt,
788                            " format"
789                        )),
790                    });
791                }
792                f.to_bits().pack_int::<E>(data);
793                Ok(())
794            }
795
796            fn unpack<E: ByteOrder>(vm: &VirtualMachine, rdr: &[u8]) -> PyObjectRef {
797                let i = PackInt::unpack_int::<E>(rdr);
798                <$T>::from_bits(i).to_pyobject(vm)
799            }
800        }
801    };
802}
803
804make_pack_float!(f32, "f");
805make_pack_float!(f64, "d");
806
807#[repr(C)]
808struct PackFloatComplex(f32, f32);
809
810#[repr(C)]
811struct PackDoubleComplex(f64, f64);
812
813macro_rules! make_pack_complex {
814    ($T:ty, $Elem:ty, $Bits:ty, $fmt:literal) => {
815        impl Packable for $T {
816            fn pack<E: ByteOrder>(
817                vm: &VirtualMachine,
818                _code: FormatType,
819                arg: PyObjectRef,
820                data: &mut [u8],
821            ) -> Result<(), PackError> {
822                let c = if let Some(value) = arg.downcast_ref::<PyComplex>() {
823                    value.as_complex()
824                } else {
825                    ArgIntoComplex::try_from_object(vm, arg)
826                        .map_err(|_| {
827                            PackError::new(
828                                PackErrorKind::Type,
829                                vm,
830                                "required argument is not a complex",
831                            )
832                        })?
833                        .into_complex()
834                };
835                for (component, bytes) in [c.re, c.im]
836                    .into_iter()
837                    .zip(data.chunks_exact_mut(size_of::<$Elem>()))
838                {
839                    let narrowed = component as $Elem;
840                    // CPython uses native casts for matching-endian complex formats.
841                    if E::convert(1u16) != 1 && narrowed.is_infinite() != component.is_infinite() {
842                        return Err(PackError {
843                            kind: PackErrorKind::Value,
844                            exception: vm.new_overflow_error(concat!(
845                                "float too large to pack with ",
846                                $fmt,
847                                " format"
848                            )),
849                        });
850                    }
851                    narrowed.to_bits().pack_int::<E>(bytes);
852                }
853                Ok(())
854            }
855
856            fn unpack<E: ByteOrder>(vm: &VirtualMachine, rdr: &[u8]) -> PyObjectRef {
857                let half = size_of::<$Elem>();
858                let re = <$Elem>::from_bits(<$Bits>::unpack_int::<E>(&rdr[..half])) as f64;
859                let im = <$Elem>::from_bits(<$Bits>::unpack_int::<E>(&rdr[half..half * 2])) as f64;
860                vm.ctx.new_complex(Complex64::new(re, im)).into()
861            }
862        }
863    };
864}
865
866make_pack_complex!(PackFloatComplex, f32, u32, "f");
867make_pack_complex!(PackDoubleComplex, f64, u64, "d");
868
869impl Packable for f16 {
870    fn pack<E: ByteOrder>(
871        vm: &VirtualMachine,
872        _code: FormatType,
873        arg: PyObjectRef,
874        data: &mut [u8],
875    ) -> Result<(), PackError> {
876        let f_64 = ArgIntoFloat::try_from_object(vm, arg)
877            .map_err(|e| PackError::from_exception(e, vm))?
878            .into_float();
879        // "from_f64 should be preferred in any non-`const` context" except it gives the wrong result :/
880        let f_16 = Self::from_f64_const(f_64);
881        if f_16.is_infinite() != f_64.is_infinite() {
882            return Err(PackError {
883                kind: PackErrorKind::Value,
884                exception: vm.new_overflow_error("float too large to pack with e format"),
885            });
886        }
887        f_16.to_bits().pack_int::<E>(data);
888        Ok(())
889    }
890
891    fn unpack<E: ByteOrder>(vm: &VirtualMachine, rdr: &[u8]) -> PyObjectRef {
892        let i = PackInt::unpack_int::<E>(rdr);
893        Self::from_bits(i).to_f64().to_pyobject(vm)
894    }
895}
896
897impl Packable for *mut raw::c_void {
898    fn pack<E: ByteOrder>(
899        vm: &VirtualMachine,
900        code: FormatType,
901        arg: PyObjectRef,
902        data: &mut [u8],
903    ) -> Result<(), PackError> {
904        usize::pack::<E>(vm, code, arg, data)
905    }
906
907    fn unpack<E: ByteOrder>(vm: &VirtualMachine, rdr: &[u8]) -> PyObjectRef {
908        usize::unpack::<E>(vm, rdr)
909    }
910}
911
912impl Packable for bool {
913    fn pack<E: ByteOrder>(
914        vm: &VirtualMachine,
915        _code: FormatType,
916        arg: PyObjectRef,
917        data: &mut [u8],
918    ) -> Result<(), PackError> {
919        let v = ArgIntoBool::try_from_object(vm, arg)
920            .map_err(PackError::raised)?
921            .into_bool() as u8;
922        v.pack_int::<E>(data);
923        Ok(())
924    }
925
926    fn unpack<E: ByteOrder>(vm: &VirtualMachine, rdr: &[u8]) -> PyObjectRef {
927        let i = u8::unpack_int::<E>(rdr);
928        vm.ctx.new_bool(i != 0).into()
929    }
930}
931
932fn pack_char(
933    vm: &VirtualMachine,
934    _code: FormatType,
935    arg: PyObjectRef,
936    data: &mut [u8],
937) -> Result<(), PackError> {
938    let v = PyBytesRef::try_from_object(vm, arg).map_err(|e| PackError::from_exception(e, vm))?;
939    let ch = *v.as_bytes().iter().exactly_one().map_err(|_| {
940        PackError::new(
941            PackErrorKind::Value,
942            vm,
943            "char format requires a bytes object of length 1",
944        )
945    })?;
946    data[0] = ch;
947    Ok(())
948}
949
950fn pack_string(vm: &VirtualMachine, arg: PyObjectRef, buf: &mut [u8]) -> PyResult<()> {
951    let b = ArgBytesLike::try_from_object(vm, arg)?;
952    b.with_ref(|data| write_string(buf, data));
953    Ok(())
954}
955
956fn pack_pascal(vm: &VirtualMachine, arg: PyObjectRef, buf: &mut [u8]) -> PyResult<()> {
957    if buf.is_empty() {
958        return Ok(());
959    }
960    let b = ArgBytesLike::try_from_object(vm, arg)?;
961    b.with_ref(|data| {
962        let string_length = core::cmp::min(core::cmp::min(data.len(), 255), buf.len() - 1);
963        buf[0] = string_length as u8;
964        write_string(&mut buf[1..], data);
965    });
966    Ok(())
967}
968
969fn write_string(buf: &mut [u8], data: &[u8]) {
970    let len_from_data = core::cmp::min(data.len(), buf.len());
971    buf[..len_from_data].copy_from_slice(&data[..len_from_data]);
972    for byte in &mut buf[len_from_data..] {
973        *byte = 0
974    }
975}
976
977fn unpack_char(vm: &VirtualMachine, data: &[u8]) -> PyObjectRef {
978    vm.ctx.new_bytes(vec![data[0]]).into()
979}
980
981fn unpack_pascal(vm: &VirtualMachine, data: &[u8]) -> PyObjectRef {
982    let (&len, data) = match data.split_first() {
983        Some(x) => x,
984        None => {
985            // cpython throws an internal SystemError here
986            return vm.ctx.new_bytes(vec![]).into();
987        }
988    };
989    let len = core::cmp::min(len as usize, data.len());
990    vm.ctx.new_bytes(data[..len].to_vec()).into()
991}
992
993// XXX: are those functions expected to be placed here?
994pub fn struct_error_type(vm: &VirtualMachine) -> &'static Py<PyType> {
995    static_cell! {
996        static INSTANCE: PyTypeRef;
997    }
998    INSTANCE.get_or_init(|| vm.ctx.new_exception_type("struct", "error", None))
999}
1000
1001pub fn new_struct_error<T: Into<Wtf8Buf>>(vm: &VirtualMachine, msg: T) -> PyBaseExceptionRef {
1002    // can't just STRUCT_ERROR.get().unwrap() cause this could be called before from buffer
1003    // machinery, independent of whether _struct was ever imported
1004    vm.new_exception_msg(struct_error_type(vm).to_owned(), msg.into())
1005}
1006
1007#[cfg(test)]
1008mod tests {
1009    use super::*;
1010    use crate::Interpreter;
1011
1012    #[test]
1013    fn format_cache_reuses_specs_and_releases_evicted_entries() {
1014        Interpreter::without_stdlib(Default::default()).enter(|vm| {
1015            let cache = &vm.state.struct_format_cache;
1016            let spec = cache.get_or_parse(b"<IH", vm).unwrap();
1017            assert!(PyRc::ptr_eq(
1018                &spec,
1019                &cache.get_or_parse(b"<IH", vm).unwrap()
1020            ));
1021            let weak = PyRc::downgrade(&spec);
1022            drop(spec);
1023
1024            for padding in 0..100 {
1025                let format = format!("{padding}x");
1026                cache.get_or_parse(format.as_bytes(), vm).unwrap();
1027            }
1028            assert!(weak.upgrade().is_none());
1029
1030            let spec = cache.get_or_parse(b"<IH", vm).unwrap();
1031            let weak = PyRc::downgrade(&spec);
1032            drop(spec);
1033            cache.clear();
1034            assert!(weak.upgrade().is_none());
1035        });
1036    }
1037}