Skip to main content

rustpython_vm/protocol/
iter.rs

1use crate::{
2    AsObject, PyObject, PyObjectRef, PyPayload, PyResult, TryFromObject, VirtualMachine,
3    builtins::iter::PySequenceIterator,
4    convert::{ToPyObject, ToPyResult},
5    object::{Traverse, TraverseFn},
6};
7use core::borrow::Borrow;
8use core::ops::Deref;
9
10/// [Iterator Protocol](https://docs.python.org/3/c-api/iter.html).
11#[derive(Debug, Clone)]
12#[repr(transparent)]
13pub struct PyIter<O = PyObjectRef>(O)
14where
15    O: Borrow<PyObject>;
16
17unsafe impl<O: Borrow<PyObject>> Traverse for PyIter<O> {
18    fn traverse(&self, tracer_fn: &mut TraverseFn<'_>) {
19        // Report the iterator itself, not its referents: an owner holding a
20        // `PyIter` owns the iterator object, and reporting what the iterator
21        // points at instead leaves the iterator's own reference unaccounted
22        // for, so a cycle running through it is never collected.
23        tracer_fn(self.0.borrow());
24    }
25}
26
27impl PyIter<PyObjectRef> {
28    pub fn check(obj: &PyObject) -> bool {
29        obj.class().slots().iternext.load().is_some()
30    }
31}
32
33impl<O> PyIter<O>
34where
35    O: Borrow<PyObject>,
36{
37    #[must_use]
38    pub const fn new(obj: O) -> Self {
39        Self(obj)
40    }
41
42    pub fn next(&self, vm: &VirtualMachine) -> PyResult<PyIterReturn> {
43        let iternext = self
44            .0
45            .borrow()
46            .class()
47            .slots
48            .iternext
49            .load()
50            .ok_or_else(|| {
51                vm.new_type_error(format!(
52                    "'{}' object is not an iterator",
53                    self.0.borrow().class().slot_name()
54                ))
55            })?;
56        iternext(self.0.borrow(), vm)
57    }
58
59    /// Walks the iterator without asking it how long it is. Almost nothing
60    /// asks: a loop over an iterator takes no room up front, so what the
61    /// object would have answered -- slowly, or by raising -- never runs.
62    pub fn iter<'a, 'b, U>(
63        &'b self,
64        vm: &'a VirtualMachine,
65    ) -> PyResult<PyIterIter<'a, U, &'b PyObject>> {
66        Ok(PyIterIter::new(vm, self.0.borrow(), None))
67    }
68}
69
70impl PyIter<PyObjectRef> {
71    /// Returns an iterator over this sequence of objects. See [`Self::iter`]
72    /// for why it does not ask how long the iterator is.
73    pub fn into_iter<U>(self, vm: &VirtualMachine) -> PyIterIter<'_, U, PyObjectRef> {
74        PyIterIter::new(vm, self.0, None)
75    }
76
77    /// [`Self::into_iter`] for a caller that fills a sized container from the
78    /// iterator, the way `PySequence_Fast()` does. It asks how much room that
79    /// takes and answers with whatever asking raised.
80    pub fn into_iter_sized<U>(
81        self,
82        vm: &VirtualMachine,
83    ) -> PyResult<PyIterIter<'_, U, PyObjectRef>> {
84        let length_hint = vm.length_hint_opt(self.as_object().to_owned())?;
85        Ok(PyIterIter::new(vm, self.0, length_hint))
86    }
87}
88
89impl From<PyIter<Self>> for PyObjectRef {
90    fn from(value: PyIter<Self>) -> Self {
91        value.0
92    }
93}
94
95impl<O> Borrow<PyObject> for PyIter<O>
96where
97    O: Borrow<PyObject>,
98{
99    #[inline(always)]
100    fn borrow(&self) -> &PyObject {
101        self.0.borrow()
102    }
103}
104
105impl<O> AsRef<PyObject> for PyIter<O>
106where
107    O: Borrow<PyObject>,
108{
109    #[inline(always)]
110    fn as_ref(&self) -> &PyObject {
111        self.0.borrow()
112    }
113}
114
115impl<O> Deref for PyIter<O>
116where
117    O: Borrow<PyObject>,
118{
119    type Target = PyObject;
120
121    #[inline(always)]
122    fn deref(&self) -> &Self::Target {
123        self.0.borrow()
124    }
125}
126
127impl ToPyObject for PyIter<PyObjectRef> {
128    #[inline(always)]
129    fn to_pyobject(self, _vm: &VirtualMachine) -> PyObjectRef {
130        self.into()
131    }
132}
133
134impl TryFromObject for PyIter<PyObjectRef> {
135    // This helper function is called at multiple places. First, it is called
136    // in the vm when a for loop is entered. Next, it is used when the builtin
137    // function 'iter' is called.
138    fn try_from_object(vm: &VirtualMachine, iter_target: PyObjectRef) -> PyResult<Self> {
139        let get_iter = iter_target.class().slots().iter.load();
140        if let Some(get_iter) = get_iter {
141            let iter = get_iter(iter_target, vm)?;
142            if Self::check(&iter) {
143                Ok(Self(iter))
144            } else {
145                Err(vm.new_type_error(format!(
146                    "iter() returned non-iterator of type '{}'",
147                    iter.class().slot_name()
148                )))
149            }
150        } else if let Ok(seq_iter) = PySequenceIterator::new(iter_target.clone(), vm) {
151            Ok(Self(seq_iter.into_pyobject(vm)))
152        } else {
153            Err(vm.new_type_error(format!(
154                "'{}' object is not iterable",
155                iter_target.class().slot_name()
156            )))
157        }
158    }
159}
160
161#[derive(result_like::ResultLike)]
162pub enum PyIterReturn<T = PyObjectRef> {
163    Return(T),
164    StopIteration(Option<PyObjectRef>),
165}
166
167unsafe impl<T: Traverse> Traverse for PyIterReturn<T> {
168    fn traverse(&self, tracer_fn: &mut TraverseFn<'_>) {
169        match self {
170            Self::Return(r) => r.traverse(tracer_fn),
171            Self::StopIteration(Some(obj)) => obj.traverse(tracer_fn),
172            _ => (),
173        }
174    }
175}
176
177impl PyIterReturn {
178    pub fn from_pyresult(result: PyResult, vm: &VirtualMachine) -> PyResult<Self> {
179        match result {
180            Ok(obj) => Ok(Self::Return(obj)),
181            Err(err) if err.fast_isinstance(vm.ctx.exceptions.stop_iteration) => {
182                let args = err.get_arg(0);
183                Ok(Self::StopIteration(args))
184            }
185            Err(err) => Err(err),
186        }
187    }
188
189    pub fn from_getitem_result(result: PyResult, vm: &VirtualMachine) -> PyResult<Self> {
190        match result {
191            Ok(obj) => Ok(Self::Return(obj)),
192            Err(err) if err.fast_isinstance(vm.ctx.exceptions.index_error) => {
193                Ok(Self::StopIteration(None))
194            }
195            Err(err) if err.fast_isinstance(vm.ctx.exceptions.stop_iteration) => {
196                let args = err.get_arg(0);
197                Ok(Self::StopIteration(args))
198            }
199            Err(err) => Err(err),
200        }
201    }
202
203    pub fn into_async_pyresult(self, vm: &VirtualMachine) -> PyResult {
204        match self {
205            Self::Return(obj) => Ok(obj),
206            Self::StopIteration(v) => Err({
207                let args = v.map_or_else(Vec::new, |v| vec![v]);
208                vm.new_exception(vm.ctx.exceptions.stop_async_iteration.to_owned(), args)
209            }),
210        }
211    }
212}
213
214impl ToPyResult for PyIterReturn {
215    fn to_pyresult(self, vm: &VirtualMachine) -> PyResult {
216        match self {
217            Self::Return(obj) => Ok(obj),
218            Self::StopIteration(v) => Err(vm.new_stop_iteration(v)),
219        }
220    }
221}
222
223impl ToPyResult for PyResult<PyIterReturn> {
224    fn to_pyresult(self, vm: &VirtualMachine) -> PyResult {
225        self?.to_pyresult(vm)
226    }
227}
228
229// Typical rust `Iter` object for `PyIter`
230pub struct PyIterIter<'a, T, O = PyObjectRef>
231where
232    O: Borrow<PyObject>,
233{
234    vm: &'a VirtualMachine,
235    obj: O, // creating PyIter<O> is zero-cost
236    length_hint: Option<usize>,
237    _phantom: core::marker::PhantomData<T>,
238}
239
240unsafe impl<T, O> Traverse for PyIterIter<'_, T, O>
241where
242    O: Traverse + Borrow<PyObject>,
243{
244    fn traverse(&self, tracer_fn: &mut TraverseFn<'_>) {
245        self.obj.traverse(tracer_fn)
246    }
247}
248
249impl<'a, T, O> PyIterIter<'a, T, O>
250where
251    O: Borrow<PyObject>,
252{
253    pub const fn new(vm: &'a VirtualMachine, obj: O, length_hint: Option<usize>) -> Self {
254        Self {
255            vm,
256            obj,
257            length_hint,
258            _phantom: core::marker::PhantomData,
259        }
260    }
261}
262
263impl<T, O> Iterator for PyIterIter<'_, T, O>
264where
265    T: TryFromObject,
266    O: Borrow<PyObject>,
267{
268    type Item = PyResult<T>;
269
270    fn next(&mut self) -> Option<Self::Item> {
271        let imp = |next: PyResult<PyIterReturn>| -> PyResult<Option<T>> {
272            let Some(obj) = next?.into_result().ok() else {
273                return Ok(None);
274            };
275            Ok(Some(T::try_from_object(self.vm, obj)?))
276        };
277        let next = PyIter::new(self.obj.borrow()).next(self.vm);
278        imp(next).transpose()
279    }
280
281    #[inline]
282    fn size_hint(&self) -> (usize, Option<usize>) {
283        (self.length_hint.unwrap_or(0), self.length_hint)
284    }
285}
286
287/// Macro to handle `PyIterReturn` values in iterator implementations.
288///
289/// Extracts the object from `PyIterReturn::Return(obj)` or performs early return
290/// for `PyIterReturn::StopIteration(v)`. This macro should only be used within
291/// functions that return `PyResult<PyIterReturn>`.
292#[macro_export]
293macro_rules! raise_if_stop {
294    ($input:expr) => {
295        match $input {
296            $crate::protocol::PyIterReturn::Return(obj) => obj,
297            $crate::protocol::PyIterReturn::StopIteration(v) => {
298                return Ok($crate::protocol::PyIterReturn::StopIteration(v))
299            }
300        }
301    };
302}