diff --git a/Cargo.lock b/Cargo.lock index 706d78b..576db28 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1857,6 +1857,7 @@ version = "0.1.0" dependencies = [ "aead", "bluer", + "cobs-acc", "env_logger", "ergot", "futures-concurrency", @@ -1880,9 +1881,11 @@ dependencies = [ "aead", "bbqueue", "cobs 0.5.1", + "cobs-acc", "defmt", "ergot", "jiff", + "log", "postcard", "postcard-schema", "rand_core 0.10.1", diff --git a/Cargo.toml b/Cargo.toml index 6bc7294..14d7d13 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -15,7 +15,7 @@ edition.workspace = true [dependencies] aead = { version = "0.6.0", features = ["alloc"] } -water-pod-common = { path = "./water-pod-common" } +water-pod-common = { path = "./water-pod-common", features = ["log"] } tokio = { version = "1", features = ["io-std", "io-util", "net", "rt", "macros", "signal", "sync"] } wharrgarbl = { git = "https://tangled.org/sachy.dev/wharrgarbl", package = "wharrgarbl" } getrandom = { version = "0.4.2", features = ["sys_rng"] } @@ -26,6 +26,7 @@ winnow = "1" futures-lite = "2.6.1" jiff = "0.2.32" ergot = { git = "https://github.com/jamesmunns/ergot", features = ["tokio-std"] } +cobs-acc = { git = "https://github.com/jamesmunns/ergot", package = "cobs-acc" } rand = { version = "0.10" } futures-concurrency = "7.7.1" postcard = { version = "1", default-features = false, features = ["alloc"] } diff --git a/src/main.rs b/src/main.rs index 3128363..24a0ba7 100644 --- a/src/main.rs +++ b/src/main.rs @@ -6,13 +6,12 @@ use std::{ }; use bluer::l2cap::{SocketAddr, Stream, stream::OwnedReadHalf}; -use postcard::accumulator::{CobsAccumulator, FeedResult}; use rand::rng; use tokio::{ io::{AsyncReadExt, AsyncWriteExt}, time::sleep, }; -use water_pod_common::{Request, Response}; +use water_pod_common::{Accumulator, AccumulatorResult, Encrypted, Pairing, Request, Response}; use wharrgarbl::{ Neko128, NekoServerHandshake128, transport::{RecvState, SendState}, @@ -56,12 +55,12 @@ async fn main() -> Result<(), std::io::Error> { let (mut reader, mut writer) = conn.into_split(); let mut raw_buf = vec![0u8; 251]; - let mut pair = cobs_read(&mut reader, &mut raw_buf).await?; + let pair: Pairing = cobs_read(&mut reader, &mut raw_buf).await?; let mut rng = rng(); - let (transport, pair) = Response::respond_to_pairing(server, &mut rng, &mut pair) - .map_err(std::io::Error::other)?; + let (transport, pair) = + Response::respond_to_pairing(server, &mut rng, pair).map_err(std::io::Error::other)?; let mut pair = Cursor::new(pair); @@ -71,7 +70,7 @@ async fn main() -> Result<(), std::io::Error> { let (mut send_state, mut recv_state) = transport.split(); - let req = cobs_read(&mut reader, &mut raw_buf).await?; + let req: Encrypted = cobs_read(&mut reader, &mut raw_buf).await?; log::info!("RESPONDING!"); @@ -97,11 +96,11 @@ async fn main() -> Result<(), std::io::Error> { } async fn req_handler<'a>( - request: Vec, + request: Encrypted, send_state: &mut SendState<'a, Neko128>, recv_state: &mut RecvState<'a, Neko128>, ) -> aead::Result> { - let request = Request::from_bytes(request, recv_state)?; + let request = request.from_bytes(recv_state)?; match request { Request::TimeSync(mut sync) => { @@ -109,38 +108,27 @@ async fn req_handler<'a>( sleep(Duration::from_millis(1)).await; sync.set_t2(); - Response::TimeSync(sync).to_bytes(send_state) + Encrypted::new(Response::TimeSync(sync)).to_bytes(send_state) } - _ => Response::Error.to_bytes(send_state), + _ => Encrypted::new(Response::Error).to_bytes(send_state), } } -async fn cobs_read(input: &mut OwnedReadHalf, raw_buf: &mut [u8]) -> std::io::Result> { - let mut cobs_buf = Box::new(CobsAccumulator::<251>::new()); +async fn cobs_read(input: &mut OwnedReadHalf, raw_buf: &mut [u8]) -> std::io::Result +where + T: for<'de> serde::Deserialize<'de>, +{ + let mut cobs_buf = Accumulator::new(); loop { let ct = input.read(raw_buf).await?; - if ct == 0 { - break; - } - - let buf = &raw_buf[..ct]; - let mut window = &buf[..]; - - 'cobs: while !window.is_empty() { - window = match cobs_buf.feed::>(&window) { - FeedResult::Consumed => break 'cobs, - FeedResult::OverFull(new_wind) => new_wind, - FeedResult::DeserError(new_wind) => new_wind, - FeedResult::Success { data, .. } => { - // Do something with `data: MyData` here. - - return Ok(data); - } - }; + match cobs_buf.accumulate(&mut raw_buf[..ct]) { + AccumulatorResult::Continue => continue, + AccumulatorResult::Error => break, + AccumulatorResult::Success(payload) => return Ok(payload), } } - return Err(ErrorKind::InvalidData.into()); + Err(ErrorKind::InvalidData.into()) } diff --git a/src/scan.rs b/src/scan.rs index 4b3e3c2..9a0bc1c 100644 --- a/src/scan.rs +++ b/src/scan.rs @@ -6,17 +6,14 @@ pub async fn scan_for_water(adapter: &Adapter) -> bluer::Result> log::info!("SEARCHING FOR WATER"); while let Some(b) = device_events.next().await { - match b { - bluer::AdapterEvent::DeviceAdded(address) => { - let device = adapter.device(address)?; - let device_name = device.name().await?; + if let bluer::AdapterEvent::DeviceAdded(address) = b { + let device = adapter.device(address)?; + let device_name = device.name().await?; - if device_name.is_some_and(|name| name == "SachyWater") { - log::info!("FOUND {address}"); - return Ok(Some(device)); - } + if device_name.is_some_and(|name| name == "SachyWater") { + log::info!("FOUND {address}"); + return Ok(Some(device)); } - _ => {} } } diff --git a/water-pod-common/Cargo.toml b/water-pod-common/Cargo.toml index 3953bc7..3cdc4d6 100644 --- a/water-pod-common/Cargo.toml +++ b/water-pod-common/Cargo.toml @@ -16,9 +16,12 @@ tokio = { version = "1", default-features = false, optional = true } bbqueue = "0.7.0" postcard-schema = { version = "0.2.5", features = ["alloc"] } ergot = { git = "https://github.com/jamesmunns/ergot" } +cobs-acc = { git = "https://github.com/jamesmunns/ergot", package = "cobs-acc" } cobs = { version = "0.5.1", default-features = false, features = ["alloc"] } +log = { version = "0.4", optional = true } [features] default = ["server"] server = ["jiff/std", "dep:tokio"] defmt = ["dep:defmt", "postcard/use-defmt", "jiff/defmt"] +log = ["dep:log"] diff --git a/water-pod-common/src/lib.rs b/water-pod-common/src/lib.rs index af99772..a08cf96 100644 --- a/water-pod-common/src/lib.rs +++ b/water-pod-common/src/lib.rs @@ -2,10 +2,12 @@ extern crate alloc; -use alloc::vec::Vec; +use alloc::{boxed::Box, vec::Vec}; use ergot::{endpoint, topic}; use jiff::{SignedDuration, Timestamp}; +use postcard::accumulator::{CobsAccumulator, FeedResult}; use rand_core::CryptoRng; +use serde::{Deserialize, Serialize}; use wharrgarbl::{ Neko128, NekoClientHandshake128, NekoServerHandshake128, transport::{AeadTransport, RecvState, SendState}, @@ -29,14 +31,95 @@ pub enum Response { Error, } +#[derive(Debug, serde::Serialize, serde::Deserialize, postcard_schema::Schema)] +#[cfg_attr(feature = "defmt", derive(defmt::Format))] +#[repr(transparent)] +pub struct Pairing( + #[serde( + serialize_with = "serialize_bytes", + deserialize_with = "deserialize_bytes" + )] + pub Vec, +); + +impl Pairing { + pub fn to_bytes(self) -> Result, postcard::Error> { + postcard::to_allocvec_cobs(&self) + } +} + +#[derive(Debug, serde::Serialize, serde::Deserialize, postcard_schema::Schema)] +#[cfg_attr(feature = "defmt", derive(defmt::Format))] +#[serde(bound = "T: Serialize + for<'a> Deserialize<'a>")] +pub struct Encrypted { + #[serde( + serialize_with = "serialize_bytes", + deserialize_with = "deserialize_bytes" + )] + pub payload: Vec, + _kind: core::marker::PhantomData, +} + +fn serialize_bytes(value: &[u8], serializer: S) -> Result +where + S: serde::Serializer, +{ + serializer.serialize_bytes(value) +} + +fn deserialize_bytes<'de, D>(deserializer: D) -> Result, D::Error> +where + D: serde::Deserializer<'de>, +{ + struct ByteArrayVisitor; + + impl<'de> serde::de::Visitor<'de> for ByteArrayVisitor { + type Value = Vec; + + fn expecting(&self, formatter: &mut core::fmt::Formatter) -> core::fmt::Result { + write!(formatter, "Expected an sequence of bytes") + } + + fn visit_bytes(self, v: &[u8]) -> Result + where + E: serde::de::Error, + { + v.try_into() + .map_err(|_| serde::de::Error::invalid_length(v.len(), &self)) + } + } + + deserializer.deserialize_bytes(ByteArrayVisitor) +} + +impl Deserialize<'de>> Encrypted { + pub fn new(payload: T) -> Self { + let buf = Vec::new(); + Self { + payload: postcard::to_extend(&payload, buf).unwrap(), + _kind: core::marker::PhantomData, + } + } + + pub fn to_bytes<'b>(mut self, send: &mut SendState<'b, Neko128>) -> aead::Result> { + send.encrypt(&mut self.payload, b"water-pod")?; + postcard::to_allocvec_cobs(&self).map_err(|_| aead::Error) + } + + pub fn from_bytes<'b>(mut self, recv: &mut RecvState<'b, Neko128>) -> aead::Result { + recv.decrypt(&mut self.payload, b"water-pod")?; + postcard::from_bytes(&self.payload).map_err(|_| aead::Error) + } +} + impl Request { pub fn request_pairing( client: &mut NekoClientHandshake128, rng: &mut impl CryptoRng, - buf: &mut Vec, + mut buf: Vec, ) -> aead::Result> { - client.send(rng, buf)?; - Ok(cobs::encode_vec_including_sentinels(buf)) + client.send(rng, &mut buf)?; + postcard::to_allocvec_cobs(&Pairing(buf)).map_err(|_| aead::Error) } pub fn finish_pairing( @@ -45,50 +128,18 @@ impl Request { ) -> aead::Result> { client.receive(ciphertext) } - - pub fn to_bytes<'b>(self, send: &mut SendState<'b, Neko128>) -> aead::Result> { - let mut payload = postcard::to_allocvec(&self).map_err(|_| aead::Error)?; - - send.encrypt(&mut payload, b"water-pod")?; - - Ok(cobs::encode_vec_including_sentinels(&payload)) - } - - pub fn from_bytes<'a>( - mut buf: Vec, - recv: &mut RecvState<'a, Neko128>, - ) -> aead::Result { - recv.decrypt(&mut buf, b"water-pod")?; - postcard::from_bytes(&buf).map_err(|_| aead::Error) - } } impl Response { pub fn respond_to_pairing( mut server: NekoServerHandshake128, rng: &mut impl CryptoRng, - buf: &mut Vec, + mut buf: Pairing, ) -> aead::Result<(AeadTransport, Vec)> { - let transport = server.respond(rng, buf)?; - let encoded = cobs::encode_vec_including_sentinels(buf); + let transport = server.respond(rng, &mut buf.0)?; + let encoded = Pairing(buf.0).to_bytes().map_err(|_| aead::Error)?; Ok((transport, encoded)) } - - pub fn to_bytes<'b>(self, send: &mut SendState<'b, Neko128>) -> aead::Result> { - let mut payload = postcard::to_allocvec(&self).map_err(|_| aead::Error)?; - - send.encrypt(&mut payload, b"water-pod")?; - - Ok(cobs::encode_vec_including_sentinels(&payload)) - } - - pub fn from_bytes<'a>( - mut buf: Vec, - recv: &mut RecvState<'a, Neko128>, - ) -> aead::Result { - recv.decrypt(&mut buf, b"water-pod")?; - postcard::from_bytes(&buf).map_err(|_| aead::Error) - } } #[derive(Debug, Clone, Copy, serde::Serialize, serde::Deserialize, postcard_schema::Schema)] @@ -149,3 +200,43 @@ pub struct Command { pub time: i64, pub duration: u8, } + +pub struct Accumulator { + cobs: Box>, +} + +impl Accumulator { + pub fn new() -> Self { + let cobs = Box::>::new_zeroed(); + + Self { + // SAFETY: Zeroed CobsAccumulator is fine. + cobs: unsafe { cobs.assume_init() }, + } + } + + pub fn accumulate(&mut self, mut input: &[u8]) -> AccumulatorResult + where + T: for<'de> Deserialize<'de>, + { + if input.is_empty() { + return AccumulatorResult::Error; + } + + loop { + input = match self.cobs.feed::(input) { + FeedResult::Consumed => return AccumulatorResult::Continue, + FeedResult::OverFull(items) => items, + FeedResult::DeserError(items) => items, + FeedResult::Success { data, .. } => return AccumulatorResult::Success(data), + }; + } + } +} + +#[derive(Debug, PartialEq, Eq)] +pub enum AccumulatorResult { + Continue, + Error, + Success(T), +}