use crate::fst::VectorFst;
use crate::semiring::TropicalWeight;
use crate::Result;
use byteorder::{LittleEndian, WriteBytesExt};
use std::collections::HashMap;
use std::io::{Read, Seek, Write};
#[derive(Debug)]
pub struct FarReader<R: Read + Seek> {
reader: R,
entries: HashMap<String, usize>, }
impl<R: Read + Seek> FarReader<R> {
pub fn new(mut reader: R) -> Result<Self> {
use byteorder::{LittleEndian, ReadBytesExt};
let file_size = reader.seek(std::io::SeekFrom::End(0))?;
if file_size < 8 {
return Err(crate::Error::Io(std::io::Error::new(
std::io::ErrorKind::InvalidData,
"FAR file too small to contain index",
)));
}
reader.seek(std::io::SeekFrom::End(-8))?;
let index_pos = reader.read_u64::<LittleEndian>()?;
reader.seek(std::io::SeekFrom::Start(index_pos))?;
let num_entries = reader.read_u32::<LittleEndian>()? as usize;
let mut entries = HashMap::with_capacity(num_entries);
for _ in 0..num_entries {
let name_len = reader.read_u32::<LittleEndian>()? as usize;
let mut name_bytes = vec![0u8; name_len];
reader.read_exact(&mut name_bytes)?;
let name = String::from_utf8(name_bytes).map_err(|e| {
crate::Error::Io(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!("Invalid UTF-8 in FAR entry name: {}", e),
))
})?;
let offset = reader.read_u64::<LittleEndian>()? as usize;
entries.insert(name, offset);
}
Ok(Self { reader, entries })
}
pub fn list(&self) -> Vec<&String> {
self.entries.keys().collect()
}
pub fn read(&mut self, name: &str) -> Result<Option<VectorFst<TropicalWeight>>> {
let offset = match self.entries.get(name) {
Some(&offset) => offset,
None => return Ok(None),
};
self.reader.seek(std::io::SeekFrom::Start(offset as u64))?;
use crate::io::read_openfst;
match read_openfst(&mut self.reader) {
Ok(fst) => Ok(Some(fst)),
Err(e) => Err(e),
}
}
}
#[derive(Debug)]
pub struct FarWriter<W: Write + Seek> {
writer: W,
entries: Vec<(String, usize)>, }
impl<W: Write + Seek> FarWriter<W> {
pub fn new(writer: W) -> Self {
Self {
writer,
entries: Vec::new(),
}
}
pub fn add(&mut self, name: &str, fst: &VectorFst<TropicalWeight>) -> Result<()> {
use crate::io::write_openfst;
let pos = self.writer.stream_position()?;
write_openfst(fst, &mut self.writer)?;
self.entries.push((name.to_string(), pos as usize));
Ok(())
}
pub fn finish(mut self) -> Result<()> {
let index_pos = self.writer.stream_position()?;
self.writer
.write_u32::<LittleEndian>(self.entries.len() as u32)?;
for (name, offset) in &self.entries {
let name_bytes = name.as_bytes();
self.writer
.write_u32::<LittleEndian>(name_bytes.len() as u32)?;
self.writer.write_all(name_bytes)?;
self.writer.write_u64::<LittleEndian>(*offset as u64)?;
}
self.writer.write_u64::<LittleEndian>(index_pos)?;
self.writer.flush()?;
Ok(())
}
}
pub fn open_far<P: AsRef<std::path::Path>>(
path: P,
) -> Result<FarReader<std::io::BufReader<std::fs::File>>> {
use std::fs::File;
use std::io::BufReader;
let file = File::open(path)?;
let reader = BufReader::new(file);
FarReader::new(reader)
}
pub fn create_far<P: AsRef<std::path::Path>>(
path: P,
) -> Result<FarWriter<std::io::BufWriter<std::fs::File>>> {
use std::io::BufWriter;
let file = std::fs::File::create(path)?;
let writer = BufWriter::new(file);
Ok(FarWriter::new(writer))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::prelude::*;
use std::io::Cursor;
#[test]
fn test_far_writer_new() {
let cursor = Cursor::new(Vec::new());
let writer = FarWriter::new(cursor);
assert_eq!(writer.entries.len(), 0);
}
#[test]
fn test_far_writer_add() {
let mut fst = VectorFst::<TropicalWeight>::new();
let s0 = fst.add_state();
fst.set_start(s0);
fst.set_final(s0, TropicalWeight::one());
let cursor = Cursor::new(Vec::new());
let mut writer = FarWriter::new(cursor);
writer.add("test_fst", &fst).unwrap();
assert_eq!(writer.entries.len(), 1);
assert_eq!(writer.entries[0].0, "test_fst");
}
#[test]
fn test_far_writer_finish() {
let mut fst = VectorFst::<TropicalWeight>::new();
let s0 = fst.add_state();
fst.set_start(s0);
fst.set_final(s0, TropicalWeight::one());
let cursor = Cursor::new(Vec::new());
let mut writer = FarWriter::new(cursor);
writer.add("test_fst", &fst).unwrap();
writer.finish().unwrap();
}
#[test]
fn test_far_reader_list() {
let mut buffer = Vec::new();
{
let mut fst = VectorFst::<TropicalWeight>::new();
let s0 = fst.add_state();
fst.set_start(s0);
fst.set_final(s0, TropicalWeight::one());
let cursor = Cursor::new(&mut buffer);
let mut writer = FarWriter::new(cursor);
writer.add("fst1", &fst).unwrap();
writer.add("fst2", &fst).unwrap();
writer.finish().unwrap();
}
let cursor = Cursor::new(buffer);
let reader = FarReader::new(cursor).unwrap();
let entries = reader.list();
assert_eq!(entries.len(), 2);
}
}