1use 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#[derive(Debug, Clone)]
18pub enum IterStatus<T> {
19 Active(T),
21 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 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 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 Err(_) => false,
113 };
114 let released = if done { self.exhaust() } else { None };
115 (ret, released)
116 }
117
118 #[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 #[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 #[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
182pub(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
197pub 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
208pub 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 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 {
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}