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//! Stream classes wrap the persistent Rust stream types from
12//! `msrtc_rans::stream`, matching Microsoft's `RansEncoderStream` /
13//! `RansDecoderStream` semantics:
14//!
15//! - `RansEncoderStream` keeps a single persistent raw rANS encoder state
16//!   across `push()` calls and flushes it once (`Flush(abort=false)`),
17//!   exactly like Microsoft's `RansEncoderStreamImpl`.
18//! - `RansDecoderStream` owns the encoded message and keeps a persistent
19//!   decode cursor (unit position + rANS state) across `decode()` calls,
20//!   matching Microsoft's `RansDecoderStreamImpl`.
21
22#![allow(missing_docs)]
23#![allow(unsafe_op_in_unsafe_fn)]
24
25use pyo3::exceptions::PyValueError;
26use pyo3::ffi;
27use pyo3::prelude::*;
28use pyo3::types::PyByteArray;
29use pyo3::types::PyBytes;
30
31use msrtc_rans::entropy::{EntropyDecoder, EntropyEncoder};
32use msrtc_rans::stream::RansDecoderStream as CoreDecoderStream;
33use msrtc_rans::stream::RansEncoderStream as CoreEncoderStream;
34use msrtc_rans::variant::Rans64;
35use msrtc_rans::variant::RansByte;
36
37// ---------------------------------------------------------------------------
38// FFI-based buffer helpers using PyObject_GetBuffer (part of Python 3.11+ API)
39//
40// PyBUF flags (numeric for safety across PyO3 versions):
41//   READ  = PyBUF_ND | PyBUF_FORMAT = 8 | 4 = 12  (PyBUF_CONTIG)
42//   WRITE = READ | PyBUF_WRITABLE  = 12 | 1 = 13
43// ---------------------------------------------------------------------------
44
45const BUF_READ: i32 = 12; // PyBUF_ND | PyBUF_FORMAT = PyBUF_CONTIG
46const BUF_WRITE: i32 = 13; // PyBUF_CONTIG | PyBUF_WRITABLE
47
48unsafe fn buffer_to_i32_slice(buf: &ffi::Py_buffer) -> &[i32] {
49    let n = (buf.len as usize) / 4;
50    std::slice::from_raw_parts(buf.buf as *const i32, n)
51}
52
53unsafe fn buffer_to_u8_slice(buf: &ffi::Py_buffer) -> &[u8] {
54    let n = buf.len as usize;
55    std::slice::from_raw_parts(buf.buf as *const u8, n)
56}
57
58fn get_i32_buffer(obj: &Bound<'_, PyAny>) -> PyResult<Vec<i32>> {
59    let mut buf: ffi::Py_buffer = unsafe { std::mem::zeroed() };
60    let ret = unsafe { ffi::PyObject_GetBuffer(obj.as_ptr(), &mut buf, BUF_READ) };
61    if ret != 0 {
62        return Err(PyValueError::new_err("cannot get i32 buffer from object"));
63    }
64    let vec = unsafe { buffer_to_i32_slice(&buf).to_vec() };
65    unsafe { ffi::PyBuffer_Release(&mut buf) };
66    Ok(vec)
67}
68
69fn get_u8_buffer(obj: &Bound<'_, PyAny>) -> PyResult<Vec<u8>> {
70    if let Ok(bytes) = obj.downcast::<PyBytes>() {
71        return Ok(bytes.as_bytes().to_vec());
72    }
73    if let Ok(ba) = obj.downcast::<PyByteArray>() {
74        let slice = unsafe { ba.as_bytes() };
75        return Ok(slice.to_vec());
76    }
77    let mut buf: ffi::Py_buffer = unsafe { std::mem::zeroed() };
78    let ret = unsafe { ffi::PyObject_GetBuffer(obj.as_ptr(), &mut buf, BUF_READ) };
79    if ret != 0 {
80        return Err(PyValueError::new_err("cannot get buffer from object"));
81    }
82    let vec = unsafe { buffer_to_u8_slice(&buf).to_vec() };
83    unsafe { ffi::PyBuffer_Release(&mut buf) };
84    Ok(vec)
85}
86
87fn write_i32_buffer(obj: &Bound<'_, PyAny>, data: &[i32]) -> PyResult<()> {
88    let mut buf: ffi::Py_buffer = unsafe { std::mem::zeroed() };
89    let ret = unsafe { ffi::PyObject_GetBuffer(obj.as_ptr(), &mut buf, BUF_WRITE) };
90    if ret != 0 {
91        return Err(PyValueError::new_err(
92            "cannot get writable i32 buffer from object",
93        ));
94    }
95    let n = data.len() * 4;
96    let copy_len = n.min(buf.len as usize);
97    unsafe {
98        let dst = buf.buf as *mut u8;
99        let src = data.as_ptr() as *const u8;
100        std::ptr::copy_nonoverlapping(src, dst, copy_len);
101        ffi::PyBuffer_Release(&mut buf);
102    }
103    Ok(())
104}
105
106// ---------------------------------------------------------------------------
107// Module-level constants
108// ---------------------------------------------------------------------------
109
110#[pyfunction]
111fn rans_byte() -> i32 {
112    1
113}
114
115#[pyfunction]
116fn rans_64() -> i32 {
117    0
118}
119
120// ---------------------------------------------------------------------------
121// RansEncoderStream
122// ---------------------------------------------------------------------------
123
124/// Persistent raw encoder held by the Python `RansEncoderStream`.
125///
126/// Matches Microsoft's `RawRansEncoderStream`: one raw rANS encoder state
127/// persists across `push()` calls; `flush()` finalizes it once.
128enum PyEncoderStream {
129    None,
130    Byte(CoreEncoderStream<RansByte>),
131    S64(CoreEncoderStream<Rans64>),
132}
133
134#[pyclass(name = "RansEncoderStream")]
135struct RansEncoderStream {
136    stream: PyEncoderStream,
137    #[allow(dead_code)]
138    variant: i32,
139    #[allow(dead_code)]
140    _initial_size: usize,
141    #[allow(dead_code)]
142    _max_size_step: usize,
143}
144
145impl RansEncoderStream {
146    /// Push a batch encoded with a RansByte entropy encoder.
147    fn push_byte(
148        &mut self,
149        encoder: &EntropyEncoder<RansByte>,
150        indices: &[i32],
151        values: &[i32],
152    ) -> PyResult<()> {
153        match &mut self.stream {
154            PyEncoderStream::Byte(s) => s
155                .push(encoder, indices, values)
156                .map_err(|e| PyValueError::new_err(format!("encode failed: {}", e))),
157            PyEncoderStream::S64(_) => Err(PyValueError::new_err(
158                "encoder stream variant mismatch: stream is Rans64, encoder is RansByte",
159            )),
160            PyEncoderStream::None => Err(PyValueError::new_err("invalid state")),
161        }
162    }
163
164    /// Push a batch encoded with a Rans64 entropy encoder.
165    fn push_64(
166        &mut self,
167        encoder: &EntropyEncoder<Rans64>,
168        indices: &[i32],
169        values: &[i32],
170    ) -> PyResult<()> {
171        match &mut self.stream {
172            PyEncoderStream::S64(s) => s
173                .push(encoder, indices, values)
174                .map_err(|e| PyValueError::new_err(format!("encode failed: {}", e))),
175            PyEncoderStream::Byte(_) => Err(PyValueError::new_err(
176                "encoder stream variant mismatch: stream is RansByte, encoder is Rans64",
177            )),
178            PyEncoderStream::None => Err(PyValueError::new_err("invalid state")),
179        }
180    }
181}
182
183#[pymethods]
184impl RansEncoderStream {
185    #[new]
186    #[pyo3(signature = (variant=1, *, initialSize=4096, maxSizeStep=1048576))]
187    fn new(variant: i32, initialSize: usize, maxSizeStep: usize) -> PyResult<Self> {
188        let stream = match variant {
189            1 => PyEncoderStream::Byte(CoreEncoderStream::new()),
190            0 => PyEncoderStream::S64(CoreEncoderStream::new()),
191            _ => {
192                return Err(PyValueError::new_err(format!(
193                    "unknown rANS variant value: {}",
194                    variant
195                )));
196            }
197        };
198        Ok(Self {
199            stream,
200            variant,
201            _initial_size: initialSize,
202            _max_size_step: maxSizeStep,
203        })
204    }
205
206    fn flush(&mut self, py: Python<'_>) -> PyResult<Py<PyAny>> {
207        let data: Vec<u8> = match &mut self.stream {
208            PyEncoderStream::Byte(s) => s
209                .flush()
210                .map_err(|e| PyValueError::new_err(format!("flush failed: {}", e)))?,
211            PyEncoderStream::S64(s) => s
212                .flush()
213                .map_err(|e| PyValueError::new_err(format!("flush failed: {}", e)))?,
214            PyEncoderStream::None => {
215                return Err(PyValueError::new_err(
216                    "invalid state: stream not initialized",
217                ));
218            }
219        };
220
221        if data.is_empty() {
222            return Err(PyValueError::new_err("invalid state: empty output"));
223        }
224
225        let ptr = unsafe {
226            ffi::PyBytes_FromStringAndSize(
227                data.as_ptr() as *const ffi::Py_ssize_t as *const i8,
228                data.len() as ffi::Py_ssize_t,
229            )
230        };
231        if ptr.is_null() {
232            return Err(PyValueError::new_err("failed to create PyBytes"));
233        }
234        let obj: Py<PyAny> = unsafe { Bound::from_owned_ptr(py, ptr).unbind() };
235        Ok(obj)
236    }
237
238    fn reset(&mut self) {
239        match &mut self.stream {
240            PyEncoderStream::Byte(s) => s.reset(),
241            PyEncoderStream::S64(s) => s.reset(),
242            PyEncoderStream::None => {}
243        }
244    }
245}
246
247// ---------------------------------------------------------------------------
248// RansDecoderStream
249// ---------------------------------------------------------------------------
250
251/// Persistent decoder held by the Python `RansDecoderStream`.
252///
253/// Matches Microsoft's `RansDecoderStreamImpl`: the raw decoder is
254/// initialized on the first `decode()` and its cursor persists across
255/// subsequent calls.
256enum PyDecoderStream {
257    None,
258    Byte(CoreDecoderStream<RansByte>),
259    S64(CoreDecoderStream<Rans64>),
260}
261
262#[pyclass(name = "RansDecoderStream")]
263struct RansDecoderStream {
264    stream: PyDecoderStream,
265    #[allow(dead_code)]
266    variant: i32,
267}
268
269impl RansDecoderStream {
270    fn byte_stream_mut(&mut self) -> PyResult<&mut CoreDecoderStream<RansByte>> {
271        match &mut self.stream {
272            PyDecoderStream::Byte(s) => Ok(s),
273            PyDecoderStream::S64(_) => Err(PyValueError::new_err(
274                "decoder stream variant mismatch: stream is Rans64",
275            )),
276            PyDecoderStream::None => Err(PyValueError::new_err("decoder stream is not open")),
277        }
278    }
279
280    fn s64_stream_mut(&mut self) -> PyResult<&mut CoreDecoderStream<Rans64>> {
281        match &mut self.stream {
282            PyDecoderStream::S64(s) => Ok(s),
283            PyDecoderStream::Byte(_) => Err(PyValueError::new_err(
284                "decoder stream variant mismatch: stream is RansByte",
285            )),
286            PyDecoderStream::None => Err(PyValueError::new_err("decoder stream is not open")),
287        }
288    }
289}
290
291#[pymethods]
292impl RansDecoderStream {
293    #[new]
294    #[pyo3(signature = (data=None, *, variant=1))]
295    fn new(data: Option<Bound<'_, PyAny>>, variant: i32) -> PyResult<Self> {
296        let stream = match variant {
297            1 => match data {
298                Some(ref obj) => {
299                    PyDecoderStream::Byte(CoreDecoderStream::open_on(&get_u8_buffer(obj)?))
300                }
301                None => PyDecoderStream::Byte(CoreDecoderStream::new()),
302            },
303            0 => match data {
304                Some(ref obj) => {
305                    PyDecoderStream::S64(CoreDecoderStream::open_on(&get_u8_buffer(obj)?))
306                }
307                None => PyDecoderStream::S64(CoreDecoderStream::new()),
308            },
309            _ => {
310                return Err(PyValueError::new_err(format!(
311                    "unknown rANS variant value: {}",
312                    variant
313                )));
314            }
315        };
316        Ok(Self { stream, variant })
317    }
318
319    fn open(&mut self, data: Bound<'_, PyAny>) -> PyResult<()> {
320        let bytes = get_u8_buffer(&data)?;
321        match self.variant {
322            1 => {
323                self.stream = PyDecoderStream::Byte(CoreDecoderStream::open_on(&bytes));
324            }
325            0 => {
326                self.stream = PyDecoderStream::S64(CoreDecoderStream::open_on(&bytes));
327            }
328            _ => return Err(PyValueError::new_err("unknown rANS variant value")),
329        }
330        Ok(())
331    }
332
333    fn close(&mut self) {
334        self.stream = PyDecoderStream::None;
335    }
336
337    #[pyo3(name = "isOpen")]
338    fn is_open(&self) -> bool {
339        !matches!(self.stream, PyDecoderStream::None)
340    }
341
342    #[pyo3(name = "decodeEOF")]
343    fn decode_eof(&mut self) -> PyResult<()> {
344        let result = match &mut self.stream {
345            PyDecoderStream::Byte(s) => s.decode_eof(),
346            PyDecoderStream::S64(s) => s.decode_eof(),
347            PyDecoderStream::None => {
348                return Err(PyValueError::new_err("decoder stream is not open"));
349            }
350        };
351        result.map_err(|e| PyValueError::new_err(format!("decodeEOF failed: {}", e)))?;
352        self.stream = PyDecoderStream::None;
353        Ok(())
354    }
355}
356
357// ---------------------------------------------------------------------------
358// EntropyEncoder
359// ---------------------------------------------------------------------------
360
361#[pyclass(name = "EntropyEncoder")]
362struct PyEntropyEncoder {
363    byte_encoder: Option<EntropyEncoder<RansByte>>,
364    _64_encoder: Option<EntropyEncoder<Rans64>>,
365    variant: i32,
366}
367
368#[pymethods]
369impl PyEntropyEncoder {
370    #[new]
371    #[pyo3(signature = (*, pmfLengths, pmfOffsets, pmfTable, variant=1, symbolBits=16, bypassBits=4))]
372    fn new(
373        pmfLengths: Bound<'_, PyAny>,
374        pmfOffsets: Bound<'_, PyAny>,
375        pmfTable: Bound<'_, PyAny>,
376        variant: i32,
377        symbolBits: u32,
378        bypassBits: u32,
379    ) -> PyResult<Self> {
380        let lengths = get_i32_buffer(&pmfLengths)?;
381        let offsets = get_i32_buffer(&pmfOffsets)?;
382        let table = get_i32_buffer(&pmfTable)?;
383
384        match variant {
385            1 => {
386                let mut encoder = EntropyEncoder::<RansByte>::new();
387                encoder
388                    .initialize(&lengths, &offsets, &table, symbolBits, bypassBits)
389                    .map_err(|e| PyValueError::new_err(format!("encoder init failed: {}", e)))?;
390                Ok(Self {
391                    byte_encoder: Some(encoder),
392                    _64_encoder: None,
393                    variant,
394                })
395            }
396            0 => {
397                let mut encoder = EntropyEncoder::<Rans64>::new();
398                encoder
399                    .initialize(&lengths, &offsets, &table, symbolBits, bypassBits)
400                    .map_err(|e| PyValueError::new_err(format!("encoder init failed: {}", e)))?;
401                Ok(Self {
402                    byte_encoder: None,
403                    _64_encoder: Some(encoder),
404                    variant,
405                })
406            }
407            _ => Err(PyValueError::new_err(format!(
408                "invalid variant: {}",
409                variant
410            ))),
411        }
412    }
413
414    #[pyo3(signature = (stream, indices, values))]
415    fn encode(
416        &self,
417        stream: &mut RansEncoderStream,
418        indices: Bound<'_, PyAny>,
419        values: Bound<'_, PyAny>,
420    ) -> PyResult<()> {
421        let indices_vec = get_i32_buffer(&indices)?;
422        let values_vec = get_i32_buffer(&values)?;
423
424        if indices_vec.len() != values_vec.len() {
425            return Err(PyValueError::new_err(
426                "indices and values must have the same length",
427            ));
428        }
429
430        match self.variant {
431            1 => {
432                if let Some(ref encoder) = self.byte_encoder {
433                    stream.push_byte(encoder, &indices_vec, &values_vec)
434                } else {
435                    Err(PyValueError::new_err("byte encoder not initialized"))
436                }
437            }
438            0 => {
439                if let Some(ref encoder) = self._64_encoder {
440                    stream.push_64(encoder, &indices_vec, &values_vec)
441                } else {
442                    Err(PyValueError::new_err("64 encoder not initialized"))
443                }
444            }
445            _ => Err(PyValueError::new_err("invalid variant")),
446        }
447    }
448}
449
450// ---------------------------------------------------------------------------
451// EntropyDecoder
452// ---------------------------------------------------------------------------
453
454#[pyclass(name = "EntropyDecoder")]
455struct PyEntropyDecoder {
456    byte_decoder: Option<EntropyDecoder<RansByte>>,
457    _64_decoder: Option<EntropyDecoder<Rans64>>,
458    variant: i32,
459}
460
461#[pymethods]
462impl PyEntropyDecoder {
463    #[new]
464    #[pyo3(signature = (*, pmfLengths, pmfOffsets, pmfTable, variant=1, symbolBits=16, bypassBits=4))]
465    fn new(
466        pmfLengths: Bound<'_, PyAny>,
467        pmfOffsets: Bound<'_, PyAny>,
468        pmfTable: Bound<'_, PyAny>,
469        variant: i32,
470        symbolBits: u32,
471        bypassBits: u32,
472    ) -> PyResult<Self> {
473        let lengths = get_i32_buffer(&pmfLengths)?;
474        let offsets = get_i32_buffer(&pmfOffsets)?;
475        let table = get_i32_buffer(&pmfTable)?;
476
477        match variant {
478            1 => {
479                let mut decoder = EntropyDecoder::<RansByte>::new();
480                decoder
481                    .initialize(&lengths, &offsets, &table, symbolBits, bypassBits)
482                    .map_err(|e| PyValueError::new_err(format!("decoder init failed: {}", e)))?;
483                Ok(Self {
484                    byte_decoder: Some(decoder),
485                    _64_decoder: None,
486                    variant,
487                })
488            }
489            0 => {
490                let mut decoder = EntropyDecoder::<Rans64>::new();
491                decoder
492                    .initialize(&lengths, &offsets, &table, symbolBits, bypassBits)
493                    .map_err(|e| PyValueError::new_err(format!("decoder init failed: {}", e)))?;
494                Ok(Self {
495                    byte_decoder: None,
496                    _64_decoder: Some(decoder),
497                    variant,
498                })
499            }
500            _ => Err(PyValueError::new_err(format!(
501                "invalid variant: {}",
502                variant
503            ))),
504        }
505    }
506
507    #[pyo3(signature = (values, indices, data))]
508    fn decode(
509        &self,
510        py: Python<'_>,
511        values: Bound<'_, PyAny>,
512        indices: Bound<'_, PyAny>,
513        data: Bound<'_, PyAny>,
514    ) -> PyResult<()> {
515        let indices_vec = get_i32_buffer(&indices)?;
516        let num_values = indices_vec.len();
517
518        // Try stream-based decode
519        if let Ok(py_stream) = data.extract::<Py<RansDecoderStream>>() {
520            let mut stream_ref = py_stream.borrow_mut(py);
521            let mut decoded = vec![0i32; num_values];
522
523            match self.variant {
524                1 => {
525                    if let Some(ref decoder) = self.byte_decoder {
526                        let core = stream_ref.byte_stream_mut()?;
527                        core.decode(decoder, &mut decoded, &indices_vec)
528                            .map_err(|e| PyValueError::new_err(format!("decode failed: {}", e)))?;
529                    } else {
530                        return Err(PyValueError::new_err("byte decoder not initialized"));
531                    }
532                }
533                0 => {
534                    if let Some(ref decoder) = self._64_decoder {
535                        let core = stream_ref.s64_stream_mut()?;
536                        core.decode(decoder, &mut decoded, &indices_vec)
537                            .map_err(|e| PyValueError::new_err(format!("decode failed: {}", e)))?;
538                    } else {
539                        return Err(PyValueError::new_err("64 decoder not initialized"));
540                    }
541                }
542                _ => return Err(PyValueError::new_err("invalid variant")),
543            }
544
545            drop(stream_ref);
546            write_i32_buffer(&values, &decoded)?;
547            return Ok(());
548        }
549
550        // Buffer-based decode
551        let data_vec = get_u8_buffer(&data)?;
552        let mut decoded = vec![0i32; num_values];
553
554        match self.variant {
555            1 => {
556                if let Some(ref decoder) = self.byte_decoder {
557                    decoder
558                        .decode(&mut decoded, &indices_vec, &data_vec)
559                        .map_err(|e| PyValueError::new_err(format!("decode failed: {}", e)))?;
560                } else {
561                    return Err(PyValueError::new_err("byte decoder not initialized"));
562                }
563            }
564            0 => {
565                if let Some(ref decoder) = self._64_decoder {
566                    decoder
567                        .decode(&mut decoded, &indices_vec, &data_vec)
568                        .map_err(|e| PyValueError::new_err(format!("decode failed: {}", e)))?;
569                } else {
570                    return Err(PyValueError::new_err("64 decoder not initialized"));
571                }
572            }
573            _ => return Err(PyValueError::new_err("invalid variant")),
574        }
575
576        write_i32_buffer(&values, &decoded)
577    }
578}
579
580// ---------------------------------------------------------------------------
581// Module definition
582// ---------------------------------------------------------------------------
583
584#[pymodule]
585fn _msrtc_rans(_py: Python, m: &Bound<'_, PyModule>) -> PyResult<()> {
586    m.add("__version__", env!("CARGO_PKG_VERSION"))?;
587    // Module-level constants matching C++ enum names used by types.py
588    m.add("RansByte", 1)?; // RansVariant.RansByte = 1
589    m.add("Rans64", 0)?; // RansVariant.Rans64 = 0
590    // Also add as functions for alternative access
591    m.add_function(wrap_pyfunction!(rans_byte, m)?)?;
592    m.add_function(wrap_pyfunction!(rans_64, m)?)?;
593    m.add_class::<RansEncoderStream>()?;
594    m.add_class::<RansDecoderStream>()?;
595    m.add_class::<PyEntropyEncoder>()?;
596    m.add_class::<PyEntropyDecoder>()?;
597    Ok(())
598}