use std::io::{Read, Write};
use alloc::boxed::Box;
use alloc::vec::Vec;
use core::marker::PhantomData;
use crate::{Config, Result};
pub trait Codec {
fn serialize_frame<T: nextjson::NsonSerialize + ?Sized>(&self, value: &T) -> Result<Vec<u8>>;
fn deserialize_frame<T: for<'a> nextjson::NsonDeserialize<'a>>(
&self,
input: &[u8],
) -> Result<T>;
}
impl Codec for Config {
fn serialize_frame<T: nextjson::NsonSerialize + ?Sized>(&self, value: &T) -> Result<Vec<u8>> {
self.serialize(value)
}
fn deserialize_frame<T: for<'a> nextjson::NsonDeserialize<'a>>(
&self,
input: &[u8],
) -> Result<T> {
self.deserialize(input)
}
}
#[cfg(feature = "cbor")]
impl Codec for crate::CborConfig {
fn serialize_frame<T: nextjson::NsonSerialize + ?Sized>(&self, value: &T) -> Result<Vec<u8>> {
self.serialize(value)
}
fn deserialize_frame<T: for<'a> nextjson::NsonDeserialize<'a>>(
&self,
input: &[u8],
) -> Result<T> {
self.deserialize(input)
}
}
#[cfg(feature = "compression")]
impl Codec for crate::CompressedConfig {
fn serialize_frame<T: nextjson::NsonSerialize + ?Sized>(&self, value: &T) -> Result<Vec<u8>> {
self.serialize(value)
}
fn deserialize_frame<T: for<'a> nextjson::NsonDeserialize<'a>>(
&self,
input: &[u8],
) -> Result<T> {
self.deserialize(input)
}
}
#[cfg(feature = "encryption")]
impl Codec for crate::EncryptedConfig {
fn serialize_frame<T: nextjson::NsonSerialize + ?Sized>(&self, value: &T) -> Result<Vec<u8>> {
self.serialize(value)
}
fn deserialize_frame<T: for<'a> nextjson::NsonDeserialize<'a>>(
&self,
input: &[u8],
) -> Result<T> {
self.deserialize(input)
}
}
#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
pub struct Untrusted;
#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
pub struct Authenticated;
#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
pub struct Closed;
pub trait AuthLevel {}
impl AuthLevel for Untrusted {}
impl AuthLevel for Authenticated {}
#[derive(Debug)]
pub struct Verified<T>(T);
impl<T> Verified<T> {
pub fn into_inner(self) -> T {
self.0
}
}
type VerifyFn = dyn Fn(&[u8]) -> Result<()>;
pub struct Verifier(Box<VerifyFn>);
impl Verifier {
pub fn new(verify: impl Fn(&[u8]) -> Result<()> + 'static) -> Self {
Self(Box::new(verify))
}
fn check(&self, bytes: &[u8]) -> Result<()> {
(self.0)(bytes)
}
}
pub struct TrustedConfig<C, A: AuthLevel> {
inner: C,
verifier: Option<Verifier>,
marker: PhantomData<A>,
}
impl<C: Codec, A: AuthLevel> TrustedConfig<C, A> {
pub fn serialize<T: nextjson::NsonSerialize + ?Sized>(&self, value: &T) -> Result<Vec<u8>> {
self.inner.serialize_frame(value)
}
}
impl<C: Codec> TrustedConfig<C, Untrusted> {
pub fn unauthenticated(config: C) -> Self {
Self {
inner: config,
verifier: None,
marker: PhantomData,
}
}
pub fn deserialize_untrusted<T: for<'a> nextjson::NsonDeserialize<'a>>(
&self,
input: &[u8],
) -> Result<T> {
self.inner.deserialize_frame(input)
}
pub fn authenticate(self, verifier: Verifier) -> TrustedConfig<C, Authenticated> {
TrustedConfig {
inner: self.inner,
verifier: Some(verifier),
marker: PhantomData,
}
}
}
impl<C: Codec> TrustedConfig<C, Authenticated> {
pub fn deserialize<T: for<'a> nextjson::NsonDeserialize<'a>>(&self, input: &[u8]) -> Result<T> {
self.verifier
.as_ref()
.ok_or(crate::Error::Trust(
"authenticated config lost its verifier",
))?
.check(input)?;
self.inner.deserialize_frame(input)
}
pub fn deserialize_verified<T: for<'a> nextjson::NsonDeserialize<'a>>(
&self,
input: &[u8],
) -> Result<Verified<T>> {
self.deserialize(input).map(Verified)
}
pub fn into_inner(self) -> C {
self.inner
}
}
pub trait SessionState {}
impl SessionState for Handshake {}
impl SessionState for Authenticated {}
impl SessionState for Closed {}
#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
pub struct Handshake;
pub struct Session<C: Codec, S: SessionState, R> {
codec: C,
max_frame_len: Option<u64>,
verifier: Option<Verifier>,
reader: R,
marker: PhantomData<S>,
}
impl<C: Codec, R: Read> Session<C, Handshake, R> {
pub fn new(codec: C, reader: R) -> Self {
Self {
codec,
max_frame_len: Some(crate::DEFAULT_SIZE_LIMIT),
verifier: None,
reader,
marker: PhantomData,
}
}
pub fn with_max_frame_len(mut self, limit: Option<u64>) -> Self {
self.max_frame_len = limit;
self
}
pub fn authenticate(self, verifier: Verifier) -> Session<C, Authenticated, R> {
Session {
codec: self.codec,
max_frame_len: self.max_frame_len,
verifier: Some(verifier),
reader: self.reader,
marker: PhantomData,
}
}
}
impl<C: Codec, R: Read> Session<C, Authenticated, R> {
pub fn recv<T: for<'a> nextjson::NsonDeserialize<'a>>(&mut self) -> Result<T> {
self.recv_verified().map(Verified::into_inner)
}
pub fn recv_verified<T: for<'a> nextjson::NsonDeserialize<'a>>(
&mut self,
) -> Result<Verified<T>> {
let mut length_bytes = [0_u8; 8];
self.reader.read_exact(&mut length_bytes)?;
let length = u64::from_le_bytes(length_bytes);
let length = usize::try_from(length)
.map_err(|_| crate::Error::Trust("session frame length does not fit usize"))?;
if let Some(limit) = self.max_frame_len {
if length as u64 > limit {
return Err(crate::Error::SizeLimit { limit });
}
}
let mut frame = vec![0_u8; length];
self.reader.read_exact(&mut frame)?;
self.verifier
.as_ref()
.ok_or(crate::Error::Trust(
"authenticated session lost its verifier",
))?
.check(&frame)?;
self.codec.deserialize_frame(&frame).map(Verified)
}
pub fn send<W: Write, T: nextjson::NsonSerialize + ?Sized>(
&self,
writer: &mut W,
value: &T,
) -> Result<()> {
let payload = self.codec.serialize_frame(value)?;
let length = u64::try_from(payload.len())
.map_err(|_| crate::Error::Trust("session frame too large"))?;
writer.write_all(&length.to_le_bytes())?;
writer.write_all(&payload)?;
Ok(())
}
pub fn close(self) -> Session<C, Closed, R> {
Session {
codec: self.codec,
max_frame_len: self.max_frame_len,
verifier: None,
reader: self.reader,
marker: PhantomData,
}
}
}
impl<C: Codec, R> Session<C, Closed, R> {}
#[cfg(test)]
mod tests {
use super::*;
use std::io::Cursor;
#[derive(Debug, PartialEq, nextjson::NsonSerialize, nextjson::NsonDeserialize)]
struct Message {
id: u64,
body: String,
}
#[test]
fn unauthenticated_config_requires_explicit_call() {
let config = TrustedConfig::<Config, Untrusted>::unauthenticated(Config::standard());
let value = Message {
id: 1,
body: "hello".into(),
};
let frame = config.serialize(&value).unwrap();
let decoded: Message = config.deserialize_untrusted(&frame).unwrap();
assert_eq!(decoded, value);
}
fn frame_verifier(known_good: Vec<u8>) -> Verifier {
Verifier::new(move |bytes| {
if bytes == known_good.as_slice() {
Ok(())
} else {
Err(crate::Error::Trust(
"frame does not match the authenticated original",
))
}
})
}
#[test]
fn authenticated_config_verifies_before_deserializing() {
let message = Message {
id: 7,
body: "verified".into(),
};
let frame = {
let unauthenticated =
TrustedConfig::<Config, Untrusted>::unauthenticated(Config::standard());
unauthenticated.serialize(&message).unwrap()
};
let authenticated = TrustedConfig::<Config, Untrusted>::unauthenticated(Config::standard())
.authenticate(frame_verifier(frame.clone()));
let decoded: Message = authenticated.deserialize(&frame).unwrap();
assert_eq!(decoded.body, "verified");
let verified: Verified<Message> = authenticated.deserialize_verified(&frame).unwrap();
assert_eq!(verified.into_inner().id, 7);
let rejecting = TrustedConfig::<Config, Untrusted>::unauthenticated(Config::standard())
.authenticate(Verifier::new(|_| {
Err(crate::Error::Trust("rejected by policy"))
}));
assert!(rejecting.deserialize::<Message>(&frame).is_err());
let mut corrupted = frame.clone();
corrupted[0] ^= 1;
assert!(authenticated.deserialize::<Message>(&corrupted).is_err());
}
#[test]
fn session_receives_only_after_authentication() {
let codec = Config::standard();
let message = Message {
id: 3,
body: "stream".into(),
};
let payload = codec.serialize(&message).unwrap();
let mut stream = Vec::new();
stream.extend_from_slice(&(payload.len() as u64).to_le_bytes());
stream.extend_from_slice(&payload);
stream.extend_from_slice(&(payload.len() as u64).to_le_bytes());
stream.extend_from_slice(&payload);
let mut session =
Session::new(codec, Cursor::new(stream)).authenticate(Verifier::new(|_| Ok(())));
let first: Verified<Message> = session.recv_verified().unwrap();
assert_eq!(first.into_inner().body, "stream");
let second: Message = session.recv().unwrap();
assert_eq!(second.body, "stream");
let _closed = session.close();
}
#[test]
#[cfg(feature = "encryption")]
fn encrypted_config_is_an_authenticated_codec() {
let key = crate::EncryptionKey::new([0x42; 32]);
let encrypted = crate::options().with_encryption(key);
let value = Message {
id: 9,
body: "aead".into(),
};
let frame = encrypted.serialize(&value).unwrap();
let decoded: Message = encrypted.deserialize(&frame).unwrap();
assert_eq!(decoded, value);
}
}