use std::fmt;
use std::io::{Read, Seek, Write};
use crate::error::RiegeliError;
use crate::record_reader::RecordReader;
use crate::record_writer::RecordWriter;
use super::field_iter::{FieldValue, ProtoFieldIter, copy_fields};
use super::handler::HandleField;
use super::wire::{WireType, is_parseable_proto_message, is_proto_message};
use super::writer::SerializedMessageWriter;
#[derive(Debug)]
pub struct StreamError {
pub record_index: usize,
pub source: RiegeliError,
}
impl fmt::Display for StreamError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(
f,
"error at record index {}: {}",
self.record_index, self.source
)
}
}
impl std::error::Error for StreamError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
Some(&self.source)
}
}
pub fn for_each_proto_record<R, H, F>(
reader: &mut RecordReader<R>,
handlers: &mut H,
mut fallback: Option<&mut F>,
) -> Result<(), StreamError>
where
R: Read + Seek,
H: HandleField,
F: FnMut(usize, &[u8]),
{
let mut record_index: usize = 0;
loop {
let record = reader.read_record().map_err(|e| StreamError {
record_index,
source: e,
})?;
let record = match record {
Some(r) => r,
None => break, };
if is_proto_message(&record) {
super::handler::read_message(&record, handlers).map_err(|e| StreamError {
record_index,
source: e,
})?;
} else if let Some(fb) = fallback.as_deref_mut() {
fb(record_index, &record);
}
record_index += 1;
}
Ok(())
}
pub fn extract_varint_column<R: Read + Seek>(
reader: &mut RecordReader<R>,
field_number: u32,
) -> Result<Vec<u64>, StreamError> {
let mut values = Vec::new();
let mut record_index: usize = 0;
loop {
let record = reader.read_record().map_err(|e| StreamError {
record_index,
source: e,
})?;
let record = match record {
Some(r) => r,
None => break,
};
if is_proto_message(&record) {
let iter = ProtoFieldIter::new(&record);
let mut group_depth: usize = 0;
for result in iter {
let field = result.map_err(|e| StreamError {
record_index,
source: e,
})?;
match field.wire_type {
WireType::StartGroup => group_depth += 1,
WireType::EndGroup => group_depth = group_depth.saturating_sub(1),
_ => {
if group_depth == 0
&& field.field_number == field_number
&& let FieldValue::Varint(v) = field.value
{
values.push(v);
}
}
}
}
}
record_index += 1;
}
Ok(values)
}
pub fn filter_fields_to_writer<R, W>(
reader: &mut RecordReader<R>,
writer: &mut RecordWriter<W>,
field_numbers: &[u32],
) -> Result<(), StreamError>
where
R: Read + Seek,
W: Write,
{
let mut record_index: usize = 0;
loop {
let record = reader.read_record().map_err(|e| StreamError {
record_index,
source: e,
})?;
let record = match record {
Some(r) => r,
None => break,
};
if is_proto_message(&record) {
let mut msg_writer = SerializedMessageWriter::new();
copy_fields(&record, field_numbers, &mut msg_writer).map_err(|e| StreamError {
record_index,
source: e,
})?;
let filtered = msg_writer.finish().map_err(|e| StreamError {
record_index,
source: e,
})?;
writer.write_record(&filtered).map_err(|e| StreamError {
record_index,
source: e,
})?;
} else if is_parseable_proto_message(&record) {
return Err(StreamError {
record_index,
source: RiegeliError::MalformedData(
"record parses as a proto message but uses non-canonical varint \
encoding; refusing to pass it through unfiltered"
.into(),
),
});
} else {
writer.write_record(&record).map_err(|e| StreamError {
record_index,
source: e,
})?;
}
record_index += 1;
}
Ok(())
}
#[cfg(test)]
mod tests {
use std::io::Cursor;
use super::{StreamError, extract_varint_column, filter_fields_to_writer};
use crate::error::RiegeliError;
use crate::record_reader::{ReaderOptions, RecordReader};
use crate::record_writer::{RecordWriter, WriterOptions};
fn write_records_file(records: &[Vec<u8>]) -> Vec<u8> {
let mut buf = Cursor::new(Vec::new());
{
let mut w = RecordWriter::new(&mut buf, WriterOptions::new()).unwrap();
for rec in records {
w.write_record(rec).unwrap();
}
w.flush().unwrap();
}
buf.into_inner()
}
fn read_records_file(data: &[u8]) -> Vec<Vec<u8>> {
let mut reader = RecordReader::new(Cursor::new(data), ReaderOptions::new()).unwrap();
let mut out = Vec::new();
while let Some(rec) = reader.read_record().unwrap() {
out.push(rec);
}
out
}
#[test]
fn extract_varint_column_ignores_group_scoped_fields() {
let group_record = vec![0x13, 0x08, 42, 0x14, 0x08, 0x07];
let file = write_records_file(&[group_record]);
let mut reader = RecordReader::new(Cursor::new(file), ReaderOptions::new()).unwrap();
let values = extract_varint_column(&mut reader, 1).unwrap();
assert_eq!(values, vec![7]);
}
#[test]
fn filter_refuses_to_pass_noncanonical_proto_record_through() {
let mut record = vec![0x08, 0x80, 0x00, 0x12, 0x06];
record.extend_from_slice(b"secret");
let file = write_records_file(&[record]);
let mut reader = RecordReader::new(Cursor::new(file), ReaderOptions::new()).unwrap();
let mut out = Cursor::new(Vec::new());
let mut writer = RecordWriter::new(&mut out, WriterOptions::new()).unwrap();
let err = filter_fields_to_writer(&mut reader, &mut writer, &[1])
.expect_err("a non-canonical proto record must not bypass field filtering");
assert_eq!(err.record_index, 0);
}
#[test]
fn filter_still_passes_genuinely_non_proto_records_through() {
let non_proto = vec![0x0F, 0xFF, 0x00];
let proto = vec![0x08, 0x07]; let file = write_records_file(&[non_proto.clone(), proto.clone()]);
let mut reader = RecordReader::new(Cursor::new(file), ReaderOptions::new()).unwrap();
let mut out = Cursor::new(Vec::new());
{
let mut writer = RecordWriter::new(&mut out, WriterOptions::new()).unwrap();
filter_fields_to_writer(&mut reader, &mut writer, &[1]).unwrap();
writer.flush().unwrap();
}
let records = read_records_file(&out.into_inner());
assert_eq!(records, vec![non_proto, proto]);
}
#[test]
fn extract_varint_column_skips_record_with_field_only_in_group() {
let group_only_record = vec![0x13, 0x08, 42, 0x14];
let file = write_records_file(&[group_only_record]);
let mut reader = RecordReader::new(Cursor::new(file), ReaderOptions::new()).unwrap();
let values = extract_varint_column(&mut reader, 1).unwrap();
assert!(values.is_empty());
}
#[test]
fn stream_error_display_and_source() {
use std::error::Error;
let err = StreamError {
record_index: 7,
source: RiegeliError::MalformedData("inner error".into()),
};
let display = format!("{err}");
assert!(display.contains("record index 7"));
assert!(display.contains("inner error"));
let source = err.source().unwrap();
let source_display = format!("{source}");
assert!(source_display.contains("inner error"));
}
}