use spate_core::checkpoint::AckRef;
use spate_core::deser::{Deserializer, EmitRecord, Owned};
use spate_core::error::DeserError;
use spate_core::record::{Flow, RawPayload, Record};
#[derive(Clone, Debug)]
enum Mode {
Passthrough,
SplitOn(u8),
FailOnPrefix(Vec<u8>),
}
#[derive(Clone, Debug)]
pub struct TestDeserializer {
mode: Mode,
}
impl TestDeserializer {
#[must_use]
pub fn passthrough() -> Self {
TestDeserializer {
mode: Mode::Passthrough,
}
}
#[must_use]
pub fn split_on(delimiter: u8) -> Self {
TestDeserializer {
mode: Mode::SplitOn(delimiter),
}
}
#[must_use]
pub fn fail_on_prefix(marker: impl Into<Vec<u8>>) -> Self {
TestDeserializer {
mode: Mode::FailOnPrefix(marker.into()),
}
}
}
impl Deserializer<Owned<Vec<u8>>> for TestDeserializer {
fn deserialize<'buf>(
&mut self,
raw: &RawPayload<'buf>,
ack: &AckRef,
out: &mut dyn EmitRecord<'buf, Vec<u8>>,
) -> Result<(), DeserError> {
let meta = raw.meta();
let mut emit = |payload: Vec<u8>| {
out.emit(Record {
payload,
meta,
ack: ack.clone(),
})
};
match &self.mode {
Mode::Passthrough => {
let _ = emit(raw.bytes.to_vec());
}
Mode::SplitOn(delim) => {
for segment in raw.bytes.split(|b| b == delim) {
if emit(segment.to_vec()) == Flow::Blocked {
break;
}
}
}
Mode::FailOnPrefix(marker) => {
if raw.bytes.starts_with(marker) {
return Err(DeserError::Malformed {
reason: format!("payload starts with test marker {marker:?}"),
});
}
let _ = emit(raw.bytes.to_vec());
}
}
Ok(())
}
}
#[derive(Debug, Default)]
pub struct EmitCollector<T> {
pub records: Vec<Record<T>>,
pub rejected: usize,
block_after: Option<usize>,
}
impl<T> EmitCollector<T> {
#[must_use]
pub fn new() -> Self {
EmitCollector {
records: Vec::new(),
rejected: 0,
block_after: None,
}
}
#[must_use]
pub fn blocking_after(n: usize) -> Self {
EmitCollector {
records: Vec::new(),
rejected: 0,
block_after: Some(n),
}
}
#[must_use]
pub fn payloads(&self) -> Vec<T>
where
T: Clone,
{
self.records.iter().map(|r| r.payload.clone()).collect()
}
}
impl<'buf, T> EmitRecord<'buf, T> for EmitCollector<T> {
fn emit(&mut self, rec: Record<T>) -> Flow {
if let Some(cap) = self.block_after
&& self.records.len() >= cap
{
self.rejected += 1;
return Flow::Blocked;
}
self.records.push(rec);
Flow::Continue
}
}