use std::borrow::Cow;
use mkit_core::protocol::PackKey;
use crate::error::{Code, ServerError};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct UploadLimits {
pub max_total_bytes: u64,
pub max_chunks: u32,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Progress {
pub complete: bool,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct UploadDone {
pub key: PackKey,
pub total: u64,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum UploadError {
HeaderMissing {
stream_empty: bool,
},
UnexpectedMessage {
header: bool,
},
BadPackId {
chunk: bool,
len: Option<usize>,
},
PackIdMismatch,
TotalMissing,
TotalTooLarge {
total: u64,
cap: u64,
},
TooManyChunks,
OffsetMissing,
OffsetGap {
offset: u64,
expected: u64,
},
ByteCountOverflow,
Overrun,
AfterLast,
NoLast,
LengthMismatch {
received: u64,
declared: u64,
},
DigestMismatch,
}
impl UploadError {
#[must_use]
pub const fn code(self) -> Code {
match self {
Self::TotalTooLarge { .. } => Code::ResourceExhausted,
_ => Code::InvalidArgument,
}
}
#[must_use]
pub const fn ssh_message(self) -> &'static str {
match self {
Self::HeaderMissing { .. } => "PackChunk arrived without UploadPack header",
Self::UnexpectedMessage { .. } => "expected PackChunk after UploadPack",
Self::BadPackId { len: None, .. } => "pack_id missing",
Self::BadPackId { len: Some(_), .. } => "pack_id must be 32 bytes",
Self::PackIdMismatch => "PackChunk.pack_id does not match UploadPack",
Self::TotalMissing => "UploadPack.total_bytes is required",
Self::TotalTooLarge { .. } => "UploadPack.total_bytes exceeds server cap",
Self::TooManyChunks => "too many PackChunk frames before last=true",
Self::OffsetMissing => "PackChunk.offset is required",
Self::OffsetGap { .. } => "PackChunk.offset is not the expected next offset",
Self::ByteCountOverflow => "PackChunk byte count overflow",
Self::Overrun => "PackChunk data exceeds declared total_bytes",
Self::AfterLast => "PackChunk after last=true",
Self::NoLast => "pack chunk read failed",
Self::LengthMismatch { .. } => "PackChunk stream ended before declared total_bytes",
Self::DigestMismatch => "uploaded pack bytes do not match UploadPack.pack_id",
}
}
#[must_use]
pub fn connect_message(self) -> Cow<'static, str> {
Cow::Borrowed(match self {
Self::HeaderMissing { stream_empty: true } => "UploadPack: empty request stream",
Self::HeaderMissing {
stream_empty: false,
} => "UploadPack: first message MUST be `header`",
Self::UnexpectedMessage { header: true } => "UploadPack: saw a second `header` message",
Self::UnexpectedMessage { header: false } => {
"UploadPack: message with neither `header` nor `chunk` set"
}
Self::BadPackId { chunk: false, len } => {
return format!(
"expected a 32-byte digest, got {} bytes",
len.unwrap_or_default()
)
.into();
}
Self::BadPackId { chunk: true, .. } | Self::PackIdMismatch => {
"UploadPack: chunk.pack_id does not match header.pack_id"
}
Self::TotalMissing => "UploadPack: header.total_bytes is required",
Self::TotalTooLarge { total, cap } => {
return format!("UploadPack: total_bytes {total} exceeds the {cap}-byte cap")
.into();
}
Self::TooManyChunks => {
"UploadPack: too many `chunk` messages before `chunk.last = true`"
}
Self::OffsetMissing => "UploadPack: chunk.offset is required",
Self::OffsetGap { offset, expected } => {
return format!(
"UploadPack: chunk.offset {offset} does not match the expected offset {expected}"
)
.into();
}
Self::ByteCountOverflow | Self::Overrun => {
"UploadPack: received bytes exceed header.total_bytes"
}
Self::AfterLast => "UploadPack: message after `chunk.last = true`",
Self::NoLast => "UploadPack: stream ended without a `chunk.last = true` message",
Self::LengthMismatch { received, declared } => {
return format!(
"UploadPack: received {received} bytes, header declared {declared}"
)
.into();
}
Self::DigestMismatch => {
"UploadPack: BLAKE3(received bytes) does not equal header.pack_id"
}
})
}
}
impl From<UploadError> for ServerError {
fn from(err: UploadError) -> Self {
Self::new(err.code(), err.connect_message())
}
}
#[derive(Debug, Clone)]
pub struct UploadValidator {
key: PackKey,
declared: u64,
received: u64,
chunks: u32,
max_chunks: u32,
complete: bool,
failed: Option<UploadError>,
}
impl UploadValidator {
pub fn new(
pack_id: Option<&[u8]>,
total_bytes: Option<u64>,
limits: UploadLimits,
) -> Result<Self, UploadError> {
let key = pack_key(pack_id, false)?;
let declared = total_bytes.ok_or(UploadError::TotalMissing)?;
if declared > limits.max_total_bytes {
return Err(UploadError::TotalTooLarge {
total: declared,
cap: limits.max_total_bytes,
});
}
Ok(Self {
key,
declared,
received: 0,
chunks: 0,
max_chunks: limits.max_chunks,
complete: false,
failed: None,
})
}
pub fn push(
&mut self,
chunk_pack_id: Option<&[u8]>,
offset: Option<u64>,
data_len: usize,
last: bool,
) -> Result<Progress, UploadError> {
if let Some(err) = self.failed {
return Err(err);
}
let result = self.accept(chunk_pack_id, offset, data_len, last);
if let Err(err) = result {
self.failed = Some(err);
}
result
}
fn accept(
&mut self,
chunk_pack_id: Option<&[u8]>,
offset: Option<u64>,
data_len: usize,
last: bool,
) -> Result<Progress, UploadError> {
if self.complete {
return Err(UploadError::AfterLast);
}
self.chunks = self.chunks.saturating_add(1);
if self.chunks > self.max_chunks {
return Err(UploadError::TooManyChunks);
}
if pack_key(chunk_pack_id, true)? != self.key {
return Err(UploadError::PackIdMismatch);
}
let offset = offset.ok_or(UploadError::OffsetMissing)?;
if offset != self.received {
return Err(UploadError::OffsetGap {
offset,
expected: self.received,
});
}
let received = u64::try_from(data_len)
.ok()
.and_then(|len| self.received.checked_add(len))
.ok_or(UploadError::ByteCountOverflow)?;
if received > self.declared {
return Err(UploadError::Overrun);
}
if last && received != self.declared {
return Err(UploadError::LengthMismatch {
received,
declared: self.declared,
});
}
self.received = received;
self.complete = last;
Ok(Progress { complete: last })
}
pub fn finish(self) -> Result<UploadDone, UploadError> {
if let Some(err) = self.failed {
return Err(err);
}
if !self.complete {
return Err(UploadError::NoLast);
}
Ok(UploadDone {
key: self.key,
total: self.received,
})
}
#[must_use]
pub const fn key(&self) -> PackKey {
self.key
}
#[must_use]
pub const fn declared(&self) -> u64 {
self.declared
}
#[must_use]
pub const fn received(&self) -> u64 {
self.received
}
}
fn pack_key(id: Option<&[u8]>, chunk: bool) -> Result<PackKey, UploadError> {
let id = id.ok_or(UploadError::BadPackId { chunk, len: None })?;
<[u8; 32]>::try_from(id)
.map(PackKey::new)
.map_err(|_| UploadError::BadPackId {
chunk,
len: Some(id.len()),
})
}
#[cfg(test)]
mod tests;
pub(crate) mod marker;
pub(crate) mod receipt;
pub(crate) mod ticket_auth;
pub mod token;