use std::time::Duration;
use efema_proto::wire::{Batch, ProblemCode, Watch};
use efema_proto::{Cursor, Epoch, StreamId, StreamName, chain};
use lacodda_seal::{Key, LockedKey};
use crate::envelope::{self, ENTRY_OVERHEAD, KEY_CONTEXT, Opened};
use crate::state::State;
use crate::transport::{Transport, TransportError};
use crate::{DeviceId, Error};
pub enum Secret<'a> {
Passphrase(&'a [u8]),
Key(Key),
}
pub struct Client<T> {
transport: T,
state: State,
stream: StreamName,
id: StreamId,
epoch: Epoch,
key: Key,
lock: LockedKey,
}
impl<T: Transport> std::fmt::Debug for Client<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Client")
.field("transport", &self.transport.describe())
.field("stream", &self.stream)
.field("stream_id", &self.id)
.field("epoch", &self.epoch)
.field("key", &self.key)
.field("device", &self.state.device())
.finish()
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct Received {
pub seq: u64,
pub epoch: Epoch,
pub device: DeviceId,
pub own: bool,
pub data: Vec<u8>,
}
#[derive(Clone, Debug)]
pub struct Pulled {
pub entries: Vec<Received>,
through: Option<Cursor>,
head: u64,
}
impl Pulled {
pub fn more(&self) -> bool {
self.through.is_some_and(|cursor| cursor.seq < self.head)
}
pub fn cursor(&self) -> Option<&Cursor> {
self.through.as_ref()
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct Pushed {
pub positions: Vec<u64>,
}
impl<T: Transport> Client<T> {
pub async fn open(
transport: T,
state: &State,
stream: StreamName,
epoch: Epoch,
secret: Secret<'_>,
) -> Result<Self, Error> {
let known = blocking({
let state = state.clone();
let stream = stream.clone();
move || state.known(&stream)
})
.await?;
let (id, lock, key) = match known {
Some(known) => {
let lock = LockedKey::from_bytes(&known.lock)?;
let key = take_key(&lock, secret, &stream).await?;
(known.stream, lock, key)
}
None => {
let (id, lock, key) = establish(&transport, &stream, epoch, secret).await?;
blocking({
let state = state.clone();
let stream = stream.clone();
let lock = lock.to_bytes();
move || state.join(&stream, id, &lock)
})
.await?;
(id, lock, key)
}
};
Ok(Self { transport, state: state.clone(), stream, id, epoch, key, lock })
}
pub async fn push<I>(&self, items: I) -> Result<Pushed, Error>
where
I: IntoIterator,
I::Item: AsRef<[u8]>,
{
let limits = self.transport.limits().await?;
let sizer = Sizer::new(self.epoch, limits.max_body);
let mut batches: Vec<Vec<Vec<u8>>> = Vec::new();
let mut body = usize::MAX;
for item in items {
let item = item.as_ref();
if item.len() > sizer.max_item() {
return Err(Error::ItemTooLarge { size: item.len(), limit: sizer.max_item() });
}
let entry = envelope::seal(&self.key, self.epoch, self.state.device(), item)?;
match batches.last_mut() {
Some(batch) if sizer.fits(body, batch.len(), entry.len()) => {
body = sizer.grow(body, batch.len(), entry.len());
batch.push(entry);
}
_ => {
body = sizer.grow(sizer.empty(), 0, entry.len());
batches.push(vec![entry]);
}
}
}
let mut positions = Vec::new();
for entries in batches {
let count = entries.len() as u64;
let batch = Batch { epoch: self.epoch, entries, stream: Some(self.id) };
let landed = match self.transport.append(&self.stream, &batch).await {
Ok(written) if written.stream != self.id => Err(Error::StreamReplaced { stream: self.stream.clone() }),
Ok(written) if written.last + 1 - written.first != count => Err(Error::Transport(not_protocol(
&self.transport,
format!("{count} entries were sent and {}..{} came back", written.first, written.last),
))),
Ok(written) => Ok(written),
Err(e) => Err(write_error(e, &self.stream, self.epoch)),
};
match landed {
Ok(written) => positions.extend(written.first..=written.last),
Err(e) if positions.is_empty() => return Err(e),
Err(e) => return Err(Error::Interrupted { written: positions.len(), source: Box::new(e) }),
}
}
Ok(Pushed { positions })
}
pub async fn pull(&self) -> Result<Pulled, Error> {
let limits = self.transport.limits().await?;
let cursor = self.cursor().await?;
let page = self
.transport
.read(&self.stream, Some(&cursor), limits.page_entries)
.await
.map_err(|e| classify_read(e, &self.stream, &cursor))?;
if page.stream != self.id {
return Err(Error::StreamReplaced { stream: self.stream.clone() });
}
let mut entries = Vec::new();
let mut previous = cursor;
let mut through: Option<Cursor> = None;
for entry in &page.entries {
let seq = entry.seq;
let stop = match self.check(&previous, entry) {
Ok(Some(received)) => {
entries.push(received);
None
}
Ok(None) => None,
Err(stop) => Some(stop),
};
if let Some(stop) = stop {
if through.is_none() {
return Err(stop);
}
break;
}
previous = Cursor { stream: page.stream, seq, hash: entry.hash };
through = Some(previous);
}
Ok(Pulled { entries, through, head: page.head })
}
fn check(&self, previous: &Cursor, entry: &efema_proto::wire::Entry) -> Result<Option<Received>, Error> {
let stream = || self.stream.clone();
if entry.seq != previous.seq + 1
|| chain::link(&previous.hash, entry.seq, entry.epoch, &entry.data) != entry.hash
{
return Err(Error::BrokenChain { stream: stream(), seq: entry.seq });
}
if entry.epoch > self.epoch {
return Err(Error::NewerEpoch { stream: stream(), seq: entry.seq, ours: self.epoch, entry: entry.epoch });
}
if lacodda_seal::is_locked_key(&entry.data) {
return Ok(None);
}
if !lacodda_seal::is_sealed(&entry.data) {
return Err(Error::ForeignEntry { stream: stream(), seq: entry.seq });
}
let newer = |version| Error::NewerEntryFormat { stream: stream(), seq: entry.seq, version };
match lacodda_seal::sealed_key_id(&entry.data) {
Ok(key) if key != self.key.id() => {
return Err(Error::UnknownKey { stream: stream(), seq: entry.seq, key });
}
Ok(_) => {}
Err(lacodda_seal::Error::NewerFormat { found, .. }) => return Err(newer(found)),
Err(e) => return Err(e.into()),
}
match envelope::open(&self.key, entry.epoch, &entry.data) {
Ok(Opened::Entry { device, data }) => Ok(Some(Received {
seq: entry.seq,
epoch: entry.epoch,
device,
own: device == self.state.device(),
data,
})),
Ok(Opened::Newer(version)) => Err(newer(version)),
Err(lacodda_seal::Error::Inauthentic) => Err(Error::Inauthentic { stream: stream(), seq: entry.seq }),
Err(e) => Err(e.into()),
}
}
pub async fn ack(&self, pulled: &Pulled) -> Result<(), Error> {
let Some(cursor) = pulled.through else {
return Ok(());
};
let state = self.state.clone();
let stream = self.stream.clone();
blocking(move || state.advance(&stream, &cursor)).await
}
pub async fn wait(&self, timeout: Duration) -> Result<bool, Error> {
let limits = self.transport.limits().await?;
let cursor = self.cursor().await?;
let watch = Watch { stream: self.stream.clone(), cursor: Some(cursor) };
let heads = self.transport.wait(&[watch], timeout.min(limits.wait_max)).await?;
Ok(!heads.heads.is_empty())
}
pub async fn max_item_len(&self) -> Result<usize, Error> {
let limits = self.transport.limits().await?;
Ok(Sizer::new(self.epoch, limits.max_body).max_item())
}
pub fn key(&self) -> &Key {
&self.key
}
pub fn locked_key(&self) -> &LockedKey {
&self.lock
}
pub fn stream(&self) -> &StreamName {
&self.stream
}
pub fn stream_id(&self) -> StreamId {
self.id
}
pub fn epoch(&self) -> Epoch {
self.epoch
}
pub fn device(&self) -> DeviceId {
self.state.device()
}
pub fn transport(&self) -> &T {
&self.transport
}
pub async fn cursor(&self) -> Result<Cursor, Error> {
let state = self.state.clone();
let stream = self.stream.clone();
let known = blocking(move || state.known(&stream)).await?;
known.map(|known| known.cursor).ok_or_else(|| Error::StreamGone { stream: self.stream.clone() })
}
}
async fn establish<T: Transport>(
transport: &T,
stream: &StreamName,
epoch: Epoch,
secret: Secret<'_>,
) -> Result<(StreamId, LockedKey, Key), Error> {
if let Some((id, lock)) = first_lock(transport, stream).await? {
let key = take_key(&lock, secret, stream).await?;
return Ok((id, lock, key));
}
let Secret::Passphrase(passphrase) = secret else {
return Err(Error::NoStreamKey { stream: stream.clone() });
};
let key = Key::generate()?;
let lock = blocking({
let key = key.clone();
let passphrase = passphrase.to_vec();
move || {
let lock = LockedKey::lock(&key, &passphrase, KEY_CONTEXT);
wipe(passphrase);
lock
}
})
.await?;
let batch = Batch { epoch, entries: vec![lock.to_bytes()], stream: None };
transport.append(stream, &batch).await.map_err(|e| write_error(e, stream, epoch))?;
let (id, first) = first_lock(transport, stream)
.await?
.ok_or_else(|| not_protocol(transport, "a key was written to the stream and is not there".into()))?;
if first.key_id() == key.id() {
return Ok((id, first, key));
}
let theirs = unlock(&first, passphrase, stream).await?;
Ok((id, first, theirs))
}
async fn first_lock<T: Transport>(transport: &T, stream: &StreamName) -> Result<Option<(StreamId, LockedKey)>, Error> {
let limits = transport.limits().await?;
let mut after: Option<Cursor> = None;
let mut limit = 1;
loop {
let page = match transport.read(stream, after.as_ref(), limit).await {
Ok(page) => page,
Err(e) if e.code() == Some(&ProblemCode::StreamNotFound) => return Ok(None),
Err(e) => return Err(e.into()),
};
for entry in &page.entries {
if lacodda_seal::is_locked_key(&entry.data) {
return Ok(Some((page.stream, LockedKey::from_bytes(&entry.data)?)));
}
if !lacodda_seal::is_sealed(&entry.data) {
return Err(Error::ForeignEntry { stream: stream.clone(), seq: entry.seq });
}
}
match page.cursor() {
Some(cursor) if cursor.seq < page.head => after = Some(cursor),
_ => return Ok(None),
}
limit = limits.page_entries;
}
}
async fn take_key(lock: &LockedKey, secret: Secret<'_>, stream: &StreamName) -> Result<Key, Error> {
match secret {
Secret::Passphrase(passphrase) => unlock(lock, passphrase, stream).await,
Secret::Key(key) if key.id() == lock.key_id() => Ok(key),
Secret::Key(key) => {
Err(Error::KeyMismatch { stream: stream.clone(), ours: key.id(), stream_key: lock.key_id() })
}
}
}
async fn unlock(lock: &LockedKey, passphrase: &[u8], stream: &StreamName) -> Result<Key, Error> {
let lock = lock.clone();
let passphrase = passphrase.to_vec();
let unlocked = blocking(move || {
let key = lock.unlock(&passphrase, KEY_CONTEXT);
wipe(passphrase);
key
})
.await;
unlocked.map_err(|e| match e {
lacodda_seal::Error::WrongPassphrase => Error::WrongPassphrase { stream: stream.clone() },
other => other.into(),
})
}
fn wipe(mut bytes: Vec<u8>) {
bytes.fill(0);
std::hint::black_box(&bytes);
}
pub(crate) fn classify_read(e: TransportError, stream: &StreamName, cursor: &Cursor) -> Error {
let stream = stream.clone();
match e.code() {
Some(ProblemCode::StreamNotFound) => Error::StreamGone { stream },
Some(ProblemCode::StreamReplaced) => Error::StreamReplaced { stream },
Some(ProblemCode::CursorAhead) => {
let head = match &e {
TransportError::Refused { problem, .. } => problem.head.unwrap_or_default(),
_ => 0,
};
Error::CursorAhead { stream, cursor: cursor.seq, head }
}
Some(ProblemCode::CursorDiverged) => Error::CursorDiverged { stream, seq: cursor.seq },
_ => Error::Transport(e),
}
}
fn write_error(e: TransportError, stream: &StreamName, ours: Epoch) -> Error {
let stream = stream.clone();
match &e {
TransportError::Refused { problem, .. } => match problem.code {
ProblemCode::EpochBehind => {
Error::EpochBehind { stream, ours, current: problem.epoch.unwrap_or(Epoch(ours.0.saturating_add(1))) }
}
ProblemCode::StreamNotFound => Error::StreamGone { stream },
ProblemCode::StreamReplaced => Error::StreamReplaced { stream },
_ => Error::Transport(e),
},
_ => Error::Transport(e),
}
}
fn not_protocol<T: Transport>(transport: &T, reason: String) -> TransportError {
TransportError::NotProtocol { at: transport.describe(), reason }
}
async fn blocking<R: Send + 'static>(work: impl FnOnce() -> R + Send + 'static) -> R {
match tokio::task::spawn_blocking(work).await {
Ok(result) => result,
Err(e) => std::panic::resume_unwind(e.into_panic()),
}
}
struct Sizer {
max_body: usize,
fixed: usize,
}
impl Sizer {
fn new(epoch: Epoch, max_body: usize) -> Self {
let fixed = 1 + 1 + head_len(u64::from(epoch.0)) + 1 + 1 + head_len(16) + 16;
Self { max_body, fixed }
}
fn empty(&self) -> usize {
self.fixed + head_len(0)
}
fn grow(&self, body: usize, count: usize, len: usize) -> usize {
body - head_len(count as u64) + head_len(count as u64 + 1) + head_len(len as u64) + len
}
fn fits(&self, body: usize, count: usize, len: usize) -> bool {
self.grow(body, count, len) <= self.max_body
}
fn max_item(&self) -> usize {
let room = self.max_body.saturating_sub(self.fixed + head_len(1));
[9, 5, 3, 2, 1]
.into_iter()
.filter_map(|head| room.checked_sub(head).filter(|entry| head_len(*entry as u64) == head))
.max()
.map_or(0, |entry| entry.saturating_sub(ENTRY_OVERHEAD))
}
}
fn head_len(value: u64) -> usize {
match value {
0..=23 => 1,
24..=0xff => 2,
0x100..=0xffff => 3,
0x1_0000..=0xffff_ffff => 5,
_ => 9,
}
}
#[cfg(test)]
mod tests {
use super::*;
use efema_proto::wire;
#[test]
fn the_sizer_measures_a_batch_as_the_encoder_writes_it() {
let stream = Some(StreamId::from_bytes([1; 16]));
for epoch in [Epoch(0), Epoch(24), Epoch(70_000)] {
let sizer = Sizer::new(epoch, usize::MAX);
let mut body = sizer.empty();
let mut entries: Vec<Vec<u8>> = Vec::new();
for len in [0, 1, 23, 24, 255, 256, 65_535, 65_536, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3] {
body = sizer.grow(body, entries.len(), len);
entries.push(vec![0; len]);
let encoded = wire::encode(&Batch { epoch, entries: entries.clone(), stream });
assert_eq!(body, encoded.len(), "{} entries in epoch {epoch}", entries.len());
}
}
}
#[test]
fn the_largest_item_fits_and_one_byte_more_does_not() {
for max_body in [200, 1000, 70_000, 16 * 1024 * 1024] {
let sizer = Sizer::new(Epoch(1), max_body);
let item = sizer.max_item();
let fits = |len: usize| sizer.grow(sizer.empty(), 0, len + ENTRY_OVERHEAD) <= max_body;
assert!(fits(item), "the largest item does not fit {max_body}");
assert!(!fits(item + 1), "one byte more still fits {max_body}");
}
assert_eq!(Sizer::new(Epoch(1), 10).max_item(), 0);
}
}