use crate::ThreadBound;
use crate::foundation::{Error, metal_error};
use crate::metal::generated_object_types::metal::Tensor;
use crate::metal::generated_value_types::BufferSparseTier;
use crate::metal::{CheckedTensorDescriptor, Device, StorageMode, Texture, TextureDescriptor};
use objc2::rc::Retained;
use objc2::runtime::{AnyObject, ProtocolObject};
use objc2::{msg_send, sel};
use objc2_foundation::{NSRange, NSString};
use objc2_metal::{MTLBuffer, MTLResource, MTLTensorDescriptor};
#[derive(Clone)]
pub struct Buffer {
pub(super) inner: Retained<ProtocolObject<dyn MTLBuffer>>,
storage: StorageMode,
_thread_bound: ThreadBound,
}
impl Buffer {
pub(super) const fn new(
inner: Retained<ProtocolObject<dyn MTLBuffer>>,
storage: StorageMode,
) -> Self {
Self {
inner,
storage,
_thread_bound: ThreadBound::new(),
}
}
pub(crate) fn from_any_object(inner: Retained<AnyObject>) -> Result<Self, Error> {
let inner: Retained<ProtocolObject<dyn MTLBuffer>> =
unsafe { Retained::cast_unchecked(inner) };
let storage = StorageMode::try_from_system_raw(inner.storageMode().0)
.ok_or_else(|| Error::unsupported("Metal returned an unknown buffer storage mode"))?;
Ok(Self::new(inner, storage))
}
pub(crate) fn as_any_object(&self) -> &AnyObject {
unsafe { &*(std::ptr::from_ref(&*self.inner).cast::<AnyObject>()) }
}
#[must_use]
pub fn length(&self) -> usize {
self.inner.length()
}
#[must_use]
pub const fn storage_mode(&self) -> StorageMode {
self.storage
}
pub fn add_debug_marker(
&self,
marker: &str,
range: std::ops::Range<usize>,
) -> Result<(), Error> {
if range.start > range.end || range.end > self.length() {
return Err(Error::invalid_argument(
"buffer debug marker range is out of bounds",
));
}
let marker = NSString::from_str(marker);
self.inner
.addDebugMarker_range(&marker, NSRange::new(range.start, range.len()));
Ok(())
}
pub fn remove_all_debug_markers(&self) {
self.inner.removeAllDebugMarkers();
}
#[must_use]
pub fn gpu_address(&self) -> u64 {
self.inner.gpuAddress()
}
#[must_use]
pub fn sparse_buffer_tier(&self) -> BufferSparseTier {
BufferSparseTier::from_system_raw(self.inner.sparseBufferTier().0)
}
#[must_use]
pub fn remote_storage_buffer(&self) -> Option<Self> {
self.inner.remoteStorageBuffer().map(|inner| {
let storage = storage_mode(&inner);
Self::new(inner, storage)
})
}
pub fn new_remote_view(&self, device: &Device) -> Result<Self, Error> {
self.inner
.newRemoteBufferViewForDevice(&device.inner)
.map(|inner| {
let storage = storage_mode(&inner);
Self::new(inner, storage)
})
.ok_or_else(|| Error::unsupported("Metal could not create a remote buffer view"))
}
pub fn new_texture(
&self,
descriptor: &TextureDescriptor,
offset: usize,
bytes_per_row: usize,
) -> Result<Texture, Error> {
if offset >= self.length() || bytes_per_row == 0 {
return Err(Error::invalid_argument(
"buffer texture offset must be in range and row stride non-zero",
));
}
self.inner
.newTextureWithDescriptor_offset_bytesPerRow(&descriptor.inner, offset, bytes_per_row)
.map(Texture::new)
.ok_or_else(|| Error::unsupported("Metal rejected the buffer texture view"))
}
pub fn new_tensor(
&self,
offset: usize,
descriptor: &CheckedTensorDescriptor,
) -> Result<Tensor, Error> {
let _checked_range = descriptor.checked_buffer_range(self.length(), offset)?;
let available: bool = unsafe {
msg_send![
&*self.inner,
respondsToSelector: sel!(newTensorWithDescriptor:offset:error:)
]
};
if !available {
return Err(Error::unsupported(
"buffer-backed tensor creation is unavailable on this system",
));
}
let descriptor: &MTLTensorDescriptor =
unsafe { &*(std::ptr::from_ref(descriptor.as_inner()).cast::<MTLTensorDescriptor>()) };
let tensor = unsafe {
self.inner
.newTensorWithDescriptor_offset_error(descriptor, offset)
}
.map_err(|error| metal_error(&error))?;
let tensor = unsafe { Retained::cast_unchecked(tensor) };
Ok(Tensor::from_inner(tensor))
}
pub fn write(&self, offset: usize, bytes: &[u8]) -> Result<(), Error> {
if matches!(self.storage, StorageMode::Private | StorageMode::Memoryless) {
return Err(Error::unsupported(
"private and memoryless buffers are not CPU writable",
));
}
let end = offset
.checked_add(bytes.len())
.ok_or_else(|| Error::invalid_argument("buffer write range overflow"))?;
if end > self.length() {
return Err(Error::invalid_argument(
"buffer write range is out of bounds",
));
}
if bytes.is_empty() {
return Ok(());
}
let destination = self.inner.contents().as_ptr().cast::<u8>();
unsafe {
std::ptr::copy_nonoverlapping(bytes.as_ptr(), destination.add(offset), bytes.len());
}
if self.storage == StorageMode::Managed {
self.inner.didModifyRange(NSRange::new(offset, bytes.len()));
}
Ok(())
}
pub(crate) fn completed_bytes(&self, length: usize) -> Result<Vec<u8>, Error> {
if length > self.length() {
return Err(Error::invalid_argument("readback length is out of bounds"));
}
if matches!(self.storage, StorageMode::Private | StorageMode::Memoryless) {
return Err(Error::unsupported("staging buffer is not CPU visible"));
}
if length == 0 {
return Ok(Vec::new());
}
let source = self.inner.contents().as_ptr().cast::<u8>();
let mut bytes = vec![0_u8; length];
unsafe {
std::ptr::copy_nonoverlapping(source, bytes.as_mut_ptr(), length);
}
Ok(bytes)
}
}
fn storage_mode(buffer: &ProtocolObject<dyn MTLBuffer>) -> StorageMode {
match buffer.storageMode().0 {
1 => StorageMode::Managed,
2 => StorageMode::Private,
3 => StorageMode::Memoryless,
_ => StorageMode::Shared,
}
}