Skip to main content

rustpython_vm/types/
structseq.rs

1use crate::common::lock::LazyLock;
2use crate::{
3    AsObject, Py, PyObject, PyObjectRef, PyPayload, PyRef, PyResult, VirtualMachine, atomic_func,
4    builtins::{
5        PyBaseExceptionRef, PyDict, PyStr, PyStrRef, PyTuple, PyTupleRef, PyType, PyTypeRef,
6    },
7    class::{PyClassImpl, StaticType, class_attr_item_doc},
8    function::{
9        Either, FuncArgs, KwArgs, NameChanges, OptionalArg, PyComparisonValue, PyMethodDef,
10        PyMethodFlags,
11    },
12    iter::PyExactSizeIterator,
13    protocol::{PyMappingMethods, PySequenceMethods},
14    sliceable::{SequenceIndex, SliceableSequenceOp},
15    types::PyComparisonOp,
16    vm::Context,
17};
18
19const DEFAULT_STRUCTSEQ_REDUCE: PyMethodDef = PyMethodDef::new_const(
20    "__reduce__",
21    |zelf: PyRef<PyTuple>, vm: &VirtualMachine| -> PyTupleRef {
22        vm.new_tuple((
23            zelf.class().to_owned(),
24            (vm.ctx.new_tuple(zelf.as_slice().to_vec()),),
25        ))
26    },
27    PyMethodFlags::METHOD,
28    crate::function::ItemDoc::static_text("__reduce__($self, /)\n--\n\n"),
29);
30
31/// Text signature `(iterable=(), /)` shared by every struct sequence.
32pub const STRUCT_SEQUENCE_PARAMS: Option<&'static [crate::function::Param]> =
33    Some(&[crate::function::Param {
34        name: "iterable",
35        kind: crate::function::ParamKind::PositionalOnly,
36        default: Some(crate::function::DefaultRepr::Raw("()")),
37    }]);
38
39/// The arguments every struct sequence constructor takes.
40#[derive(FromArgs)]
41pub struct StructSequenceNewArgs {
42    #[pyarg(any)]
43    pub sequence: PyObjectRef,
44    #[pyarg(any, optional, py_default = "{}")]
45    pub dict: OptionalArg<PyObjectRef>,
46}
47
48/// Create a new struct sequence instance from a sequence.
49///
50/// `dict` supplies the hidden fields — the ones past `n_sequence_fields`, named
51/// by `hidden_field_names` in order — that the sequence itself did not cover. It
52/// may not name a field the sequence already supplied, nor one that does not
53/// exist.
54///
55/// The class must have `n_sequence_fields` and `n_fields` attributes set
56/// (done automatically by `PyStructSequence::extend_pyclass`).
57pub fn struct_sequence_new(
58    cls: PyTypeRef,
59    args: StructSequenceNewArgs,
60    hidden_field_names: &[&str],
61    vm: &VirtualMachine,
62) -> PyResult {
63    // = structseq_new
64    let StructSequenceNewArgs {
65        sequence: seq,
66        dict,
67    } = args;
68
69    #[cold]
70    fn length_error(
71        tp_name: &str,
72        min_len: usize,
73        max_len: usize,
74        len: usize,
75        vm: &VirtualMachine,
76    ) -> PyBaseExceptionRef {
77        if min_len == max_len {
78            vm.new_type_error(format!(
79                "{tp_name}() takes a {min_len}-sequence ({len}-sequence given)"
80            ))
81        } else if len < min_len {
82            vm.new_type_error(format!(
83                "{tp_name}() takes an at least {min_len}-sequence ({len}-sequence given)"
84            ))
85        } else {
86            vm.new_type_error(format!(
87                "{tp_name}() takes an at most {max_len}-sequence ({len}-sequence given)"
88            ))
89        }
90    }
91
92    let min_len: usize = cls
93        .get_attr(identifier!(vm.ctx, n_sequence_fields))
94        .ok_or_else(|| vm.new_type_error("missing n_sequence_fields attribute"))?
95        .try_into_value(vm)?;
96    let max_len: usize = cls
97        .get_attr(identifier!(vm.ctx, n_fields))
98        .ok_or_else(|| vm.new_type_error("missing n_fields attribute"))?
99        .try_into_value(vm)?;
100
101    let dict = match dict {
102        OptionalArg::Missing => None,
103        OptionalArg::Present(dict) => Some(dict.downcast::<PyDict>().map_err(|_| {
104            vm.new_type_error(format!(
105                "{}() takes a dict as second arg, if any",
106                cls.slot_name()
107            ))
108        })?),
109    };
110
111    let seq: Vec<PyObjectRef> = seq.try_into_value(vm)?;
112    let len = seq.len();
113
114    if len < min_len || len > max_len {
115        return Err(length_error(&cls.slot_name(), min_len, max_len, len, vm));
116    }
117
118    // Copy items and pad the hidden fields the sequence did not cover with None.
119    let mut items = seq;
120    items.resize_with(max_len, || vm.ctx.none());
121
122    // Fill those padded slots from `dict`. Every key has to land in one of them:
123    // a key naming a field the sequence already supplied, or no field at all,
124    // would otherwise be silently dropped.
125    if let Some(dict) = dict.filter(|dict| !dict.is_empty()) {
126        let mut found = 0;
127        let names = hidden_field_names.get(len - min_len..).unwrap_or(&[]);
128        for (item, name) in items[len..].iter_mut().zip(names) {
129            if let Some(value) = dict.get_item_opt(*name, vm)? {
130                *item = value;
131                found += 1;
132            }
133        }
134        if found != dict.__len__() {
135            return Err(vm.new_type_error(format!(
136                "{}() got duplicate or unexpected field name(s)",
137                cls.slot_name()
138            )));
139        }
140    }
141
142    PyTuple::new_unchecked(items.into_boxed_slice())
143        .into_ref_with_type(vm, cls)
144        .map(Into::into)
145}
146
147fn get_visible_len(obj: &PyObject, vm: &VirtualMachine) -> PyResult<usize> {
148    obj.class()
149        .get_attr(identifier!(vm.ctx, n_sequence_fields))
150        .ok_or_else(|| vm.new_type_error("missing n_sequence_fields"))?
151        .try_into_value(vm)
152}
153
154/// Sequence methods for struct sequences.
155/// Uses n_sequence_fields to determine visible length.
156static STRUCT_SEQUENCE_AS_SEQUENCE: LazyLock<PySequenceMethods> =
157    LazyLock::new(|| PySequenceMethods {
158        length: atomic_func!(|seq, vm| get_visible_len(seq.obj, vm)),
159        concat: atomic_func!(|seq, other, vm| {
160            // Convert to visible-only tuple, then use regular tuple concat
161            let n_seq = get_visible_len(seq.obj, vm)?;
162            let tuple = seq.obj.downcast_ref::<PyTuple>().unwrap();
163            let visible: Vec<_> = tuple.as_slice().iter().take(n_seq).cloned().collect();
164            let visible_tuple = PyTuple::new_ref(visible, &vm.ctx);
165            // Use tuple's concat implementation
166            visible_tuple
167                .as_object()
168                .sequence_unchecked()
169                .concat(other, vm)
170        }),
171        repeat: atomic_func!(|seq, n, vm| {
172            // Convert to visible-only tuple, then use regular tuple repeat
173            let n_seq = get_visible_len(seq.obj, vm)?;
174            let tuple = seq.obj.downcast_ref::<PyTuple>().unwrap();
175            let visible: Vec<_> = tuple.as_slice().iter().take(n_seq).cloned().collect();
176            let visible_tuple = PyTuple::new_ref(visible, &vm.ctx);
177            // Use tuple's repeat implementation
178            visible_tuple.as_object().sequence_unchecked().repeat(n, vm)
179        }),
180        item: atomic_func!(|seq, i, vm| {
181            let n_seq = get_visible_len(seq.obj, vm)?;
182            let tuple = seq.obj.downcast_ref::<PyTuple>().unwrap();
183            let idx = if i < 0 {
184                let pos_i = n_seq as isize + i;
185                if pos_i < 0 {
186                    return Err(vm.new_index_error("tuple index out of range"));
187                }
188                pos_i as usize
189            } else {
190                i as usize
191            };
192            if idx >= n_seq {
193                return Err(vm.new_index_error("tuple index out of range"));
194            }
195            Ok(tuple.as_slice()[idx].clone())
196        }),
197        contains: atomic_func!(|seq, needle, vm| {
198            let n_seq = get_visible_len(seq.obj, vm)?;
199            let tuple = seq.obj.downcast_ref::<PyTuple>().unwrap();
200            for item in tuple.as_slice().iter().take(n_seq) {
201                if item.rich_compare_bool(needle, PyComparisonOp::Eq, vm)? {
202                    return Ok(true);
203                }
204            }
205            Ok(false)
206        }),
207        ..PySequenceMethods::NOT_IMPLEMENTED
208    });
209
210/// Mapping methods for struct sequences.
211/// Handles subscript (indexing) with visible length bounds.
212static STRUCT_SEQUENCE_AS_MAPPING: LazyLock<PyMappingMethods> =
213    LazyLock::new(|| PyMappingMethods {
214        length: atomic_func!(|mapping, vm| get_visible_len(mapping.obj, vm)),
215        subscript: atomic_func!(|mapping, needle, vm| {
216            let n_seq = get_visible_len(mapping.obj, vm)?;
217            let tuple = mapping.obj.downcast_ref::<PyTuple>().unwrap();
218            let visible_elements = &tuple.as_slice()[..n_seq];
219
220            match SequenceIndex::try_from_borrowed_object(vm, needle, "tuple")? {
221                SequenceIndex::Int(i) => visible_elements.getitem_by_index(vm, i),
222                SequenceIndex::Slice(slice) => visible_elements
223                    .getitem_by_slice(vm, slice)
224                    .map(|x| vm.ctx.new_tuple(x).into()),
225            }
226        }),
227        ..PyMappingMethods::NOT_IMPLEMENTED
228    });
229
230/// Trait for Data structs that back a PyStructSequence.
231///
232/// This trait is implemented by `#[pystruct_sequence_data]` on the Data struct.
233/// It provides field information, tuple conversion, and element parsing.
234pub trait PyStructSequenceData: Sized {
235    /// Names of required fields (in order). Shown in repr.
236    const REQUIRED_FIELD_NAMES: &'static [&'static str];
237
238    /// Names of optional/skipped fields (in order, after required fields).
239    const OPTIONAL_FIELD_NAMES: &'static [&'static str];
240
241    /// Number of unnamed fields (visible but index-only access).
242    const UNNAMED_FIELDS_LEN: usize = 0;
243
244    /// Convert this Data struct into a PyTuple.
245    fn into_tuple(self, vm: &VirtualMachine) -> PyTuple;
246
247    /// Construct this Data struct from tuple elements.
248    /// Default implementation returns an error.
249    /// Override with `#[pystruct_sequence_data(try_from_object)]` to enable.
250    fn try_from_elements(_elements: Vec<PyObjectRef>, vm: &VirtualMachine) -> PyResult<Self> {
251        Err(vm.new_type_error("This struct sequence does not support construction from elements"))
252    }
253}
254
255/// Trait for Python struct sequence types.
256///
257/// This trait is implemented by the `#[pystruct_sequence]` macro on the Python type struct.
258/// It connects to the Data struct and provides Python-level functionality.
259#[pyclass]
260pub trait PyStructSequence: StaticType + PyClassImpl + Sized + 'static {
261    /// The Data struct that provides field definitions.
262    type Data: PyStructSequenceData;
263
264    #[pyslot]
265    fn slot_new(cls: PyTypeRef, args: FuncArgs, vm: &VirtualMachine) -> PyResult {
266        struct_sequence_new(
267            cls,
268            args.bind_for(vm, Self::NAME)?,
269            Self::Data::OPTIONAL_FIELD_NAMES,
270            vm,
271        )
272    }
273
274    /// Convert a Data struct into a PyStructSequence instance.
275    fn from_data(data: Self::Data, vm: &VirtualMachine) -> PyTupleRef {
276        let tuple =
277            <Self::Data as ::rustpython_vm::types::PyStructSequenceData>::into_tuple(data, vm);
278        let typ = Self::static_type();
279        tuple
280            .into_ref_with_type(vm, typ.to_owned())
281            .expect("Every PyStructSequence must be a valid tuple. This is a RustPython bug.")
282    }
283
284    #[pyslot]
285    fn slot_repr(zelf: &PyObject, vm: &VirtualMachine) -> PyResult<PyStrRef> {
286        let zelf = zelf
287            .downcast_ref::<PyTuple>()
288            .ok_or_else(|| vm.new_type_error("unexpected payload for __repr__"))?;
289
290        let field_names = Self::Data::REQUIRED_FIELD_NAMES;
291        let format_field = |(value, name): (&PyObject, _)| {
292            let s = value.repr(vm)?;
293            Ok(format!("{name}={s}"))
294        };
295        let (body, suffix) =
296            if let Some(_guard) = rustpython_vm::recursion::ReprGuard::enter(vm, zelf.as_ref()) {
297                let fields: PyResult<Vec<_>> = zelf
298                    .as_slice()
299                    .iter()
300                    .map(|value| value.as_ref())
301                    .zip(field_names.iter().copied())
302                    .map(format_field)
303                    .collect();
304                (fields?.join(", "), "")
305            } else {
306                (String::new(), "...")
307            };
308        // Build qualified name: if MODULE_NAME is already in TP_NAME, use it directly.
309        // Otherwise, check __module__ attribute (set by #[pymodule] at runtime).
310        let type_name = if Self::MODULE_NAME.is_some() {
311            alloc::borrow::Cow::Borrowed(Self::TP_NAME)
312        } else {
313            let typ = zelf.class();
314            match typ.get_attr(identifier!(vm.ctx, __module__)) {
315                Some(module) if module.downcastable::<PyStr>() => {
316                    let module_str = module.downcast_ref::<PyStr>().unwrap();
317                    alloc::borrow::Cow::Owned(format!("{}.{}", module_str.as_wtf8(), Self::NAME))
318                }
319                _ => alloc::borrow::Cow::Borrowed(Self::TP_NAME),
320            }
321        };
322        let repr_str = format!("{type_name}({body}{suffix})");
323        Ok(vm.ctx.new_str(repr_str))
324    }
325
326    // Return a copy of the structure with new values for the specified fields.
327    #[pymethod]
328    fn __replace__(
329        zelf: PyRef<PyTuple>,
330        changes: KwArgs<PyObjectRef, NameChanges>,
331        vm: &VirtualMachine,
332    ) -> PyResult {
333        if Self::Data::UNNAMED_FIELDS_LEN > 0 {
334            return Err(vm.new_type_error(format!(
335                "__replace__() is not supported for {} because it has unnamed field(s)",
336                zelf.class().slot_name()
337            )));
338        }
339
340        let n_fields =
341            Self::Data::REQUIRED_FIELD_NAMES.len() + Self::Data::OPTIONAL_FIELD_NAMES.len();
342        let mut items: Vec<PyObjectRef> = zelf.as_slice()[..n_fields].to_vec();
343
344        let mut kwargs = changes;
345
346        // Replace fields from kwargs
347        let all_field_names: Vec<&str> = Self::Data::REQUIRED_FIELD_NAMES
348            .iter()
349            .chain(Self::Data::OPTIONAL_FIELD_NAMES.iter())
350            .copied()
351            .collect();
352        for (i, &name) in all_field_names.iter().enumerate() {
353            if let Some(val) = kwargs.shift_remove(name) {
354                items[i] = val;
355            }
356        }
357
358        if !kwargs.is_empty() {
359            let names = vm.ctx.new_list(
360                kwargs
361                    .keys()
362                    .map(|k| vm.ctx.new_str(k.to_owned()).into())
363                    .collect(),
364            );
365            let names_repr = names.as_object().repr(vm)?;
366            return Err(vm.new_type_error(format!("Got unexpected field name(s): {names_repr}")));
367        }
368
369        PyTuple::new_unchecked(items.into_boxed_slice())
370            .into_ref_with_type(vm, zelf.class().to_owned())
371            .map(Into::into)
372    }
373
374    #[pymethod]
375    fn __getitem__(zelf: PyRef<PyTuple>, needle: PyObjectRef, vm: &VirtualMachine) -> PyResult {
376        let n_seq = get_visible_len(zelf.as_ref(), vm)?;
377        let visible_elements = &zelf.as_slice()[..n_seq];
378
379        match SequenceIndex::try_from_borrowed_object(vm, &needle, "tuple")? {
380            SequenceIndex::Int(i) => visible_elements.getitem_by_index(vm, i),
381            SequenceIndex::Slice(slice) => visible_elements
382                .getitem_by_slice(vm, slice)
383                .map(|x| vm.ctx.new_tuple(x).into()),
384        }
385    }
386
387    #[extend_class]
388    fn extend_pyclass(ctx: &Context, class: &'static Py<PyType>) {
389        // Getters for named visible fields (indices 0 to REQUIRED_FIELD_NAMES.len() - 1)
390        for (i, &name) in Self::Data::REQUIRED_FIELD_NAMES.iter().enumerate() {
391            class.set_attr(
392                ctx.intern_str(name),
393                ctx.new_readonly_tuple_member(name, class, i, class_attr_item_doc::<Self>(name))
394                    .into(),
395            );
396        }
397
398        // Getters for hidden/skipped fields (indices after visible fields)
399        let visible_count = Self::Data::REQUIRED_FIELD_NAMES.len() + Self::Data::UNNAMED_FIELDS_LEN;
400        for (i, &name) in Self::Data::OPTIONAL_FIELD_NAMES.iter().enumerate() {
401            class.set_attr(
402                ctx.intern_str(name),
403                ctx.new_readonly_tuple_member(
404                    name,
405                    class,
406                    visible_count + i,
407                    class_attr_item_doc::<Self>(name),
408                )
409                .into(),
410            );
411        }
412
413        class.set_attr(
414            identifier!(ctx, __match_args__),
415            ctx.new_tuple(
416                Self::Data::REQUIRED_FIELD_NAMES
417                    .iter()
418                    .map(|&name| ctx.new_str(name).into())
419                    .collect::<Vec<_>>(),
420            )
421            .into(),
422        );
423
424        // special fields:
425        // n_sequence_fields = visible fields (named + unnamed)
426        // n_fields = all fields (visible + hidden/skipped)
427        // n_unnamed_fields
428        let n_unnamed_fields = Self::Data::UNNAMED_FIELDS_LEN;
429        let n_sequence_fields = Self::Data::REQUIRED_FIELD_NAMES.len() + n_unnamed_fields;
430        let n_fields = n_sequence_fields + Self::Data::OPTIONAL_FIELD_NAMES.len();
431        class.set_attr(
432            identifier!(ctx, n_sequence_fields),
433            ctx.new_int(n_sequence_fields).into(),
434        );
435        class.set_attr(identifier!(ctx, n_fields), ctx.new_int(n_fields).into());
436        class.set_attr(
437            identifier!(ctx, n_unnamed_fields),
438            ctx.new_int(n_unnamed_fields).into(),
439        );
440
441        // Override as_sequence and as_mapping slots to use visible length
442        class
443            .slots
444            .as_sequence
445            .copy_from(&STRUCT_SEQUENCE_AS_SEQUENCE);
446        class
447            .slots
448            .as_mapping
449            .copy_from(&STRUCT_SEQUENCE_AS_MAPPING);
450
451        // Override iter slot to return only visible elements
452        class.slots.iter.store(Some(struct_sequence_iter));
453
454        // Override hash slot to hash only visible elements
455        class.slots.hash.store(Some(struct_sequence_hash));
456
457        // Override richcompare slot to compare only visible elements
458        class
459            .slots
460            .richcompare
461            .store(Some(struct_sequence_richcompare));
462
463        // Default __reduce__: only set if not already overridden by the impl's extend_class.
464        // This allows struct sequences like sched_param to provide a custom __reduce__
465        // (equivalent to METH_COEXIST in structseq.c).
466        if !class.attributes.contains(ctx.intern_str("__reduce__")) {
467            class.set_attr(
468                ctx.intern_str("__reduce__"),
469                DEFAULT_STRUCTSEQ_REDUCE.to_proper_method(class, ctx),
470            );
471        }
472    }
473}
474
475/// Iterator function for struct sequences - returns only visible elements
476fn struct_sequence_iter(zelf: PyObjectRef, vm: &VirtualMachine) -> PyResult {
477    let tuple = zelf
478        .downcast_ref::<PyTuple>()
479        .ok_or_else(|| vm.new_type_error("expected tuple"))?;
480    let n_seq = get_visible_len(&zelf, vm)?;
481    let visible: Vec<_> = tuple.as_slice().iter().take(n_seq).cloned().collect();
482    let visible_tuple = PyTuple::new_ref(visible, &vm.ctx);
483    visible_tuple
484        .as_object()
485        .to_owned()
486        .get_iter(vm)
487        .map(Into::into)
488}
489
490/// Hash function for struct sequences - hashes only visible elements
491fn struct_sequence_hash(
492    zelf: &PyObject,
493    vm: &VirtualMachine,
494) -> PyResult<crate::common::hash::PyHash> {
495    let tuple = zelf
496        .downcast_ref::<PyTuple>()
497        .ok_or_else(|| vm.new_type_error("expected tuple"))?;
498    let n_seq = get_visible_len(zelf, vm)?;
499    // Create a visible-only tuple and hash it
500    let visible: Vec<_> = tuple.as_slice().iter().take(n_seq).cloned().collect();
501    let visible_tuple = PyTuple::new_ref(visible, &vm.ctx);
502    visible_tuple.as_object().hash(vm)
503}
504
505/// Rich comparison for struct sequences - compares only visible elements
506fn struct_sequence_richcompare(
507    zelf: &PyObject,
508    other: &PyObject,
509    op: PyComparisonOp,
510    vm: &VirtualMachine,
511) -> PyResult<Either<PyObjectRef, PyComparisonValue>> {
512    let zelf_tuple = zelf
513        .downcast_ref::<PyTuple>()
514        .ok_or_else(|| vm.new_type_error("expected tuple"))?;
515
516    // If other is not a tuple, return NotImplemented
517    let Some(other_tuple) = other.downcast_ref::<PyTuple>() else {
518        return Ok(Either::B(PyComparisonValue::NotImplemented));
519    };
520
521    let zelf_len = get_visible_len(zelf, vm)?;
522    // For other, try to get visible len; if it fails (not a struct sequence), use full length
523    let other_len = get_visible_len(other, vm).unwrap_or(other_tuple.as_slice().len());
524
525    let zelf_visible = &zelf_tuple.as_slice()[..zelf_len];
526    let other_visible = &other_tuple.as_slice()[..other_len];
527
528    // Use the same comparison logic as regular tuples
529    zelf_visible
530        .iter()
531        .map(|o| &**o)
532        .richcompare(other_visible.iter().map(|o| &**o), op, vm)
533        .map(|v| Either::B(PyComparisonValue::Implemented(v)))
534}