use std::ffi::c_void;
use crate::ffi::{DLManagedTensor, DLManagedTensorVersioned, DLPackVersion, DLTensor};
use bitflags::bitflags;
bitflags! {
#[repr(transparent)]
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct DlpackFlags: u64 {
const READ_ONLY = 1 << 0;
const IS_COPIED = 1 << 1;
const IS_SUBBYTE_TYPE_PADDED = 1 << 2;
}
}
impl DlpackFlags {
pub(crate) fn newly_asserts_is_copied(self, current: DlpackFlags) -> bool {
self.contains(DlpackFlags::IS_COPIED) && !current.contains(DlpackFlags::IS_COPIED)
}
}
pub unsafe trait ManagedTensorBase {
fn from_parts(
tensor: DLTensor,
manager_ctx: *mut c_void,
deleter: Option<unsafe extern "C" fn(self_: *mut Self)>,
) -> Self;
fn tensor(&self) -> &DLTensor;
fn tensor_mut(&mut self) -> &mut DLTensor;
fn manager_ctx(&self) -> *mut c_void;
fn deleter(&self) -> Option<unsafe extern "C" fn(self_: *mut Self)>;
#[inline]
fn version(&self) -> Option<DLPackVersion> {
None
}
#[inline]
fn flags(&self) -> DlpackFlags {
DlpackFlags::empty()
}
unsafe fn set_flags_unchecked(&mut self, _flags: crate::DlpackFlags) {}
#[inline]
unsafe fn drop_raw(ptr: *mut Self) {
if let Some(deleter) = unsafe { (*ptr).deleter() } {
unsafe { deleter(ptr) };
}
}
}
unsafe impl ManagedTensorBase for DLManagedTensor {
#[inline]
fn from_parts(
tensor: DLTensor,
manager_ctx: *mut c_void,
deleter: Option<unsafe extern "C" fn(self_: *mut Self)>,
) -> Self {
Self {
dl_tensor: tensor,
manager_ctx,
deleter,
}
}
#[inline]
fn tensor(&self) -> &DLTensor {
&self.dl_tensor
}
#[inline]
fn tensor_mut(&mut self) -> &mut DLTensor {
&mut self.dl_tensor
}
#[inline]
fn manager_ctx(&self) -> *mut c_void {
self.manager_ctx
}
#[inline]
fn deleter(&self) -> Option<unsafe extern "C" fn(self_: *mut Self)> {
self.deleter
}
}
unsafe impl ManagedTensorBase for DLManagedTensorVersioned {
#[inline]
fn from_parts(
tensor: DLTensor,
manager_ctx: *mut c_void,
deleter: Option<unsafe extern "C" fn(self_: *mut Self)>,
) -> Self {
Self {
version: DLPackVersion::default(),
manager_ctx,
deleter,
flags: DlpackFlags::empty(),
dl_tensor: tensor,
}
}
#[inline]
fn tensor(&self) -> &DLTensor {
&self.dl_tensor
}
#[inline]
fn tensor_mut(&mut self) -> &mut DLTensor {
&mut self.dl_tensor
}
#[inline]
fn manager_ctx(&self) -> *mut c_void {
self.manager_ctx
}
#[inline]
fn deleter(&self) -> Option<unsafe extern "C" fn(self_: *mut Self)> {
self.deleter
}
#[inline]
fn version(&self) -> Option<DLPackVersion> {
Some(self.version)
}
#[inline]
unsafe fn set_flags_unchecked(&mut self, flags: crate::DlpackFlags) {
self.flags = flags;
}
#[inline]
fn flags(&self) -> DlpackFlags {
self.flags
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn version_comparisons_are_semantic() {
let current = DLPackVersion::CURRENT;
let newer_minor = DLPackVersion {
major: current.major,
minor: current.minor + 1,
};
let other_major = DLPackVersion {
major: current.major + 1,
minor: 0,
};
assert!(current.is_compatible_with(newer_minor));
assert!(newer_minor.supports(current));
assert!(!current.supports(newer_minor));
assert!(!current.is_compatible_with(other_major));
assert!(!current.supports(other_major));
assert_eq!(DLPackVersion::default(), current);
}
}