Skip to main content

rustpython_vm/builtins/
map.rs

1use super::{PyType, PyTypeRef};
2use crate::{
3    AsObject, Context, Py, PyObjectRef, PyPayload, PyRef, PyResult, TryFromObject, VirtualMachine,
4    builtins::PyTupleRef,
5    class::PyClassImpl,
6    function::{ArgIntoBool, FuncArgs, PosArgs},
7    protocol::{PyIter, PyIterReturn},
8    types::{Constructor, IterNext, Iterable, SelfIter},
9};
10use rustpython_common::atomic::{self, PyAtomic, Radium};
11
12#[pyclass(module = false, name = "map", traverse)]
13#[derive(Debug)]
14pub struct PyMap {
15    mapper: PyObjectRef,
16    iterators: Vec<PyIter>,
17    #[pytraverse(skip)]
18    strict: PyAtomic<bool>,
19}
20
21impl PyPayload for PyMap {
22    #[inline]
23    fn class(ctx: &Context) -> &'static Py<PyType> {
24        ctx.types.map_type
25    }
26}
27
28#[derive(FromArgs)]
29pub struct PyMapNewArgs {
30    #[pyarg(positional)]
31    function: PyObjectRef,
32    #[pyarg(positional)]
33    iterable: PyIter,
34    #[pyarg(flatten)]
35    iterables: PosArgs<PyIter, crate::function::NameIterables>,
36    #[pyarg(named, default)]
37    strict: bool,
38}
39
40#[derive(FromArgs)]
41struct MapCallArgs {
42    #[pyarg(positional)]
43    function: PyObjectRef,
44    #[pyarg(flatten)]
45    iterables: PosArgs<PyIter, crate::function::NameIterables>,
46    #[pyarg(named, default)]
47    strict: bool,
48}
49
50impl Constructor for PyMap {
51    type Args = PyMapNewArgs;
52
53    fn slot_new(cls: PyTypeRef, args: FuncArgs, vm: &VirtualMachine) -> PyResult {
54        let MapCallArgs {
55            function: mapper,
56            iterables,
57            strict,
58        } = args.bind_for(vm, "map")?;
59        let iterators = iterables.into_vec();
60        if iterators.is_empty() {
61            return Err(vm.new_type_error("map() must have at least two arguments."));
62        }
63        let payload = Self {
64            mapper,
65            iterators,
66            strict: Radium::new(strict),
67        };
68        payload.into_ref_with_type(vm, cls).map(Into::into)
69    }
70
71    fn py_new(_cls: &Py<PyType>, args: Self::Args, vm: &VirtualMachine) -> PyResult<Self> {
72        let PyMapNewArgs {
73            function,
74            iterable,
75            iterables,
76            strict,
77        } = args;
78        let _ = (function, iterable, iterables, strict);
79        Err(vm.new_type_error("use slot_new"))
80    }
81}
82
83#[pyclass(with(IterNext, Iterable, Constructor), flags(BASETYPE))]
84impl PyMap {
85    #[pymethod]
86    fn __reduce__(zelf: PyRef<Self>, vm: &VirtualMachine) -> PyTupleRef {
87        let cls = zelf.class().to_owned();
88        let mut vec = vec![zelf.mapper.clone()];
89        vec.extend(zelf.iterators.iter().map(|o| o.clone().into()));
90        let tuple_args = vm.ctx.new_tuple(vec);
91        if zelf.strict.load(atomic::Ordering::Acquire) {
92            vm.new_tuple((cls, tuple_args, true))
93        } else {
94            vm.new_tuple((cls, tuple_args))
95        }
96    }
97
98    #[pymethod]
99    fn __setstate__(zelf: PyRef<Self>, object: PyObjectRef, vm: &VirtualMachine) {
100        if let Ok(obj) = ArgIntoBool::try_from_object(vm, object) {
101            zelf.strict.store(obj.into(), atomic::Ordering::Release);
102        }
103    }
104}
105
106impl SelfIter for PyMap {}
107
108impl IterNext for PyMap {
109    fn next(zelf: &Py<Self>, vm: &VirtualMachine) -> PyResult<PyIterReturn> {
110        let mut next_objs = Vec::new();
111        for (idx, iterator) in zelf.iterators.iter().enumerate() {
112            let item = match iterator.next(vm)? {
113                PyIterReturn::Return(obj) => obj,
114                PyIterReturn::StopIteration(v) => {
115                    if zelf.strict.load(atomic::Ordering::Acquire) {
116                        if idx > 0 {
117                            let plural = if idx == 1 { " " } else { "s 1-" };
118                            return Err(vm.new_value_error(format!(
119                                "map() argument {} is shorter than argument{}{}",
120                                idx + 1,
121                                plural,
122                                idx,
123                            )));
124                        }
125                        for (idx, iterator) in zelf.iterators[1..].iter().enumerate() {
126                            if let PyIterReturn::Return(_) = iterator.next(vm)? {
127                                let plural = if idx == 0 { " " } else { "s 1-" };
128                                return Err(vm.new_value_error(format!(
129                                    "map() argument {} is longer than argument{}{}",
130                                    idx + 2,
131                                    plural,
132                                    idx + 1,
133                                )));
134                            }
135                        }
136                    }
137                    return Ok(PyIterReturn::StopIteration(v));
138                }
139            };
140            next_objs.push(item);
141        }
142
143        // the mapper itself can raise StopIteration which does stop the map iteration
144        PyIterReturn::from_pyresult(zelf.mapper.call(next_objs, vm), vm)
145    }
146}
147
148pub(crate) fn init(context: &'static Context) {
149    PyMap::extend_class(context, context.types.map_type);
150}