Skip to main content

rustpython_vm/builtins/
enumerate.rs

1use super::{
2    IterStatus, PositionIterInternal, PyGenericAlias, PyIntRef, PyTupleRef, PyType, PyTypeRef,
3    iter::builtins_reversed, locked_rev_next,
4};
5use crate::common::lock::{PyMutex, PyRwLock};
6use crate::{
7    AsObject, Context, Py, PyObjectRef, PyPayload, PyResult, VirtualMachine,
8    class::PyClassImpl,
9    protocol::{PyIter, PyIterReturn},
10    raise_if_stop,
11    types::{Constructor, IterNext, Iterable, SelfIter},
12};
13use malachite_bigint::BigInt;
14use num_traits::ToPrimitive;
15
16/// Fast-path counter for `enumerate`.
17///
18/// malachite `BigInt` already keeps limb-sized values inline (`Natural::Small`)
19/// and only heap-allocates beyond that, so this enum looks redundant. The
20/// small/large tag is crate-private, so we cannot inspect or bump that limb
21/// ourselves; storing a `BigInt` would still clone it and go through
22/// `+= 1` on every `next()`. Keep a `usize` until it overflows, or when
23/// `start` does not fit in `usize`.
24#[derive(Debug, Clone)]
25enum Counter {
26    Small(usize),
27    Big(BigInt),
28}
29
30impl Counter {
31    fn to_bigint(&self) -> BigInt {
32        match self {
33            Self::Small(n) => BigInt::from(*n),
34            Self::Big(b) => b.clone(),
35        }
36    }
37}
38
39#[pyclass(module = false, name = "enumerate", traverse)]
40#[derive(Debug)]
41pub struct PyEnumerate {
42    #[pytraverse(skip)]
43    counter: PyRwLock<Counter>,
44    iterable: PyIter,
45}
46
47impl PyPayload for PyEnumerate {
48    #[inline]
49    fn class(ctx: &Context) -> &'static Py<PyType> {
50        ctx.types.enumerate_type
51    }
52}
53
54#[derive(FromArgs)]
55pub struct EnumerateArgs {
56    #[pyarg(any)]
57    iterable: PyIter,
58    #[pyarg(any, default = 0)]
59    start: PyIntRef,
60}
61
62impl Constructor for PyEnumerate {
63    type Args = EnumerateArgs;
64
65    fn py_new(
66        _cls: &Py<PyType>,
67        Self::Args { iterable, start }: Self::Args,
68        _vm: &VirtualMachine,
69    ) -> PyResult<Self> {
70        let counter = match start.as_bigint().to_usize() {
71            Some(n) => Counter::Small(n),
72            None => Counter::Big(start.as_bigint().clone()),
73        };
74        Ok(Self {
75            counter: PyRwLock::new(counter),
76            iterable,
77        })
78    }
79}
80
81#[pyclass(with(IterNext, Iterable, Constructor), flags(BASETYPE))]
82impl Py<PyEnumerate> {
83    #[pyclassmethod]
84    fn __class_getitem__(
85        cls: PyTypeRef,
86        object: PyObjectRef,
87        vm: &VirtualMachine,
88    ) -> PyResult<PyGenericAlias> {
89        PyGenericAlias::from_args(cls, object, vm)
90    }
91
92    #[pymethod]
93    fn __reduce__(&self) -> (PyTypeRef, (PyIter, BigInt)) {
94        (
95            self.class().to_owned(),
96            (self.iterable.clone(), self.counter.read().to_bigint()),
97        )
98    }
99}
100
101impl SelfIter for PyEnumerate {}
102
103impl IterNext for PyEnumerate {
104    fn next(zelf: &Py<Self>, vm: &VirtualMachine) -> PyResult<PyIterReturn> {
105        let next_obj = raise_if_stop!(zelf.iterable.next(vm)?);
106        let mut counter = zelf.counter.write();
107        let position = match &mut *counter {
108            Counter::Small(n) => {
109                let cur = *n;
110                match cur.checked_add(1) {
111                    Some(next_n) => {
112                        *n = next_n;
113                        vm.ctx.new_int(cur)
114                    }
115                    None => {
116                        let cur_int = vm.ctx.new_int(cur);
117                        *counter = Counter::Big(BigInt::from(cur) + 1);
118                        cur_int
119                    }
120                }
121            }
122            Counter::Big(b) => {
123                let position = b.clone();
124                *b += 1;
125                vm.ctx.new_bigint(&position)
126            }
127        };
128        drop(counter);
129        Ok(PyIterReturn::Return(
130            vm.new_tuple((position, next_obj)).into(),
131        ))
132    }
133}
134
135#[pyclass(module = false, name = "reversed", traverse)]
136#[derive(Debug)]
137pub(crate) struct PyReverseSequenceIterator {
138    internal: PyMutex<PositionIterInternal<PyObjectRef>>,
139}
140
141impl PyPayload for PyReverseSequenceIterator {
142    #[inline]
143    fn class(ctx: &Context) -> &'static Py<PyType> {
144        ctx.types.reverse_iter_type
145    }
146}
147
148impl PyReverseSequenceIterator {
149    pub(crate) const fn new(obj: PyObjectRef, len: usize) -> Self {
150        let position = len.saturating_sub(1);
151        Self {
152            internal: PyMutex::new(PositionIterInternal::new(obj, position)),
153        }
154    }
155}
156
157#[pyclass(with(IterNext, Iterable))]
158impl Py<PyReverseSequenceIterator> {
159    #[pymethod]
160    fn __length_hint__(&self, vm: &VirtualMachine) -> PyResult<usize> {
161        let internal = self.internal.lock();
162        if let IterStatus::Active(obj) = &internal.status
163            && internal.position <= obj.length(vm)?
164        {
165            return Ok(internal.position + 1);
166        }
167        Ok(0)
168    }
169
170    #[pymethod]
171    fn __setstate__(&self, state: PyObjectRef, vm: &VirtualMachine) -> PyResult<()> {
172        self.internal.lock().set_state(&state, |_, pos| pos, vm)
173    }
174
175    #[pymethod]
176    fn __reduce__(&self, vm: &VirtualMachine) -> PyResult<PyTupleRef> {
177        let func = builtins_reversed(vm)?;
178        Ok(self.internal.lock().reduce(
179            func,
180            |x| x.clone(),
181            |vm| vm.ctx.empty_tuple.clone().into(),
182            vm,
183        ))
184    }
185}
186
187impl SelfIter for PyReverseSequenceIterator {}
188impl IterNext for PyReverseSequenceIterator {
189    fn next(zelf: &Py<Self>, vm: &VirtualMachine) -> PyResult<PyIterReturn> {
190        locked_rev_next(&zelf.internal, |obj, pos| {
191            PyIterReturn::from_getitem_result(obj.get_item(&pos, vm), vm)
192        })
193    }
194}
195
196pub(crate) fn init(context: &'static Context) {
197    PyEnumerate::extend_class(context, context.types.enumerate_type);
198    PyReverseSequenceIterator::extend_class(context, context.types.reverse_iter_type);
199}