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; unsafe 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#[pyfunction]
111fn rans_byte() -> i32 {
112 1
113}
114
115#[pyfunction]
116fn rans_64() -> i32 {
117 0
118}
119
120enum 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 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 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
247enum 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#[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#[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 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 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#[pymodule]
585fn _msrtc_rans(_py: Python, m: &Bound<'_, PyModule>) -> PyResult<()> {
586 m.add("__version__", env!("CARGO_PKG_VERSION"))?;
587 m.add("RansByte", 1)?; m.add("Rans64", 0)?; 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}