Skip to main content

rustpython_vm/builtins/
mappingproxy.rs

1use super::{PyDict, PyDictRef, PyGenericAlias, PyList, PyTuple, PyType, PyTypeRef};
2use crate::{
3    AsObject, Context, Py, PyObject, PyObjectRef, PyPayload, PyRef, PyResult, VirtualMachine,
4    atomic_func,
5    class::PyClassImpl,
6    common::{hash, lock::LazyLock},
7    convert::ToPyObject,
8    function::{ArgMapping, OptionalArg, PyArithmeticValue, PyComparisonValue},
9    object::{Traverse, TraverseFn},
10    protocol::{PyMappingMethods, PyNumberMethods, PySequenceMethods},
11    types::{
12        AsMapping, AsNumber, AsSequence, Comparable, Constructor, Hashable, Iterable,
13        PyComparisonOp, Representable,
14    },
15};
16use rustpython_common::wtf8::{Wtf8Buf, wtf8_concat};
17
18#[pyclass(module = false, name = "mappingproxy", traverse)]
19#[derive(Debug)]
20pub struct PyMappingProxy {
21    mapping: MappingProxyInner,
22}
23
24#[derive(Debug)]
25enum MappingProxyInner {
26    Class(PyTypeRef),
27    Mapping(ArgMapping),
28}
29
30unsafe impl Traverse for MappingProxyInner {
31    fn traverse(&self, tracer_fn: &mut TraverseFn<'_>) {
32        match self {
33            Self::Class(r) => r.traverse(tracer_fn),
34            Self::Mapping(arg) => arg.traverse(tracer_fn),
35        }
36    }
37}
38
39impl PyPayload for PyMappingProxy {
40    #[inline]
41    fn class(ctx: &Context) -> &'static Py<PyType> {
42        ctx.types.mappingproxy_type
43    }
44}
45
46impl From<PyTypeRef> for PyMappingProxy {
47    fn from(dict: PyTypeRef) -> Self {
48        Self {
49            mapping: MappingProxyInner::Class(dict),
50        }
51    }
52}
53
54impl From<PyDictRef> for PyMappingProxy {
55    fn from(dict: PyDictRef) -> Self {
56        Self {
57            mapping: MappingProxyInner::Mapping(ArgMapping::from_dict_exact(dict)),
58        }
59    }
60}
61
62impl Constructor for PyMappingProxy {
63    type Args = PyObjectRef;
64
65    fn py_new(_cls: &Py<PyType>, mapping: Self::Args, vm: &VirtualMachine) -> PyResult<Self> {
66        Self::from_object(mapping, vm)
67    }
68}
69
70impl PyMappingProxy {
71    pub fn from_object(mapping: PyObjectRef, vm: &VirtualMachine) -> PyResult<Self> {
72        if mapping.mapping_unchecked().check()
73            && !mapping.downcastable::<PyList>()
74            && !mapping.downcastable::<PyTuple>()
75        {
76            return Ok(Self {
77                mapping: MappingProxyInner::Mapping(ArgMapping::new(mapping)),
78            });
79        }
80        Err(vm.new_type_error(format!(
81            "mappingproxy() argument must be a mapping, not {}",
82            mapping.class()
83        )))
84    }
85
86    fn get_inner(&self, key: &PyObject, vm: &VirtualMachine) -> PyResult<Option<PyObjectRef>> {
87        match &self.mapping {
88            MappingProxyInner::Class(class) => Self::class_get(class, key, vm),
89            MappingProxyInner::Mapping(mapping) => mapping.mapping().subscript(key, vm).map(Some),
90        }
91    }
92
93    fn class_get(
94        class: &Py<PyType>,
95        key: &PyObject,
96        vm: &VirtualMachine,
97    ) -> PyResult<Option<PyObjectRef>> {
98        match class.attributes.as_dict() {
99            Some(dict) => dict.get_item_opt(key, vm),
100            None => Ok(key
101                .as_interned_str(vm)
102                .and_then(|key| class.attributes.get(key))),
103        }
104    }
105
106    pub fn __getitem__(&self, key: PyObjectRef, vm: &VirtualMachine) -> PyResult {
107        self.get_inner(&key, vm)?
108            .ok_or_else(|| vm.new_key_error(key))
109    }
110
111    fn _contains(&self, key: &PyObject, vm: &VirtualMachine) -> PyResult<bool> {
112        match &self.mapping {
113            MappingProxyInner::Class(class) => Ok(Self::class_contains(class, key, vm)),
114            MappingProxyInner::Mapping(mapping) => {
115                mapping.obj().sequence_unchecked().contains(key, vm)
116            }
117        }
118    }
119
120    fn class_contains(class: &Py<PyType>, key: &PyObject, vm: &VirtualMachine) -> bool {
121        match class.attributes.as_dict() {
122            Some(dict) => dict.contains_key(key, vm),
123            None => key
124                .as_interned_str(vm)
125                .is_some_and(|key| class.attributes.contains(key)),
126        }
127    }
128
129    pub fn __contains__(&self, key: &PyObject, vm: &VirtualMachine) -> PyResult<bool> {
130        self._contains(key, vm)
131    }
132
133    fn to_object(&self, vm: &VirtualMachine) -> PyResult {
134        Ok(match &self.mapping {
135            MappingProxyInner::Mapping(d) => d.as_ref().to_owned(),
136            MappingProxyInner::Class(c) => Self::class_to_dict(c, vm)?,
137        })
138    }
139
140    fn class_to_dict(class: &Py<PyType>, vm: &VirtualMachine) -> PyResult {
141        if let Some(dict) = class.attributes.as_dict() {
142            return Ok(dict.copy().to_pyobject(vm));
143        }
144        Ok(PyDict::from_attributes(class.attributes.attributes(&vm.ctx), vm)?.to_pyobject(vm))
145    }
146
147    fn __len__(&self, vm: &VirtualMachine) -> PyResult<usize> {
148        let obj = self.to_object(vm)?;
149        obj.length(vm)
150    }
151
152    fn __ior__(&self, _args: PyObjectRef, vm: &VirtualMachine) -> PyResult {
153        Err(vm.new_type_error(format!(
154            r#""'|=' is not supported by {}; use '|' instead""#,
155            Self::class(&vm.ctx)
156        )))
157    }
158
159    fn __or__(&self, args: &PyObject, vm: &VirtualMachine) -> PyResult {
160        vm._or(self.copy(vm)?.as_ref(), args)
161    }
162
163    pub fn copy(&self, vm: &VirtualMachine) -> PyResult {
164        match &self.mapping {
165            MappingProxyInner::Mapping(d) => {
166                vm.call_method(d.obj(), identifier!(vm, copy).as_str(), ())
167            }
168            MappingProxyInner::Class(c) => Self::class_to_dict(c, vm),
169        }
170    }
171}
172
173#[pyclass(with(
174    AsMapping,
175    Iterable,
176    Constructor,
177    AsSequence,
178    Comparable,
179    Hashable,
180    AsNumber,
181    Representable
182))]
183impl Py<PyMappingProxy> {
184    #[pymethod]
185    fn get(
186        &self,
187        key: PyObjectRef,
188        default: OptionalArg,
189        vm: &VirtualMachine,
190    ) -> PyResult<Option<PyObjectRef>> {
191        let obj = self.to_object(vm)?;
192        Ok(Some(vm.call_method(
193            &obj,
194            "get",
195            (key, default.unwrap_or_none(vm)),
196        )?))
197    }
198
199    #[pymethod]
200    pub fn items(&self, vm: &VirtualMachine) -> PyResult {
201        let obj = self.to_object(vm)?;
202        vm.call_method(&obj, identifier!(vm, items).as_str(), ())
203    }
204
205    #[pymethod]
206    pub fn keys(&self, vm: &VirtualMachine) -> PyResult {
207        let obj = self.to_object(vm)?;
208        vm.call_method(&obj, identifier!(vm, keys).as_str(), ())
209    }
210
211    #[pymethod]
212    pub fn values(&self, vm: &VirtualMachine) -> PyResult {
213        let obj = self.to_object(vm)?;
214        vm.call_method(&obj, identifier!(vm, values).as_str(), ())
215    }
216
217    #[pymethod]
218    pub fn copy(&self, vm: &VirtualMachine) -> PyResult {
219        self.payload.copy(vm)
220    }
221
222    #[pyclassmethod]
223    fn __class_getitem__(
224        cls: PyTypeRef,
225        args: PyObjectRef,
226        vm: &VirtualMachine,
227    ) -> PyResult<PyGenericAlias> {
228        PyGenericAlias::from_args(cls, args, vm)
229    }
230
231    #[pymethod]
232    fn __reversed__(&self, vm: &VirtualMachine) -> PyResult {
233        vm.call_method(
234            self.to_object(vm)?.as_object(),
235            identifier!(vm, __reversed__).as_str(),
236            (),
237        )
238    }
239}
240
241impl Comparable for PyMappingProxy {
242    fn cmp(
243        zelf: &Py<Self>,
244        other: &PyObject,
245        op: PyComparisonOp,
246        vm: &VirtualMachine,
247    ) -> PyResult<PyComparisonValue> {
248        let obj = zelf.to_object(vm)?;
249        // CPython parity (Objects/descrobject.c::mappingproxy_richcompare):
250        // delegate to PyObject_RichCompare on the underlying mapping.
251        let res = obj.rich_compare(other.to_owned(), op, vm)?;
252        PyArithmeticValue::from_object(vm, res)
253            .map(|o| o.try_to_bool(vm))
254            .transpose()
255    }
256}
257
258impl Hashable for PyMappingProxy {
259    #[inline]
260    fn hash(zelf: &Py<Self>, vm: &VirtualMachine) -> PyResult<hash::PyHash> {
261        // Delegate hash to the underlying mapping
262        let obj = zelf.to_object(vm)?;
263        obj.hash(vm)
264    }
265}
266
267impl AsMapping for PyMappingProxy {
268    fn as_mapping() -> &'static PyMappingMethods {
269        static AS_MAPPING: LazyLock<PyMappingMethods> = LazyLock::new(|| PyMappingMethods {
270            length: atomic_func!(
271                |mapping, vm| PyMappingProxy::mapping_downcast(mapping).__len__(vm)
272            ),
273            subscript: atomic_func!(|mapping, needle, vm| {
274                PyMappingProxy::mapping_downcast(mapping).__getitem__(needle.to_owned(), vm)
275            }),
276            ..PyMappingMethods::NOT_IMPLEMENTED
277        });
278        &AS_MAPPING
279    }
280}
281
282impl AsSequence for PyMappingProxy {
283    fn as_sequence() -> &'static PySequenceMethods {
284        static AS_SEQUENCE: LazyLock<PySequenceMethods> = LazyLock::new(|| PySequenceMethods {
285            length: atomic_func!(|seq, vm| PyMappingProxy::sequence_downcast(seq).__len__(vm)),
286            contains: atomic_func!(
287                |seq, target, vm| PyMappingProxy::sequence_downcast(seq)._contains(target, vm)
288            ),
289            ..PySequenceMethods::NOT_IMPLEMENTED
290        });
291        &AS_SEQUENCE
292    }
293}
294
295impl AsNumber for PyMappingProxy {
296    fn as_number() -> &'static PyNumberMethods {
297        static AS_NUMBER: PyNumberMethods = PyNumberMethods {
298            or: Some(|a, b, vm| {
299                // Mirror CPython's mappingproxy_or: when either side is a
300                // mappingproxy, unwrap to its underlying mapping and delegate
301                // to PyNumber_Or so `dict | mp`, `mp | dict`, and `mp | mp`
302                // all produce a `dict` result.
303                let a_obj = match a.downcast_ref::<PyMappingProxy>() {
304                    Some(mp) => mp.copy(vm)?,
305                    None => a.to_pyobject(vm),
306                };
307                let b_obj = match b.downcast_ref::<PyMappingProxy>() {
308                    Some(mp) => mp.copy(vm)?,
309                    None => b.to_pyobject(vm),
310                };
311                vm._or(a_obj.as_ref(), b_obj.as_ref())
312            }),
313            inplace_or: Some(|a, b, vm| {
314                if let Some(a) = a.downcast_ref::<PyMappingProxy>() {
315                    a.__ior__(b.to_pyobject(vm), vm)
316                } else {
317                    Ok(vm.ctx.not_implemented())
318                }
319            }),
320            ..PyNumberMethods::NOT_IMPLEMENTED
321        };
322        &AS_NUMBER
323    }
324}
325
326impl Iterable for PyMappingProxy {
327    fn iter(zelf: PyRef<Self>, vm: &VirtualMachine) -> PyResult {
328        let obj = zelf.to_object(vm)?;
329        let iter = obj.get_iter(vm)?;
330        Ok(iter.into())
331    }
332}
333
334impl Representable for PyMappingProxy {
335    #[inline]
336    fn repr_wtf8(zelf: &Py<Self>, vm: &VirtualMachine) -> PyResult<Wtf8Buf> {
337        let obj = zelf.to_object(vm)?;
338        Ok(wtf8_concat!("mappingproxy(", obj.repr(vm)?.as_wtf8(), ')'))
339    }
340}
341
342pub(crate) fn init(context: &'static Context) {
343    PyMappingProxy::extend_class(context, context.types.mappingproxy_type)
344}