use std::collections::HashMap;
use super::field_iter::{FieldValue, ProtoField, ProtoFieldIter};
use super::wire::WireType;
pub trait HandleField {
fn dispatch(&mut self, field: &ProtoField<'_>) -> Result<bool, crate::RiegeliError>;
}
pub trait FieldHandler {
const FIELD_NUMBER: u32;
fn handle_varint(&mut self, _value: u64) -> Result<(), crate::RiegeliError> {
Ok(())
}
fn handle_fixed32(&mut self, _value: u32) -> Result<(), crate::RiegeliError> {
Ok(())
}
fn handle_fixed64(&mut self, _value: u64) -> Result<(), crate::RiegeliError> {
Ok(())
}
fn handle_length_delimited(&mut self, _data: &[u8]) -> Result<(), crate::RiegeliError> {
Ok(())
}
fn handle_start_group(&mut self) -> Result<(), crate::RiegeliError> {
Ok(())
}
fn handle_end_group(&mut self) -> Result<(), crate::RiegeliError> {
Ok(())
}
}
fn dispatch_to_handler<H: FieldHandler>(
handler: &mut H,
field: &ProtoField<'_>,
) -> Result<bool, crate::RiegeliError> {
if field.field_number != H::FIELD_NUMBER {
return Ok(false);
}
match &field.value {
FieldValue::Varint(v) => handler.handle_varint(*v)?,
FieldValue::Fixed32(v) => handler.handle_fixed32(*v)?,
FieldValue::Fixed64(v) => handler.handle_fixed64(*v)?,
FieldValue::LengthDelimited(data) => handler.handle_length_delimited(data)?,
FieldValue::StartGroup => handler.handle_start_group()?,
FieldValue::EndGroup => handler.handle_end_group()?,
}
Ok(true)
}
pub struct StaticHandlerSet1<H1: FieldHandler> {
pub h1: H1,
}
impl<H1: FieldHandler> HandleField for StaticHandlerSet1<H1> {
fn dispatch(&mut self, field: &ProtoField<'_>) -> Result<bool, crate::RiegeliError> {
dispatch_to_handler(&mut self.h1, field)
}
}
pub struct StaticHandlerSet2<H1: FieldHandler, H2: FieldHandler> {
pub h1: H1,
pub h2: H2,
}
impl<H1: FieldHandler, H2: FieldHandler> HandleField for StaticHandlerSet2<H1, H2> {
fn dispatch(&mut self, field: &ProtoField<'_>) -> Result<bool, crate::RiegeliError> {
if dispatch_to_handler(&mut self.h1, field)? {
return Ok(true);
}
dispatch_to_handler(&mut self.h2, field)
}
}
pub struct StaticHandlerSet3<H1: FieldHandler, H2: FieldHandler, H3: FieldHandler> {
pub h1: H1,
pub h2: H2,
pub h3: H3,
}
impl<H1: FieldHandler, H2: FieldHandler, H3: FieldHandler> HandleField
for StaticHandlerSet3<H1, H2, H3>
{
fn dispatch(&mut self, field: &ProtoField<'_>) -> Result<bool, crate::RiegeliError> {
if dispatch_to_handler(&mut self.h1, field)? {
return Ok(true);
}
if dispatch_to_handler(&mut self.h2, field)? {
return Ok(true);
}
dispatch_to_handler(&mut self.h3, field)
}
}
pub struct StaticHandlerSet;
impl StaticHandlerSet {
#[allow(clippy::new_ret_no_self)] pub fn new<H1: FieldHandler>(h1: H1) -> StaticHandlerSet1<H1> {
StaticHandlerSet1 { h1 }
}
}
impl<H1: FieldHandler> StaticHandlerSet1<H1> {
pub fn and<H2: FieldHandler>(self, h2: H2) -> StaticHandlerSet2<H1, H2> {
StaticHandlerSet2 { h1: self.h1, h2 }
}
}
impl<H1: FieldHandler, H2: FieldHandler> StaticHandlerSet2<H1, H2> {
pub fn and<H3: FieldHandler>(self, h3: H3) -> StaticHandlerSet3<H1, H2, H3> {
StaticHandlerSet3 {
h1: self.h1,
h2: self.h2,
h3,
}
}
}
type DynamicHandlerKey = (u32, u8);
pub struct DynamicHandlerSet<'a> {
handlers: HashMap<
DynamicHandlerKey,
Box<dyn FnMut(&FieldValue<'_>) -> Result<(), crate::RiegeliError> + 'a>,
>,
}
impl<'a> DynamicHandlerSet<'a> {
pub fn new() -> Self {
DynamicHandlerSet {
handlers: HashMap::new(),
}
}
pub fn on_varint<F>(&mut self, field_number: u32, mut f: F)
where
F: FnMut(u64) -> Result<(), crate::RiegeliError> + 'a,
{
self.handlers.insert(
(field_number, WireType::Varint as u8),
Box::new(move |value| {
if let FieldValue::Varint(v) = value {
f(*v)
} else {
Ok(())
}
}),
);
}
pub fn on_fixed32<F>(&mut self, field_number: u32, mut f: F)
where
F: FnMut(u32) -> Result<(), crate::RiegeliError> + 'a,
{
self.handlers.insert(
(field_number, WireType::Fixed32 as u8),
Box::new(move |value| {
if let FieldValue::Fixed32(v) = value {
f(*v)
} else {
Ok(())
}
}),
);
}
pub fn on_fixed64<F>(&mut self, field_number: u32, mut f: F)
where
F: FnMut(u64) -> Result<(), crate::RiegeliError> + 'a,
{
self.handlers.insert(
(field_number, WireType::Fixed64 as u8),
Box::new(move |value| {
if let FieldValue::Fixed64(v) = value {
f(*v)
} else {
Ok(())
}
}),
);
}
pub fn on_length_delimited<F>(&mut self, field_number: u32, mut f: F)
where
F: FnMut(&[u8]) -> Result<(), crate::RiegeliError> + 'a,
{
self.handlers.insert(
(field_number, WireType::LengthDelimited as u8),
Box::new(move |value| {
if let FieldValue::LengthDelimited(data) = value {
f(data)
} else {
Ok(())
}
}),
);
}
pub fn on_start_group<F>(&mut self, field_number: u32, mut f: F)
where
F: FnMut() -> Result<(), crate::RiegeliError> + 'a,
{
self.handlers.insert(
(field_number, WireType::StartGroup as u8),
Box::new(move |value| {
if let FieldValue::StartGroup = value {
f()
} else {
Ok(())
}
}),
);
}
pub fn on_end_group<F>(&mut self, field_number: u32, mut f: F)
where
F: FnMut() -> Result<(), crate::RiegeliError> + 'a,
{
self.handlers.insert(
(field_number, WireType::EndGroup as u8),
Box::new(move |value| {
if let FieldValue::EndGroup = value {
f()
} else {
Ok(())
}
}),
);
}
}
impl<'a> Default for DynamicHandlerSet<'a> {
fn default() -> Self {
Self::new()
}
}
impl<'a> HandleField for DynamicHandlerSet<'a> {
fn dispatch(&mut self, field: &ProtoField<'_>) -> Result<bool, crate::RiegeliError> {
let wire_type_raw = field.wire_type as u8;
let key = (field.field_number, wire_type_raw);
if let Some(handler) = self.handlers.get_mut(&key) {
handler(&field.value)?;
Ok(true)
} else {
Ok(false)
}
}
}
pub struct EmptyHandlerSet;
impl HandleField for EmptyHandlerSet {
fn dispatch(&mut self, _field: &ProtoField<'_>) -> Result<bool, crate::RiegeliError> {
Ok(false)
}
}
pub fn read_message<H: HandleField>(
data: &[u8],
handlers: &mut H,
) -> Result<(), crate::RiegeliError> {
let iter = ProtoFieldIter::new(data);
for result in iter {
let field = result?;
handlers.dispatch(&field)?;
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::proto::SerializedMessageWriter;
struct VarintOnlyField1 {
values: Vec<u64>,
}
impl FieldHandler for VarintOnlyField1 {
const FIELD_NUMBER: u32 = 1;
fn handle_varint(&mut self, value: u64) -> Result<(), crate::RiegeliError> {
self.values.push(value);
Ok(())
}
}
#[test]
fn default_handler_methods_ignore_other_wire_types() {
let mut w = SerializedMessageWriter::new();
w.write_fixed32(1, 0xABCD).unwrap();
let data = w.finish().unwrap();
let mut handler = StaticHandlerSet::new(VarintOnlyField1 { values: vec![] });
read_message(&data, &mut handler).unwrap();
assert!(handler.h1.values.is_empty());
}
#[test]
fn reregistering_dynamic_handler_replaces_previous() {
let mut w = SerializedMessageWriter::new();
w.write_uint64(1, 42).unwrap();
let data = w.finish().unwrap();
let result = std::cell::RefCell::new(0u64);
{
let r = &result;
let mut handlers = DynamicHandlerSet::new();
handlers.on_varint(1, |_| {
panic!("first handler should have been replaced");
});
handlers.on_varint(1, move |v| {
*r.borrow_mut() = v;
Ok(())
});
read_message(&data, &mut handlers).unwrap();
}
assert_eq!(*result.borrow(), 42);
}
}