Skip to main content

rustpython_vm/protocol/
mapping.rs

1use crossbeam_utils::atomic::AtomicCell;
2
3use crate::{
4    AsObject, PyObject, PyObjectRef, PyResult, VirtualMachine,
5    builtins::{
6        PyDict, PyStrInterned,
7        dict::{PyDictItems, PyDictKeys, PyDictValues},
8    },
9    convert::ToPyResult,
10    object::{Traverse, TraverseFn},
11};
12
13/// [Mapping protocol](https://docs.python.org/3/c-api/mapping.html)
14#[expect(clippy::type_complexity)]
15#[derive(Default)]
16pub struct PyMappingSlots {
17    pub length: AtomicCell<Option<fn(PyMapping<'_>, &VirtualMachine) -> PyResult<usize>>>,
18    pub subscript: AtomicCell<Option<fn(PyMapping<'_>, &PyObject, &VirtualMachine) -> PyResult>>,
19    pub ass_subscript: AtomicCell<
20        Option<fn(PyMapping<'_>, &PyObject, Option<PyObjectRef>, &VirtualMachine) -> PyResult<()>>,
21    >,
22}
23
24impl core::fmt::Debug for PyMappingSlots {
25    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
26        f.write_str("PyMappingSlots")
27    }
28}
29
30impl PyMappingSlots {
31    #[must_use]
32    pub fn has_subscript(&self) -> bool {
33        self.subscript.load().is_some()
34    }
35
36    /// Copy from static [`PyMappingMethods`].
37    pub fn copy_from(&self, methods: &PyMappingMethods) {
38        if let Some(f) = methods.length {
39            self.length.store(Some(f));
40        }
41
42        if let Some(f) = methods.subscript {
43            self.subscript.store(Some(f));
44        }
45
46        if let Some(f) = methods.ass_subscript {
47            self.ass_subscript.store(Some(f));
48        }
49    }
50}
51
52#[expect(clippy::type_complexity)]
53#[derive(Default)]
54pub struct PyMappingMethods {
55    pub length: Option<fn(PyMapping<'_>, &VirtualMachine) -> PyResult<usize>>,
56    pub subscript: Option<fn(PyMapping<'_>, &PyObject, &VirtualMachine) -> PyResult>,
57    pub ass_subscript:
58        Option<fn(PyMapping<'_>, &PyObject, Option<PyObjectRef>, &VirtualMachine) -> PyResult<()>>,
59}
60
61impl core::fmt::Debug for PyMappingMethods {
62    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
63        f.write_str("PyMappingMethods")
64    }
65}
66
67impl PyMappingMethods {
68    pub const NOT_IMPLEMENTED: Self = Self {
69        length: None,
70        subscript: None,
71        ass_subscript: None,
72    };
73}
74
75impl PyObject {
76    #[must_use]
77    pub const fn mapping_unchecked(&self) -> PyMapping<'_> {
78        PyMapping { obj: self }
79    }
80
81    pub fn try_mapping(&self, vm: &VirtualMachine) -> PyResult<PyMapping<'_>> {
82        let mapping = self.mapping_unchecked();
83        if mapping.check() {
84            Ok(mapping)
85        } else {
86            Err(vm.new_type_error(format!(
87                "{} is not a mapping object",
88                self.class().slot_name()
89            )))
90        }
91    }
92}
93
94#[derive(Copy, Clone)]
95pub struct PyMapping<'a> {
96    pub obj: &'a PyObject,
97}
98
99unsafe impl Traverse for PyMapping<'_> {
100    fn traverse(&self, tracer_fn: &mut TraverseFn<'_>) {
101        self.obj.traverse(tracer_fn)
102    }
103}
104
105impl AsRef<PyObject> for PyMapping<'_> {
106    #[inline(always)]
107    fn as_ref(&self) -> &PyObject {
108        self.obj
109    }
110}
111
112impl PyMapping<'_> {
113    #[inline]
114    #[must_use]
115    pub fn slots(&self) -> &PyMappingSlots {
116        &self.obj.class().slots().as_mapping
117    }
118
119    #[inline]
120    #[must_use]
121    pub fn check(&self) -> bool {
122        self.slots().has_subscript()
123    }
124
125    pub fn length_opt(self, vm: &VirtualMachine) -> Option<PyResult<usize>> {
126        self.slots().length.load().map(|f| f(self, vm))
127    }
128
129    // Py_ssize_t PyMapping_Size(PyObject *o)
130    pub fn length(self, vm: &VirtualMachine) -> PyResult<usize> {
131        self.length_opt(vm).ok_or_else(|| {
132            let name = self.obj.class().slot_name();
133            // Something that measures itself as a sequence is no mapping at all.
134            let msg = if self
135                .obj
136                .sequence_unchecked()
137                .slots()
138                .length
139                .load()
140                .is_some()
141            {
142                format!("{name} is not a mapping")
143            } else {
144                format!("object of type '{name}' has no len()")
145            };
146            vm.new_type_error(msg)
147        })?
148    }
149
150    pub fn subscript(self, needle: &impl AsObject, vm: &VirtualMachine) -> PyResult {
151        self._subscript(needle.as_object(), vm)
152    }
153
154    pub fn ass_subscript(
155        self,
156        needle: &impl AsObject,
157        value: Option<PyObjectRef>,
158        vm: &VirtualMachine,
159    ) -> PyResult<()> {
160        self._ass_subscript(needle.as_object(), value, vm)
161    }
162
163    fn _subscript(self, needle: &PyObject, vm: &VirtualMachine) -> PyResult {
164        let f = self.slots().subscript.load().ok_or_else(|| {
165            vm.new_type_error(format!("{} is not a mapping", self.obj.class().slot_name()))
166        })?;
167        f(self, needle, vm)
168    }
169
170    fn _ass_subscript(
171        self,
172        needle: &PyObject,
173        value: Option<PyObjectRef>,
174        vm: &VirtualMachine,
175    ) -> PyResult<()> {
176        let f = self.slots().ass_subscript.load().ok_or_else(|| {
177            vm.new_type_error(format!(
178                "'{}' object does not support item assignment",
179                self.obj.class().slot_name()
180            ))
181        })?;
182        f(self, needle, value, vm)
183    }
184
185    pub fn keys(self, vm: &VirtualMachine) -> PyResult {
186        if let Some(dict) = self.obj.downcast_ref_if_exact::<PyDict>(vm) {
187            PyDictKeys::new(dict.to_owned()).to_pyresult(vm)
188        } else {
189            self.method_output_as_list(identifier!(vm, keys), vm)
190        }
191    }
192
193    pub fn values(self, vm: &VirtualMachine) -> PyResult {
194        if let Some(dict) = self.obj.downcast_ref_if_exact::<PyDict>(vm) {
195            PyDictValues::new(dict.to_owned()).to_pyresult(vm)
196        } else {
197            self.method_output_as_list(identifier!(vm, values), vm)
198        }
199    }
200
201    pub fn items(self, vm: &VirtualMachine) -> PyResult {
202        if let Some(dict) = self.obj.downcast_ref_if_exact::<PyDict>(vm) {
203            PyDictItems::new(dict.to_owned()).to_pyresult(vm)
204        } else {
205            self.method_output_as_list(identifier!(vm, items), vm)
206        }
207    }
208
209    fn method_output_as_list(
210        self,
211        method_name: &'static PyStrInterned,
212        vm: &VirtualMachine,
213    ) -> PyResult {
214        let meth_output = vm.call_method(self.obj, method_name.as_str(), ())?;
215        if meth_output.is(vm.ctx.types.list_type) {
216            return Ok(meth_output);
217        }
218
219        let iter = meth_output.get_iter(vm).map_err(|_| {
220            vm.new_type_error(format!(
221                "{}.{}() returned a non-iterable (type {})",
222                self.obj.class().slot_name(),
223                method_name.as_str(),
224                meth_output.class().slot_name()
225            ))
226        })?;
227
228        // TODO
229        // PySequence::from(&iter).list(vm).map(|x| x.into())
230        vm.ctx.new_list(iter.try_to_value(vm)?).to_pyresult(vm)
231    }
232}