1#![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
37const BUF_READ: i32 = 12; const BUF_WRITE: i32 = 13; fn validate_i32_buffer(buf: &ffi::Py_buffer) -> Result<(), String> {
54 if buf.ndim != 1 {
55 return Err(format!("expected 1-d array, got ndim={}", buf.ndim));
56 }
57 if buf.itemsize != 4 {
58 return Err(format!("expected int32 (itemsize 4), got {}", buf.itemsize));
59 }
60 if buf.format.is_null() {
61 return Err("buffer has no format string".into());
62 }
63 let fmt = unsafe { std::ffi::CStr::from_ptr(buf.format) }.to_string_lossy();
64 let fmt = fmt.trim_end_matches(|c: char| c.is_ascii_digit()); if fmt != "i" && fmt != "l" {
66 return Err(format!("expected int32 format, got '{}'", fmt));
67 }
68 if buf.len as usize % 4 != 0 {
69 return Err(format!(
70 "int32 array length must be multiple of 4, got {}",
71 buf.len
72 ));
73 }
74 if (buf.buf as usize) % 4 != 0 {
75 return Err("int32 buffer is not 4-byte aligned".into());
76 }
77 if buf.shape.is_null() || unsafe { *buf.shape } != buf.len / buf.itemsize {
78 return Err("buffer shape does not match length".into());
79 }
80 Ok(())
81}
82
83unsafe fn buffer_to_i32_slice(buf: &ffi::Py_buffer) -> &[i32] {
84 let n = (buf.len as usize) / 4;
85 std::slice::from_raw_parts(buf.buf as *const i32, n)
86}
87
88unsafe fn buffer_to_u8_slice(buf: &ffi::Py_buffer) -> &[u8] {
89 let n = buf.len as usize;
90 std::slice::from_raw_parts(buf.buf as *const u8, n)
91}
92
93fn get_i32_buffer(obj: &Bound<'_, PyAny>) -> PyResult<Vec<i32>> {
94 let mut buf: ffi::Py_buffer = unsafe { std::mem::zeroed() };
95 let ret = unsafe { ffi::PyObject_GetBuffer(obj.as_ptr(), &mut buf, BUF_READ) };
96 if ret != 0 {
97 return Err(PyValueError::new_err(
98 "indices/values/pmf must be an int32 1-d array (buffer protocol)",
99 ));
100 }
101 let result = validate_i32_buffer(&buf)
102 .map_err(PyValueError::new_err)
103 .and_then(|_| Ok(unsafe { buffer_to_i32_slice(&buf).to_vec() }));
104 unsafe { ffi::PyBuffer_Release(&mut buf) };
105 result
106}
107
108fn get_u8_buffer(obj: &Bound<'_, PyAny>) -> PyResult<Vec<u8>> {
109 if let Ok(bytes) = obj.downcast::<PyBytes>() {
110 return Ok(bytes.as_bytes().to_vec());
111 }
112 if let Ok(ba) = obj.downcast::<PyByteArray>() {
113 let slice = unsafe { ba.as_bytes() };
114 return Ok(slice.to_vec());
115 }
116 let mut buf: ffi::Py_buffer = unsafe { std::mem::zeroed() };
117 let ret = unsafe { ffi::PyObject_GetBuffer(obj.as_ptr(), &mut buf, BUF_READ) };
118 if ret != 0 {
119 return Err(PyValueError::new_err("cannot get buffer from object"));
120 }
121 let result = (|| -> Result<Vec<u8>, String> {
122 if buf.ndim != 1 {
123 return Err(format!("expected 1-d buffer, got ndim={}", buf.ndim));
124 }
125 if buf.itemsize != 1 {
126 return Err(format!(
127 "expected byte buffer (itemsize 1), got {}",
128 buf.itemsize
129 ));
130 }
131 if buf.shape.is_null() || unsafe { *buf.shape } != buf.len {
132 return Err("buffer shape does not match length".into());
133 }
134 Ok(unsafe { buffer_to_u8_slice(&buf).to_vec() })
135 })()
136 .map_err(PyValueError::new_err);
137 unsafe { ffi::PyBuffer_Release(&mut buf) };
138 result
139}
140
141fn write_i32_buffer(obj: &Bound<'_, PyAny>, data: &[i32]) -> PyResult<()> {
142 let mut buf: ffi::Py_buffer = unsafe { std::mem::zeroed() };
143 let ret = unsafe { ffi::PyObject_GetBuffer(obj.as_ptr(), &mut buf, BUF_WRITE) };
144 if ret != 0 {
145 return Err(PyValueError::new_err(
146 "values must be a writable int32 1-d array",
147 ));
148 }
149 let result = (|| -> Result<(), String> {
150 validate_i32_buffer(&buf)?;
151 let n = data.len() * 4;
152 if buf.len as usize != n {
154 return Err(format!(
155 "output int32 array has {} bytes, expected exactly {} for {} values",
156 buf.len,
157 n,
158 data.len()
159 ));
160 }
161 unsafe {
162 let dst = buf.buf as *mut u8;
163 let src = data.as_ptr() as *const u8;
164 std::ptr::copy_nonoverlapping(src, dst, n);
165 }
166 Ok(())
167 })()
168 .map_err(PyValueError::new_err);
169 unsafe { ffi::PyBuffer_Release(&mut buf) };
170 result
171}
172
173#[pyfunction]
178fn rans_byte() -> i32 {
179 1
180}
181
182#[pyfunction]
183fn rans_64() -> i32 {
184 0
185}
186
187enum PyEncoderStream {
196 None,
197 Byte(CoreEncoderStream<RansByte>),
198 S64(CoreEncoderStream<Rans64>),
199}
200
201#[pyclass(name = "RansEncoderStream")]
202struct RansEncoderStream {
203 stream: PyEncoderStream,
204 #[allow(dead_code)]
205 variant: i32,
206 #[allow(dead_code)]
207 _initial_size: usize,
208 #[allow(dead_code)]
209 _max_size_step: usize,
210}
211
212impl RansEncoderStream {
213 fn push_byte(
215 &mut self,
216 encoder: &EntropyEncoder<RansByte>,
217 indices: &[i32],
218 values: &[i32],
219 ) -> PyResult<()> {
220 match &mut self.stream {
221 PyEncoderStream::Byte(s) => s
222 .push(encoder, indices, values)
223 .map_err(|e| PyValueError::new_err(format!("encode failed: {}", e))),
224 PyEncoderStream::S64(_) => Err(PyValueError::new_err(
225 "encoder stream variant mismatch: stream is Rans64, encoder is RansByte",
226 )),
227 PyEncoderStream::None => Err(PyValueError::new_err("invalid state")),
228 }
229 }
230
231 fn push_64(
233 &mut self,
234 encoder: &EntropyEncoder<Rans64>,
235 indices: &[i32],
236 values: &[i32],
237 ) -> PyResult<()> {
238 match &mut self.stream {
239 PyEncoderStream::S64(s) => s
240 .push(encoder, indices, values)
241 .map_err(|e| PyValueError::new_err(format!("encode failed: {}", e))),
242 PyEncoderStream::Byte(_) => Err(PyValueError::new_err(
243 "encoder stream variant mismatch: stream is RansByte, encoder is Rans64",
244 )),
245 PyEncoderStream::None => Err(PyValueError::new_err("invalid state")),
246 }
247 }
248}
249
250#[pymethods]
251impl RansEncoderStream {
252 #[new]
253 #[pyo3(signature = (variant=1, *, initialSize=4096, maxSizeStep=1048576))]
254 fn new(variant: i32, initialSize: usize, maxSizeStep: usize) -> PyResult<Self> {
255 let stream = match variant {
256 1 => PyEncoderStream::Byte(CoreEncoderStream::new()),
257 0 => PyEncoderStream::S64(CoreEncoderStream::new()),
258 _ => {
259 return Err(PyValueError::new_err(format!(
260 "unknown rANS variant value: {}",
261 variant
262 )));
263 }
264 };
265 Ok(Self {
266 stream,
267 variant,
268 _initial_size: initialSize,
269 _max_size_step: maxSizeStep,
270 })
271 }
272
273 fn flush(&mut self, py: Python<'_>) -> PyResult<Py<PyAny>> {
274 let data: Vec<u8> = match &mut self.stream {
275 PyEncoderStream::Byte(s) => s
276 .flush()
277 .map_err(|e| PyValueError::new_err(format!("flush failed: {}", e)))?,
278 PyEncoderStream::S64(s) => s
279 .flush()
280 .map_err(|e| PyValueError::new_err(format!("flush failed: {}", e)))?,
281 PyEncoderStream::None => {
282 return Err(PyValueError::new_err(
283 "invalid state: stream not initialized",
284 ));
285 }
286 };
287
288 if data.is_empty() {
289 return Err(PyValueError::new_err("invalid state: empty output"));
290 }
291
292 let ptr = unsafe {
293 ffi::PyBytes_FromStringAndSize(
294 data.as_ptr() as *const ffi::Py_ssize_t as *const i8,
295 data.len() as ffi::Py_ssize_t,
296 )
297 };
298 if ptr.is_null() {
299 return Err(PyValueError::new_err("failed to create PyBytes"));
300 }
301 let obj: Py<PyAny> = unsafe { Bound::from_owned_ptr(py, ptr).unbind() };
302 Ok(obj)
303 }
304
305 fn reset(&mut self) {
306 match &mut self.stream {
307 PyEncoderStream::Byte(s) => s.reset(),
308 PyEncoderStream::S64(s) => s.reset(),
309 PyEncoderStream::None => {}
310 }
311 }
312}
313
314enum PyDecoderStream {
324 None,
325 Byte(CoreDecoderStream<RansByte>),
326 S64(CoreDecoderStream<Rans64>),
327}
328
329#[pyclass(name = "RansDecoderStream")]
330struct RansDecoderStream {
331 stream: PyDecoderStream,
332 #[allow(dead_code)]
333 variant: i32,
334}
335
336impl RansDecoderStream {
337 fn byte_stream_mut(&mut self) -> PyResult<&mut CoreDecoderStream<RansByte>> {
338 match &mut self.stream {
339 PyDecoderStream::Byte(s) => Ok(s),
340 PyDecoderStream::S64(_) => Err(PyValueError::new_err(
341 "decoder stream variant mismatch: stream is Rans64",
342 )),
343 PyDecoderStream::None => Err(PyValueError::new_err("decoder stream is not open")),
344 }
345 }
346
347 fn s64_stream_mut(&mut self) -> PyResult<&mut CoreDecoderStream<Rans64>> {
348 match &mut self.stream {
349 PyDecoderStream::S64(s) => Ok(s),
350 PyDecoderStream::Byte(_) => Err(PyValueError::new_err(
351 "decoder stream variant mismatch: stream is RansByte",
352 )),
353 PyDecoderStream::None => Err(PyValueError::new_err("decoder stream is not open")),
354 }
355 }
356}
357
358#[pymethods]
359impl RansDecoderStream {
360 #[new]
361 #[pyo3(signature = (data=None, *, variant=1))]
362 fn new(data: Option<Bound<'_, PyAny>>, variant: i32) -> PyResult<Self> {
363 let stream = match variant {
364 1 => match data {
365 Some(ref obj) => {
366 PyDecoderStream::Byte(CoreDecoderStream::open_on(&get_u8_buffer(obj)?))
367 }
368 None => PyDecoderStream::Byte(CoreDecoderStream::new()),
369 },
370 0 => match data {
371 Some(ref obj) => {
372 PyDecoderStream::S64(CoreDecoderStream::open_on(&get_u8_buffer(obj)?))
373 }
374 None => PyDecoderStream::S64(CoreDecoderStream::new()),
375 },
376 _ => {
377 return Err(PyValueError::new_err(format!(
378 "unknown rANS variant value: {}",
379 variant
380 )));
381 }
382 };
383 Ok(Self { stream, variant })
384 }
385
386 fn open(&mut self, data: Bound<'_, PyAny>) -> PyResult<()> {
387 let bytes = get_u8_buffer(&data)?;
388 match self.variant {
389 1 => {
390 self.stream = PyDecoderStream::Byte(CoreDecoderStream::open_on(&bytes));
391 }
392 0 => {
393 self.stream = PyDecoderStream::S64(CoreDecoderStream::open_on(&bytes));
394 }
395 _ => return Err(PyValueError::new_err("unknown rANS variant value")),
396 }
397 Ok(())
398 }
399
400 fn close(&mut self) {
401 self.stream = PyDecoderStream::None;
402 }
403
404 #[pyo3(name = "isOpen")]
405 fn is_open(&self) -> bool {
406 !matches!(self.stream, PyDecoderStream::None)
407 }
408
409 #[pyo3(name = "decodeEOF")]
410 fn decode_eof(&mut self) -> PyResult<()> {
411 let result = match &mut self.stream {
412 PyDecoderStream::Byte(s) => s.decode_eof(),
413 PyDecoderStream::S64(s) => s.decode_eof(),
414 PyDecoderStream::None => {
415 return Err(PyValueError::new_err("decoder stream is not open"));
416 }
417 };
418 result.map_err(|e| PyValueError::new_err(format!("decodeEOF failed: {}", e)))?;
419 self.stream = PyDecoderStream::None;
420 Ok(())
421 }
422}
423
424#[pyclass(name = "EntropyEncoder")]
429struct PyEntropyEncoder {
430 byte_encoder: Option<EntropyEncoder<RansByte>>,
431 _64_encoder: Option<EntropyEncoder<Rans64>>,
432 variant: i32,
433}
434
435#[pymethods]
436impl PyEntropyEncoder {
437 #[new]
438 #[pyo3(signature = (*, pmfLengths, pmfOffsets, pmfTable, variant=1, symbolBits=16, bypassBits=4))]
439 fn new(
440 pmfLengths: Bound<'_, PyAny>,
441 pmfOffsets: Bound<'_, PyAny>,
442 pmfTable: Bound<'_, PyAny>,
443 variant: i32,
444 symbolBits: u32,
445 bypassBits: u32,
446 ) -> PyResult<Self> {
447 let lengths = get_i32_buffer(&pmfLengths)?;
448 let offsets = get_i32_buffer(&pmfOffsets)?;
449 let table = get_i32_buffer(&pmfTable)?;
450
451 match variant {
452 1 => {
453 let mut encoder = EntropyEncoder::<RansByte>::new();
454 encoder
455 .initialize(&lengths, &offsets, &table, symbolBits, bypassBits)
456 .map_err(|e| PyValueError::new_err(format!("encoder init failed: {}", e)))?;
457 Ok(Self {
458 byte_encoder: Some(encoder),
459 _64_encoder: None,
460 variant,
461 })
462 }
463 0 => {
464 let mut encoder = EntropyEncoder::<Rans64>::new();
465 encoder
466 .initialize(&lengths, &offsets, &table, symbolBits, bypassBits)
467 .map_err(|e| PyValueError::new_err(format!("encoder init failed: {}", e)))?;
468 Ok(Self {
469 byte_encoder: None,
470 _64_encoder: Some(encoder),
471 variant,
472 })
473 }
474 _ => Err(PyValueError::new_err(format!(
475 "invalid variant: {}",
476 variant
477 ))),
478 }
479 }
480
481 #[pyo3(signature = (stream, indices, values))]
482 fn encode(
483 &self,
484 stream: &mut RansEncoderStream,
485 indices: Bound<'_, PyAny>,
486 values: Bound<'_, PyAny>,
487 ) -> PyResult<()> {
488 let indices_vec = get_i32_buffer(&indices)?;
489 let values_vec = get_i32_buffer(&values)?;
490
491 if indices_vec.len() != values_vec.len() {
492 return Err(PyValueError::new_err(
493 "indices and values must have the same length",
494 ));
495 }
496
497 match self.variant {
498 1 => {
499 if let Some(ref encoder) = self.byte_encoder {
500 stream.push_byte(encoder, &indices_vec, &values_vec)
501 } else {
502 Err(PyValueError::new_err("byte encoder not initialized"))
503 }
504 }
505 0 => {
506 if let Some(ref encoder) = self._64_encoder {
507 stream.push_64(encoder, &indices_vec, &values_vec)
508 } else {
509 Err(PyValueError::new_err("64 encoder not initialized"))
510 }
511 }
512 _ => Err(PyValueError::new_err("invalid variant")),
513 }
514 }
515}
516
517#[pyclass(name = "EntropyDecoder")]
522struct PyEntropyDecoder {
523 byte_decoder: Option<EntropyDecoder<RansByte>>,
524 _64_decoder: Option<EntropyDecoder<Rans64>>,
525 variant: i32,
526}
527
528#[pymethods]
529impl PyEntropyDecoder {
530 #[new]
531 #[pyo3(signature = (*, pmfLengths, pmfOffsets, pmfTable, variant=1, symbolBits=16, bypassBits=4))]
532 fn new(
533 pmfLengths: Bound<'_, PyAny>,
534 pmfOffsets: Bound<'_, PyAny>,
535 pmfTable: Bound<'_, PyAny>,
536 variant: i32,
537 symbolBits: u32,
538 bypassBits: u32,
539 ) -> PyResult<Self> {
540 let lengths = get_i32_buffer(&pmfLengths)?;
541 let offsets = get_i32_buffer(&pmfOffsets)?;
542 let table = get_i32_buffer(&pmfTable)?;
543
544 match variant {
545 1 => {
546 let mut decoder = EntropyDecoder::<RansByte>::new();
547 decoder
548 .initialize(&lengths, &offsets, &table, symbolBits, bypassBits)
549 .map_err(|e| PyValueError::new_err(format!("decoder init failed: {}", e)))?;
550 Ok(Self {
551 byte_decoder: Some(decoder),
552 _64_decoder: None,
553 variant,
554 })
555 }
556 0 => {
557 let mut decoder = EntropyDecoder::<Rans64>::new();
558 decoder
559 .initialize(&lengths, &offsets, &table, symbolBits, bypassBits)
560 .map_err(|e| PyValueError::new_err(format!("decoder init failed: {}", e)))?;
561 Ok(Self {
562 byte_decoder: None,
563 _64_decoder: Some(decoder),
564 variant,
565 })
566 }
567 _ => Err(PyValueError::new_err(format!(
568 "invalid variant: {}",
569 variant
570 ))),
571 }
572 }
573
574 #[pyo3(signature = (values, indices, data))]
575 fn decode(
576 &self,
577 py: Python<'_>,
578 values: Bound<'_, PyAny>,
579 indices: Bound<'_, PyAny>,
580 data: Bound<'_, PyAny>,
581 ) -> PyResult<()> {
582 let indices_vec = get_i32_buffer(&indices)?;
583 let num_values = indices_vec.len();
584
585 if let Ok(py_stream) = data.extract::<Py<RansDecoderStream>>() {
587 let mut stream_ref = py_stream.borrow_mut(py);
588 let mut decoded = vec![0i32; num_values];
589
590 match self.variant {
591 1 => {
592 if let Some(ref decoder) = self.byte_decoder {
593 let core = stream_ref.byte_stream_mut()?;
594 core.decode(decoder, &mut decoded, &indices_vec)
595 .map_err(|e| PyValueError::new_err(format!("decode failed: {}", e)))?;
596 } else {
597 return Err(PyValueError::new_err("byte decoder not initialized"));
598 }
599 }
600 0 => {
601 if let Some(ref decoder) = self._64_decoder {
602 let core = stream_ref.s64_stream_mut()?;
603 core.decode(decoder, &mut decoded, &indices_vec)
604 .map_err(|e| PyValueError::new_err(format!("decode failed: {}", e)))?;
605 } else {
606 return Err(PyValueError::new_err("64 decoder not initialized"));
607 }
608 }
609 _ => return Err(PyValueError::new_err("invalid variant")),
610 }
611
612 drop(stream_ref);
613 write_i32_buffer(&values, &decoded)?;
614 return Ok(());
615 }
616
617 let data_vec = get_u8_buffer(&data)?;
619 let mut decoded = vec![0i32; num_values];
620
621 match self.variant {
622 1 => {
623 if let Some(ref decoder) = self.byte_decoder {
624 decoder
625 .decode(&mut decoded, &indices_vec, &data_vec)
626 .map_err(|e| PyValueError::new_err(format!("decode failed: {}", e)))?;
627 } else {
628 return Err(PyValueError::new_err("byte decoder not initialized"));
629 }
630 }
631 0 => {
632 if let Some(ref decoder) = self._64_decoder {
633 decoder
634 .decode(&mut decoded, &indices_vec, &data_vec)
635 .map_err(|e| PyValueError::new_err(format!("decode failed: {}", e)))?;
636 } else {
637 return Err(PyValueError::new_err("64 decoder not initialized"));
638 }
639 }
640 _ => return Err(PyValueError::new_err("invalid variant")),
641 }
642
643 write_i32_buffer(&values, &decoded)
644 }
645}
646
647#[pymodule]
652fn _msrtc_rans(_py: Python, m: &Bound<'_, PyModule>) -> PyResult<()> {
653 m.add("__version__", env!("CARGO_PKG_VERSION"))?;
654 m.add("RansByte", 1)?; m.add("Rans64", 0)?; m.add_function(wrap_pyfunction!(rans_byte, m)?)?;
659 m.add_function(wrap_pyfunction!(rans_64, m)?)?;
660 m.add_class::<RansEncoderStream>()?;
661 m.add_class::<RansDecoderStream>()?;
662 m.add_class::<PyEntropyEncoder>()?;
663 m.add_class::<PyEntropyDecoder>()?;
664 Ok(())
665}