use std::io::{self, Read, Write};
use super::Border;
pub type Frequency = u16;
pub struct Model {
total: Border,
table: Vec<Frequency>,
cut_threshold: Border,
cut_shift: usize,
}
impl Model {
pub fn new_custom<F>(num_values: usize, threshold: Border,
mut fn_init: F) -> Model
where F: FnMut(usize) -> Frequency
{
let freq: Vec<Frequency> = (0..num_values).map(|i| fn_init(i)).collect();
let total = freq.iter().fold(0 as Border, |u,&f| u+(f as Border));
let mut ft = Model {
total: total,
table: freq,
cut_threshold: threshold,
cut_shift: 1,
};
while ft.total >= threshold {
ft.downscale();
}
ft
}
pub fn new_flat(num_values: usize, threshold: Border) -> Model {
Model::new_custom(num_values, threshold, |_| 1)
}
pub fn reset_flat(&mut self) {
for freq in self.table.iter_mut() {
*freq = 1;
}
self.total = self.table.len() as Border;
}
pub fn update(&mut self, value: usize, add_log: usize, add_const: Border) {
let add = (self.total>>add_log) + add_const;
assert!(add < 2*self.cut_threshold);
debug!("\tUpdating by adding {} to value {}", add, value);
self.table[value] += add as Frequency;
self.total += add;
if self.total >= self.cut_threshold {
self.downscale();
assert!(self.total < self.cut_threshold);
}
}
pub fn downscale(&mut self) {
debug!("\tDownscaling frequencies");
let roundup = (1<<self.cut_shift) - 1;
self.total = 0;
for freq in self.table.iter_mut() {
*freq = (*freq+roundup) >> self.cut_shift;
self.total += *freq as Border;
}
}
pub fn get_frequencies<'a>(&'a self) -> &'a [Frequency] {
&self.table[..]
}
}
impl super::Model<usize> for Model {
fn get_range(&self, value: usize) -> (Border,Border) {
let lo = self.table[..value].iter().fold(0, |u,&f| u+(f as Border));
(lo, lo + (self.table[value] as Border))
}
fn find_value(&self, offset: Border) -> (usize,Border,Border) {
assert!(offset < self.total,
"Invalid frequency offset {} requested under total {}",
offset, self.total);
let mut value = 0;
let mut lo = 0 as Border;
let mut hi;
while {hi=lo+(self.table[value] as Border); hi} <= offset {
lo = hi;
value += 1;
}
(value, lo, hi)
}
fn get_denominator(&self) -> Border {
self.total
}
}
pub struct SumProxy<'a> {
first: &'a Model,
second: &'a Model,
w_first: Border,
w_second: Border,
w_shift: Border,
}
impl<'a> SumProxy<'a> {
pub fn new(wa: Border, fa: &'a Model, wb: Border, fb: &'a Model, shift: Border) -> SumProxy<'a> {
assert_eq!(fa.get_frequencies().len(), fb.get_frequencies().len());
SumProxy {
first: fa,
second: fb,
w_first: wa,
w_second: wb,
w_shift: shift,
}
}
}
impl<'a> super::Model<usize> for SumProxy<'a> {
fn get_range(&self, value: usize) -> (Border,Border) {
let (lo0, hi0) = self.first.get_range(value);
let (lo1, hi1) = self.second.get_range(value);
let (wa, wb, ws) = (self.w_first, self.w_second, self.w_shift as usize);
((wa*lo0 + wb*lo1)>>ws, (wa*hi0 + wb*hi1)>>ws)
}
fn find_value(&self, offset: Border) -> (usize,Border,Border) {
assert!(offset < self.get_denominator(),
"Invalid frequency offset {} requested under total {}",
offset, self.get_denominator());
let mut value = 0;
let mut lo = 0 as Border;
let mut hi;
while { hi = lo +
(self.w_first * (self.first.get_frequencies()[value] as Border) +
self.w_second * (self.second.get_frequencies()[value] as Border)) >>
(self.w_shift as usize);
hi <= offset } {
lo = hi;
value += 1;
}
(value, lo, hi)
}
fn get_denominator(&self) -> Border {
(self.w_first * self.first.get_denominator() +
self.w_second * self.second.get_denominator()) >>
(self.w_shift as usize)
}
}
pub struct ByteEncoder<W> {
pub encoder: super::Encoder<W>,
pub freq: Model,
}
impl<W: Write> ByteEncoder<W> {
pub fn new(w: W) -> ByteEncoder<W> {
let freq_max = super::RANGE_DEFAULT_THRESHOLD >> 2;
ByteEncoder {
encoder: super::Encoder::new(w),
freq: Model::new_flat(super::SYMBOL_TOTAL+1, freq_max),
}
}
pub fn finish(mut self) -> (W, io::Result<()>) {
let ret = self.encoder.encode(super::SYMBOL_TOTAL, &self.freq);
let (w,r2) = self.encoder.finish();
(w, ret.and(r2))
}
}
impl<W: Write> Write for ByteEncoder<W> {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
for byte in buf.iter() {
let value = *byte as usize;
try!(self.encoder.encode(value, &self.freq));
self.freq.update(value, 10, 1);
}
Ok(buf.len())
}
fn flush(&mut self) -> io::Result<()> {
self.encoder.flush()
}
}
pub struct ByteDecoder<R> {
pub decoder: super::Decoder<R>,
pub freq: Model,
is_eof: bool,
}
impl<R: Read> ByteDecoder<R> {
pub fn new(r: R) -> ByteDecoder<R> {
let freq_max = super::RANGE_DEFAULT_THRESHOLD >> 2;
ByteDecoder {
decoder: super::Decoder::new(r),
freq: Model::new_flat(super::SYMBOL_TOTAL+1, freq_max),
is_eof: false,
}
}
pub fn finish(self) -> (R, io::Result<()>) {
self.decoder.finish()
}
}
impl<R: Read> Read for ByteDecoder<R> {
fn read(&mut self, dst: &mut [u8]) -> io::Result<usize> {
if self.is_eof {
return Ok(0)
}
let mut amount = 0;
for out_byte in dst.iter_mut() {
let value = try!(self.decoder.decode(&self.freq));
if value == super::SYMBOL_TOTAL {
self.is_eof = true;
break
}
self.freq.update(value, 10, 1);
*out_byte = value as u8;
amount += 1;
}
Ok(amount)
}
}