use std::any::type_name;
use std::cell::UnsafeCell;
use std::collections::hash_map::Entry;
use std::fmt;
use std::mem::size_of;
use std::ops::Deref;
use ahash::AHashMap;
use ndarray::{
ArrayView, ArrayViewMut, Dimension, IntoDimension, Ix1, Ix2, Ix3, Ix4, Ix5, Ix6, IxDyn,
};
use num_integer::gcd;
use pyo3::{FromPyObject, PyAny, PyResult, Python};
use crate::array::PyArray;
use crate::cold;
use crate::convert::NpyIndex;
use crate::dtype::Element;
use crate::error::{BorrowError, NotContiguousError};
use crate::npyffi::{self, PyArrayObject, NPY_ARRAY_WRITEABLE};
#[derive(Clone, Copy, PartialEq, Eq, Hash)]
struct BorrowKey {
range: (*mut u8, *mut u8),
data_ptr: *mut u8,
gcd_strides: isize,
}
impl BorrowKey {
fn from_array<T, D>(array: &PyArray<T, D>) -> Self
where
T: Element,
D: Dimension,
{
let range = data_range(array);
let data_ptr = array.data() as *mut u8;
let gcd_strides = gcd_strides(array.strides());
Self {
range,
data_ptr,
gcd_strides,
}
}
fn conflicts(&self, other: &Self) -> bool {
debug_assert!(self.range.0 <= self.range.1);
debug_assert!(other.range.0 <= other.range.1);
if other.range.0 >= self.range.1 || self.range.0 >= other.range.1 {
return false;
}
let ptr_diff = unsafe { self.data_ptr.offset_from(other.data_ptr).abs() };
let gcd_strides = gcd(self.gcd_strides, other.gcd_strides);
if ptr_diff % gcd_strides != 0 {
return false;
}
true
}
}
type BorrowFlagsInner = AHashMap<*mut u8, AHashMap<BorrowKey, isize>>;
struct BorrowFlags(UnsafeCell<Option<BorrowFlagsInner>>);
unsafe impl Sync for BorrowFlags {}
impl BorrowFlags {
const fn new() -> Self {
Self(UnsafeCell::new(None))
}
#[allow(clippy::mut_from_ref)]
unsafe fn get(&self) -> &mut BorrowFlagsInner {
(*self.0.get()).get_or_insert_with(AHashMap::new)
}
fn acquire(&self, _py: Python, address: *mut u8, key: BorrowKey) -> Result<(), BorrowError> {
let borrow_flags = unsafe { BORROW_FLAGS.get() };
match borrow_flags.entry(address) {
Entry::Occupied(entry) => {
let same_base_arrays = entry.into_mut();
if let Some(readers) = same_base_arrays.get_mut(&key) {
assert_ne!(*readers, 0);
let new_readers = readers.wrapping_add(1);
if new_readers <= 0 {
cold();
return Err(BorrowError::AlreadyBorrowed);
}
*readers = new_readers;
} else {
if same_base_arrays
.iter()
.any(|(other, readers)| key.conflicts(other) && *readers < 0)
{
cold();
return Err(BorrowError::AlreadyBorrowed);
}
same_base_arrays.insert(key, 1);
}
}
Entry::Vacant(entry) => {
let mut same_base_arrays = AHashMap::with_capacity(1);
same_base_arrays.insert(key, 1);
entry.insert(same_base_arrays);
}
}
Ok(())
}
fn release(&self, _py: Python, address: *mut u8, key: BorrowKey) {
let borrow_flags = unsafe { BORROW_FLAGS.get() };
let same_base_arrays = borrow_flags.get_mut(&address).unwrap();
let readers = same_base_arrays.get_mut(&key).unwrap();
*readers -= 1;
if *readers == 0 {
if same_base_arrays.len() > 1 {
same_base_arrays.remove(&key).unwrap();
} else {
borrow_flags.remove(&address).unwrap();
}
}
}
fn acquire_mut(
&self,
_py: Python,
address: *mut u8,
key: BorrowKey,
) -> Result<(), BorrowError> {
let borrow_flags = unsafe { BORROW_FLAGS.get() };
match borrow_flags.entry(address) {
Entry::Occupied(entry) => {
let same_base_arrays = entry.into_mut();
if let Some(writers) = same_base_arrays.get_mut(&key) {
assert_ne!(*writers, 0);
cold();
return Err(BorrowError::AlreadyBorrowed);
} else {
if same_base_arrays
.iter()
.any(|(other, writers)| key.conflicts(other) && *writers != 0)
{
cold();
return Err(BorrowError::AlreadyBorrowed);
}
same_base_arrays.insert(key, -1);
}
}
Entry::Vacant(entry) => {
let mut same_base_arrays = AHashMap::with_capacity(1);
same_base_arrays.insert(key, -1);
entry.insert(same_base_arrays);
}
}
Ok(())
}
fn release_mut(&self, _py: Python, address: *mut u8, key: BorrowKey) {
let borrow_flags = unsafe { BORROW_FLAGS.get() };
let same_base_arrays = borrow_flags.get_mut(&address).unwrap();
if same_base_arrays.len() > 1 {
same_base_arrays.remove(&key).unwrap();
} else {
borrow_flags.remove(&address);
}
}
}
static BORROW_FLAGS: BorrowFlags = BorrowFlags::new();
#[repr(C)]
pub struct PyReadonlyArray<'py, T, D>
where
T: Element,
D: Dimension,
{
array: &'py PyArray<T, D>,
address: *mut u8,
key: BorrowKey,
}
pub type PyReadonlyArray1<'py, T> = PyReadonlyArray<'py, T, Ix1>;
pub type PyReadonlyArray2<'py, T> = PyReadonlyArray<'py, T, Ix2>;
pub type PyReadonlyArray3<'py, T> = PyReadonlyArray<'py, T, Ix3>;
pub type PyReadonlyArray4<'py, T> = PyReadonlyArray<'py, T, Ix4>;
pub type PyReadonlyArray5<'py, T> = PyReadonlyArray<'py, T, Ix5>;
pub type PyReadonlyArray6<'py, T> = PyReadonlyArray<'py, T, Ix6>;
pub type PyReadonlyArrayDyn<'py, T> = PyReadonlyArray<'py, T, IxDyn>;
impl<'py, T, D> Deref for PyReadonlyArray<'py, T, D>
where
T: Element,
D: Dimension,
{
type Target = PyArray<T, D>;
fn deref(&self) -> &Self::Target {
self.array
}
}
impl<'py, T: Element, D: Dimension> FromPyObject<'py> for PyReadonlyArray<'py, T, D> {
fn extract(obj: &'py PyAny) -> PyResult<Self> {
let array: &'py PyArray<T, D> = obj.extract()?;
Ok(array.readonly())
}
}
impl<'py, T, D> PyReadonlyArray<'py, T, D>
where
T: Element,
D: Dimension,
{
pub(crate) fn try_new(array: &'py PyArray<T, D>) -> Result<Self, BorrowError> {
let address = base_address(array);
let key = BorrowKey::from_array(array);
BORROW_FLAGS.acquire(array.py(), address, key)?;
Ok(Self {
array,
address,
key,
})
}
#[inline(always)]
pub fn as_array(&self) -> ArrayView<T, D> {
unsafe { self.array.as_array() }
}
#[inline(always)]
pub fn as_slice(&self) -> Result<&[T], NotContiguousError> {
unsafe { self.array.as_slice() }
}
#[inline(always)]
pub fn get<I>(&self, index: I) -> Option<&T>
where
I: NpyIndex<Dim = D>,
{
unsafe { self.array.get(index) }
}
}
impl<'a, T, D> Clone for PyReadonlyArray<'a, T, D>
where
T: Element,
D: Dimension,
{
fn clone(&self) -> Self {
BORROW_FLAGS
.acquire(self.array.py(), self.address, self.key)
.unwrap();
Self {
array: self.array,
address: self.address,
key: self.key,
}
}
}
impl<'a, T, D> Drop for PyReadonlyArray<'a, T, D>
where
T: Element,
D: Dimension,
{
fn drop(&mut self) {
BORROW_FLAGS.release(self.array.py(), self.address, self.key);
}
}
impl<'py, T, D> fmt::Debug for PyReadonlyArray<'py, T, D>
where
T: Element,
D: Dimension,
{
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let name = format!(
"PyReadonlyArray<{}, {}>",
type_name::<T>(),
type_name::<D>()
);
f.debug_struct(&name).finish()
}
}
#[repr(C)]
pub struct PyReadwriteArray<'py, T, D>
where
T: Element,
D: Dimension,
{
array: &'py PyArray<T, D>,
address: *mut u8,
key: BorrowKey,
}
pub type PyReadwriteArray1<'py, T> = PyReadwriteArray<'py, T, Ix1>;
pub type PyReadwriteArray2<'py, T> = PyReadwriteArray<'py, T, Ix2>;
pub type PyReadwriteArray3<'py, T> = PyReadwriteArray<'py, T, Ix3>;
pub type PyReadwriteArray4<'py, T> = PyReadwriteArray<'py, T, Ix4>;
pub type PyReadwriteArray5<'py, T> = PyReadwriteArray<'py, T, Ix5>;
pub type PyReadwriteArray6<'py, T> = PyReadwriteArray<'py, T, Ix6>;
pub type PyReadwriteArrayDyn<'py, T> = PyReadwriteArray<'py, T, IxDyn>;
impl<'py, T, D> Deref for PyReadwriteArray<'py, T, D>
where
T: Element,
D: Dimension,
{
type Target = PyReadonlyArray<'py, T, D>;
fn deref(&self) -> &Self::Target {
unsafe { &*(self as *const Self as *const Self::Target) }
}
}
impl<'py, T: Element, D: Dimension> FromPyObject<'py> for PyReadwriteArray<'py, T, D> {
fn extract(obj: &'py PyAny) -> PyResult<Self> {
let array: &'py PyArray<T, D> = obj.extract()?;
Ok(array.readwrite())
}
}
impl<'py, T, D> PyReadwriteArray<'py, T, D>
where
T: Element,
D: Dimension,
{
pub(crate) fn try_new(array: &'py PyArray<T, D>) -> Result<Self, BorrowError> {
if !array.check_flags(NPY_ARRAY_WRITEABLE) {
return Err(BorrowError::NotWriteable);
}
let address = base_address(array);
let key = BorrowKey::from_array(array);
BORROW_FLAGS.acquire_mut(array.py(), address, key)?;
Ok(Self {
array,
address,
key,
})
}
#[inline(always)]
pub fn as_array_mut(&mut self) -> ArrayViewMut<T, D> {
unsafe { self.array.as_array_mut() }
}
#[inline(always)]
pub fn as_slice_mut(&mut self) -> Result<&mut [T], NotContiguousError> {
unsafe { self.array.as_slice_mut() }
}
#[inline(always)]
pub fn get_mut<I>(&mut self, index: I) -> Option<&mut T>
where
I: NpyIndex<Dim = D>,
{
unsafe { self.array.get_mut(index) }
}
}
impl<'py, T> PyReadwriteArray<'py, T, Ix1>
where
T: Element,
{
pub fn resize<ID: IntoDimension>(self, dims: ID) -> PyResult<Self> {
let array = self.array;
unsafe {
array.resize(dims)?;
}
drop(self);
Ok(Self::try_new(array).unwrap())
}
}
impl<'a, T, D> Drop for PyReadwriteArray<'a, T, D>
where
T: Element,
D: Dimension,
{
fn drop(&mut self) {
BORROW_FLAGS.release_mut(self.array.py(), self.address, self.key);
}
}
impl<'py, T, D> fmt::Debug for PyReadwriteArray<'py, T, D>
where
T: Element,
D: Dimension,
{
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let name = format!(
"PyReadwriteArray<{}, {}>",
type_name::<T>(),
type_name::<D>()
);
f.debug_struct(&name).finish()
}
}
fn base_address<T, D>(array: &PyArray<T, D>) -> *mut u8 {
fn inner(py: Python, mut array: *mut PyArrayObject) -> *mut u8 {
loop {
let base = unsafe { (*array).base };
if base.is_null() {
return array as *mut u8;
} else if unsafe { npyffi::PyArray_Check(py, base) } != 0 {
array = base as *mut PyArrayObject;
} else {
return base as *mut u8;
}
}
}
inner(array.py(), array.as_array_ptr())
}
fn data_range<T, D>(array: &PyArray<T, D>) -> (*mut u8, *mut u8)
where
T: Element,
D: Dimension,
{
fn inner(
shape: &[usize],
strides: &[isize],
itemsize: isize,
data: *mut u8,
) -> (*mut u8, *mut u8) {
let mut start = 0;
let mut end = 0;
if shape.iter().all(|dim| *dim != 0) {
for (&dim, &stride) in shape.iter().zip(strides) {
let offset = (dim - 1) as isize * stride;
if offset >= 0 {
end += offset;
} else {
start += offset;
}
}
end += itemsize;
}
let start = unsafe { data.offset(start) };
let end = unsafe { data.offset(end) };
(start, end)
}
inner(
array.shape(),
array.strides(),
size_of::<T>() as isize,
array.data() as *mut u8,
)
}
fn gcd_strides(strides: &[isize]) -> isize {
reduce(strides.iter().copied(), gcd).unwrap_or(1)
}
fn reduce<I, F>(mut iter: I, f: F) -> Option<I::Item>
where
I: Iterator,
F: FnMut(I::Item, I::Item) -> I::Item,
{
let first = iter.next()?;
Some(iter.fold(first, f))
}
#[cfg(test)]
mod tests {
use super::*;
use ndarray::Array;
use pyo3::{types::IntoPyDict, Python};
use crate::array::{PyArray1, PyArray2, PyArray3};
use crate::convert::IntoPyArray;
#[test]
fn without_base_object() {
Python::with_gil(|py| {
let array = PyArray::<f64, _>::zeros(py, (1, 2, 3), false);
let base = unsafe { (*array.as_array_ptr()).base };
assert!(base.is_null());
let base_address = base_address(array);
assert_eq!(base_address, array as *const _ as *mut u8);
let data_range = data_range(array);
assert_eq!(data_range.0, array.data() as *mut u8);
assert_eq!(data_range.1, unsafe { array.data().add(6) } as *mut u8);
});
}
#[test]
fn with_base_object() {
Python::with_gil(|py| {
let array = Array::<f64, _>::zeros((1, 2, 3)).into_pyarray(py);
let base = unsafe { (*array.as_array_ptr()).base };
assert!(!base.is_null());
let base_address = base_address(array);
assert_ne!(base_address, array as *const _ as *mut u8);
assert_eq!(base_address, base as *mut u8);
let data_range = data_range(array);
assert_eq!(data_range.0, array.data() as *mut u8);
assert_eq!(data_range.1, unsafe { array.data().add(6) } as *mut u8);
});
}
#[test]
fn view_without_base_object() {
Python::with_gil(|py| {
let array = PyArray::<f64, _>::zeros(py, (1, 2, 3), false);
let locals = [("array", array)].into_py_dict(py);
let view = py
.eval("array[:,:,0]", None, Some(locals))
.unwrap()
.downcast::<PyArray2<f64>>()
.unwrap();
assert_ne!(view as *const _ as *mut u8, array as *const _ as *mut u8);
let base = unsafe { (*view.as_array_ptr()).base };
assert_eq!(base as *mut u8, array as *const _ as *mut u8);
let base_address = base_address(view);
assert_ne!(base_address, view as *const _ as *mut u8);
assert_eq!(base_address, base as *mut u8);
let data_range = data_range(view);
assert_eq!(data_range.0, array.data() as *mut u8);
assert_eq!(data_range.1, unsafe { array.data().add(4) } as *mut u8);
});
}
#[test]
fn view_with_base_object() {
Python::with_gil(|py| {
let array = Array::<f64, _>::zeros((1, 2, 3)).into_pyarray(py);
let locals = [("array", array)].into_py_dict(py);
let view = py
.eval("array[:,:,0]", None, Some(locals))
.unwrap()
.downcast::<PyArray2<f64>>()
.unwrap();
assert_ne!(view as *const _ as *mut u8, array as *const _ as *mut u8);
let base = unsafe { (*view.as_array_ptr()).base };
assert_eq!(base as *mut u8, array as *const _ as *mut u8);
let base = unsafe { (*array.as_array_ptr()).base };
assert!(!base.is_null());
let base_address = base_address(view);
assert_ne!(base_address, view as *const _ as *mut u8);
assert_ne!(base_address, array as *const _ as *mut u8);
assert_eq!(base_address, base as *mut u8);
let data_range = data_range(view);
assert_eq!(data_range.0, array.data() as *mut u8);
assert_eq!(data_range.1, unsafe { array.data().add(4) } as *mut u8);
});
}
#[test]
fn view_of_view_without_base_object() {
Python::with_gil(|py| {
let array = PyArray::<f64, _>::zeros(py, (1, 2, 3), false);
let locals = [("array", array)].into_py_dict(py);
let view1 = py
.eval("array[:,:,0]", None, Some(locals))
.unwrap()
.downcast::<PyArray2<f64>>()
.unwrap();
assert_ne!(view1 as *const _ as *mut u8, array as *const _ as *mut u8);
let locals = [("view1", view1)].into_py_dict(py);
let view2 = py
.eval("view1[:,0]", None, Some(locals))
.unwrap()
.downcast::<PyArray1<f64>>()
.unwrap();
assert_ne!(view2 as *const _ as *mut u8, array as *const _ as *mut u8);
assert_ne!(view2 as *const _ as *mut u8, view1 as *const _ as *mut u8);
let base = unsafe { (*view2.as_array_ptr()).base };
assert_eq!(base as *mut u8, array as *const _ as *mut u8);
let base = unsafe { (*view1.as_array_ptr()).base };
assert_eq!(base as *mut u8, array as *const _ as *mut u8);
let base_address = base_address(view2);
assert_ne!(base_address, view2 as *const _ as *mut u8);
assert_ne!(base_address, view1 as *const _ as *mut u8);
assert_eq!(base_address, base as *mut u8);
let data_range = data_range(view2);
assert_eq!(data_range.0, array.data() as *mut u8);
assert_eq!(data_range.1, unsafe { array.data().add(1) } as *mut u8);
});
}
#[test]
fn view_of_view_with_base_object() {
Python::with_gil(|py| {
let array = Array::<f64, _>::zeros((1, 2, 3)).into_pyarray(py);
let locals = [("array", array)].into_py_dict(py);
let view1 = py
.eval("array[:,:,0]", None, Some(locals))
.unwrap()
.downcast::<PyArray2<f64>>()
.unwrap();
assert_ne!(view1 as *const _ as *mut u8, array as *const _ as *mut u8);
let locals = [("view1", view1)].into_py_dict(py);
let view2 = py
.eval("view1[:,0]", None, Some(locals))
.unwrap()
.downcast::<PyArray1<f64>>()
.unwrap();
assert_ne!(view2 as *const _ as *mut u8, array as *const _ as *mut u8);
assert_ne!(view2 as *const _ as *mut u8, view1 as *const _ as *mut u8);
let base = unsafe { (*view2.as_array_ptr()).base };
assert_eq!(base as *mut u8, array as *const _ as *mut u8);
let base = unsafe { (*view1.as_array_ptr()).base };
assert_eq!(base as *mut u8, array as *const _ as *mut u8);
let base = unsafe { (*array.as_array_ptr()).base };
assert!(!base.is_null());
let base_address = base_address(view2);
assert_ne!(base_address, view2 as *const _ as *mut u8);
assert_ne!(base_address, view1 as *const _ as *mut u8);
assert_ne!(base_address, array as *const _ as *mut u8);
assert_eq!(base_address, base as *mut u8);
let data_range = data_range(view2);
assert_eq!(data_range.0, array.data() as *mut u8);
assert_eq!(data_range.1, unsafe { array.data().add(1) } as *mut u8);
});
}
#[test]
fn view_with_negative_strides() {
Python::with_gil(|py| {
let array = PyArray::<f64, _>::zeros(py, (1, 2, 3), false);
let locals = [("array", array)].into_py_dict(py);
let view = py
.eval("array[::-1,:,::-1]", None, Some(locals))
.unwrap()
.downcast::<PyArray3<f64>>()
.unwrap();
assert_ne!(view as *const _ as *mut u8, array as *const _ as *mut u8);
let base = unsafe { (*view.as_array_ptr()).base };
assert_eq!(base as *mut u8, array as *const _ as *mut u8);
let base_address = base_address(view);
assert_ne!(base_address, view as *const _ as *mut u8);
assert_eq!(base_address, base as *mut u8);
let data_range = data_range(view);
assert_eq!(view.data(), unsafe { array.data().offset(2) });
assert_eq!(data_range.0, unsafe { view.data().offset(-2) } as *mut u8);
assert_eq!(data_range.1, unsafe { view.data().offset(4) } as *mut u8);
});
}
#[test]
fn array_with_zero_dimensions() {
Python::with_gil(|py| {
let array = PyArray::<f64, _>::zeros(py, (1, 0, 3), false);
let base = unsafe { (*array.as_array_ptr()).base };
assert!(base.is_null());
let base_address = base_address(array);
assert_eq!(base_address, array as *const _ as *mut u8);
let data_range = data_range(array);
assert_eq!(data_range.0, array.data() as *mut u8);
assert_eq!(data_range.1, array.data() as *mut u8);
});
}
#[test]
fn view_with_non_dividing_strides() {
Python::with_gil(|py| {
let array = PyArray::<f64, _>::zeros(py, (10, 10), false);
let locals = [("array", array)].into_py_dict(py);
let view1 = py
.eval("array[:,::3]", None, Some(locals))
.unwrap()
.downcast::<PyArray2<f64>>()
.unwrap();
let key1 = BorrowKey::from_array(view1);
assert_eq!(view1.strides(), &[80, 24]);
assert_eq!(key1.gcd_strides, 8);
let view2 = py
.eval("array[:,1::3]", None, Some(locals))
.unwrap()
.downcast::<PyArray2<f64>>()
.unwrap();
let key2 = BorrowKey::from_array(view2);
assert_eq!(view2.strides(), &[80, 24]);
assert_eq!(key2.gcd_strides, 8);
let view3 = py
.eval("array[:,::2]", None, Some(locals))
.unwrap()
.downcast::<PyArray2<f64>>()
.unwrap();
let key3 = BorrowKey::from_array(view3);
assert_eq!(view3.strides(), &[80, 16]);
assert_eq!(key3.gcd_strides, 16);
let view4 = py
.eval("array[:,1::2]", None, Some(locals))
.unwrap()
.downcast::<PyArray2<f64>>()
.unwrap();
let key4 = BorrowKey::from_array(view4);
assert_eq!(view4.strides(), &[80, 16]);
assert_eq!(key4.gcd_strides, 16);
assert!(!key3.conflicts(&key4));
assert!(key1.conflicts(&key3));
assert!(key2.conflicts(&key4));
assert!(key1.conflicts(&key2));
});
}
#[test]
fn borrow_multiple_arrays() {
Python::with_gil(|py| {
let array1 = PyArray::<f64, _>::zeros(py, 10, false);
let array2 = PyArray::<f64, _>::zeros(py, 10, false);
let base1 = base_address(array1);
let base2 = base_address(array2);
let key1 = BorrowKey::from_array(array1);
let _exclusive1 = array1.readwrite();
{
let borrow_flags = unsafe { BORROW_FLAGS.get() };
assert_eq!(borrow_flags.len(), 1);
let same_base_arrays = &borrow_flags[&base1];
assert_eq!(same_base_arrays.len(), 1);
let flag = same_base_arrays[&key1];
assert_eq!(flag, -1);
}
let key2 = BorrowKey::from_array(array2);
let _shared2 = array2.readonly();
{
let borrow_flags = unsafe { BORROW_FLAGS.get() };
assert_eq!(borrow_flags.len(), 2);
let same_base_arrays = &borrow_flags[&base1];
assert_eq!(same_base_arrays.len(), 1);
let flag = same_base_arrays[&key1];
assert_eq!(flag, -1);
let same_base_arrays = &borrow_flags[&base2];
assert_eq!(same_base_arrays.len(), 1);
let flag = same_base_arrays[&key2];
assert_eq!(flag, 1);
}
});
}
#[test]
fn borrow_multiple_views() {
Python::with_gil(|py| {
let array = PyArray::<f64, _>::zeros(py, 10, false);
let base = base_address(array);
let locals = [("array", array)].into_py_dict(py);
let view1 = py
.eval("array[:5]", None, Some(locals))
.unwrap()
.downcast::<PyArray1<f64>>()
.unwrap();
let key1 = BorrowKey::from_array(view1);
let exclusive1 = view1.readwrite();
{
let borrow_flags = unsafe { BORROW_FLAGS.get() };
assert_eq!(borrow_flags.len(), 1);
let same_base_arrays = &borrow_flags[&base];
assert_eq!(same_base_arrays.len(), 1);
let flag = same_base_arrays[&key1];
assert_eq!(flag, -1);
}
let view2 = py
.eval("array[5:]", None, Some(locals))
.unwrap()
.downcast::<PyArray1<f64>>()
.unwrap();
let key2 = BorrowKey::from_array(view2);
let shared2 = view2.readonly();
{
let borrow_flags = unsafe { BORROW_FLAGS.get() };
assert_eq!(borrow_flags.len(), 1);
let same_base_arrays = &borrow_flags[&base];
assert_eq!(same_base_arrays.len(), 2);
let flag = same_base_arrays[&key1];
assert_eq!(flag, -1);
let flag = same_base_arrays[&key2];
assert_eq!(flag, 1);
}
let view3 = py
.eval("array[5:]", None, Some(locals))
.unwrap()
.downcast::<PyArray1<f64>>()
.unwrap();
let key3 = BorrowKey::from_array(view3);
let shared3 = view3.readonly();
{
let borrow_flags = unsafe { BORROW_FLAGS.get() };
assert_eq!(borrow_flags.len(), 1);
let same_base_arrays = &borrow_flags[&base];
assert_eq!(same_base_arrays.len(), 2);
let flag = same_base_arrays[&key1];
assert_eq!(flag, -1);
let flag = same_base_arrays[&key2];
assert_eq!(flag, 2);
let flag = same_base_arrays[&key3];
assert_eq!(flag, 2);
}
let view4 = py
.eval("array[7:]", None, Some(locals))
.unwrap()
.downcast::<PyArray1<f64>>()
.unwrap();
let key4 = BorrowKey::from_array(view4);
let shared4 = view4.readonly();
{
let borrow_flags = unsafe { BORROW_FLAGS.get() };
assert_eq!(borrow_flags.len(), 1);
let same_base_arrays = &borrow_flags[&base];
assert_eq!(same_base_arrays.len(), 3);
let flag = same_base_arrays[&key1];
assert_eq!(flag, -1);
let flag = same_base_arrays[&key2];
assert_eq!(flag, 2);
let flag = same_base_arrays[&key3];
assert_eq!(flag, 2);
let flag = same_base_arrays[&key4];
assert_eq!(flag, 1);
}
drop(shared2);
{
let borrow_flags = unsafe { BORROW_FLAGS.get() };
assert_eq!(borrow_flags.len(), 1);
let same_base_arrays = &borrow_flags[&base];
assert_eq!(same_base_arrays.len(), 3);
let flag = same_base_arrays[&key1];
assert_eq!(flag, -1);
let flag = same_base_arrays[&key2];
assert_eq!(flag, 1);
let flag = same_base_arrays[&key3];
assert_eq!(flag, 1);
let flag = same_base_arrays[&key4];
assert_eq!(flag, 1);
}
drop(shared3);
{
let borrow_flags = unsafe { BORROW_FLAGS.get() };
assert_eq!(borrow_flags.len(), 1);
let same_base_arrays = &borrow_flags[&base];
assert_eq!(same_base_arrays.len(), 2);
let flag = same_base_arrays[&key1];
assert_eq!(flag, -1);
assert!(!same_base_arrays.contains_key(&key2));
assert!(!same_base_arrays.contains_key(&key3));
let flag = same_base_arrays[&key4];
assert_eq!(flag, 1);
}
drop(exclusive1);
{
let borrow_flags = unsafe { BORROW_FLAGS.get() };
assert_eq!(borrow_flags.len(), 1);
let same_base_arrays = &borrow_flags[&base];
assert_eq!(same_base_arrays.len(), 1);
assert!(!same_base_arrays.contains_key(&key1));
assert!(!same_base_arrays.contains_key(&key2));
assert!(!same_base_arrays.contains_key(&key3));
let flag = same_base_arrays[&key4];
assert_eq!(flag, 1);
}
drop(shared4);
{
let borrow_flags = unsafe { BORROW_FLAGS.get() };
assert_eq!(borrow_flags.len(), 0);
}
});
}
#[test]
fn test_debug_formatting() {
Python::with_gil(|py| {
let array = PyArray::<f64, _>::zeros(py, (1, 2, 3), false);
{
let shared = array.readonly();
assert_eq!(
format!("{:?}", shared),
"PyReadonlyArray<f64, ndarray::dimension::dim::Dim<[usize; 3]>>"
);
}
{
let exclusive = array.readwrite();
assert_eq!(
format!("{:?}", exclusive),
"PyReadwriteArray<f64, ndarray::dimension::dim::Dim<[usize; 3]>>"
);
}
});
}
#[test]
#[should_panic(expected = "AlreadyBorrowed")]
fn cannot_clone_exclusive_borrow_via_deref() {
Python::with_gil(|py| {
let array = PyArray::<f64, _>::zeros(py, (3, 2, 1), false);
let exclusive = array.readwrite();
let _shared = exclusive.clone();
});
}
#[test]
fn failed_resize_does_not_double_release() {
Python::with_gil(|py| {
let array = PyArray::<f64, _>::zeros(py, 10, false);
let locals = [("array", array)].into_py_dict(py);
let _view = py
.eval("array[:]", None, Some(locals))
.unwrap()
.downcast::<PyArray1<f64>>()
.unwrap();
let exclusive = array.readwrite();
assert!(exclusive.resize(100).is_err());
});
}
#[test]
fn ineffective_resize_does_not_conflict() {
Python::with_gil(|py| {
let array = PyArray::<f64, _>::zeros(py, 10, false);
let exclusive = array.readwrite();
assert!(exclusive.resize(10).is_ok());
});
}
}