Skip to main content

rustpython_vm/protocol/
sequence.rs

1//! [Sequence Protocol](https://docs.python.org/3/c-api/sequence.html)
2
3use crossbeam_utils::atomic::AtomicCell;
4use itertools::Itertools;
5
6use crate::{
7    AsObject, PyObject, PyObjectRef, PyPayload, PyResult, VirtualMachine,
8    builtins::{PyList, PyListRef, PySlice, PyTuple, PyTupleRef},
9    convert::ToPyObject,
10    function::PyArithmeticValue,
11    object::{Traverse, TraverseFn},
12    protocol::PyNumberBinaryOp,
13};
14
15#[expect(clippy::type_complexity)]
16#[derive(Default)]
17pub struct PySequenceSlots {
18    pub length: AtomicCell<Option<fn(PySequence<'_>, &VirtualMachine) -> PyResult<usize>>>,
19    pub concat: AtomicCell<Option<fn(PySequence<'_>, &PyObject, &VirtualMachine) -> PyResult>>,
20    pub repeat: AtomicCell<Option<fn(PySequence<'_>, isize, &VirtualMachine) -> PyResult>>,
21    pub item: AtomicCell<Option<fn(PySequence<'_>, isize, &VirtualMachine) -> PyResult>>,
22    pub ass_item: AtomicCell<
23        Option<fn(PySequence<'_>, isize, Option<PyObjectRef>, &VirtualMachine) -> PyResult<()>>,
24    >,
25    pub contains:
26        AtomicCell<Option<fn(PySequence<'_>, &PyObject, &VirtualMachine) -> PyResult<bool>>>,
27    pub inplace_concat:
28        AtomicCell<Option<fn(PySequence<'_>, &PyObject, &VirtualMachine) -> PyResult>>,
29    pub inplace_repeat: AtomicCell<Option<fn(PySequence<'_>, isize, &VirtualMachine) -> PyResult>>,
30}
31
32impl core::fmt::Debug for PySequenceSlots {
33    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
34        f.write_str("PySequenceSlots")
35    }
36}
37
38impl PySequenceSlots {
39    pub fn has_item(&self) -> bool {
40        self.item.load().is_some()
41    }
42
43    /// Whether any slot is filled, which is what a non-null `tp_as_sequence`
44    /// amounts to for a statically declared type.
45    pub fn has_any(&self) -> bool {
46        self.length.load().is_some()
47            || self.concat.load().is_some()
48            || self.repeat.load().is_some()
49            || self.item.load().is_some()
50            || self.ass_item.load().is_some()
51            || self.contains.load().is_some()
52            || self.inplace_concat.load().is_some()
53            || self.inplace_repeat.load().is_some()
54    }
55
56    /// Copy from static PySequenceMethods
57    pub fn copy_from(&self, methods: &PySequenceMethods) {
58        if let Some(f) = methods.length {
59            self.length.store(Some(f));
60        }
61
62        if let Some(f) = methods.concat {
63            self.concat.store(Some(f));
64        }
65
66        if let Some(f) = methods.repeat {
67            self.repeat.store(Some(f));
68        }
69
70        if let Some(f) = methods.item {
71            self.item.store(Some(f));
72        }
73
74        if let Some(f) = methods.ass_item {
75            self.ass_item.store(Some(f));
76        }
77
78        if let Some(f) = methods.contains {
79            self.contains.store(Some(f));
80        }
81
82        if let Some(f) = methods.inplace_concat {
83            self.inplace_concat.store(Some(f));
84        }
85
86        if let Some(f) = methods.inplace_repeat {
87            self.inplace_repeat.store(Some(f));
88        }
89    }
90}
91
92#[expect(clippy::type_complexity)]
93#[derive(Default)]
94pub struct PySequenceMethods {
95    pub length: Option<fn(PySequence<'_>, &VirtualMachine) -> PyResult<usize>>,
96    pub concat: Option<fn(PySequence<'_>, &PyObject, &VirtualMachine) -> PyResult>,
97    pub repeat: Option<fn(PySequence<'_>, isize, &VirtualMachine) -> PyResult>,
98    pub item: Option<fn(PySequence<'_>, isize, &VirtualMachine) -> PyResult>,
99    pub ass_item:
100        Option<fn(PySequence<'_>, isize, Option<PyObjectRef>, &VirtualMachine) -> PyResult<()>>,
101    pub contains: Option<fn(PySequence<'_>, &PyObject, &VirtualMachine) -> PyResult<bool>>,
102    pub inplace_concat: Option<fn(PySequence<'_>, &PyObject, &VirtualMachine) -> PyResult>,
103    pub inplace_repeat: Option<fn(PySequence<'_>, isize, &VirtualMachine) -> PyResult>,
104}
105
106impl core::fmt::Debug for PySequenceMethods {
107    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
108        f.write_str("PySequenceMethods")
109    }
110}
111
112impl PySequenceMethods {
113    pub const NOT_IMPLEMENTED: Self = Self {
114        length: None,
115        concat: None,
116        repeat: None,
117        item: None,
118        ass_item: None,
119        contains: None,
120        inplace_concat: None,
121        inplace_repeat: None,
122    };
123}
124
125impl PyObject {
126    #[inline]
127    pub const fn sequence_unchecked(&self) -> PySequence<'_> {
128        PySequence { obj: self }
129    }
130
131    pub fn try_sequence(&self, vm: &VirtualMachine) -> PyResult<PySequence<'_>> {
132        let seq = self.sequence_unchecked();
133        if seq.check() {
134            Ok(seq)
135        } else {
136            Err(vm.new_type_error(format!("{} is not a sequence", self.class().slot_name())))
137        }
138    }
139}
140
141#[derive(Copy, Clone)]
142pub struct PySequence<'a> {
143    pub obj: &'a PyObject,
144}
145
146unsafe impl Traverse for PySequence<'_> {
147    fn traverse(&self, tracer_fn: &mut TraverseFn<'_>) {
148        self.obj.traverse(tracer_fn)
149    }
150}
151
152impl PySequence<'_> {
153    #[inline]
154    #[must_use]
155    pub fn slots(&self) -> &PySequenceSlots {
156        &self.obj.class().slots().as_sequence
157    }
158
159    #[must_use]
160    pub fn check(&self) -> bool {
161        self.slots().has_item()
162    }
163
164    pub fn length_opt(self, vm: &VirtualMachine) -> Option<PyResult<usize>> {
165        self.slots().length.load().map(|f| f(self, vm))
166    }
167
168    // Py_ssize_t PySequence_Size(PyObject *s)
169    pub fn length(self, vm: &VirtualMachine) -> PyResult<usize> {
170        self.length_opt(vm).ok_or_else(|| {
171            let name = self.obj.class().slot_name();
172            // Something that measures itself as a mapping is no sequence at all.
173            let msg = if self.obj.mapping_unchecked().slots().length.load().is_some() {
174                format!("{name} is not a sequence")
175            } else {
176                format!("object of type '{name}' has no len()")
177            };
178            vm.new_type_error(msg)
179        })?
180    }
181
182    pub fn concat(self, other: &PyObject, vm: &VirtualMachine) -> PyResult {
183        if let Some(f) = self.slots().concat.load() {
184            return f(self, other, vm);
185        }
186
187        // if both arguments appear to be sequences, try fallback to __add__
188        if self.check() && other.sequence_unchecked().check() {
189            let ret = vm.binary_op1(self.obj, other, PyNumberBinaryOp::Add)?;
190            if let PyArithmeticValue::Implemented(ret) = PyArithmeticValue::from_object(vm, ret) {
191                return Ok(ret);
192            }
193        }
194
195        Err(vm.new_type_error(format!(
196            "'{}' object can't be concatenated",
197            self.obj.class().slot_name()
198        )))
199    }
200
201    pub fn repeat(self, n: isize, vm: &VirtualMachine) -> PyResult {
202        if let Some(f) = self.slots().repeat.load() {
203            return f(self, n, vm);
204        }
205
206        // fallback to __mul__
207        if self.check() {
208            let ret = vm.binary_op1(self.obj, &n.to_pyobject(vm), PyNumberBinaryOp::Multiply)?;
209            if let PyArithmeticValue::Implemented(ret) = PyArithmeticValue::from_object(vm, ret) {
210                return Ok(ret);
211            }
212        }
213
214        Err(vm.new_type_error(format!(
215            "'{}' object can't be repeated",
216            self.obj.class().slot_name()
217        )))
218    }
219
220    pub fn inplace_concat(self, other: &PyObject, vm: &VirtualMachine) -> PyResult {
221        if let Some(f) = self.slots().inplace_concat.load() {
222            return f(self, other, vm);
223        }
224        if let Some(f) = self.slots().concat.load() {
225            return f(self, other, vm);
226        }
227
228        // if both arguments appear to be sequences, try fallback to __iadd__
229        if self.check() && other.sequence_unchecked().check() {
230            let ret = vm.binary_iop1(
231                self.obj,
232                other,
233                PyNumberBinaryOp::InplaceAdd,
234                PyNumberBinaryOp::Add,
235            )?;
236            if let PyArithmeticValue::Implemented(ret) = PyArithmeticValue::from_object(vm, ret) {
237                return Ok(ret);
238            }
239        }
240
241        Err(vm.new_type_error(format!(
242            "'{}' object can't be concatenated",
243            self.obj.class().slot_name()
244        )))
245    }
246
247    pub fn inplace_repeat(self, n: isize, vm: &VirtualMachine) -> PyResult {
248        if let Some(f) = self.slots().inplace_repeat.load() {
249            return f(self, n, vm);
250        }
251
252        if let Some(f) = self.slots().repeat.load() {
253            return f(self, n, vm);
254        }
255
256        if self.check() {
257            let ret = vm.binary_iop1(
258                self.obj,
259                &n.to_pyobject(vm),
260                PyNumberBinaryOp::InplaceMultiply,
261                PyNumberBinaryOp::Multiply,
262            )?;
263            if let PyArithmeticValue::Implemented(ret) = PyArithmeticValue::from_object(vm, ret) {
264                return Ok(ret);
265            }
266        }
267
268        Err(vm.new_type_error(format!(
269            "'{}' object can't be repeated",
270            self.obj.class().slot_name()
271        )))
272    }
273
274    pub fn get_item(self, i: isize, vm: &VirtualMachine) -> PyResult {
275        if let Some(f) = self.slots().item.load() {
276            return f(self, i, vm);
277        }
278
279        let name = self.obj.class().slot_name();
280        // Something that subscripts itself as a mapping is no sequence at all.
281        let msg = if self
282            .obj
283            .mapping_unchecked()
284            .slots()
285            .subscript
286            .load()
287            .is_some()
288        {
289            format!("{name} is not a sequence")
290        } else {
291            format!("'{name}' object does not support indexing")
292        };
293        Err(vm.new_type_error(msg))
294    }
295
296    fn _ass_item(self, i: isize, value: Option<PyObjectRef>, vm: &VirtualMachine) -> PyResult<()> {
297        if let Some(f) = self.slots().ass_item.load() {
298            return f(self, i, value, vm);
299        }
300
301        let name = self.obj.class().slot_name();
302        let msg = if self
303            .obj
304            .mapping_unchecked()
305            .slots()
306            .ass_subscript
307            .load()
308            .is_some()
309        {
310            format!("{name} is not a sequence")
311        } else if value.is_some() {
312            format!("'{name}' object does not support item assignment")
313        } else {
314            format!("'{name}' object doesn't support item deletion")
315        };
316        Err(vm.new_type_error(msg))
317    }
318
319    pub fn set_item(self, i: isize, value: PyObjectRef, vm: &VirtualMachine) -> PyResult<()> {
320        self._ass_item(i, Some(value), vm)
321    }
322
323    pub fn del_item(self, i: isize, vm: &VirtualMachine) -> PyResult<()> {
324        self._ass_item(i, None, vm)
325    }
326
327    pub fn get_slice(&self, start: isize, stop: isize, vm: &VirtualMachine) -> PyResult {
328        if let Ok(mapping) = self.obj.try_mapping(vm) {
329            let slice = PySlice {
330                start: Some(start.to_pyobject(vm)),
331                stop: stop.to_pyobject(vm),
332                step: None,
333            };
334            mapping.subscript(&slice.into_pyobject(vm), vm)
335        } else {
336            Err(vm.new_type_error(format!(
337                "'{}' object is unsliceable",
338                self.obj.class().slot_name()
339            )))
340        }
341    }
342
343    fn _ass_slice(
344        self,
345        start: isize,
346        stop: isize,
347        value: Option<PyObjectRef>,
348        vm: &VirtualMachine,
349    ) -> PyResult<()> {
350        let mapping = self.obj.mapping_unchecked();
351        if let Some(f) = mapping.slots().ass_subscript.load() {
352            let slice = PySlice {
353                start: Some(start.to_pyobject(vm)),
354                stop: stop.to_pyobject(vm),
355                step: None,
356            };
357            f(mapping, &slice.into_pyobject(vm), value, vm)
358        } else {
359            Err(vm.new_type_error(format!(
360                "'{}' object doesn't support slice {}",
361                self.obj.class().slot_name(),
362                if value.is_some() {
363                    "assignment"
364                } else {
365                    "deletion"
366                }
367            )))
368        }
369    }
370
371    pub fn set_slice(
372        &self,
373        start: isize,
374        stop: isize,
375        value: PyObjectRef,
376        vm: &VirtualMachine,
377    ) -> PyResult<()> {
378        self._ass_slice(start, stop, Some(value), vm)
379    }
380
381    pub fn del_slice(&self, start: isize, stop: isize, vm: &VirtualMachine) -> PyResult<()> {
382        self._ass_slice(start, stop, None, vm)
383    }
384
385    pub fn tuple(&self, vm: &VirtualMachine) -> PyResult<PyTupleRef> {
386        if let Some(tuple) = self.obj.downcast_ref_if_exact::<PyTuple>(vm) {
387            Ok(tuple.to_owned())
388        } else if let Some(list) = self.obj.downcast_ref_if_exact::<PyList>(vm) {
389            Ok(vm.ctx.new_tuple(list.borrow_vec().to_vec()))
390        } else {
391            let iter = self.obj.to_owned().get_iter(vm)?;
392            let iter = iter.iter(vm)?;
393            Ok(vm.ctx.new_tuple(iter.try_collect()?))
394        }
395    }
396
397    pub fn list(&self, vm: &VirtualMachine) -> PyResult<PyListRef> {
398        Ok(vm.ctx.new_list(self.obj.try_to_value(vm)?))
399    }
400
401    pub fn count(&self, target: &PyObject, vm: &VirtualMachine) -> PyResult<usize> {
402        let mut n = 0;
403
404        let iter = self.obj.to_owned().get_iter(vm)?;
405        let iter = iter.iter::<PyObjectRef>(vm)?;
406
407        for elem in iter {
408            let elem = elem?;
409            if vm.bool_eq(&elem, target)? {
410                if n == isize::MAX as usize {
411                    return Err(vm.new_overflow_error("index exceeds C integer size"));
412                }
413                n += 1;
414            }
415        }
416
417        Ok(n)
418    }
419
420    pub fn index(&self, target: &PyObject, vm: &VirtualMachine) -> PyResult<usize> {
421        let iter = self.obj.to_owned().get_iter(vm)?;
422        let iter = iter.iter::<PyObjectRef>(vm)?;
423
424        for (index, elem) in iter.enumerate() {
425            if isize::try_from(index).is_err() {
426                return Err(vm.new_overflow_error("index exceeds C integer size"));
427            }
428
429            let elem = elem?;
430            if vm.bool_eq(&elem, target)? {
431                return Ok(index);
432            }
433        }
434
435        Err(vm.new_value_error("sequence.index(x): x not in sequence"))
436    }
437
438    pub fn extract<F, R>(&self, mut f: F, vm: &VirtualMachine) -> PyResult<Vec<R>>
439    where
440        F: FnMut(&PyObject) -> PyResult<R>,
441    {
442        let mut v = Vec::new();
443        if let Some(tuple) = self.obj.downcast_ref_if_exact::<PyTuple>(vm) {
444            v.try_reserve_exact(tuple.len())
445                .map_err(|_| vm.no_memory_error())?;
446            for x in tuple.as_slice() {
447                v.push(f(x.as_ref())?);
448            }
449        } else if let Some(list) = self.obj.downcast_ref_if_exact::<PyList>(vm) {
450            let elements = list.borrow_vec();
451            v.try_reserve_exact(elements.len())
452                .map_err(|_| vm.no_memory_error())?;
453            for x in elements.iter() {
454                v.push(f(x.as_ref())?);
455            }
456        } else {
457            let iter = self.obj.to_owned().get_iter(vm)?;
458            let iter = iter.iter::<PyObjectRef>(vm)?;
459            let len = self.length(vm).unwrap_or(0);
460            v.try_reserve_exact(len).map_err(|_| vm.no_memory_error())?;
461            for x in iter {
462                let item = f(x?.as_ref())?;
463                if v.len() == v.capacity() {
464                    v.try_reserve(1).map_err(|_| vm.no_memory_error())?;
465                }
466                v.push(item);
467            }
468        }
469        Ok(v)
470    }
471
472    pub fn contains(self, target: &PyObject, vm: &VirtualMachine) -> PyResult<bool> {
473        if let Some(f) = self.slots().contains.load() {
474            return f(self, target, vm);
475        }
476
477        // CPython parity: when neither __contains__ nor __iter__ is available,
478        // `PySequence_Contains` rewords the get_iter TypeError into the
479        // membership-test wording. Other exception types propagate unchanged.
480        let iter = self.obj.to_owned().get_iter(vm).map_err(|e| {
481            if e.fast_isinstance(vm.ctx.exceptions.type_error) {
482                vm.new_type_error(format!(
483                    "argument of type '{}' is not a container or iterable",
484                    self.obj.class().slot_name()
485                ))
486            } else {
487                e
488            }
489        })?;
490        let iter = iter.iter::<PyObjectRef>(vm)?;
491
492        for elem in iter {
493            let elem = elem?;
494            if vm.bool_eq(&elem, target)? {
495                return Ok(true);
496            }
497        }
498        Ok(false)
499    }
500}
501
502#[cfg(test)]
503mod tests {
504    use super::*;
505    use crate::{Interpreter, builtins::PyRange};
506    use core::cell::Cell;
507
508    #[derive(Debug)]
509    struct Converted<'a>(&'a Cell<usize>);
510
511    impl Drop for Converted<'_> {
512        fn drop(&mut self) {
513            self.0.set(self.0.get() + 1);
514        }
515    }
516
517    #[test]
518    fn unsupported_inplace_operations_keep_sequence_errors() {
519        Interpreter::without_stdlib(Default::default()).enter(|vm| {
520            let range = PyRange {
521                start: vm.ctx.new_int(0),
522                stop: vm.ctx.new_int(3),
523                step: vm.ctx.new_int(1),
524            }
525            .into_ref(&vm.ctx);
526            let sequence = range.as_object().sequence_unchecked();
527            for (result, message) in [
528                (
529                    sequence.inplace_repeat(2, vm),
530                    "'range' object can't be repeated",
531                ),
532                (
533                    sequence.inplace_concat(range.as_object(), vm),
534                    "'range' object can't be concatenated",
535                ),
536            ] {
537                let error = result.unwrap_err();
538                assert!(error.fast_isinstance(vm.ctx.exceptions.type_error));
539                let actual: String = error.args().as_slice()[0].try_to_value(vm).unwrap();
540                assert_eq!(actual, message);
541            }
542        });
543    }
544
545    #[test]
546    fn conversion_and_partial_result_cleanup() {
547        Interpreter::without_stdlib(Default::default()).enter(|vm| {
548            let elements: Vec<PyObjectRef> =
549                (0..20).map(|value| vm.ctx.new_int(value).into()).collect();
550            let list = vm.ctx.new_list(elements.clone());
551            let iterator = list.as_object().get_iter(vm).unwrap();
552            assert!(iterator.sequence_unchecked().length_opt(vm).is_none());
553            let sequences: [PyObjectRef; 3] = [
554                vm.ctx.new_tuple(elements).into(),
555                list.clone().into(),
556                iterator.into(),
557            ];
558            for sequence in sequences {
559                let values: Vec<i32> = sequence
560                    .sequence_unchecked()
561                    .extract(|item| item.try_to_value(vm), vm)
562                    .unwrap();
563                assert_eq!(values, (0..20).collect::<Vec<_>>());
564            }
565
566            let iterator = list.as_object().get_iter(vm).unwrap();
567            let error = vm.new_value_error("conversion failed");
568            let dropped = Cell::new(0);
569            let mut calls = 0;
570            let raised = iterator
571                .sequence_unchecked()
572                .extract(
573                    |_| {
574                        calls += 1;
575                        if calls == 3 {
576                            Err(error.clone())
577                        } else {
578                            Ok(Converted(&dropped))
579                        }
580                    },
581                    vm,
582                )
583                .unwrap_err();
584            assert!(raised.is(&error));
585            assert_eq!(calls, 3);
586            assert_eq!(dropped.get(), 2);
587        });
588    }
589
590    #[test]
591    fn capacity_overflow_precedes_conversion() {
592        Interpreter::without_stdlib(Default::default()).enter(|vm| {
593            let range = PyRange {
594                start: vm.ctx.new_int(0),
595                stop: vm.ctx.new_int(isize::MAX),
596                step: vm.ctx.new_int(1),
597            }
598            .into_ref(&vm.ctx);
599            let called = Cell::new(false);
600            // The byte capacity exceeds isize::MAX; no large allocation is attempted.
601            let error = range
602                .as_object()
603                .sequence_unchecked()
604                .extract(
605                    |_| {
606                        called.set(true);
607                        Ok(0u64)
608                    },
609                    vm,
610                )
611                .unwrap_err();
612            assert!(error.fast_isinstance(vm.ctx.exceptions.memory_error));
613            assert!(!called.get());
614        });
615    }
616}