Skip to main content

rustpython_vm/builtins/
filter.rs

1use super::{PyType, PyTypeRef};
2use crate::{
3    Context, Py, PyObjectRef, PyPayload, PyResult, VirtualMachine,
4    class::PyClassImpl,
5    protocol::{PyIter, PyIterReturn},
6    raise_if_stop,
7    types::{Constructor, IterNext, Iterable, SelfIter},
8};
9
10#[pyclass(module = false, name = "filter", traverse)]
11#[derive(Debug)]
12pub struct PyFilter {
13    predicate: PyObjectRef,
14    iterator: PyIter,
15}
16
17impl PyPayload for PyFilter {
18    #[inline]
19    fn class(ctx: &Context) -> &'static Py<PyType> {
20        ctx.types.filter_type
21    }
22}
23
24#[derive(FromArgs)]
25pub struct FilterArgs {
26    #[pyarg(positional)]
27    function: PyObjectRef,
28    #[pyarg(positional)]
29    iterable: PyIter,
30}
31
32impl Constructor for PyFilter {
33    type Args = FilterArgs;
34    const DROP_KWARGS_WHEN_INIT_OVERRIDDEN: bool = true;
35
36    fn py_new(
37        _cls: &Py<PyType>,
38        Self::Args {
39            function,
40            iterable: iterator,
41        }: Self::Args,
42        _vm: &VirtualMachine,
43    ) -> PyResult<Self> {
44        Ok(Self {
45            predicate: function,
46            iterator,
47        })
48    }
49}
50
51#[pyclass(with(IterNext, Iterable, Constructor), flags(BASETYPE))]
52impl Py<PyFilter> {
53    #[pymethod]
54    fn __reduce__(&self, vm: &VirtualMachine) -> (PyTypeRef, (PyObjectRef, PyIter)) {
55        (
56            vm.ctx.types.filter_type.to_owned(),
57            (self.predicate.clone(), self.iterator.clone()),
58        )
59    }
60}
61
62impl SelfIter for PyFilter {}
63
64impl IterNext for PyFilter {
65    fn next(zelf: &Py<Self>, vm: &VirtualMachine) -> PyResult<PyIterReturn> {
66        let predicate = &zelf.predicate;
67        loop {
68            let next_obj = raise_if_stop!(zelf.iterator.next(vm)?);
69            let predicate_value = if vm.is_none(predicate) {
70                next_obj.clone()
71            } else {
72                // the predicate itself can raise StopIteration which does stop the filter iteration
73                raise_if_stop!(PyIterReturn::from_pyresult(
74                    predicate.call((next_obj.clone(),), vm),
75                    vm
76                )?)
77            };
78
79            if predicate_value.try_to_bool(vm)? {
80                return Ok(PyIterReturn::Return(next_obj));
81            }
82        }
83    }
84}
85
86pub(crate) fn init(context: &'static Context) {
87    PyFilter::extend_class(context, context.types.filter_type);
88}