1use crate::error::{IoError, Result};
8
9pub trait Codec: std::fmt::Debug + Send + Sync {
11 fn name(&self) -> &str;
13
14 fn encode(&self, data: &[u8]) -> Result<Vec<u8>>;
16
17 fn decode(&self, data: &[u8]) -> Result<Vec<u8>>;
19}
20
21#[derive(Debug, Clone, Copy)]
23pub struct BytesCodec {
24 pub endian: Endian,
26 pub element_size: usize,
28}
29
30#[derive(Debug, Clone, Copy, PartialEq, Eq)]
32pub enum Endian {
33 Little,
35 Big,
37}
38
39impl BytesCodec {
40 pub fn new(endian: Endian, element_size: usize) -> Self {
42 Self {
43 endian,
44 element_size,
45 }
46 }
47
48 fn needs_swap(&self) -> bool {
49 match self.endian {
50 Endian::Little => cfg!(target_endian = "big"),
51 Endian::Big => cfg!(target_endian = "little"),
52 }
53 }
54
55 fn swap_bytes(data: &[u8], elem_size: usize) -> Vec<u8> {
56 if elem_size <= 1 {
57 return data.to_vec();
58 }
59 let mut out = data.to_vec();
60 for chunk in out.chunks_exact_mut(elem_size) {
61 chunk.reverse();
62 }
63 out
64 }
65}
66
67impl Codec for BytesCodec {
68 fn name(&self) -> &str {
69 "bytes"
70 }
71
72 fn encode(&self, data: &[u8]) -> Result<Vec<u8>> {
73 if self.needs_swap() {
74 Ok(Self::swap_bytes(data, self.element_size))
75 } else {
76 Ok(data.to_vec())
77 }
78 }
79
80 fn decode(&self, data: &[u8]) -> Result<Vec<u8>> {
81 if self.needs_swap() {
83 Ok(Self::swap_bytes(data, self.element_size))
84 } else {
85 Ok(data.to_vec())
86 }
87 }
88}
89
90#[derive(Debug, Clone)]
95pub struct TransposeCodec {
96 shape: Vec<usize>,
98 element_size: usize,
100}
101
102impl TransposeCodec {
103 pub fn new(shape: Vec<usize>, element_size: usize) -> Self {
105 Self {
106 shape,
107 element_size,
108 }
109 }
110
111 fn total_elements(&self) -> usize {
112 self.shape.iter().product()
113 }
114
115 fn c_to_f_index(&self, c_linear: usize) -> usize {
117 let ndim = self.shape.len();
118 if ndim == 0 {
119 return 0;
120 }
121 let mut indices = vec![0usize; ndim];
123 let mut rem = c_linear;
124 for d in (0..ndim).rev() {
125 indices[d] = rem % self.shape[d];
126 rem /= self.shape[d];
127 }
128 let mut f_linear = 0usize;
130 let mut stride = 1usize;
131 for d in 0..ndim {
132 f_linear += indices[d] * stride;
133 stride *= self.shape[d];
134 }
135 f_linear
136 }
137
138 fn f_to_c_index(&self, f_linear: usize) -> usize {
140 let ndim = self.shape.len();
141 if ndim == 0 {
142 return 0;
143 }
144 let mut indices = vec![0usize; ndim];
146 let mut rem = f_linear;
147 for d in 0..ndim {
148 indices[d] = rem % self.shape[d];
149 rem /= self.shape[d];
150 }
151 let mut c_linear = 0usize;
153 let mut stride = 1usize;
154 for d in (0..ndim).rev() {
155 c_linear += indices[d] * stride;
156 stride *= self.shape[d];
157 }
158 c_linear
159 }
160}
161
162impl Codec for TransposeCodec {
163 fn name(&self) -> &str {
164 "transpose"
165 }
166
167 fn encode(&self, data: &[u8]) -> Result<Vec<u8>> {
168 let n = self.total_elements();
169 let expected = n * self.element_size;
170 if data.len() != expected {
171 return Err(IoError::FormatError(format!(
172 "Transpose encode: expected {} bytes, got {}",
173 expected,
174 data.len()
175 )));
176 }
177 let mut out = vec![0u8; expected];
178 for c_idx in 0..n {
179 let f_idx = self.c_to_f_index(c_idx);
180 let src = c_idx * self.element_size;
181 let dst = f_idx * self.element_size;
182 out[dst..dst + self.element_size].copy_from_slice(&data[src..src + self.element_size]);
183 }
184 Ok(out)
185 }
186
187 fn decode(&self, data: &[u8]) -> Result<Vec<u8>> {
188 let n = self.total_elements();
189 let expected = n * self.element_size;
190 if data.len() != expected {
191 return Err(IoError::FormatError(format!(
192 "Transpose decode: expected {} bytes, got {}",
193 expected,
194 data.len()
195 )));
196 }
197 let mut out = vec![0u8; expected];
198 for f_idx in 0..n {
199 let c_idx = self.f_to_c_index(f_idx);
200 let src = f_idx * self.element_size;
201 let dst = c_idx * self.element_size;
202 out[dst..dst + self.element_size].copy_from_slice(&data[src..src + self.element_size]);
203 }
204 Ok(out)
205 }
206}
207
208#[derive(Debug, Clone, Copy)]
213pub struct ZstdCodec {
214 pub level: i32,
216}
217
218impl ZstdCodec {
219 pub fn new(level: i32) -> Self {
221 Self { level }
222 }
223}
224
225impl Default for ZstdCodec {
226 fn default() -> Self {
227 Self { level: 3 }
228 }
229}
230
231impl Codec for ZstdCodec {
232 fn name(&self) -> &str {
233 "zstd"
234 }
235
236 fn encode(&self, data: &[u8]) -> Result<Vec<u8>> {
237 oxiarc_zstd::compress(data)
238 .map_err(|e| IoError::CompressionError(format!("Zstd compression failed: {e}")))
239 }
240
241 fn decode(&self, data: &[u8]) -> Result<Vec<u8>> {
242 oxiarc_zstd::decompress(data)
243 .map_err(|e| IoError::DecompressionError(format!("Zstd decompression failed: {e}")))
244 }
245}
246
247#[derive(Debug, Clone, Copy)]
252pub struct ShuffleCodec {
253 pub element_size: usize,
255}
256
257impl ShuffleCodec {
258 pub fn new(element_size: usize) -> Self {
260 Self { element_size }
261 }
262}
263
264impl Codec for ShuffleCodec {
265 fn name(&self) -> &str {
266 "shuffle"
267 }
268
269 fn encode(&self, data: &[u8]) -> Result<Vec<u8>> {
270 if self.element_size <= 1 {
271 return Ok(data.to_vec());
272 }
273 let n_elements = data.len() / self.element_size;
274 if !data.len().is_multiple_of(self.element_size) {
275 return Err(IoError::FormatError(format!(
276 "Shuffle encode: data length {} not divisible by element size {}",
277 data.len(),
278 self.element_size
279 )));
280 }
281 let mut out = vec![0u8; data.len()];
282 for elem_idx in 0..n_elements {
283 for byte_idx in 0..self.element_size {
284 let src = elem_idx * self.element_size + byte_idx;
285 let dst = byte_idx * n_elements + elem_idx;
286 out[dst] = data[src];
287 }
288 }
289 Ok(out)
290 }
291
292 fn decode(&self, data: &[u8]) -> Result<Vec<u8>> {
293 if self.element_size <= 1 {
294 return Ok(data.to_vec());
295 }
296 let n_elements = data.len() / self.element_size;
297 if !data.len().is_multiple_of(self.element_size) {
298 return Err(IoError::FormatError(format!(
299 "Shuffle decode: data length {} not divisible by element size {}",
300 data.len(),
301 self.element_size
302 )));
303 }
304 let mut out = vec![0u8; data.len()];
305 for elem_idx in 0..n_elements {
306 for byte_idx in 0..self.element_size {
307 let src = byte_idx * n_elements + elem_idx;
308 let dst = elem_idx * self.element_size + byte_idx;
309 out[dst] = data[src];
310 }
311 }
312 Ok(out)
313 }
314}
315
316#[derive(Debug)]
318pub struct CodecPipeline {
319 codecs: Vec<Box<dyn Codec>>,
320}
321
322impl CodecPipeline {
323 pub fn new() -> Self {
325 Self { codecs: Vec::new() }
326 }
327
328 pub fn push<C: Codec + 'static>(&mut self, codec: C) {
330 self.codecs.push(Box::new(codec));
331 }
332
333 pub fn len(&self) -> usize {
335 self.codecs.len()
336 }
337
338 pub fn is_empty(&self) -> bool {
340 self.codecs.is_empty()
341 }
342
343 pub fn encode(&self, data: &[u8]) -> Result<Vec<u8>> {
345 let mut buf = data.to_vec();
346 for codec in &self.codecs {
347 buf = codec.encode(&buf)?;
348 }
349 Ok(buf)
350 }
351
352 pub fn decode(&self, data: &[u8]) -> Result<Vec<u8>> {
354 let mut buf = data.to_vec();
355 for codec in self.codecs.iter().rev() {
356 buf = codec.decode(&buf)?;
357 }
358 Ok(buf)
359 }
360}
361
362impl Default for CodecPipeline {
363 fn default() -> Self {
364 Self::new()
365 }
366}
367
368#[cfg(test)]
369mod tests {
370 use super::*;
371
372 #[test]
373 fn test_bytes_codec_no_swap() {
374 let codec = BytesCodec::new(
375 if cfg!(target_endian = "little") {
376 Endian::Little
377 } else {
378 Endian::Big
379 },
380 4,
381 );
382 let data = vec![1, 2, 3, 4, 5, 6, 7, 8];
383 let encoded = codec.encode(&data).expect("encode");
384 assert_eq!(encoded, data);
385 let decoded = codec.decode(&encoded).expect("decode");
386 assert_eq!(decoded, data);
387 }
388
389 #[test]
390 fn test_bytes_codec_swap() {
391 let non_native = if cfg!(target_endian = "little") {
392 Endian::Big
393 } else {
394 Endian::Little
395 };
396 let codec = BytesCodec::new(non_native, 2);
397 let data = vec![0x01, 0x02, 0x03, 0x04];
398 let encoded = codec.encode(&data).expect("encode");
399 assert_eq!(encoded, vec![0x02, 0x01, 0x04, 0x03]);
400 let decoded = codec.decode(&encoded).expect("decode");
401 assert_eq!(decoded, data);
402 }
403
404 #[test]
405 fn test_transpose_codec_roundtrip() {
406 let codec = TransposeCodec::new(vec![2, 3], 4);
408 let mut data = Vec::new();
410 for val in [1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0] {
411 data.extend_from_slice(&val.to_ne_bytes());
412 }
413 let encoded = codec.encode(&data).expect("encode");
414 assert_ne!(encoded, data);
416 let decoded = codec.decode(&encoded).expect("decode");
417 assert_eq!(decoded, data);
418 }
419
420 #[test]
421 fn test_zstd_codec_roundtrip() {
422 let codec = ZstdCodec::new(3);
423 let data: Vec<u8> = vec![42u8; 4096];
425 let compressed = codec.encode(&data).expect("compress");
426 assert!(compressed.len() < data.len());
428 let decompressed = codec.decode(&compressed).expect("decompress");
429 assert_eq!(decompressed, data);
430 }
431
432 #[test]
433 fn test_shuffle_codec_roundtrip() {
434 let codec = ShuffleCodec::new(4);
435 let data: Vec<u8> = (0..32).collect();
436 let encoded = codec.encode(&data).expect("encode");
437 assert_ne!(encoded, data);
438 let decoded = codec.decode(&encoded).expect("decode");
439 assert_eq!(decoded, data);
440 }
441
442 #[test]
443 fn test_shuffle_single_byte_passthrough() {
444 let codec = ShuffleCodec::new(1);
445 let data = vec![10, 20, 30];
446 let encoded = codec.encode(&data).expect("encode");
447 assert_eq!(encoded, data);
448 }
449
450 #[test]
451 fn test_codec_pipeline_chain() {
452 let mut pipeline = CodecPipeline::new();
453 pipeline.push(ShuffleCodec::new(8));
454 pipeline.push(ZstdCodec::new(1));
455 assert_eq!(pipeline.len(), 2);
456
457 let data: Vec<u8> = (0..800).map(|i| (i % 256) as u8).collect();
458 let encoded = pipeline.encode(&data).expect("pipeline encode");
459 let decoded = pipeline.decode(&encoded).expect("pipeline decode");
460 assert_eq!(decoded, data);
461 }
462
463 #[test]
464 fn test_codec_pipeline_empty() {
465 let pipeline = CodecPipeline::new();
466 assert!(pipeline.is_empty());
467 let data = vec![1, 2, 3];
468 let encoded = pipeline.encode(&data).expect("encode");
469 assert_eq!(encoded, data);
470 let decoded = pipeline.decode(&data).expect("decode");
471 assert_eq!(decoded, data);
472 }
473
474 #[test]
475 fn test_transpose_codec_1d() {
476 let codec = TransposeCodec::new(vec![8], 4);
478 let data: Vec<u8> = (0..32).collect();
479 let encoded = codec.encode(&data).expect("encode");
480 assert_eq!(encoded, data);
481 }
482}