1#![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
24const BUF_READ: i32 = 12; const BUF_WRITE: i32 = 13; unsafe 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#[pyfunction]
98fn rans_byte() -> i32 {
99 1
100}
101
102#[pyfunction]
103fn rans_64() -> i32 {
104 0
105}
106
107#[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 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#[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#[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#[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 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 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#[pymodule]
462fn _msrtc_rans(_py: Python, m: &Bound<'_, PyModule>) -> PyResult<()> {
463 m.add("__version__", env!("CARGO_PKG_VERSION"))?;
464 m.add("RansByte", 1)?; m.add("Rans64", 0)?; 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}