1use crate::{
5 Py, PyObject, PyObjectRef, PyPayload, PyRef, PyResult, TryFromBorrowedObject, VirtualMachine,
6 common::{
7 borrow::{BorrowedValue, BorrowedValueMut},
8 lock::{MapImmutable, PyMutex, PyMutexGuard},
9 rc::PyRc,
10 },
11 object::PyObjectPayload,
12 sliceable::SequenceIndexOp,
13};
14use alloc::borrow::Cow;
15use bitflags::bitflags;
16use core::{fmt::Debug, ops::Range};
17use crossbeam_utils::atomic::AtomicCell;
18use itertools::Itertools;
19
20bitflags! {
21 #[derive(Copy, Clone, Debug, PartialEq, Eq)]
28 pub struct BufferFlags: u32 {
29 const WRITABLE = 0x0001;
30 const FORMAT = 0x0004;
31 const ND = 0x0008;
32 const STRIDES = 0x0010 | Self::ND.bits();
33 const C_CONTIGUOUS = 0x0020 | Self::STRIDES.bits();
34 const F_CONTIGUOUS = 0x0040 | Self::STRIDES.bits();
35 const ANY_CONTIGUOUS = 0x0080 | Self::STRIDES.bits();
36 const INDIRECT = 0x0100 | Self::STRIDES.bits();
37 }
38}
39
40impl BufferFlags {
41 pub const SIMPLE: Self = Self::empty();
43 pub const CONTIG: Self = Self::ND.union(Self::WRITABLE);
45 pub const CONTIG_RO: Self = Self::ND;
47 pub const STRIDED: Self = Self::STRIDES.union(Self::WRITABLE);
49 pub const STRIDED_RO: Self = Self::STRIDES;
51 pub const RECORDS: Self = Self::STRIDED.union(Self::FORMAT);
53 pub const RECORDS_RO: Self = Self::STRIDED_RO.union(Self::FORMAT);
55 pub const FULL: Self = Self::INDIRECT.union(Self::WRITABLE).union(Self::FORMAT);
57 pub const FULL_RO: Self = Self::INDIRECT.union(Self::FORMAT);
59
60 const MEMORY_READ: Self = Self::from_bits_retain(0x100);
62 const MEMORY_WRITE: Self = Self::from_bits_retain(0x200);
64
65 #[must_use]
68 pub const fn is_memory_access_mode(self) -> bool {
69 self.bits() == Self::MEMORY_READ.bits() || self.bits() == Self::MEMORY_WRITE.bits()
70 }
71
72 #[must_use]
74 pub const fn is_writable(self) -> bool {
75 self.intersects(Self::WRITABLE)
76 }
77
78 pub fn fill_info_check(self, readonly: bool, vm: &VirtualMachine) -> PyResult<()> {
81 if self == Self::SIMPLE {
82 return Ok(());
83 }
84 if self.is_memory_access_mode() {
85 return Err(vm.new_system_error("bad argument to internal function"));
86 }
87 self.check_writable(readonly, "Object is not writable.", vm)
88 }
89
90 pub fn check_writable(
92 self,
93 readonly: bool,
94 message: &str,
95 vm: &VirtualMachine,
96 ) -> PyResult<()> {
97 if self.is_writable() && readonly {
98 return Err(vm.new_buffer_error(message.to_owned()));
99 }
100 Ok(())
101 }
102}
103
104pub struct BufferMethods {
105 pub obj_bytes: fn(&PyBuffer) -> BorrowedValue<'_, [u8]>,
106 pub obj_bytes_mut: fn(&PyBuffer) -> BorrowedValueMut<'_, [u8]>,
107 pub release: fn(&PyBuffer),
108 pub retain: fn(&PyBuffer),
109}
110
111impl Debug for BufferMethods {
112 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
113 f.debug_struct("BufferMethods")
114 .field("obj_bytes", &(self.obj_bytes as usize))
115 .field("obj_bytes_mut", &(self.obj_bytes_mut as usize))
116 .field("release", &(self.release as usize))
117 .field("retain", &(self.retain as usize))
118 .finish()
119 }
120}
121
122#[derive(Debug)]
125struct BufferExport {
126 shares: AtomicCell<usize>,
128 released: AtomicCell<bool>,
131}
132
133#[derive(Debug, Traverse)]
134pub struct PyBuffer {
135 pub obj: PyObjectRef,
136 #[pytraverse(skip)]
137 pub desc: BufferDescriptor,
138 #[pytraverse(skip)]
139 methods: &'static BufferMethods,
140 #[pytraverse(skip)]
141 export: PyRc<BufferExport>,
142 #[pytraverse(skip)]
144 owns_share: AtomicCell<bool>,
145}
146
147impl Clone for PyBuffer {
151 fn clone(&self) -> Self {
152 debug_assert!(!self.export.released.load());
153 self.export.shares.fetch_add(1);
154 Self {
155 obj: self.obj.clone(),
156 desc: self.desc.clone(),
157 methods: self.methods,
158 export: self.export.clone(),
159 owns_share: AtomicCell::new(true),
160 }
161 }
162}
163
164impl PyBuffer {
165 #[must_use]
166 pub fn new(obj: PyObjectRef, desc: BufferDescriptor, methods: &'static BufferMethods) -> Self {
167 #[cfg(debug_assertions)]
168 let desc = desc.validate();
169
170 let zelf = Self {
171 obj,
172 desc,
173 methods,
174 export: PyRc::new(BufferExport {
175 shares: AtomicCell::new(1),
176 released: AtomicCell::new(false),
177 }),
178 owns_share: AtomicCell::new(true),
179 };
180 (zelf.methods.retain)(&zelf);
181 zelf
182 }
183
184 #[must_use]
185 pub fn as_contiguous(&self) -> Option<BorrowedValue<'_, [u8]>> {
186 self.desc
187 .is_contiguous()
188 .then(|| unsafe { self.contiguous_unchecked() })
189 }
190
191 #[must_use]
192 pub fn as_contiguous_mut(&self) -> Option<BorrowedValueMut<'_, [u8]>> {
193 (!self.desc.readonly && self.desc.is_contiguous())
194 .then(|| unsafe { self.contiguous_mut_unchecked() })
195 }
196
197 pub fn from_byte_vector(bytes: Vec<u8>, vm: &VirtualMachine) -> Self {
198 let bytes_len = bytes.len();
199 Self::new(
200 PyPayload::into_pyobject(VecBuffer::from(bytes), vm),
201 BufferDescriptor::simple(bytes_len, true),
202 &VEC_BUFFER_METHODS,
203 )
204 }
205
206 #[must_use]
209 pub unsafe fn contiguous_unchecked(&self) -> BorrowedValue<'_, [u8]> {
210 let range = self.desc.contiguous_range();
211 BorrowedValue::map(self.obj_bytes(), |x| &x[range])
212 }
213
214 #[must_use]
217 pub unsafe fn contiguous_mut_unchecked(&self) -> BorrowedValueMut<'_, [u8]> {
218 let range = self.desc.contiguous_range();
219 BorrowedValueMut::map(self.obj_bytes_mut(), |x| &mut x[range])
220 }
221
222 pub fn append_to(&self, buf: &mut Vec<u8>) {
223 if let Some(bytes) = self.as_contiguous() {
224 buf.extend_from_slice(&bytes);
225 } else {
226 let bytes = &*self.obj_bytes();
227 self.desc.for_each_segment(true, |range| {
228 buf.extend_from_slice(&bytes[range.start as usize..range.end as usize])
229 });
230 }
231 }
232
233 pub fn contiguous_or_collect<R, F: FnOnce(&[u8]) -> R>(&self, f: F) -> R {
234 let borrowed;
235 let mut collected;
236 let v = if let Some(bytes) = self.as_contiguous() {
237 borrowed = bytes;
238 &*borrowed
239 } else {
240 collected = vec![];
241 self.append_to(&mut collected);
242 &collected
243 };
244 f(v)
245 }
246
247 #[must_use]
251 pub fn to_contiguous(&self, vm: &VirtualMachine) -> Self {
252 let mut data = vec![];
253 self.append_to(&mut data);
254 VecBuffer::from(data)
255 .into_ref(&vm.ctx)
256 .into_pybuffer_with_descriptor(self.desc.contiguous())
257 }
258
259 #[must_use]
260 pub fn obj_as<T: PyObjectPayload>(&self) -> &Py<T> {
261 unsafe { self.obj.downcast_unchecked_ref() }
262 }
263
264 #[must_use]
265 pub fn obj_bytes(&self) -> BorrowedValue<'_, [u8]> {
266 (self.methods.obj_bytes)(self)
267 }
268
269 #[must_use]
270 pub fn obj_bytes_mut(&self) -> BorrowedValueMut<'_, [u8]> {
271 (self.methods.obj_bytes_mut)(self)
272 }
273
274 pub fn release(&self) {
283 if self.owns_share.swap(false) {
284 self.drop_share();
285 }
286 }
287
288 pub(crate) fn retain_share(&self) {
292 self.export.shares.fetch_add(1);
293 }
294
295 pub(crate) fn release_share(&self) {
297 self.drop_share();
298 }
299
300 fn drop_share(&self) {
301 if self.export.shares.fetch_sub(1) == 1 {
302 self.finalize();
303 }
304 }
305
306 fn finalize(&self) {
308 if self.export.released.swap(true) {
311 return;
312 }
313 if self.obj.class().slots().python_release_buffer.load() {
316 crate::builtins::memory::release_buffer_call_python(self);
317 }
318 (self.methods.release)(self)
319 }
320
321 pub(crate) fn abort_acquisition(self) {
325 debug_assert_eq!(self.export.shares.load(), 1);
326 self.owns_share.store(false);
327 self.export.released.store(true);
328 (self.methods.release)(&self);
329 }
330
331 #[must_use]
335 pub fn detached(&self) -> Self {
336 Self {
337 obj: self.obj.clone(),
338 desc: self.desc.clone(),
339 methods: self.methods,
340 export: self.export.clone(),
341 owns_share: AtomicCell::new(false),
342 }
343 }
344}
345
346impl PyBuffer {
347 pub fn from_object(vm: &VirtualMachine, obj: &PyObject, flags: BufferFlags) -> PyResult<Self> {
349 if flags.is_memory_access_mode() {
350 return Err(vm.new_system_error("bad argument to internal function"));
351 }
352 let cls = obj.class();
353 if let Some(f) = cls.slots.as_buffer.load() {
354 return f(obj, flags, vm);
355 }
356 Err(vm.new_type_error(format!(
357 "a bytes-like object is required, not '{}'",
358 cls.name()
359 )))
360 }
361}
362
363impl PyObject {
364 #[must_use]
370 pub fn check_buffer(&self) -> bool {
371 self.class().slots().as_buffer.load().is_some()
372 }
373}
374
375impl<'a> TryFromBorrowedObject<'a> for PyBuffer {
378 fn try_from_borrowed_object(vm: &VirtualMachine, obj: &'a PyObject) -> PyResult<Self> {
379 Self::from_object(vm, obj, BufferFlags::FULL_RO)
380 }
381}
382
383impl Drop for PyBuffer {
384 fn drop(&mut self) {
385 self.release();
386 }
387}
388
389#[derive(Debug, Clone)]
390pub struct BufferDescriptor {
391 pub len: usize,
394 pub offset: isize,
403 pub readonly: bool,
404 pub itemsize: usize,
405 pub format: Cow<'static, str>,
406 pub dim_desc: Vec<(usize, isize, isize)>,
409 }
411
412impl BufferDescriptor {
413 #[must_use]
414 pub fn simple(bytes_len: usize, readonly: bool) -> Self {
415 Self {
416 len: bytes_len,
417 offset: 0,
418 readonly,
419 itemsize: 1,
420 format: Cow::Borrowed("B"),
421 dim_desc: vec![(bytes_len, 1, 0)],
422 }
423 }
424
425 #[must_use]
426 pub fn format(
427 bytes_len: usize,
428 readonly: bool,
429 itemsize: usize,
430 format: Cow<'static, str>,
431 ) -> Self {
432 Self {
433 len: bytes_len,
434 offset: 0,
435 readonly,
436 itemsize,
437 format,
438 dim_desc: vec![(bytes_len / itemsize, itemsize as isize, 0)],
439 }
440 }
441
442 #[must_use]
454 pub fn projected(&self, flags: BufferFlags) -> Self {
455 let mut desc = self.clone();
456 if !flags.contains(BufferFlags::FORMAT) {
457 desc.format = Cow::Borrowed("B");
458 }
459 if !flags.contains(BufferFlags::ND) {
460 let shape = desc.len.checked_div(desc.itemsize).unwrap_or(0);
463 desc.dim_desc = vec![(shape, desc.itemsize as isize, 0)];
464 } else if !flags.contains(BufferFlags::STRIDES) {
465 let mut stride = desc.itemsize as isize;
467 for (shape, dim_stride, suboffset) in desc.dim_desc.iter_mut().rev() {
468 *dim_stride = stride;
469 *suboffset = 0;
470 stride *= *shape as isize;
471 }
472 }
473 desc
474 }
475
476 #[cfg(debug_assertions)]
477 #[must_use]
478 pub fn validate(self) -> Self {
479 if self.len != 0 {
482 debug_assert!(self.offset >= 0);
483 }
484 if self.ndim() == 0 {
486 if self.len > 0 {
488 debug_assert_ne!(self.itemsize, 0);
489 }
490 debug_assert_eq!(self.itemsize, self.len);
491 } else {
492 let mut shape_product = 1;
493 let has_zero_dim = self.dim_desc.iter().any(|(s, _, _)| *s == 0);
494 for (shape, stride, suboffset) in self.dim_desc.iter().copied() {
495 shape_product *= shape;
496 debug_assert!(suboffset >= 0);
497 if !has_zero_dim {
499 debug_assert_ne!(stride, 0);
500 }
501 }
502 debug_assert_eq!(shape_product * self.itemsize, self.len);
503 }
504 self
505 }
506
507 #[must_use]
508 pub fn ndim(&self) -> usize {
509 self.dim_desc.len()
510 }
511
512 #[must_use]
514 pub fn is_contiguous(&self) -> bool {
515 if self.len == 0 {
516 return true;
517 }
518 let mut sd = self.itemsize;
519 for (shape, stride, _) in self.dim_desc.iter().copied().rev() {
520 if shape > 1 && stride != sd as isize {
521 return false;
522 }
523 sd *= shape;
524 }
525 true
526 }
527
528 #[must_use]
532 pub fn is_fortran_contiguous(&self) -> bool {
533 if self.len == 0 {
534 return true;
535 }
536 let mut sd = self.itemsize;
537 for (shape, stride, _) in self.dim_desc.iter().copied() {
538 if shape > 1 && stride != sd as isize {
539 return false;
540 }
541 sd *= shape;
542 }
543 true
544 }
545
546 #[must_use]
552 pub fn contiguous_range(&self) -> Range<usize> {
553 if self.len == 0 {
554 return 0..0;
555 }
556 debug_assert!(self.offset >= 0);
557 let start = self.offset as usize;
558 start..start + self.len
559 }
560
561 #[must_use]
563 pub fn contiguous(&self) -> Self {
564 let itemsize = self.itemsize;
565 let mut dim_desc = self.dim_desc.clone();
566 if let Some((_, stride, suboffset)) = dim_desc.last_mut() {
567 *stride = itemsize as isize;
568 *suboffset = 0;
569 }
570 for i in (1..dim_desc.len()).rev() {
571 dim_desc[i - 1].1 = dim_desc[i].1 * dim_desc[i].0 as isize;
572 dim_desc[i - 1].2 = 0;
573 }
574 Self {
575 len: self.len,
576 offset: 0,
577 readonly: self.readonly,
578 itemsize: self.itemsize,
579 format: self.format.clone(),
580 dim_desc,
581 }
582 }
583
584 #[must_use]
587 pub fn has_suboffsets(&self) -> bool {
588 self.dim_desc
589 .iter()
590 .any(|(_, _, suboffset)| *suboffset != 0)
591 }
592
593 #[must_use]
596 pub fn fast_position(&self, indices: &[usize]) -> isize {
597 let mut pos = self.offset;
598 for (i, (_, stride, suboffset)) in indices
599 .iter()
600 .copied()
601 .zip_eq(self.dim_desc.iter().copied())
602 {
603 pos += i as isize * stride + suboffset;
604 }
605 pos
606 }
607
608 pub fn position(&self, indices: &[isize], vm: &VirtualMachine) -> PyResult<isize> {
610 let mut pos = self.offset;
611 for (dim, (i, (shape, stride, suboffset))) in indices
612 .iter()
613 .copied()
614 .zip_eq(self.dim_desc.iter().copied())
615 .enumerate()
616 {
617 let i = i.wrapped_at(shape).ok_or_else(|| {
619 vm.new_index_error(format!("index out of bounds on dimension {}", dim + 1))
620 })?;
621 pos += i as isize * stride + suboffset;
622 }
623 Ok(pos)
624 }
625
626 pub fn for_each_segment<F>(&self, try_contiguous: bool, mut f: F)
627 where
628 F: FnMut(Range<isize>),
629 {
630 if self.len == 0 {
633 return;
634 }
635 if self.ndim() == 0 {
636 f(self.offset..self.offset + self.itemsize as isize);
637 return;
638 }
639 if try_contiguous && self.is_last_dim_contiguous() {
640 self._for_each_segment::<_, true>(self.offset, 0, &mut f);
641 } else {
642 self._for_each_segment::<_, false>(self.offset, 0, &mut f);
643 }
644 }
645
646 pub fn for_each_segment_fortran<F>(&self, mut f: F)
652 where
653 F: FnMut(Range<isize>),
654 {
655 if self.len == 0 {
656 return;
657 }
658 if self.ndim() == 0 {
659 f(self.offset..self.offset + self.itemsize as isize);
660 return;
661 }
662 let mut indices = vec![0usize; self.ndim()];
663 loop {
664 let pos = self.offset
665 + indices
666 .iter()
667 .zip_eq(self.dim_desc.iter())
668 .map(|(&i, &(_, stride, suboffset))| i as isize * stride + suboffset)
669 .sum::<isize>();
670 f(pos..pos + self.itemsize as isize);
671
672 let mut dim = 0;
673 loop {
674 indices[dim] += 1;
675 if indices[dim] < self.dim_desc[dim].0 {
676 break;
677 }
678 indices[dim] = 0;
679 dim += 1;
680 if dim == self.ndim() {
681 return;
682 }
683 }
684 }
685 }
686
687 fn _for_each_segment<F, const CONTIGUOUS: bool>(&self, mut index: isize, dim: usize, f: &mut F)
688 where
689 F: FnMut(Range<isize>),
690 {
691 let (shape, stride, suboffset) = self.dim_desc[dim];
692 if dim + 1 == self.ndim() {
693 if CONTIGUOUS {
694 f(index..index + (shape * self.itemsize) as isize);
695 } else {
696 for _ in 0..shape {
697 let pos = index + suboffset;
698 f(pos..pos + self.itemsize as isize);
699 index += stride;
700 }
701 }
702 return;
703 }
704 for _ in 0..shape {
705 self._for_each_segment::<F, CONTIGUOUS>(index + suboffset, dim + 1, f);
706 index += stride;
707 }
708 }
709
710 pub fn zip_eq<F>(&self, other: &Self, try_contiguous: bool, mut f: F)
712 where
713 F: FnMut(Range<isize>, Range<isize>) -> bool,
714 {
715 if self.len == 0 {
716 return;
717 }
718 if self.ndim() == 0 {
719 f(
720 self.offset..self.offset + self.itemsize as isize,
721 other.offset..other.offset + other.itemsize as isize,
722 );
723 return;
724 }
725 let run_at_once =
728 try_contiguous && self.is_last_dim_contiguous() && other.is_last_dim_contiguous();
729 if run_at_once {
730 self._zip_eq::<_, true>(other, self.offset, other.offset, 0, &mut f);
731 } else {
732 self._zip_eq::<_, false>(other, self.offset, other.offset, 0, &mut f);
733 }
734 }
735
736 fn _zip_eq<F, const CONTIGUOUS: bool>(
737 &self,
738 other: &Self,
739 mut a_index: isize,
740 mut b_index: isize,
741 dim: usize,
742 f: &mut F,
743 ) where
744 F: FnMut(Range<isize>, Range<isize>) -> bool,
745 {
746 let (shape, a_stride, a_suboffset) = self.dim_desc[dim];
747 let (_b_shape, b_stride, b_suboffset) = other.dim_desc[dim];
748 debug_assert_eq!(shape, _b_shape);
749 if dim + 1 == self.ndim() {
750 if CONTIGUOUS {
751 if f(
752 a_index..a_index + (shape * self.itemsize) as isize,
753 b_index..b_index + (shape * other.itemsize) as isize,
754 ) {
755 return;
756 }
757 } else {
758 for _ in 0..shape {
759 let a_pos = a_index + a_suboffset;
760 let b_pos = b_index + b_suboffset;
761 if f(
762 a_pos..a_pos + self.itemsize as isize,
763 b_pos..b_pos + other.itemsize as isize,
764 ) {
765 return;
766 }
767 a_index += a_stride;
768 b_index += b_stride;
769 }
770 }
771 return;
772 }
773
774 for _ in 0..shape {
775 self._zip_eq::<F, CONTIGUOUS>(
776 other,
777 a_index + a_suboffset,
778 b_index + b_suboffset,
779 dim + 1,
780 f,
781 );
782 a_index += a_stride;
783 b_index += b_stride;
784 }
785 }
786
787 #[must_use]
788 fn is_last_dim_contiguous(&self) -> bool {
789 let (_, stride, suboffset) = self.dim_desc[self.ndim() - 1];
790 suboffset == 0 && stride == self.itemsize as isize
791 }
792
793 #[must_use]
794 pub fn is_zero_in_shape(&self) -> bool {
795 self.dim_desc.iter().any(|(shape, _, _)| *shape == 0)
796 }
797
798 }
800
801pub trait BufferResizeGuard {
802 type Resizable<'a>: 'a
803 where
804 Self: 'a;
805 fn try_resizable_opt(&self) -> Option<Self::Resizable<'_>>;
806 fn try_resizable(&self, vm: &VirtualMachine) -> PyResult<Self::Resizable<'_>> {
807 self.try_resizable_opt().ok_or_else(|| {
808 vm.new_buffer_error("Existing exports of data: object cannot be re-sized")
809 })
810 }
811}
812
813#[pyclass(module = false, name = "vec_buffer")]
814#[derive(Debug, PyPayload)]
815pub struct VecBuffer {
816 data: PyMutex<Vec<u8>>,
817}
818
819#[pyclass(flags(BASETYPE, DISALLOW_INSTANTIATION))]
820impl VecBuffer {
821 pub fn take(&self) -> Vec<u8> {
822 core::mem::take(&mut self.data.lock())
823 }
824}
825
826impl From<Vec<u8>> for VecBuffer {
827 fn from(data: Vec<u8>) -> Self {
828 Self {
829 data: PyMutex::new(data),
830 }
831 }
832}
833
834impl PyRef<VecBuffer> {
835 #[must_use]
836 pub fn into_pybuffer(self, readonly: bool) -> PyBuffer {
837 let len = self.data.lock().len();
838 PyBuffer::new(
839 self.into(),
840 BufferDescriptor::simple(len, readonly),
841 &VEC_BUFFER_METHODS,
842 )
843 }
844
845 #[must_use]
846 pub fn into_pybuffer_with_descriptor(self, desc: BufferDescriptor) -> PyBuffer {
847 PyBuffer::new(self.into(), desc, &VEC_BUFFER_METHODS)
848 }
849}
850
851static VEC_BUFFER_METHODS: BufferMethods = BufferMethods {
852 obj_bytes: |buffer| {
853 PyMutexGuard::map_immutable(buffer.obj_as::<VecBuffer>().data.lock(), |x| x.as_slice())
854 .into()
855 },
856 obj_bytes_mut: |buffer| {
857 PyMutexGuard::map(buffer.obj_as::<VecBuffer>().data.lock(), |x| {
858 x.as_mut_slice()
859 })
860 .into()
861 },
862 release: |_| {},
863 retain: |_| {},
864};