diff --git a/Lib/test/test_dict.py b/Lib/test/test_dict.py index 79c975946f7..e2a73773cc2 100644 --- a/Lib/test/test_dict.py +++ b/Lib/test/test_dict.py @@ -275,7 +275,6 @@ def __next__(self): self.assertRaises(ValueError, {}.update, [(1, 2, 3)]) - @unittest.expectedFailure # TODO: RUSTPYTHON def test_update_type_error(self): with self.assertRaises(TypeError) as cm: {}.update([object() for _ in range(3)]) diff --git a/crates/vm/src/builtins/dict.rs b/crates/vm/src/builtins/dict.rs index af74a259157..5db071e1d8f 100644 --- a/crates/vm/src/builtins/dict.rs +++ b/crates/vm/src/builtins/dict.rs @@ -7,7 +7,7 @@ use crate::object::{Traverse, TraverseFn}; use crate::{ AsObject, Context, Py, PyExact, PyObject, PyObjectRef, PyPayload, PyRef, PyRefExact, PyResult, TryFromObject, atomic_func, - builtins::{PyTuple, iter::builtins_iter, type_::PyAttributes}, + builtins::{PyList, PyTuple, iter::builtins_iter, type_::PyAttributes}, class::{PyClassDef, PyClassImpl}, common::ascii, dict_inner::{self, DictKey}, @@ -183,6 +183,76 @@ impl PyDict { self.merge_object_with_override(other, false, vm) } + fn add_update_sequence_note( + exc: PyBaseExceptionRef, + index: usize, + vm: &VirtualMachine, + ) -> PyBaseExceptionRef { + if !exc.fast_isinstance(vm.ctx.exceptions.type_error) { + return exc; + } + + let note = + format!("Cannot convert dictionary update sequence element #{index} to a sequence"); + match vm.call_method(exc.as_object(), "add_note", (vm.ctx.new_str(note),)) { + Ok(_) => exc, + Err(note_err) => { + note_err.set___context__(Some(exc)); + note_err + } + } + } + + fn update_sequence_pair_from_slice( + elements: &[PyObjectRef], + index: usize, + vm: &VirtualMachine, + ) -> PyResult<(PyObjectRef, PyObjectRef)> { + let [key, value] = elements else { + return Err(vm.new_value_error(format!( + "dictionary update sequence element #{index} has length {}; 2 is required", + elements.len() + ))); + }; + Ok((key.clone(), value.clone())) + } + + fn update_sequence_pair( + element: PyObjectRef, + index: usize, + vm: &VirtualMachine, + ) -> PyResult<(PyObjectRef, PyObjectRef)> { + let element = match element.downcast_exact::(vm) { + Ok(list) => { + let elements = list.borrow_vec(); + return Self::update_sequence_pair_from_slice(&elements, index, vm); + } + Err(element) => element, + }; + let element = match element.downcast_exact::(vm) { + Ok(tuple) => { + return Self::update_sequence_pair_from_slice(tuple.as_slice(), index, vm); + } + Err(element) => element, + }; + + let elements = (|| { + let elem_iter = element.get_iter(vm).map_err(|exc| { + if exc.fast_isinstance(vm.ctx.exceptions.type_error) { + vm.new_type_error("object is not iterable") + } else { + exc + } + })?; + elem_iter + .into_iter::(vm)? + .collect::>>() + })() + .map_err(|exc| Self::add_update_sequence_note(exc, index, vm))?; + + Self::update_sequence_pair_from_slice(&elements, index, vm) + } + pub fn merge_from_seq2( &self, seq2: PyObjectRef, @@ -191,20 +261,10 @@ impl PyDict { ) -> PyResult<()> { let iter = seq2.get_iter(vm)?; let dict = &self.entries; - loop { - fn err(vm: &VirtualMachine) -> PyBaseExceptionRef { - vm.new_value_error("Iterator must have exactly two elements") - } - let element = match iter.next(vm)? { - PyIterReturn::Return(obj) => obj, - PyIterReturn::StopIteration(_) => break, - }; - let elem_iter = element.get_iter(vm)?; - let key = elem_iter.next(vm)?.into_result().map_err(|_| err(vm))?; - let value = elem_iter.next(vm)?.into_result().map_err(|_| err(vm))?; - if matches!(elem_iter.next(vm)?, PyIterReturn::Return(_)) { - return Err(err(vm)); - } + + for (index, element) in iter.iter_without_hint::(vm)?.enumerate() { + let (key, value) = Self::update_sequence_pair(element?, index, vm)?; + if !override_existing && dict.contains(vm, &*key)? { continue; } diff --git a/crates/vm/src/exceptions.rs b/crates/vm/src/exceptions.rs index fe07a7e3c9e..7eaf8bafcdd 100644 --- a/crates/vm/src/exceptions.rs +++ b/crates/vm/src/exceptions.rs @@ -704,7 +704,7 @@ impl PyRef { let notes = notes .downcast::() - .map_err(|_| vm.new_type_error("__notes__ must be a list"))?; + .map_err(|_| vm.new_type_error("Cannot add note: __notes__ is not a list"))?; notes.borrow_vec_mut().push(note.into()); Ok(())