use ::base64::{
decoded_len_estimate, encoded_len,
engine::GeneralPurpose,
prelude::{BASE64_STANDARD, BASE64_URL_SAFE, Engine},
};
use serde::{Deserialize, Serialize};
use tocat_api::{
Boundaries, BuildCtx, Ctx, Needs, Plugin, PluginError, PluginFactory, Result, Stage,
};
pub const BASE64: &str = "base64";
pub const UNBASE64: &str = "unbase64";
const MISFRAMED: &str = concat!(
"message is not a whole number of base64 groups: unbase64 decodes one ",
"complete message per call, so it needs an unframe stage ahead of it",
);
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Deserialize, Serialize)]
#[serde(rename_all = "kebab-case")]
pub enum Alphabet {
#[default]
Standard,
UrlSafe,
}
impl Alphabet {
fn engine(self) -> &'static GeneralPurpose {
match self {
Self::Standard => &BASE64_STANDARD,
Self::UrlSafe => &BASE64_URL_SAFE,
}
}
}
#[derive(Debug, Clone, Default, Deserialize, Serialize)]
#[serde(rename_all = "kebab-case", deny_unknown_fields)]
pub struct Base64Config {
#[serde(default)]
pub alphabet: Alphabet,
}
#[derive(Debug, Clone, Default, Deserialize, Serialize)]
#[serde(rename_all = "kebab-case", deny_unknown_fields)]
pub struct Unbase64Config {
#[serde(default)]
pub alphabet: Alphabet,
#[serde(default)]
pub accept_unpadded: bool,
}
fn decode_error(err: ::base64::DecodeSliceError) -> PluginError {
PluginError::runtime(UNBASE64, err)
}
fn encode_error(err: ::base64::EncodeSliceError) -> PluginError {
PluginError::runtime(BASE64, err)
}
pub struct Base64 {
engine: &'static GeneralPurpose,
out: Vec<u8>,
}
impl Base64 {
fn new(alphabet: Alphabet) -> Self {
Self {
engine: alphabet.engine(),
out: Vec::new(),
}
}
}
impl Plugin for Base64 {
fn name(&self) -> &str {
BASE64
}
fn on_bytes(&mut self, ctx: &mut Ctx<'_>, input: &[u8]) -> Result<()> {
if input.is_empty() {
return Ok(());
}
let len = encoded_len(input.len(), true)
.ok_or_else(|| PluginError::runtime(BASE64, "encoded message would overflow usize"))?;
self.out.resize(len, 0);
let written = self
.engine
.encode_slice(input, &mut self.out)
.map_err(encode_error)?;
ctx.forward(&self.out[..written]);
Ok(())
}
fn boundaries(&self) -> Boundaries {
Boundaries::Preserve
}
}
pub struct Unbase64 {
engine: &'static GeneralPurpose,
accept_unpadded: bool,
repadded: Vec<u8>,
out: Vec<u8>,
}
impl Unbase64 {
fn new(config: Unbase64Config) -> Self {
Self {
engine: config.alphabet.engine(),
accept_unpadded: config.accept_unpadded,
repadded: Vec::new(),
out: Vec::new(),
}
}
fn decode_into(
engine: &GeneralPurpose,
out: &mut Vec<u8>,
ctx: &mut Ctx<'_>,
message: &[u8],
) -> Result<()> {
out.resize(decoded_len_estimate(message.len()), 0);
let written = engine.decode_slice(message, out).map_err(decode_error)?;
ctx.forward(&out[..written]);
Ok(())
}
}
impl Plugin for Unbase64 {
fn name(&self) -> &str {
UNBASE64
}
fn on_bytes(&mut self, ctx: &mut Ctx<'_>, input: &[u8]) -> Result<()> {
if input.is_empty() {
return Ok(());
}
let remainder = input.len() % 4;
if remainder == 0 {
return Self::decode_into(self.engine, &mut self.out, ctx, input);
}
if remainder == 1 || !self.accept_unpadded {
return Err(PluginError::runtime(UNBASE64, MISFRAMED));
}
self.repadded.clear();
self.repadded.extend_from_slice(input);
self.repadded.resize(input.len() + 4 - remainder, b'=');
Self::decode_into(self.engine, &mut self.out, ctx, &self.repadded)
}
fn boundaries(&self) -> Boundaries {
Boundaries::Preserve
}
fn needs(&self) -> Needs {
Needs::Upstream
}
}
pub struct Base64Factory;
impl PluginFactory for Base64Factory {
fn name(&self) -> &str {
BASE64
}
fn description(&self) -> &str {
"base64-encode this direction, one message per chunk"
}
fn build(&self, ctx: &mut BuildCtx<'_>) -> Result<Stage> {
let config: Base64Config = ctx.config()?;
Ok(Stage::filter(Base64::new(config.alphabet)))
}
}
pub struct Unbase64Factory;
impl PluginFactory for Unbase64Factory {
fn name(&self) -> &str {
UNBASE64
}
fn description(&self) -> &str {
"base64-decode this direction, one message per chunk"
}
fn build(&self, ctx: &mut BuildCtx<'_>) -> Result<Stage> {
let config: Unbase64Config = ctx.config()?;
Ok(Stage::filter(Unbase64::new(config)))
}
}
#[cfg(test)]
mod tests {
use serde_json::{Value, json};
use tocat_api::{
ChannelId, ChannelTarget, Direction, EffectSink, Emission, HostBuilder, LogLevel,
PipelineMeta, StageInfo,
};
use super::*;
const SAMPLE: &[u8] = b"the quick brown fox jumps over 13 lazy dogs";
struct NullHost;
impl HostBuilder for NullHost {
fn open_channel(&mut self, _target: ChannelTarget) -> Result<ChannelId> {
Ok(ChannelId(0))
}
}
#[derive(Default)]
struct Silent;
impl EffectSink for Silent {
fn write(&mut self, _channel: ChannelId, _bytes: &[u8]) {}
fn log(&mut self, _level: LogLevel, _stage: &str, _message: &str) {}
}
fn meta() -> PipelineMeta {
PipelineMeta::new(Direction::SourceToSink, "src", "sink")
}
fn build(factory: &dyn PluginFactory, config: Value) -> Box<dyn Plugin> {
let map = config.as_object().expect("object").clone();
let meta = meta();
let mut host = NullHost;
let stage = StageInfo {
index: 0,
total: 1,
name: factory.name(),
upstream: "src",
downstream: "sink",
};
let mut ctx = BuildCtx::new(factory.name(), &map, &meta, stage, &mut host);
match factory.build(&mut ctx).expect("build") {
Stage::Filter(plugin) => plugin,
Stage::External(_) => unreachable!("base64 stages are filters"),
}
}
fn try_feed(plugin: &mut dyn Plugin, message: &[u8]) -> Result<Vec<u8>> {
let name = plugin.name().to_owned();
let meta = meta();
let mut emission = Emission::new();
let mut sink = Silent;
{
let mut ctx = Ctx::new(&meta, &name, message, &mut emission, &mut sink);
plugin.on_bytes(&mut ctx, message)?;
}
Ok(emission.bytes().to_vec())
}
fn feed(plugin: &mut dyn Plugin, message: &[u8]) -> Vec<u8> {
try_feed(plugin, message).expect("on_bytes")
}
#[test]
fn round_trips_a_message() {
let mut encoder = build(&Base64Factory, json!({}));
let mut decoder = build(&Unbase64Factory, json!({}));
let wire = feed(encoder.as_mut(), SAMPLE);
assert_eq!(feed(decoder.as_mut(), &wire), SAMPLE);
}
#[test]
fn round_trips_every_length() {
let mut encoder = build(&Base64Factory, json!({}));
let mut decoder = build(&Unbase64Factory, json!({}));
for len in 0..=64usize {
let plain: Vec<u8> = (0..len).map(|i| (i * 7 + 3) as u8).collect();
let wire = feed(encoder.as_mut(), &plain);
assert_eq!(feed(decoder.as_mut(), &wire), plain, "length {len}");
}
}
#[test]
fn each_call_is_a_self_contained_message() {
let mut encoder = build(&Base64Factory, json!({}));
let mut decoder = build(&Unbase64Factory, json!({}));
let first = feed(encoder.as_mut(), b"one");
let second = feed(encoder.as_mut(), b"two");
assert_eq!(first, BASE64_STANDARD.encode(b"one").into_bytes());
assert_eq!(second, BASE64_STANDARD.encode(b"two").into_bytes());
assert_eq!(feed(decoder.as_mut(), &second), b"two");
assert_eq!(feed(decoder.as_mut(), &first), b"one");
}
#[test]
fn encodes_a_message_that_needs_padding() {
let mut encoder = build(&Base64Factory, json!({}));
assert_eq!(feed(encoder.as_mut(), b"ab"), b"YWI=");
}
#[test]
fn empty_message_emits_nothing() {
let mut encoder = build(&Base64Factory, json!({}));
let mut decoder = build(&Unbase64Factory, json!({}));
assert!(feed(encoder.as_mut(), b"").is_empty());
assert!(feed(decoder.as_mut(), b"").is_empty());
}
#[test]
fn decoder_rejects_a_message_cut_mid_group() {
let wire = BASE64_STANDARD.encode(SAMPLE).into_bytes();
let mut decoder = build(&Unbase64Factory, json!({}));
for cut in [1usize, 2, 3] {
let err = try_feed(decoder.as_mut(), &wire[..wire.len() - cut])
.expect_err("a partial group must not decode");
assert!(
err.to_string().contains("unframe"),
"the error must point at the framing stage: {err}",
);
}
}
#[test]
fn decoder_accepts_an_unpadded_message_when_configured() {
let mut wire = BASE64_STANDARD.encode(SAMPLE).into_bytes();
while wire.last() == Some(&b'=') {
wire.pop();
}
let mut decoder = build(&Unbase64Factory, json!({ "accept-unpadded": true }));
assert_eq!(feed(decoder.as_mut(), &wire), SAMPLE);
assert!(try_feed(decoder.as_mut(), &wire[..5]).is_err());
}
#[test]
fn decoder_rejects_characters_outside_the_alphabet() {
let mut decoder = build(&Unbase64Factory, json!({}));
assert!(try_feed(decoder.as_mut(), b"aGVs*G8=").is_err());
}
#[test]
fn url_safe_alphabet_avoids_plus_and_slash() {
let plain = [0xfb_u8, 0xff, 0xbf];
let config = json!({ "alphabet": "url-safe" });
let mut encoder = build(&Base64Factory, config.clone());
let wire = feed(encoder.as_mut(), &plain);
assert_eq!(wire, b"-_-_");
let mut decoder = build(&Unbase64Factory, config);
assert_eq!(feed(decoder.as_mut(), &wire), plain);
}
#[test]
fn both_stages_preserve_boundaries() {
let encoder = build(&Base64Factory, json!({}));
let decoder = build(&Unbase64Factory, json!({}));
assert_eq!(encoder.boundaries(), Boundaries::Preserve);
assert_eq!(decoder.boundaries(), Boundaries::Preserve);
}
#[test]
fn rejects_unknown_config_keys() {
let map = json!({ "level": 3 }).as_object().unwrap().clone();
let meta = meta();
let mut host = NullHost;
let stage = StageInfo {
index: 0,
total: 1,
name: BASE64,
upstream: "src",
downstream: "sink",
};
let mut ctx = BuildCtx::new(BASE64, &map, &meta, stage, &mut host);
assert!(Base64Factory.build(&mut ctx).is_err());
}
}