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#[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}