Skip to main content

rustpython_vm/builtins/
iter.rs

1/*
2 * iterator types
3 */
4
5use super::{PyInt, PyTupleRef, PyType};
6use crate::{
7    Context, Py, PyObject, PyObjectRef, PyPayload, PyResult, VirtualMachine,
8    class::PyClassImpl,
9    function::ArgCallable,
10    object::{Traverse, TraverseFn},
11    protocol::PyIterReturn,
12    types::{IterNext, Iterable, SelfIter},
13};
14use rustpython_common::lock::{PyMutex, PyRwLock, PyRwLockUpgradableReadGuard};
15
16/// Marks status of iterator.
17#[derive(Debug, Clone)]
18pub enum IterStatus<T> {
19    /// Iterator hasn't raised StopIteration.
20    Active(T),
21    /// Iterator has raised StopIteration.
22    Exhausted,
23}
24
25unsafe impl<T: Traverse> Traverse for IterStatus<T> {
26    fn traverse(&self, tracer_fn: &mut TraverseFn<'_>) {
27        match self {
28            Self::Active(r) => r.traverse(tracer_fn),
29            Self::Exhausted => (),
30        }
31    }
32}
33
34#[derive(Debug)]
35pub struct PositionIterInternal<T> {
36    pub status: IterStatus<T>,
37    pub position: usize,
38}
39
40unsafe impl<T: Traverse> Traverse for PositionIterInternal<T> {
41    fn traverse(&self, tracer_fn: &mut TraverseFn<'_>) {
42        self.status.traverse(tracer_fn)
43    }
44}
45
46impl<T> PositionIterInternal<T> {
47    pub const fn new(obj: T, position: usize) -> Self {
48        Self {
49            status: IterStatus::Active(obj),
50            position,
51        }
52    }
53
54    pub fn set_state<F>(&mut self, state: &PyObject, f: F, vm: &VirtualMachine) -> PyResult<()>
55    where
56        F: FnOnce(&T, usize) -> usize,
57    {
58        if let IterStatus::Active(obj) = &self.status {
59            if let Some(i) = state.downcast_ref::<PyInt>() {
60                let i = i.try_to_primitive(vm).unwrap_or(0);
61                self.position = f(obj, i);
62                Ok(())
63            } else {
64                Err(vm.new_type_error("an integer is required."))
65            }
66        } else {
67            Ok(())
68        }
69    }
70
71    /// Build a pickle-compatible reduce tuple.
72    ///
73    /// `func` must be resolved **before** acquiring any lock that guards this
74    /// `PositionIterInternal`, so that the builtins lookup cannot trigger
75    /// reentrant iterator access and deadlock.
76    pub fn reduce<F, E>(
77        &self,
78        func: PyObjectRef,
79        active: F,
80        empty: E,
81        vm: &VirtualMachine,
82    ) -> PyTupleRef
83    where
84        F: FnOnce(&T) -> PyObjectRef,
85        E: FnOnce(&VirtualMachine) -> PyObjectRef,
86    {
87        if let IterStatus::Active(obj) = &self.status {
88            vm.new_tuple((func, (active(obj),), self.position))
89        } else {
90            vm.new_tuple((func, (empty(vm),)))
91        }
92    }
93
94    /// `op` answers whether the step it took left this exhausted.
95    fn _next<F, OP>(&mut self, f: F, op: OP) -> (PyResult<PyIterReturn>, Option<T>)
96    where
97        F: FnOnce(&T, usize) -> PyResult<PyIterReturn>,
98        OP: FnOnce(&mut Self) -> bool,
99    {
100        let IterStatus::Active(obj) = &self.status else {
101            return (Ok(PyIterReturn::StopIteration(None)), None);
102        };
103        let ret = f(obj, self.position);
104        let done = match &ret {
105            Ok(PyIterReturn::Return(_)) => op(self),
106            Ok(PyIterReturn::StopIteration(_)) => true,
107            // An error belongs to the element, not to the walk, so the next
108            // call reaches for the same one again. `iter_iternext()` lets go of
109            // its sequence for `IndexError` and `StopIteration` alone, and
110            // `PyIterReturn::from_getitem_result` has already turned the first
111            // of those into the second.
112            Err(_) => false,
113        };
114        let released = if done { self.exhaust() } else { None };
115        (ret, released)
116    }
117
118    /// Mark this exhausted and hand back what it was holding, for the caller to
119    /// release once it has dropped the lock guarding this. Releasing it under
120    /// that lock would let a `__del__` that iterates again deadlock.
121    #[must_use]
122    pub fn exhaust(&mut self) -> Option<T> {
123        match core::mem::replace(&mut self.status, IterStatus::Exhausted) {
124            IterStatus::Active(obj) => Some(obj),
125            IterStatus::Exhausted => None,
126        }
127    }
128
129    /// Advance, along with what this was holding if the step exhausted it. See
130    /// [`Self::exhaust`] for why the caller is handed it rather than the drop
131    /// happening here; [`locked_next`] does the release for the common case.
132    #[must_use = "what this hands back is released after the lock, not here"]
133    pub fn next<F>(&mut self, f: F) -> (PyResult<PyIterReturn>, Option<T>)
134    where
135        F: FnOnce(&T, usize) -> PyResult<PyIterReturn>,
136    {
137        self._next(f, |zelf| {
138            zelf.position += 1;
139            false
140        })
141    }
142
143    /// [`Self::next`] walking backwards, exhausted once it steps off the front.
144    #[must_use = "what this hands back is released after the lock, not here"]
145    pub fn rev_next<F>(&mut self, f: F) -> (PyResult<PyIterReturn>, Option<T>)
146    where
147        F: FnOnce(&T, usize) -> PyResult<PyIterReturn>,
148    {
149        self._next(f, |zelf| {
150            if zelf.position == 0 {
151                return true;
152            }
153            zelf.position -= 1;
154            false
155        })
156    }
157
158    pub fn length_hint<F>(&self, f: F) -> usize
159    where
160        F: FnOnce(&T) -> usize,
161    {
162        if let IterStatus::Active(obj) = &self.status {
163            f(obj).saturating_sub(self.position)
164        } else {
165            0
166        }
167    }
168
169    pub fn rev_length_hint<F>(&self, f: F) -> usize
170    where
171        F: FnOnce(&T) -> usize,
172    {
173        if let IterStatus::Active(obj) = &self.status
174            && self.position <= f(obj)
175        {
176            return self.position + 1;
177        }
178        0
179    }
180}
181
182/// Take `step` under the lock `internal` holds, releasing whatever the step
183/// hands back only after that lock is gone. `setiter_iternext()` puts its
184/// `Py_DECREF(so)` past `Py_END_CRITICAL_SECTION()` for the same reason: a
185/// `__del__` that iterates again would otherwise wait on a lock still held here.
186pub(crate) fn locked_step<T>(
187    internal: &PyMutex<PositionIterInternal<T>>,
188    step: impl FnOnce(&mut PositionIterInternal<T>) -> (PyResult<PyIterReturn>, Option<T>),
189) -> PyResult<PyIterReturn> {
190    let mut guard = internal.lock();
191    let (ret, released) = step(&mut guard);
192    drop(guard);
193    drop(released);
194    ret
195}
196
197/// [`PositionIterInternal::next`] with the release [`locked_step`] describes.
198pub fn locked_next<T, F>(
199    internal: &PyMutex<PositionIterInternal<T>>,
200    f: F,
201) -> PyResult<PyIterReturn>
202where
203    F: FnOnce(&T, usize) -> PyResult<PyIterReturn>,
204{
205    locked_step(internal, |internal| internal.next(f))
206}
207
208/// [`locked_next`] walking backwards.
209pub fn locked_rev_next<T, F>(
210    internal: &PyMutex<PositionIterInternal<T>>,
211    f: F,
212) -> PyResult<PyIterReturn>
213where
214    F: FnOnce(&T, usize) -> PyResult<PyIterReturn>,
215{
216    locked_step(internal, |internal| internal.rev_next(f))
217}
218
219pub fn builtins_iter(vm: &VirtualMachine) -> PyResult {
220    vm.eval_get_builtin(vm.ctx.intern_str("iter"))
221}
222
223pub fn builtins_reversed(vm: &VirtualMachine) -> PyResult {
224    vm.eval_get_builtin(vm.ctx.intern_str("reversed"))
225}
226
227#[pyclass(module = false, name = "iterator", traverse)]
228#[derive(Debug)]
229pub struct PySequenceIterator {
230    internal: PyMutex<PositionIterInternal<PyObjectRef>>,
231}
232
233impl PyPayload for PySequenceIterator {
234    #[inline]
235    fn class(ctx: &Context) -> &'static Py<PyType> {
236        ctx.types.iter_type
237    }
238}
239
240impl PySequenceIterator {
241    pub fn new(obj: PyObjectRef, vm: &VirtualMachine) -> PyResult<Self> {
242        let _seq = obj.try_sequence(vm)?;
243        Ok(Self {
244            internal: PyMutex::new(PositionIterInternal::new(obj, 0)),
245        })
246    }
247}
248
249#[pyclass(with(IterNext, Iterable))]
250impl Py<PySequenceIterator> {
251    #[pymethod]
252    fn __length_hint__(&self, vm: &VirtualMachine) -> PyResult<PyObjectRef> {
253        vm.with_recursion("in __length_hint__", || {
254            let (obj, position) = {
255                let internal = self.internal.lock();
256                match &internal.status {
257                    IterStatus::Active(obj) => (Some(obj.clone()), internal.position),
258                    IterStatus::Exhausted => (None, 0),
259                }
260            };
261            if let Some(obj) = obj {
262                let seq = obj.sequence_unchecked();
263                match seq.length_opt(vm) {
264                    Some(len) => {
265                        len.map(|len| PyInt::from(len.saturating_sub(position)).into_pyobject(vm))
266                    }
267                    None => Ok(vm.ctx.not_implemented()),
268                }
269            } else {
270                Ok(PyInt::from(0).into_pyobject(vm))
271            }
272        })
273    }
274
275    #[pymethod]
276    fn __reduce__(&self, vm: &VirtualMachine) -> PyResult<PyTupleRef> {
277        let func = builtins_iter(vm)?;
278        Ok(self.internal.lock().reduce(
279            func,
280            |x| x.clone(),
281            |vm| vm.ctx.empty_tuple.clone().into(),
282            vm,
283        ))
284    }
285
286    #[pymethod]
287    fn __setstate__(&self, state: PyObjectRef, vm: &VirtualMachine) -> PyResult<()> {
288        self.internal.lock().set_state(&state, |_, pos| pos, vm)
289    }
290}
291
292impl SelfIter for PySequenceIterator {}
293impl IterNext for PySequenceIterator {
294    fn next(zelf: &Py<Self>, vm: &VirtualMachine) -> PyResult<PyIterReturn> {
295        locked_next(&zelf.internal, |obj, pos| {
296            let seq = obj.sequence_unchecked();
297            PyIterReturn::from_getitem_result(seq.get_item(pos as isize, vm), vm)
298        })
299    }
300}
301
302#[pyclass(module = false, name = "callable_iterator", traverse)]
303#[derive(Debug)]
304pub struct PyCallableIterator {
305    sentinel: PyObjectRef,
306    status: PyRwLock<IterStatus<ArgCallable>>,
307}
308
309impl PyPayload for PyCallableIterator {
310    #[inline]
311    fn class(ctx: &Context) -> &'static Py<PyType> {
312        ctx.types.callable_iterator
313    }
314}
315
316impl PyCallableIterator {
317    #[must_use]
318    pub const fn new(callable: ArgCallable, sentinel: PyObjectRef) -> Self {
319        Self {
320            sentinel,
321            status: PyRwLock::new(IterStatus::Active(callable)),
322        }
323    }
324}
325
326#[pyclass(with(IterNext, Iterable))]
327impl Py<PyCallableIterator> {
328    #[pymethod]
329    fn __reduce__(&self, vm: &VirtualMachine) -> PyResult<PyTupleRef> {
330        let func = builtins_iter(vm)?;
331        let status = self.status.read();
332        if let IterStatus::Active(callable) = &*status {
333            let callable_obj: PyObjectRef = callable.clone().into();
334            Ok(vm.new_tuple((func, (callable_obj, self.sentinel.clone()))))
335        } else {
336            Ok(vm.new_tuple((func, (vm.ctx.empty_tuple.clone(),))))
337        }
338    }
339}
340
341impl SelfIter for PyCallableIterator {}
342impl IterNext for PyCallableIterator {
343    fn next(zelf: &Py<Self>, vm: &VirtualMachine) -> PyResult<PyIterReturn> {
344        // Clone the callable and release the lock before invoking,
345        // so that reentrant next() calls don't deadlock.
346        let callable = {
347            let status = zelf.status.read();
348            match &*status {
349                IterStatus::Active(callable) => callable.clone(),
350                IterStatus::Exhausted => return Ok(PyIterReturn::StopIteration(None)),
351            }
352        };
353
354        let ret = callable.invoke((), vm)?;
355
356        // Re-check before comparing, but don't hold the lock while running
357        // sentinel equality. User __eq__ code can re-enter this iterator.
358        {
359            let status = zelf.status.read();
360            if !matches!(&*status, IterStatus::Active(_)) {
361                return Ok(PyIterReturn::StopIteration(None));
362            }
363        }
364
365        let is_sentinel = vm.identical_or_equal(&ret, &zelf.sentinel)?;
366
367        if is_sentinel {
368            let status = zelf.status.upgradable_read();
369            if !matches!(&*status, IterStatus::Active(_)) {
370                return Ok(PyIterReturn::StopIteration(None));
371            }
372            *PyRwLockUpgradableReadGuard::upgrade(status) = IterStatus::Exhausted;
373            Ok(PyIterReturn::StopIteration(None))
374        } else {
375            Ok(PyIterReturn::Return(ret))
376        }
377    }
378}
379
380pub fn init(context: &'static Context) {
381    PySequenceIterator::extend_class(context, context.types.iter_type);
382    PyCallableIterator::extend_class(context, context.types.callable_iterator);
383}