use serde::{Deserialize, Serialize};
use crate::constants::protocol_header_value;
pub const FRAME_HELLO: u8 = 0x01;
pub const FRAME_HELLO_OK: u8 = 0x02;
pub const FRAME_LIST_REQ: u8 = 0x10;
pub const FRAME_LIST_REPLY: u8 = 0x11;
pub const FRAME_META_REQ: u8 = 0x12;
pub const FRAME_META_REPLY: u8 = 0x13;
pub const FRAME_START: u8 = 0x20;
pub const FRAME_READY: u8 = 0x21;
pub const FRAME_BLOCK: u8 = 0x30;
pub const FRAME_NAK: u8 = 0x31;
pub const FRAME_REQ: u8 = 0x32;
pub const FRAME_WAVE_DONE: u8 = 0x33;
pub const FRAME_COMPLETE: u8 = 0x34;
pub const FRAME_ERROR: u8 = 0xFF;
pub fn frame_type(frame: &[u8]) -> Option<u8> {
frame.first().copied()
}
pub fn frame_payload(frame: &[u8]) -> &[u8] {
frame.get(1..).unwrap_or(&[])
}
pub fn control_frame<T: Serialize>(kind: u8, payload: &T) -> Vec<u8> {
let json = serde_json::to_vec(payload).unwrap_or_default();
let mut out = Vec::with_capacity(1 + json.len());
out.push(kind);
out.extend_from_slice(&json);
out
}
pub fn parse_control<'a, T: Deserialize<'a>>(frame: &'a [u8], kind: u8) -> Option<T> {
if frame.first() != Some(&kind) {
return None;
}
serde_json::from_slice(frame.get(1..)?).ok()
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Hello {
pub protocol: String,
pub token: String,
}
impl Hello {
pub fn new(token: &str) -> Self {
Hello {
protocol: protocol_header_value().to_string(),
token: token.to_string(),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum TransferKind {
Upload,
Download,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct StartRequest {
pub kind: TransferKind,
pub path: String,
#[serde(default)]
pub size: u64,
#[serde(default)]
pub mtime: u64,
#[serde(default)]
pub etag: String,
#[serde(default)]
pub compress: bool,
#[serde(default)]
pub mode: String,
#[serde(default)]
pub offset: u64,
#[serde(default)]
pub block_size: u64,
#[serde(default)]
pub window: u32,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ReadyReply {
pub kind: TransferKind,
pub path: String,
#[serde(default)]
pub size: u64,
#[serde(default)]
pub mtime: u64,
#[serde(default)]
pub etag: String,
#[serde(default)]
pub compress: bool,
pub block_size: u64,
pub total_blocks: u32,
#[serde(default)]
pub offset: u64,
#[serde(default)]
pub received: Vec<[u64; 2]>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CompleteMessage {
pub ok: bool,
#[serde(default)]
pub size: u64,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub error: Option<String>,
}
impl CompleteMessage {
pub fn ok(size: u64) -> Self {
CompleteMessage {
ok: true,
size,
error: None,
}
}
pub fn err(message: impl Into<String>) -> Self {
CompleteMessage {
ok: false,
size: 0,
error: Some(message.into()),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ErrorMessage {
pub code: String,
pub message: String,
}
#[derive(Debug, Clone)]
pub struct Block {
pub index: u32,
pub crc: u32,
pub raw_len: u32,
pub data: Vec<u8>,
}
pub fn block_frame(index: u32, crc: u32, raw_len: u32, data: &[u8]) -> Vec<u8> {
let mut out = Vec::with_capacity(1 + 12 + data.len());
out.push(FRAME_BLOCK);
out.extend_from_slice(&index.to_be_bytes());
out.extend_from_slice(&crc.to_be_bytes());
out.extend_from_slice(&raw_len.to_be_bytes());
out.extend_from_slice(data);
out
}
pub fn parse_block(frame: &[u8]) -> Option<Block> {
if frame.first() != Some(&FRAME_BLOCK) || frame.len() < 13 {
return None;
}
Some(Block {
index: u32::from_be_bytes(frame[1..5].try_into().ok()?),
crc: u32::from_be_bytes(frame[5..9].try_into().ok()?),
raw_len: u32::from_be_bytes(frame[9..13].try_into().ok()?),
data: frame[13..].to_vec(),
})
}
pub fn nak_frame(index: u32) -> Vec<u8> {
let mut out = vec![FRAME_NAK];
out.extend_from_slice(&index.to_be_bytes());
out
}
pub fn req_frame(indices: &[u32]) -> Vec<u8> {
let mut out = Vec::with_capacity(5 + indices.len() * 4);
out.push(FRAME_REQ);
out.extend_from_slice(&(indices.len() as u32).to_be_bytes());
for i in indices {
out.extend_from_slice(&i.to_be_bytes());
}
out
}
pub fn parse_req(frame: &[u8]) -> Option<Vec<u32>> {
if frame.first() != Some(&FRAME_REQ) || frame.len() < 5 {
return None;
}
let count = u32::from_be_bytes(frame[1..5].try_into().ok()?) as usize;
let mut out = Vec::with_capacity(count);
let mut off = 5usize;
for _ in 0..count {
if off + 4 > frame.len() {
return None;
}
out.push(u32::from_be_bytes(frame[off..off + 4].try_into().ok()?));
off += 4;
}
Some(out)
}
pub fn parse_nak(frame: &[u8]) -> Option<u32> {
if frame.first() != Some(&FRAME_NAK) || frame.len() < 5 {
return None;
}
Some(u32::from_be_bytes(frame[1..5].try_into().ok()?))
}
pub fn wave_done_frame() -> Vec<u8> {
vec![FRAME_WAVE_DONE]
}
pub fn crc32(data: &[u8]) -> u32 {
crc32fast::hash(data)
}
pub fn block_count(size: u64, block_size: u64) -> u32 {
let bs = block_size.max(1);
if size == 0 {
1
} else {
(size.div_ceil(bs)) as u32
}
}
pub fn block_bounds(index: u32, block_size: u64, size: u64) -> (u64, u64) {
let start = index as u64 * block_size;
let end = (start + block_size).min(size);
(start, end)
}
pub fn block_offset(index: u32, block_size: u64, offset: u64) -> u64 {
offset.saturating_add(index as u64 * block_size)
}
#[derive(Debug, Clone, Default)]
pub struct BlockSet {
bits: Vec<u64>,
total: u32,
count: u32,
}
impl BlockSet {
pub fn new(total: u32) -> Self {
BlockSet {
bits: vec![0; (total as usize).div_ceil(64)],
total,
count: 0,
}
}
pub fn insert(&mut self, index: u32) {
if index >= self.total {
return;
}
let word = (index / 64) as usize;
let bit = 1u64 << (index % 64);
if self.bits[word] & bit == 0 {
self.bits[word] |= bit;
self.count += 1;
}
}
pub fn contains(&self, index: u32) -> bool {
index < self.total && (self.bits[(index / 64) as usize] & (1u64 << (index % 64))) != 0
}
pub fn count(&self) -> u32 {
self.count
}
pub fn total(&self) -> u32 {
self.total
}
pub fn missing(&self) -> Vec<u32> {
let mut out = Vec::new();
for i in 0..self.total {
if !self.contains(i) {
out.push(i);
}
}
out
}
pub fn seed_from_ranges(&mut self, block_size: u64, ranges: &[(u64, u64)]) {
let bs = block_size.max(1);
for &(start, end) in ranges {
if end <= start {
continue;
}
let first = (start / bs) as u32;
let last = ((end - 1) / bs) as u32; for i in first..=last {
self.insert(i);
}
}
}
}
pub fn missing_blocks(size: u64, block_size: u64, received: &[(u64, u64)]) -> Vec<u32> {
let total = block_count(size, block_size);
let mut set = BlockSet::new(total);
set.seed_from_ranges(block_size, received);
set.missing()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn block_count_and_bounds() {
assert_eq!(block_count(0, 4), 1);
assert_eq!(block_count(10, 4), 3);
assert_eq!(block_count(8, 4), 2);
assert_eq!(block_bounds(1, 4, 10), (4, 8));
assert_eq!(block_bounds(2, 4, 10), (8, 10));
assert_eq!(block_offset(2, 4, 100), 108);
}
#[test]
fn crc_roundtrip_detects_corruption() {
let data = b"hello libfw block";
let crc = crc32(data);
let mut bad = data.to_vec();
bad[0] ^= 0xFF;
assert_ne!(crc, crc32(&bad));
assert_eq!(crc, crc32(data));
}
#[test]
fn block_frame_roundtrip() {
let data = vec![7u8; 100];
let frame = block_frame(42, 12345, 100, &data);
let parsed = parse_block(&frame).unwrap();
assert_eq!(parsed.index, 42);
assert_eq!(parsed.crc, 12345);
assert_eq!(parsed.raw_len, 100);
assert_eq!(parsed.data, data);
assert_eq!(frame_type(&frame), Some(FRAME_BLOCK));
}
#[test]
fn req_and_nak_roundtrip() {
let req = req_frame(&[0, 3, 7, 9]);
assert_eq!(parse_req(&req), Some(vec![0, 3, 7, 9]));
let nak = nak_frame(5);
assert_eq!(parse_nak(&nak), Some(5));
}
#[test]
fn block_set_tracks_verified_and_missing() {
let mut set = BlockSet::new(10);
assert_eq!(set.total(), 10);
set.insert(0);
set.insert(3);
set.insert(3); assert_eq!(set.count(), 2);
assert!(set.contains(0));
assert!(!set.contains(1));
assert_eq!(set.missing(), vec![1, 2, 4, 5, 6, 7, 8, 9]);
}
#[test]
fn seed_from_ranges_marks_overlapping_blocks() {
let mut set = BlockSet::new(3);
set.seed_from_ranges(4, &[(0, 4)]);
assert!(set.contains(0));
assert!(!set.contains(1));
let mut set = BlockSet::new(3);
set.seed_from_ranges(4, &[(4, 10)]);
assert!(set.contains(1));
assert!(set.contains(2));
assert!(!set.contains(0));
}
#[test]
fn missing_blocks_from_received_ranges() {
assert_eq!(missing_blocks(10, 4, &[(0, 4)]), vec![1, 2]);
assert_eq!(missing_blocks(10, 4, &[]), vec![0, 1, 2]);
assert_eq!(missing_blocks(10, 4, &[(0, 10)]), Vec::<u32>::new());
}
#[test]
fn control_frame_roundtrip() {
let hello = Hello::new("tok");
let frame = control_frame(FRAME_HELLO, &hello);
assert_eq!(frame_type(&frame), Some(FRAME_HELLO));
let parsed: Hello = parse_control(&frame, FRAME_HELLO).unwrap();
assert_eq!(parsed.token, "tok");
assert_eq!(parsed.protocol, protocol_header_value());
}
}