use blake3::Hasher;
use std::sync::atomic::{AtomicU32, Ordering};
use std::sync::Mutex;
use crate::binary::{ChunkFlags, StreamChunk};
use crate::stream::ring_buffer::StreamRingBuffer;
use crate::DCPError;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct RetransmitRequest {
pub start_seq: u32,
pub count: u32,
}
pub struct DcpStream {
buffer: StreamRingBuffer,
hasher: Mutex<Hasher>,
write_lock: Mutex<()>,
next_seq: AtomicU32,
send_seq: AtomicU32,
stream_id: u32,
complete: std::sync::atomic::AtomicBool,
}
impl DcpStream {
pub fn new(stream_id: u32, capacity: usize) -> Self {
Self {
buffer: StreamRingBuffer::new(capacity),
hasher: Mutex::new(Hasher::new()),
write_lock: Mutex::new(()),
next_seq: AtomicU32::new(0),
send_seq: AtomicU32::new(0),
stream_id,
complete: std::sync::atomic::AtomicBool::new(false),
}
}
pub fn stream_id(&self) -> u32 {
self.stream_id
}
pub fn buffer(&self) -> &StreamRingBuffer {
&self.buffer
}
pub fn is_complete(&self) -> bool {
self.complete.load(Ordering::Acquire)
}
pub fn write_chunk(&self, data: &[u8], is_last: bool) -> Result<StreamChunk, DCPError> {
let _write_guard = self.write_lock.lock().unwrap();
if self.is_complete() {
return Err(DCPError::ValidationFailed);
}
if data.len() > u16::MAX as usize {
return Err(DCPError::OutOfBounds);
}
let total_len = StreamChunk::SIZE
.checked_add(data.len())
.ok_or(DCPError::OutOfBounds)?;
if self.buffer.available_space() < total_len {
return Err(DCPError::Backpressure);
}
let seq = self.send_seq.fetch_add(1, Ordering::AcqRel);
let is_first = seq == 0;
let flags = if is_first && is_last {
ChunkFlags::FIRST | ChunkFlags::LAST
} else if is_first {
ChunkFlags::FIRST
} else if is_last {
ChunkFlags::LAST
} else {
ChunkFlags::CONTINUE
};
let chunk = StreamChunk::new(seq, flags, data.len() as u16);
self.buffer.push(chunk.as_bytes())?;
if !data.is_empty() {
self.buffer.push(data)?;
}
{
let mut hasher = self.hasher.lock().unwrap();
hasher.update(data);
}
if is_last {
self.complete.store(true, Ordering::Release);
}
Ok(chunk)
}
pub fn read_chunk(&self) -> Result<Option<(StreamChunk, Vec<u8>)>, DCPError> {
let mut header_buf = [0u8; StreamChunk::SIZE];
let peeked = self.buffer.peek(&mut header_buf);
if peeked < StreamChunk::SIZE {
return Ok(None);
}
let chunk = StreamChunk::from_bytes(&header_buf)?;
let chunk_len = chunk.len as usize;
let total_len = StreamChunk::SIZE + chunk_len;
if self.buffer.len() < total_len {
return Ok(None);
}
let expected_seq = self.next_seq.load(Ordering::Acquire);
if chunk.sequence != expected_seq {
return Err(DCPError::ChecksumMismatch);
}
let mut full_buf = vec![0u8; total_len];
self.buffer.pop(&mut full_buf);
let payload = full_buf[StreamChunk::SIZE..].to_vec();
self.next_seq.store(expected_seq + 1, Ordering::Release);
let result_chunk = StreamChunk::new(chunk.sequence, chunk.flags, chunk.len);
Ok(Some((result_chunk, payload)))
}
pub fn checksum(&self) -> [u8; 32] {
let hasher = self.hasher.lock().unwrap();
*hasher.finalize().as_bytes()
}
pub fn verify_checksum(&self, expected: &[u8; 32]) -> bool {
&self.checksum() == expected
}
pub fn request_retransmit(&self, missing_seq: u32) -> RetransmitRequest {
let expected = self.next_seq.load(Ordering::Acquire);
RetransmitRequest {
start_seq: expected,
count: missing_seq.saturating_sub(expected) + 1,
}
}
pub fn next_expected_seq(&self) -> u32 {
self.next_seq.load(Ordering::Acquire)
}
pub fn next_send_seq(&self) -> u32 {
self.send_seq.load(Ordering::Acquire)
}
pub fn is_backpressure(&self) -> bool {
self.buffer.backpressure().is_full()
}
pub fn available_space(&self) -> usize {
self.buffer.available_space()
}
pub fn reset(&self) {
self.buffer.clear();
self.next_seq.store(0, Ordering::Release);
self.send_seq.store(0, Ordering::Release);
self.complete.store(false, Ordering::Release);
*self.hasher.lock().unwrap() = Hasher::new();
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_stream_basic() {
let stream = DcpStream::new(1, 1024);
let chunk1 = stream.write_chunk(b"hello", false).unwrap();
assert!(chunk1.is_first());
assert!(!chunk1.is_last());
let seq1 = chunk1.sequence;
assert_eq!(seq1, 0);
let chunk2 = stream.write_chunk(b"world", true).unwrap();
assert!(!chunk2.is_first());
assert!(chunk2.is_last());
let seq2 = chunk2.sequence;
assert_eq!(seq2, 1);
assert!(stream.is_complete());
}
#[test]
fn test_stream_read_write() {
let stream = DcpStream::new(1, 1024);
stream.write_chunk(b"test", false).unwrap();
stream.write_chunk(b"data", true).unwrap();
let (chunk1, data1) = stream.read_chunk().unwrap().unwrap();
let seq1 = chunk1.sequence;
assert_eq!(seq1, 0);
assert_eq!(data1, b"test");
let (chunk2, data2) = stream.read_chunk().unwrap().unwrap();
let seq2 = chunk2.sequence;
assert_eq!(seq2, 1);
assert_eq!(data2, b"data");
assert!(stream.read_chunk().unwrap().is_none());
}
#[test]
fn test_stream_checksum() {
let stream = DcpStream::new(1, 1024);
stream.write_chunk(b"hello", false).unwrap();
let checksum1 = stream.checksum();
stream.write_chunk(b"world", true).unwrap();
let checksum2 = stream.checksum();
assert_ne!(checksum1, checksum2);
assert!(stream.verify_checksum(&checksum2));
assert!(!stream.verify_checksum(&checksum1));
}
#[test]
fn test_stream_single_chunk() {
let stream = DcpStream::new(1, 1024);
let chunk = stream.write_chunk(b"single", true).unwrap();
assert!(chunk.is_first());
assert!(chunk.is_last());
let seq = chunk.sequence;
assert_eq!(seq, 0);
}
#[test]
fn test_retransmit_request() {
let stream = DcpStream::new(1, 1024);
let req = stream.request_retransmit(5);
assert_eq!(req.start_seq, 0);
assert_eq!(req.count, 6);
}
#[test]
fn test_stream_reset() {
let stream = DcpStream::new(1, 1024);
stream.write_chunk(b"data", true).unwrap();
assert!(stream.is_complete());
assert_eq!(stream.next_send_seq(), 1);
stream.reset();
assert!(!stream.is_complete());
assert_eq!(stream.next_send_seq(), 0);
assert_eq!(stream.next_expected_seq(), 0);
}
}