#![allow(missing_docs)]
#![allow(unsafe_op_in_unsafe_fn)]
use pyo3::exceptions::PyValueError;
use pyo3::ffi;
use pyo3::prelude::*;
use pyo3::types::PyByteArray;
use pyo3::types::PyBytes;
use msrtc_rans::entropy::{EntropyDecoder, EntropyEncoder};
use msrtc_rans::stream::RansDecoderStream as CoreDecoderStream;
use msrtc_rans::stream::RansEncoderStream as CoreEncoderStream;
use msrtc_rans::variant::Rans64;
use msrtc_rans::variant::RansByte;
const BUF_READ: i32 = 12; const BUF_WRITE: i32 = 13;
fn validate_i32_buffer(buf: &ffi::Py_buffer) -> Result<(), String> {
if buf.ndim != 1 {
return Err(format!("expected 1-d array, got ndim={}", buf.ndim));
}
if buf.itemsize != 4 {
return Err(format!("expected int32 (itemsize 4), got {}", buf.itemsize));
}
if buf.format.is_null() {
return Err("buffer has no format string".into());
}
let fmt = unsafe { std::ffi::CStr::from_ptr(buf.format) }.to_string_lossy();
let fmt = fmt.trim_end_matches(|c: char| c.is_ascii_digit()); if fmt != "i" && fmt != "l" {
return Err(format!("expected int32 format, got '{}'", fmt));
}
if buf.len as usize % 4 != 0 {
return Err(format!(
"int32 array length must be multiple of 4, got {}",
buf.len
));
}
if (buf.buf as usize) % 4 != 0 {
return Err("int32 buffer is not 4-byte aligned".into());
}
if buf.shape.is_null() || unsafe { *buf.shape } != buf.len / buf.itemsize {
return Err("buffer shape does not match length".into());
}
Ok(())
}
unsafe fn buffer_to_i32_slice(buf: &ffi::Py_buffer) -> &[i32] {
let n = (buf.len as usize) / 4;
std::slice::from_raw_parts(buf.buf as *const i32, n)
}
unsafe fn buffer_to_u8_slice(buf: &ffi::Py_buffer) -> &[u8] {
let n = buf.len as usize;
std::slice::from_raw_parts(buf.buf as *const u8, n)
}
fn get_i32_buffer(obj: &Bound<'_, PyAny>) -> PyResult<Vec<i32>> {
let mut buf: ffi::Py_buffer = unsafe { std::mem::zeroed() };
let ret = unsafe { ffi::PyObject_GetBuffer(obj.as_ptr(), &mut buf, BUF_READ) };
if ret != 0 {
return Err(PyValueError::new_err(
"indices/values/pmf must be an int32 1-d array (buffer protocol)",
));
}
let result = validate_i32_buffer(&buf)
.map_err(PyValueError::new_err)
.and_then(|_| Ok(unsafe { buffer_to_i32_slice(&buf).to_vec() }));
unsafe { ffi::PyBuffer_Release(&mut buf) };
result
}
fn get_u8_buffer(obj: &Bound<'_, PyAny>) -> PyResult<Vec<u8>> {
if let Ok(bytes) = obj.downcast::<PyBytes>() {
return Ok(bytes.as_bytes().to_vec());
}
if let Ok(ba) = obj.downcast::<PyByteArray>() {
let slice = unsafe { ba.as_bytes() };
return Ok(slice.to_vec());
}
let mut buf: ffi::Py_buffer = unsafe { std::mem::zeroed() };
let ret = unsafe { ffi::PyObject_GetBuffer(obj.as_ptr(), &mut buf, BUF_READ) };
if ret != 0 {
return Err(PyValueError::new_err("cannot get buffer from object"));
}
let result = (|| -> Result<Vec<u8>, String> {
if buf.ndim != 1 {
return Err(format!("expected 1-d buffer, got ndim={}", buf.ndim));
}
if buf.itemsize != 1 {
return Err(format!(
"expected byte buffer (itemsize 1), got {}",
buf.itemsize
));
}
if buf.shape.is_null() || unsafe { *buf.shape } != buf.len {
return Err("buffer shape does not match length".into());
}
Ok(unsafe { buffer_to_u8_slice(&buf).to_vec() })
})()
.map_err(PyValueError::new_err);
unsafe { ffi::PyBuffer_Release(&mut buf) };
result
}
fn write_i32_buffer(obj: &Bound<'_, PyAny>, data: &[i32]) -> PyResult<()> {
let mut buf: ffi::Py_buffer = unsafe { std::mem::zeroed() };
let ret = unsafe { ffi::PyObject_GetBuffer(obj.as_ptr(), &mut buf, BUF_WRITE) };
if ret != 0 {
return Err(PyValueError::new_err(
"values must be a writable int32 1-d array",
));
}
let result = (|| -> Result<(), String> {
validate_i32_buffer(&buf)?;
let n = data.len() * 4;
if buf.len as usize != n {
return Err(format!(
"output int32 array has {} bytes, expected exactly {} for {} values",
buf.len,
n,
data.len()
));
}
unsafe {
let dst = buf.buf as *mut u8;
let src = data.as_ptr() as *const u8;
std::ptr::copy_nonoverlapping(src, dst, n);
}
Ok(())
})()
.map_err(PyValueError::new_err);
unsafe { ffi::PyBuffer_Release(&mut buf) };
result
}
#[pyfunction]
fn rans_byte() -> i32 {
1
}
#[pyfunction]
fn rans_64() -> i32 {
0
}
enum PyEncoderStream {
None,
Byte(CoreEncoderStream<RansByte>),
S64(CoreEncoderStream<Rans64>),
}
#[pyclass(name = "RansEncoderStream")]
struct RansEncoderStream {
stream: PyEncoderStream,
#[allow(dead_code)]
variant: i32,
#[allow(dead_code)]
_initial_size: usize,
#[allow(dead_code)]
_max_size_step: usize,
}
impl RansEncoderStream {
fn push_byte(
&mut self,
encoder: &EntropyEncoder<RansByte>,
indices: &[i32],
values: &[i32],
) -> PyResult<()> {
match &mut self.stream {
PyEncoderStream::Byte(s) => s
.push(encoder, indices, values)
.map_err(|e| PyValueError::new_err(format!("encode failed: {}", e))),
PyEncoderStream::S64(_) => Err(PyValueError::new_err(
"encoder stream variant mismatch: stream is Rans64, encoder is RansByte",
)),
PyEncoderStream::None => Err(PyValueError::new_err("invalid state")),
}
}
fn push_64(
&mut self,
encoder: &EntropyEncoder<Rans64>,
indices: &[i32],
values: &[i32],
) -> PyResult<()> {
match &mut self.stream {
PyEncoderStream::S64(s) => s
.push(encoder, indices, values)
.map_err(|e| PyValueError::new_err(format!("encode failed: {}", e))),
PyEncoderStream::Byte(_) => Err(PyValueError::new_err(
"encoder stream variant mismatch: stream is RansByte, encoder is Rans64",
)),
PyEncoderStream::None => Err(PyValueError::new_err("invalid state")),
}
}
}
#[pymethods]
impl RansEncoderStream {
#[new]
#[pyo3(signature = (variant=1, *, initialSize=4096, maxSizeStep=1048576))]
fn new(variant: i32, initialSize: usize, maxSizeStep: usize) -> PyResult<Self> {
let stream = match variant {
1 => PyEncoderStream::Byte(CoreEncoderStream::new()),
0 => PyEncoderStream::S64(CoreEncoderStream::new()),
_ => {
return Err(PyValueError::new_err(format!(
"unknown rANS variant value: {}",
variant
)));
}
};
Ok(Self {
stream,
variant,
_initial_size: initialSize,
_max_size_step: maxSizeStep,
})
}
fn flush(&mut self, py: Python<'_>) -> PyResult<Py<PyAny>> {
let data: Vec<u8> = match &mut self.stream {
PyEncoderStream::Byte(s) => s
.flush()
.map_err(|e| PyValueError::new_err(format!("flush failed: {}", e)))?,
PyEncoderStream::S64(s) => s
.flush()
.map_err(|e| PyValueError::new_err(format!("flush failed: {}", e)))?,
PyEncoderStream::None => {
return Err(PyValueError::new_err(
"invalid state: stream not initialized",
));
}
};
if data.is_empty() {
return Err(PyValueError::new_err("invalid state: empty output"));
}
let ptr = unsafe {
ffi::PyBytes_FromStringAndSize(
data.as_ptr() as *const ffi::Py_ssize_t as *const i8,
data.len() as ffi::Py_ssize_t,
)
};
if ptr.is_null() {
return Err(PyValueError::new_err("failed to create PyBytes"));
}
let obj: Py<PyAny> = unsafe { Bound::from_owned_ptr(py, ptr).unbind() };
Ok(obj)
}
fn reset(&mut self) {
match &mut self.stream {
PyEncoderStream::Byte(s) => s.reset(),
PyEncoderStream::S64(s) => s.reset(),
PyEncoderStream::None => {}
}
}
}
enum PyDecoderStream {
None,
Byte(CoreDecoderStream<RansByte>),
S64(CoreDecoderStream<Rans64>),
}
#[pyclass(name = "RansDecoderStream")]
struct RansDecoderStream {
stream: PyDecoderStream,
#[allow(dead_code)]
variant: i32,
}
impl RansDecoderStream {
fn byte_stream_mut(&mut self) -> PyResult<&mut CoreDecoderStream<RansByte>> {
match &mut self.stream {
PyDecoderStream::Byte(s) => Ok(s),
PyDecoderStream::S64(_) => Err(PyValueError::new_err(
"decoder stream variant mismatch: stream is Rans64",
)),
PyDecoderStream::None => Err(PyValueError::new_err("decoder stream is not open")),
}
}
fn s64_stream_mut(&mut self) -> PyResult<&mut CoreDecoderStream<Rans64>> {
match &mut self.stream {
PyDecoderStream::S64(s) => Ok(s),
PyDecoderStream::Byte(_) => Err(PyValueError::new_err(
"decoder stream variant mismatch: stream is RansByte",
)),
PyDecoderStream::None => Err(PyValueError::new_err("decoder stream is not open")),
}
}
}
#[pymethods]
impl RansDecoderStream {
#[new]
#[pyo3(signature = (data=None, *, variant=1))]
fn new(data: Option<Bound<'_, PyAny>>, variant: i32) -> PyResult<Self> {
let stream = match variant {
1 => match data {
Some(ref obj) => {
PyDecoderStream::Byte(CoreDecoderStream::open_on(&get_u8_buffer(obj)?))
}
None => PyDecoderStream::Byte(CoreDecoderStream::new()),
},
0 => match data {
Some(ref obj) => {
PyDecoderStream::S64(CoreDecoderStream::open_on(&get_u8_buffer(obj)?))
}
None => PyDecoderStream::S64(CoreDecoderStream::new()),
},
_ => {
return Err(PyValueError::new_err(format!(
"unknown rANS variant value: {}",
variant
)));
}
};
Ok(Self { stream, variant })
}
fn open(&mut self, data: Bound<'_, PyAny>) -> PyResult<()> {
let bytes = get_u8_buffer(&data)?;
match self.variant {
1 => {
self.stream = PyDecoderStream::Byte(CoreDecoderStream::open_on(&bytes));
}
0 => {
self.stream = PyDecoderStream::S64(CoreDecoderStream::open_on(&bytes));
}
_ => return Err(PyValueError::new_err("unknown rANS variant value")),
}
Ok(())
}
fn close(&mut self) {
self.stream = PyDecoderStream::None;
}
#[pyo3(name = "isOpen")]
fn is_open(&self) -> bool {
!matches!(self.stream, PyDecoderStream::None)
}
#[pyo3(name = "decodeEOF")]
fn decode_eof(&mut self) -> PyResult<()> {
let result = match &mut self.stream {
PyDecoderStream::Byte(s) => s.decode_eof(),
PyDecoderStream::S64(s) => s.decode_eof(),
PyDecoderStream::None => {
return Err(PyValueError::new_err("decoder stream is not open"));
}
};
result.map_err(|e| PyValueError::new_err(format!("decodeEOF failed: {}", e)))?;
self.stream = PyDecoderStream::None;
Ok(())
}
}
#[pyclass(name = "EntropyEncoder")]
struct PyEntropyEncoder {
byte_encoder: Option<EntropyEncoder<RansByte>>,
_64_encoder: Option<EntropyEncoder<Rans64>>,
variant: i32,
}
#[pymethods]
impl PyEntropyEncoder {
#[new]
#[pyo3(signature = (*, pmfLengths, pmfOffsets, pmfTable, variant=1, symbolBits=16, bypassBits=4))]
fn new(
pmfLengths: Bound<'_, PyAny>,
pmfOffsets: Bound<'_, PyAny>,
pmfTable: Bound<'_, PyAny>,
variant: i32,
symbolBits: u32,
bypassBits: u32,
) -> PyResult<Self> {
let lengths = get_i32_buffer(&pmfLengths)?;
let offsets = get_i32_buffer(&pmfOffsets)?;
let table = get_i32_buffer(&pmfTable)?;
match variant {
1 => {
let mut encoder = EntropyEncoder::<RansByte>::new();
encoder
.initialize(&lengths, &offsets, &table, symbolBits, bypassBits)
.map_err(|e| PyValueError::new_err(format!("encoder init failed: {}", e)))?;
Ok(Self {
byte_encoder: Some(encoder),
_64_encoder: None,
variant,
})
}
0 => {
let mut encoder = EntropyEncoder::<Rans64>::new();
encoder
.initialize(&lengths, &offsets, &table, symbolBits, bypassBits)
.map_err(|e| PyValueError::new_err(format!("encoder init failed: {}", e)))?;
Ok(Self {
byte_encoder: None,
_64_encoder: Some(encoder),
variant,
})
}
_ => Err(PyValueError::new_err(format!(
"invalid variant: {}",
variant
))),
}
}
#[pyo3(signature = (stream, indices, values))]
fn encode(
&self,
stream: &mut RansEncoderStream,
indices: Bound<'_, PyAny>,
values: Bound<'_, PyAny>,
) -> PyResult<()> {
let indices_vec = get_i32_buffer(&indices)?;
let values_vec = get_i32_buffer(&values)?;
if indices_vec.len() != values_vec.len() {
return Err(PyValueError::new_err(
"indices and values must have the same length",
));
}
match self.variant {
1 => {
if let Some(ref encoder) = self.byte_encoder {
stream.push_byte(encoder, &indices_vec, &values_vec)
} else {
Err(PyValueError::new_err("byte encoder not initialized"))
}
}
0 => {
if let Some(ref encoder) = self._64_encoder {
stream.push_64(encoder, &indices_vec, &values_vec)
} else {
Err(PyValueError::new_err("64 encoder not initialized"))
}
}
_ => Err(PyValueError::new_err("invalid variant")),
}
}
}
#[pyclass(name = "EntropyDecoder")]
struct PyEntropyDecoder {
byte_decoder: Option<EntropyDecoder<RansByte>>,
_64_decoder: Option<EntropyDecoder<Rans64>>,
variant: i32,
}
#[pymethods]
impl PyEntropyDecoder {
#[new]
#[pyo3(signature = (*, pmfLengths, pmfOffsets, pmfTable, variant=1, symbolBits=16, bypassBits=4))]
fn new(
pmfLengths: Bound<'_, PyAny>,
pmfOffsets: Bound<'_, PyAny>,
pmfTable: Bound<'_, PyAny>,
variant: i32,
symbolBits: u32,
bypassBits: u32,
) -> PyResult<Self> {
let lengths = get_i32_buffer(&pmfLengths)?;
let offsets = get_i32_buffer(&pmfOffsets)?;
let table = get_i32_buffer(&pmfTable)?;
match variant {
1 => {
let mut decoder = EntropyDecoder::<RansByte>::new();
decoder
.initialize(&lengths, &offsets, &table, symbolBits, bypassBits)
.map_err(|e| PyValueError::new_err(format!("decoder init failed: {}", e)))?;
Ok(Self {
byte_decoder: Some(decoder),
_64_decoder: None,
variant,
})
}
0 => {
let mut decoder = EntropyDecoder::<Rans64>::new();
decoder
.initialize(&lengths, &offsets, &table, symbolBits, bypassBits)
.map_err(|e| PyValueError::new_err(format!("decoder init failed: {}", e)))?;
Ok(Self {
byte_decoder: None,
_64_decoder: Some(decoder),
variant,
})
}
_ => Err(PyValueError::new_err(format!(
"invalid variant: {}",
variant
))),
}
}
#[pyo3(signature = (values, indices, data))]
fn decode(
&self,
py: Python<'_>,
values: Bound<'_, PyAny>,
indices: Bound<'_, PyAny>,
data: Bound<'_, PyAny>,
) -> PyResult<()> {
let indices_vec = get_i32_buffer(&indices)?;
let num_values = indices_vec.len();
if let Ok(py_stream) = data.extract::<Py<RansDecoderStream>>() {
let mut stream_ref = py_stream.borrow_mut(py);
let mut decoded = vec![0i32; num_values];
match self.variant {
1 => {
if let Some(ref decoder) = self.byte_decoder {
let core = stream_ref.byte_stream_mut()?;
core.decode(decoder, &mut decoded, &indices_vec)
.map_err(|e| PyValueError::new_err(format!("decode failed: {}", e)))?;
} else {
return Err(PyValueError::new_err("byte decoder not initialized"));
}
}
0 => {
if let Some(ref decoder) = self._64_decoder {
let core = stream_ref.s64_stream_mut()?;
core.decode(decoder, &mut decoded, &indices_vec)
.map_err(|e| PyValueError::new_err(format!("decode failed: {}", e)))?;
} else {
return Err(PyValueError::new_err("64 decoder not initialized"));
}
}
_ => return Err(PyValueError::new_err("invalid variant")),
}
drop(stream_ref);
write_i32_buffer(&values, &decoded)?;
return Ok(());
}
let data_vec = get_u8_buffer(&data)?;
let mut decoded = vec![0i32; num_values];
match self.variant {
1 => {
if let Some(ref decoder) = self.byte_decoder {
decoder
.decode(&mut decoded, &indices_vec, &data_vec)
.map_err(|e| PyValueError::new_err(format!("decode failed: {}", e)))?;
} else {
return Err(PyValueError::new_err("byte decoder not initialized"));
}
}
0 => {
if let Some(ref decoder) = self._64_decoder {
decoder
.decode(&mut decoded, &indices_vec, &data_vec)
.map_err(|e| PyValueError::new_err(format!("decode failed: {}", e)))?;
} else {
return Err(PyValueError::new_err("64 decoder not initialized"));
}
}
_ => return Err(PyValueError::new_err("invalid variant")),
}
write_i32_buffer(&values, &decoded)
}
}
#[pymodule]
fn _msrtc_rans(_py: Python, m: &Bound<'_, PyModule>) -> PyResult<()> {
m.add("__version__", env!("CARGO_PKG_VERSION"))?;
m.add("RansByte", 1)?; m.add("Rans64", 0)?; m.add_function(wrap_pyfunction!(rans_byte, m)?)?;
m.add_function(wrap_pyfunction!(rans_64, m)?)?;
m.add_class::<RansEncoderStream>()?;
m.add_class::<RansDecoderStream>()?;
m.add_class::<PyEntropyEncoder>()?;
m.add_class::<PyEntropyDecoder>()?;
Ok(())
}