Skip to main content

_msrtc_rans/
lib.rs

1// Licensed under the MIT license.
2// Author: Riaan de Beer - github.com/infinityabundance - rdebeer.infinityabundance@gmail.com
3
4//! # msrtc-rans-python
5//!
6//! Python extension module `_msrtc_rans` providing the msrtc.rans Python API.
7//!
8//! This crate uses PyO3 to create a CPython extension module that is
9//! import-path-compatible with the existing `msrtc.rans` package.
10
11#![allow(missing_docs)]
12#![allow(unsafe_op_in_unsafe_fn)]
13
14use pyo3::exceptions::PyValueError;
15use pyo3::ffi;
16use pyo3::prelude::*;
17use pyo3::types::PyByteArray;
18use pyo3::types::PyBytes;
19
20use msrtc_rans::entropy::{EntropyDecoder, EntropyEncoder};
21use msrtc_rans::variant::Rans64;
22use msrtc_rans::variant::RansByte;
23
24// ---------------------------------------------------------------------------
25// FFI-based buffer helpers using PyObject_GetBuffer (part of Python 3.11+ API)
26//
27// PyBUF flags (numeric for safety across PyO3 versions):
28//   READ  = PyBUF_ND | PyBUF_FORMAT = 8 | 4 = 12  (PyBUF_CONTIG)
29//   WRITE = READ | PyBUF_WRITABLE  = 12 | 1 = 13
30// ---------------------------------------------------------------------------
31
32const BUF_READ: i32 = 12; // PyBUF_ND | PyBUF_FORMAT = PyBUF_CONTIG
33const BUF_WRITE: i32 = 13; // PyBUF_CONTIG | PyBUF_WRITABLE
34
35unsafe fn buffer_to_i32_slice(buf: &ffi::Py_buffer) -> &[i32] {
36    let n = (buf.len as usize) / 4;
37    std::slice::from_raw_parts(buf.buf as *const i32, n)
38}
39
40unsafe fn buffer_to_u8_slice(buf: &ffi::Py_buffer) -> &[u8] {
41    let n = buf.len as usize;
42    std::slice::from_raw_parts(buf.buf as *const u8, n)
43}
44
45fn get_i32_buffer(obj: &Bound<'_, PyAny>) -> PyResult<Vec<i32>> {
46    let mut buf: ffi::Py_buffer = unsafe { std::mem::zeroed() };
47    let ret = unsafe { ffi::PyObject_GetBuffer(obj.as_ptr(), &mut buf, BUF_READ) };
48    if ret != 0 {
49        return Err(PyValueError::new_err("cannot get i32 buffer from object"));
50    }
51    let vec = unsafe { buffer_to_i32_slice(&buf).to_vec() };
52    unsafe { ffi::PyBuffer_Release(&mut buf) };
53    Ok(vec)
54}
55
56fn get_u8_buffer(obj: &Bound<'_, PyAny>) -> PyResult<Vec<u8>> {
57    if let Ok(bytes) = obj.downcast::<PyBytes>() {
58        return Ok(bytes.as_bytes().to_vec());
59    }
60    if let Ok(ba) = obj.downcast::<PyByteArray>() {
61        let slice = unsafe { ba.as_bytes() };
62        return Ok(slice.to_vec());
63    }
64    let mut buf: ffi::Py_buffer = unsafe { std::mem::zeroed() };
65    let ret = unsafe { ffi::PyObject_GetBuffer(obj.as_ptr(), &mut buf, BUF_READ) };
66    if ret != 0 {
67        return Err(PyValueError::new_err("cannot get buffer from object"));
68    }
69    let vec = unsafe { buffer_to_u8_slice(&buf).to_vec() };
70    unsafe { ffi::PyBuffer_Release(&mut buf) };
71    Ok(vec)
72}
73
74fn write_i32_buffer(obj: &Bound<'_, PyAny>, data: &[i32]) -> PyResult<()> {
75    let mut buf: ffi::Py_buffer = unsafe { std::mem::zeroed() };
76    let ret = unsafe { ffi::PyObject_GetBuffer(obj.as_ptr(), &mut buf, BUF_WRITE) };
77    if ret != 0 {
78        return Err(PyValueError::new_err(
79            "cannot get writable i32 buffer from object",
80        ));
81    }
82    let n = data.len() * 4;
83    let copy_len = n.min(buf.len as usize);
84    unsafe {
85        let dst = buf.buf as *mut u8;
86        let src = data.as_ptr() as *const u8;
87        std::ptr::copy_nonoverlapping(src, dst, copy_len);
88        ffi::PyBuffer_Release(&mut buf);
89    }
90    Ok(())
91}
92
93// ---------------------------------------------------------------------------
94// Module-level constants
95// ---------------------------------------------------------------------------
96
97#[pyfunction]
98fn rans_byte() -> i32 {
99    1
100}
101
102#[pyfunction]
103fn rans_64() -> i32 {
104    0
105}
106
107// ---------------------------------------------------------------------------
108// RansEncoderStream
109// ---------------------------------------------------------------------------
110
111#[pyclass(name = "RansEncoderStream")]
112struct RansEncoderStream {
113    segments: Vec<Vec<u8>>,
114    #[allow(dead_code)]
115    variant: i32,
116    #[allow(dead_code)]
117    _initial_size: usize,
118    #[allow(dead_code)]
119    _max_size_step: usize,
120}
121
122#[pymethods]
123impl RansEncoderStream {
124    #[new]
125    #[pyo3(signature = (variant=1, *, initialSize=4096, maxSizeStep=1048576))]
126    fn new(variant: i32, initialSize: usize, maxSizeStep: usize) -> Self {
127        Self {
128            segments: Vec::new(),
129            variant,
130            _initial_size: initialSize,
131            _max_size_step: maxSizeStep,
132        }
133    }
134
135    fn flush(&mut self, py: Python<'_>) -> PyResult<Py<PyAny>> {
136        let total_len: usize = self.segments.iter().map(|s| s.len()).sum();
137        let mut buffer = Vec::with_capacity(total_len);
138        for segment in self.segments.iter().rev() {
139            buffer.extend_from_slice(segment);
140        }
141        self.segments.clear();
142        // Use FFI to create PyBytes (PyBytes::new not available in PyO3 0.22 abi3)
143        let ptr = unsafe {
144            ffi::PyBytes_FromStringAndSize(
145                buffer.as_ptr() as *const ffi::Py_ssize_t as *const i8,
146                buffer.len() as ffi::Py_ssize_t,
147            )
148        };
149        if ptr.is_null() {
150            return Err(PyValueError::new_err("failed to create PyBytes"));
151        }
152        let obj: Py<PyAny> = unsafe { Bound::from_owned_ptr(py, ptr).unbind() };
153        Ok(obj)
154    }
155
156    fn reset(&mut self) {
157        self.segments.clear();
158    }
159}
160
161// ---------------------------------------------------------------------------
162// RansDecoderStream
163// ---------------------------------------------------------------------------
164
165#[pyclass(name = "RansDecoderStream")]
166struct RansDecoderStream {
167    data: Option<Vec<u8>>,
168    offset: usize,
169    #[allow(dead_code)]
170    _variant: i32,
171}
172
173#[pymethods]
174impl RansDecoderStream {
175    #[new]
176    #[pyo3(signature = (data=None, *, variant=1))]
177    fn new(data: Option<Bound<'_, PyAny>>, variant: i32) -> PyResult<Self> {
178        let vec = match data {
179            Some(ref obj) => Some(get_u8_buffer(obj)?),
180            None => None,
181        };
182        Ok(Self {
183            data: vec,
184            offset: 0,
185            _variant: variant,
186        })
187    }
188
189    fn open(&mut self, data: Bound<'_, PyAny>) -> PyResult<()> {
190        self.data = Some(get_u8_buffer(&data)?);
191        self.offset = 0;
192        Ok(())
193    }
194
195    fn close(&mut self) {
196        self.data = None;
197        self.offset = 0;
198    }
199
200    #[pyo3(name = "isOpen")]
201    fn is_open(&self) -> bool {
202        self.data.is_some()
203    }
204
205    #[pyo3(name = "decodeEOF")]
206    fn decode_eof(&mut self) -> PyResult<()> {
207        if let Some(ref data) = self.data {
208            if self.offset != data.len() {
209                return Err(PyValueError::new_err(format!(
210                    "decodeEOF: stream not fully consumed (offset={}, len={})",
211                    self.offset,
212                    data.len()
213                )));
214            }
215        }
216        self.close();
217        Ok(())
218    }
219}
220
221// ---------------------------------------------------------------------------
222// EntropyEncoder
223// ---------------------------------------------------------------------------
224
225#[pyclass(name = "EntropyEncoder")]
226struct PyEntropyEncoder {
227    byte_encoder: Option<EntropyEncoder<RansByte>>,
228    _64_encoder: Option<EntropyEncoder<Rans64>>,
229    variant: i32,
230}
231
232#[pymethods]
233impl PyEntropyEncoder {
234    #[new]
235    #[pyo3(signature = (*, pmfLengths, pmfOffsets, pmfTable, variant=1, symbolBits=16, bypassBits=4))]
236    fn new(
237        pmfLengths: Bound<'_, PyAny>,
238        pmfOffsets: Bound<'_, PyAny>,
239        pmfTable: Bound<'_, PyAny>,
240        variant: i32,
241        symbolBits: u32,
242        bypassBits: u32,
243    ) -> PyResult<Self> {
244        let lengths = get_i32_buffer(&pmfLengths)?;
245        let offsets = get_i32_buffer(&pmfOffsets)?;
246        let table = get_i32_buffer(&pmfTable)?;
247
248        match variant {
249            1 => {
250                let mut encoder = EntropyEncoder::<RansByte>::new();
251                encoder
252                    .initialize(&lengths, &offsets, &table, symbolBits, bypassBits)
253                    .map_err(|e| PyValueError::new_err(format!("encoder init failed: {}", e)))?;
254                Ok(Self {
255                    byte_encoder: Some(encoder),
256                    _64_encoder: None,
257                    variant,
258                })
259            }
260            0 => {
261                let mut encoder = EntropyEncoder::<Rans64>::new();
262                encoder
263                    .initialize(&lengths, &offsets, &table, symbolBits, bypassBits)
264                    .map_err(|e| PyValueError::new_err(format!("encoder init failed: {}", e)))?;
265                Ok(Self {
266                    byte_encoder: None,
267                    _64_encoder: Some(encoder),
268                    variant,
269                })
270            }
271            _ => Err(PyValueError::new_err(format!(
272                "invalid variant: {}",
273                variant
274            ))),
275        }
276    }
277
278    #[pyo3(signature = (stream, indices, values))]
279    fn encode(
280        &self,
281        stream: &mut RansEncoderStream,
282        indices: Bound<'_, PyAny>,
283        values: Bound<'_, PyAny>,
284    ) -> PyResult<()> {
285        let indices_vec = get_i32_buffer(&indices)?;
286        let values_vec = get_i32_buffer(&values)?;
287
288        match self.variant {
289            1 => {
290                if let Some(ref encoder) = self.byte_encoder {
291                    let mut buffer = Vec::new();
292                    encoder
293                        .encode(&indices_vec, &values_vec, &mut buffer)
294                        .map_err(|e| PyValueError::new_err(format!("encode failed: {}", e)))?;
295                    stream.segments.push(buffer);
296                    Ok(())
297                } else {
298                    Err(PyValueError::new_err("byte encoder not initialized"))
299                }
300            }
301            0 => {
302                if let Some(ref encoder) = self._64_encoder {
303                    let mut buffer = Vec::new();
304                    encoder
305                        .encode(&indices_vec, &values_vec, &mut buffer)
306                        .map_err(|e| PyValueError::new_err(format!("encode failed: {}", e)))?;
307                    stream.segments.push(buffer);
308                    Ok(())
309                } else {
310                    Err(PyValueError::new_err("64 encoder not initialized"))
311                }
312            }
313            _ => Err(PyValueError::new_err("invalid variant")),
314        }
315    }
316}
317
318// ---------------------------------------------------------------------------
319// EntropyDecoder
320// ---------------------------------------------------------------------------
321
322#[pyclass(name = "EntropyDecoder")]
323struct PyEntropyDecoder {
324    byte_decoder: Option<EntropyDecoder<RansByte>>,
325    _64_decoder: Option<EntropyDecoder<Rans64>>,
326    variant: i32,
327}
328
329#[pymethods]
330impl PyEntropyDecoder {
331    #[new]
332    #[pyo3(signature = (*, pmfLengths, pmfOffsets, pmfTable, variant=1, symbolBits=16, bypassBits=4))]
333    fn new(
334        pmfLengths: Bound<'_, PyAny>,
335        pmfOffsets: Bound<'_, PyAny>,
336        pmfTable: Bound<'_, PyAny>,
337        variant: i32,
338        symbolBits: u32,
339        bypassBits: u32,
340    ) -> PyResult<Self> {
341        let lengths = get_i32_buffer(&pmfLengths)?;
342        let offsets = get_i32_buffer(&pmfOffsets)?;
343        let table = get_i32_buffer(&pmfTable)?;
344
345        match variant {
346            1 => {
347                let mut decoder = EntropyDecoder::<RansByte>::new();
348                decoder
349                    .initialize(&lengths, &offsets, &table, symbolBits, bypassBits)
350                    .map_err(|e| PyValueError::new_err(format!("decoder init failed: {}", e)))?;
351                Ok(Self {
352                    byte_decoder: Some(decoder),
353                    _64_decoder: None,
354                    variant,
355                })
356            }
357            0 => {
358                let mut decoder = EntropyDecoder::<Rans64>::new();
359                decoder
360                    .initialize(&lengths, &offsets, &table, symbolBits, bypassBits)
361                    .map_err(|e| PyValueError::new_err(format!("decoder init failed: {}", e)))?;
362                Ok(Self {
363                    byte_decoder: None,
364                    _64_decoder: Some(decoder),
365                    variant,
366                })
367            }
368            _ => Err(PyValueError::new_err(format!(
369                "invalid variant: {}",
370                variant
371            ))),
372        }
373    }
374
375    #[pyo3(signature = (values, indices, data))]
376    fn decode(
377        &self,
378        py: Python<'_>,
379        values: Bound<'_, PyAny>,
380        indices: Bound<'_, PyAny>,
381        data: Bound<'_, PyAny>,
382    ) -> PyResult<()> {
383        let indices_vec = get_i32_buffer(&indices)?;
384        let num_values = indices_vec.len();
385
386        // Try stream-based decode
387        if let Ok(py_stream) = data.extract::<Py<RansDecoderStream>>() {
388            let mut stream_ref = py_stream.borrow_mut(py);
389            let stream_data = stream_ref
390                .data
391                .as_ref()
392                .ok_or_else(|| PyValueError::new_err("RansDecoderStream is not open"))?;
393            let remaining = stream_data[stream_ref.offset..].to_vec();
394            let current_offset = stream_ref.offset;
395
396            let mut decoded = vec![0i32; num_values];
397
398            let consumed = match self.variant {
399                1 => {
400                    if let Some(ref decoder) = self.byte_decoder {
401                        decoder
402                            .decode_partial(&mut decoded, &indices_vec, &remaining)
403                            .map_err(|e| PyValueError::new_err(format!("decode failed: {}", e)))?
404                    } else {
405                        return Err(PyValueError::new_err("byte decoder not initialized"));
406                    }
407                }
408                0 => {
409                    if let Some(ref decoder) = self._64_decoder {
410                        decoder
411                            .decode_partial(&mut decoded, &indices_vec, &remaining)
412                            .map_err(|e| PyValueError::new_err(format!("decode failed: {}", e)))?
413                    } else {
414                        return Err(PyValueError::new_err("64 decoder not initialized"));
415                    }
416                }
417                _ => return Err(PyValueError::new_err("invalid variant")),
418            };
419
420            stream_ref.offset = current_offset + consumed;
421            drop(stream_ref);
422
423            write_i32_buffer(&values, &decoded)?;
424            return Ok(());
425        }
426
427        // Buffer-based decode
428        let data_vec = get_u8_buffer(&data)?;
429        let mut decoded = vec![0i32; num_values];
430
431        match self.variant {
432            1 => {
433                if let Some(ref decoder) = self.byte_decoder {
434                    decoder
435                        .decode(&mut decoded, &indices_vec, &data_vec)
436                        .map_err(|e| PyValueError::new_err(format!("decode failed: {}", e)))?;
437                } else {
438                    return Err(PyValueError::new_err("byte decoder not initialized"));
439                }
440            }
441            0 => {
442                if let Some(ref decoder) = self._64_decoder {
443                    decoder
444                        .decode(&mut decoded, &indices_vec, &data_vec)
445                        .map_err(|e| PyValueError::new_err(format!("decode failed: {}", e)))?;
446                } else {
447                    return Err(PyValueError::new_err("64 decoder not initialized"));
448                }
449            }
450            _ => return Err(PyValueError::new_err("invalid variant")),
451        }
452
453        write_i32_buffer(&values, &decoded)
454    }
455}
456
457// ---------------------------------------------------------------------------
458// Module definition
459// ---------------------------------------------------------------------------
460
461#[pymodule]
462fn _msrtc_rans(_py: Python, m: &Bound<'_, PyModule>) -> PyResult<()> {
463    m.add("__version__", env!("CARGO_PKG_VERSION"))?;
464    // Module-level constants matching C++ enum names used by types.py
465    m.add("RansByte", 1)?; // RansVariant.RansByte = 1
466    m.add("Rans64", 0)?; // RansVariant.Rans64 = 0
467    // Also add as functions for alternative access
468    m.add_function(wrap_pyfunction!(rans_byte, m)?)?;
469    m.add_function(wrap_pyfunction!(rans_64, m)?)?;
470    m.add_class::<RansEncoderStream>()?;
471    m.add_class::<RansDecoderStream>()?;
472    m.add_class::<PyEntropyEncoder>()?;
473    m.add_class::<PyEntropyDecoder>()?;
474    Ok(())
475}