FazBrowse GitHub Viewer | Trending |
URL:
| Home
Tools: [Download Repo ZIP]   [Original HTTPS Page]

Fix array compare oparation · RustPython/RustPython@4890932 · GitHub

Repository navigation

Commit 4890932

Browse files
committed
Fix array compare oparation
1 parent a5d0df2 commit 4890932

3 files changed

Lines changed: 58 additions & 68 deletions

File tree

‎Lib/test/test_array.py‎

Lines changed: 0 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -509,8 +509,6 @@ def test_str(self):
509509
a = array.array(self.typecode, 2*self.example)
510510
str(a)
511511

512-
# TODO: RUSTPYTHON
513-
@unittest.expectedFailure
514512
def test_cmp(self):
515513
a = array.array(self.typecode, self.example)
516514
self.assertIs(a == 42, False)

‎tests/snippets/stdlib_array.py‎

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -54,3 +54,15 @@ def test_float_with_integer_input():
5454
assert a == array(T, [1, 2, 3, 4, 5, 6, 7, 8])
5555
del a[0:0:-9999999999]
5656
assert a == array(T, [1, 2, 3, 4, 5, 6, 7, 8])
57+
58+
def test_float_with_nan():
59+
f = float('nan')
60+
a = array('f')
61+
a.append(f)
62+
assert not (a == a)
63+
assert a != a
64+
assert not (a < a)
65+
assert not (a <= a)
66+
assert not (a > a)
67+
assert not (a >= a)
68+
test_float_with_nan()

‎vm/src/stdlib/array.rs‎

Lines changed: 46 additions & 66 deletions
Original file line numberDiff line numberDiff line change
@@ -5,11 +5,11 @@ use crate::common::cell::{
55
use crate::function::OptionalArg;
66
use crate::obj::objbytes::PyBytesRef;
77
use crate::obj::objfloat::try_float;
8+
use crate::obj::objiter;
89
use crate::obj::objsequence::{PySliceableSequence, PySliceableSequenceMut};
910
use crate::obj::objslice::PySliceRef;
1011
use crate::obj::objstr::PyStringRef;
1112
use crate::obj::objtype::PyClassRef;
12-
use crate::obj::{objbool, objiter};
1313
use crate::pyobject::{
1414
BorrowValue, Either, IdProtocol, IntoPyObject, PyArithmaticValue, PyClassImpl,
1515
PyComparisonValue, PyIterable, PyObjectRef, PyRef, PyResult, PyValue, TryFromObject,
@@ -19,6 +19,7 @@ use crate::VirtualMachine;
1919
use crossbeam_utils::atomic::AtomicCell;
2020
use itertools::Itertools;
2121
use std::fmt;
22+
use PyArithmaticValue::Implemented;
2223

2324
struct ArrayTypeSpecifierError {
2425
_priv: (),
@@ -362,7 +363,7 @@ trait ArrayElement: Sized {
362363
}
363364

364365
macro_rules! adapt_try_into_from_object {
365-
($(($t:ty, $f: path),)*) => {$(
366+
($(($t:ty, $f:path),)*) => {$(
366367
impl ArrayElement for $t {
367368
fn try_into_from_object(vm: &VirtualMachine, obj: PyObjectRef) -> PyResult<Self> {
368369
$f(vm, obj)
@@ -651,22 +652,45 @@ impl PyArray {
651652
zelf.borrow_value().repr(vm)
652653
}
653654

655+
fn cmp<L, O>(
656+
&self,
657+
other: PyArrayRef,
658+
len_cmp: L,
659+
obj_cmp: O,
660+
vm: &VirtualMachine,
661+
) -> PyResult<PyComparisonValue>
662+
where
663+
L: Fn(usize, usize) -> bool,
664+
O: Fn(PyObjectRef, PyObjectRef) -> PyResult<Option<bool>>,
665+
{
666+
let array_a = self.borrow_value();
667+
let array_b = other.borrow_value();
668+
let iter = Iterator::zip(array_a.iter(vm), array_b.iter(vm));
669+
for (a, b) in iter {
670+
if let Some(v) = obj_cmp(a, b)? {
671+
return Ok(Implemented(v));
672+
}
673+
}
674+
Ok(Implemented(len_cmp(self.len(), other.len())))
675+
}
676+
654677
#[pymethod(name = "__eq__")]
655678
fn eq(&self, other: PyObjectRef, vm: &VirtualMachine) -> PyResult<PyComparisonValue> {
656-
let lhs = self.borrow_value();
657-
let rhs = class_or_notimplemented!(vm, Self, other);
658-
let rhs = rhs.borrow_value();
659-
if lhs.len() != rhs.len() {
660-
Ok(PyArithmaticValue::Implemented(false))
661-
} else {
662-
for (a, b) in lhs.iter(vm).zip(rhs.iter(vm)) {
663-
let ne = objbool::boolval(vm, vm._ne(a, b)?)?;
664-
if ne {
665-
return Ok(PyArithmaticValue::Implemented(false));
666-
}
679+
// we cannot use zelf.is(other) for shortcut because if we contenting a
680+
// float value NaN we always return False even they are the same object.
681+
let other = class_or_notimplemented!(vm, Self, other);
682+
if self.len() != other.len() {
683+
return Ok(Implemented(false));
684+
}
685+
let array_a = self.borrow_value();
686+
let array_b = other.borrow_value();
687+
let iter = Iterator::zip(array_a.iter(vm), array_b.iter(vm));
688+
for (a, b) in iter {
689+
if !vm.bool_eq(a, b)? {
690+
return Ok(Implemented(false));
667691
}
668-
Ok(PyArithmaticValue::Implemented(true))
669692
}
693+
Ok(Implemented(true))
670694
}
671695

672696
#[pymethod(name = "__ne__")]
@@ -676,70 +700,26 @@ impl PyArray {
676700

677701
#[pymethod(name = "__lt__")]
678702
fn lt(&self, other: PyObjectRef, vm: &VirtualMachine) -> PyResult<PyComparisonValue> {
679-
let lhs = self.borrow_value();
680-
let rhs = class_or_notimplemented!(vm, Self, other);
681-
let rhs = rhs.borrow_value();
682-
683-
for (a, b) in lhs.iter(vm).zip(rhs.iter(vm)) {
684-
let lt = objbool::boolval(vm, vm._lt(a, b)?)?;
685-
686-
if lt {
687-
return Ok(PyArithmaticValue::Implemented(true));
688-
}
689-
}
690-
691-
Ok(PyArithmaticValue::Implemented(lhs.len() < rhs.len()))
703+
let other = class_or_notimplemented!(vm, Self, other);
704+
self.cmp(other, |a, b| a < b, |a, b| vm.bool_seq_lt(a, b), vm)
692705
}
693706

694707
#[pymethod(name = "__le__")]
695708
fn le(&self, other: PyObjectRef, vm: &VirtualMachine) -> PyResult<PyComparisonValue> {
696-
let lhs = self.borrow_value();
697-
let rhs = class_or_notimplemented!(vm, Self, other);
698-
let rhs = rhs.borrow_value();
699-
700-
for (a, b) in lhs.iter(vm).zip(rhs.iter(vm)) {
701-
let le = objbool::boolval(vm, vm._le(a, b)?)?;
702-
703-
if le {
704-
return Ok(PyArithmaticValue::Implemented(true));
705-
}
706-
}
707-
708-
Ok(PyArithmaticValue::Implemented(lhs.len() <= rhs.len()))
709+
let other = class_or_notimplemented!(vm, Self, other);
710+
self.cmp(other, |a, b| a <= b, |a, b| vm.bool_seq_lt(a, b), vm)
709711
}
710712

711713
#[pymethod(name = "__gt__")]
712714
fn gt(&self, other: PyObjectRef, vm: &VirtualMachine) -> PyResult<PyComparisonValue> {
713-
let lhs = self.borrow_value();
714-
let rhs = class_or_notimplemented!(vm, Self, other);
715-
let rhs = rhs.borrow_value();
716-
717-
for (a, b) in lhs.iter(vm).zip(rhs.iter(vm)) {
718-
let gt = objbool::boolval(vm, vm._gt(a, b)?)?;
719-
720-
if gt {
721-
return Ok(PyArithmaticValue::Implemented(true));
722-
}
723-
}
724-
725-
Ok(PyArithmaticValue::Implemented(lhs.len() > rhs.len()))
715+
let other = class_or_notimplemented!(vm, Self, other);
716+
self.cmp(other, |a, b| a > b, |a, b| vm.bool_seq_gt(a, b), vm)
726717
}
727718

728719
#[pymethod(name = "__ge__")]
729720
fn ge(&self, other: PyObjectRef, vm: &VirtualMachine) -> PyResult<PyComparisonValue> {
730-
let lhs = self.borrow_value();
731-
let rhs = class_or_notimplemented!(vm, Self, other);
732-
let rhs = rhs.borrow_value();
733-
734-
for (a, b) in lhs.iter(vm).zip(rhs.iter(vm)) {
735-
let ge = objbool::boolval(vm, vm._ge(a, b)?)?;
736-
737-
if ge {
738-
return Ok(PyArithmaticValue::Implemented(true));
739-
}
740-
}
741-
742-
Ok(PyArithmaticValue::Implemented(lhs.len() >= rhs.len()))
721+
let other = class_or_notimplemented!(vm, Self, other);
722+
self.cmp(other, |a, b| a >= b, |a, b| vm.bool_seq_gt(a, b), vm)
743723
}
744724

745725
#[pymethod(name = "__len__")]

0 commit comments

Comments
 (0)

Back | FazBrowse Home | New Git URL