Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 0 additions & 1 deletion Lib/test/test_dict.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)])
Expand Down
90 changes: 75 additions & 15 deletions crates/vm/src/builtins/dict.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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},
Expand Down Expand Up @@ -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::<PyList>(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::<PyTuple>(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::<PyObjectRef>(vm)?
.collect::<PyResult<Vec<_>>>()
})()
.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,
Expand All @@ -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::<PyObjectRef>(vm)?.enumerate() {
let (key, value) = Self::update_sequence_pair(element?, index, vm)?;

if !override_existing && dict.contains(vm, &*key)? {
continue;
}
Expand Down
2 changes: 1 addition & 1 deletion crates/vm/src/exceptions.rs
Original file line number Diff line number Diff line change
Expand Up @@ -704,7 +704,7 @@ impl PyRef<PyBaseException> {

let notes = notes
.downcast::<PyList>()
.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(())
Expand Down
Loading