use std::{
collections::VecDeque,
io::{BufRead, Cursor},
};
use tokio::io::AsyncWriteExt;
use tracing::trace;
use crate::{
comms::{chunker::Chunker, message_chunk::MessageChunk, secure_channel::SecureChannel},
Message,
};
use opcua_types::{Error, SimpleBinaryEncodable, StatusCode};
use super::{
sequence_number::SequenceNumberHandle,
tcp_types::{AcknowledgeMessage, ErrorMessage},
};
#[derive(Copy, Clone, Debug)]
enum SendBufferState {
Reading(usize),
Writing,
}
#[derive(Debug)]
enum PendingPayload {
Chunk(MessageChunk),
Ack(AcknowledgeMessage),
Error(ErrorMessage),
}
pub struct SendBuffer {
buffer: Cursor<Vec<u8>>,
chunk_storage: bytes::BytesMut,
chunk_scratch: Vec<MessageChunk>,
chunks: VecDeque<PendingPayload>,
last_request_id: u32,
sequence_numbers: SequenceNumberHandle,
pub max_message_size: usize,
pub max_chunk_count: usize,
pub send_buffer_size: usize,
state: SendBufferState,
}
impl SendBuffer {
pub fn new(
buffer_size: usize,
max_message_size: usize,
max_chunk_count: usize,
sequence_numbers_legacy: bool,
) -> Self {
Self {
buffer: Cursor::new(vec![0u8; buffer_size + 1024]),
chunk_storage: bytes::BytesMut::new(),
chunk_scratch: Vec::new(),
chunks: VecDeque::with_capacity(max_chunk_count),
last_request_id: 1000,
sequence_numbers: SequenceNumberHandle::new(sequence_numbers_legacy),
max_message_size,
max_chunk_count,
send_buffer_size: buffer_size,
state: SendBufferState::Writing,
}
}
pub fn encode_next_chunk(&mut self, secure_channel: &SecureChannel) -> Result<(), StatusCode> {
if matches!(self.state, SendBufferState::Reading(_)) {
return Err(StatusCode::BadInvalidState);
}
let Some(next_chunk) = self.chunks.pop_front() else {
return Ok(());
};
let size = match next_chunk {
PendingPayload::Chunk(c) => secure_channel.apply_security(&c, self.buffer.get_mut())?,
PendingPayload::Ack(a) => {
a.encode(&mut self.buffer)?;
self.buffer.position() as usize
}
PendingPayload::Error(e) => {
e.encode(&mut self.buffer)?;
self.buffer.position() as usize
}
};
self.buffer.set_position(0);
self.state = SendBufferState::Reading(size);
Ok(())
}
pub fn set_sequence_number_legacy(&mut self, is_legacy: bool) {
self.sequence_numbers.set_is_legacy(is_legacy);
}
pub fn write_error(&mut self, error: ErrorMessage) {
self.chunks.clear();
self.chunks.push_back(PendingPayload::Error(error));
}
pub fn write_ack(&mut self, ack: AcknowledgeMessage) {
self.chunks.push_back(PendingPayload::Ack(ack));
}
pub fn write(
&mut self,
request_id: u32,
message: impl Message,
secure_channel: &SecureChannel,
) -> Result<u32, Error> {
trace!("Writing request to buffer");
self.chunk_scratch.clear();
let chunk_count = Chunker::encode_into(
self.sequence_numbers.clone(),
request_id,
self.max_message_size,
self.send_buffer_size,
secure_channel,
&message,
&mut self.chunk_storage,
&mut self.chunk_scratch,
)
.map_err(|e| e.with_context(Some(request_id), Some(message.request_handle())))?;
if self.max_chunk_count > 0 && chunk_count > self.max_chunk_count {
self.chunk_scratch.clear();
Err(Error::new(
StatusCode::BadCommunicationError,
format!(
"Cannot write message since {chunk_count} chunks exceeds {} chunk limit",
self.max_chunk_count
),
)
.with_context(Some(request_id), Some(message.request_handle())))
} else {
self.sequence_numbers.increment(chunk_count as u32);
self.chunks
.extend(self.chunk_scratch.drain(..).map(PendingPayload::Chunk));
Ok(request_id)
}
}
pub fn next_request_id(&mut self) -> u32 {
self.last_request_id += 1;
self.last_request_id
}
pub async fn read_into_async(
&mut self,
write: &mut (impl tokio::io::AsyncWrite + Unpin),
) -> Result<(), tokio::io::Error> {
let end = match self.state {
SendBufferState::Writing => {
let end = self.buffer.position() as usize;
self.state = SendBufferState::Reading(end);
self.buffer.set_position(0);
end
}
SendBufferState::Reading(end) => end,
};
let pos = self.buffer.position() as usize;
let buf = &self.buffer.get_ref()[pos..end];
let written = write.write(buf).await?;
self.buffer.consume(written);
if end == self.buffer.position() as usize {
self.state = SendBufferState::Writing;
self.buffer.set_position(0);
}
Ok(())
}
pub fn should_encode_chunks(&self) -> bool {
!self.chunks.is_empty() && !self.can_read()
}
pub fn can_read(&self) -> bool {
matches!(self.state, SendBufferState::Reading(_)) || self.buffer.position() != 0
}
pub fn revise(
&mut self,
send_buffer_size: usize,
max_message_size: usize,
max_chunk_count: usize,
) {
if self.send_buffer_size > send_buffer_size {
self.buffer.get_mut().shrink_to(send_buffer_size + 1024);
self.send_buffer_size = send_buffer_size;
}
if self.max_message_size > max_message_size && max_message_size > 0 {
self.max_message_size = max_message_size;
}
if self.max_chunk_count > max_chunk_count && max_chunk_count > 0 {
self.max_chunk_count = max_chunk_count;
}
}
}
#[cfg(test)]
mod tests {
use std::io::Cursor;
use std::sync::Arc;
use parking_lot::RwLock;
use super::SendBuffer;
use crate::comms::secure_channel::{Role, SecureChannel};
use crate::RequestMessage;
use opcua_crypto::CertificateStore;
use opcua_types::StatusCode;
use opcua_types::{
DateTime, NodeId, ReadRequest, ReadValueId, RequestHeader, TimestampsToReturn,
};
fn get_buffer_and_channel() -> (SendBuffer, SecureChannel) {
let buffer = SendBuffer::new(8196, 81960, 5, true);
let channel = SecureChannel::new(
Arc::new(RwLock::new(CertificateStore::new(std::path::Path::new(
"./pki",
)))),
Role::Client,
Default::default(),
);
(buffer, channel)
}
#[tokio::test]
async fn test_buffer_simple() {
let message = ReadRequest {
request_header: RequestHeader::new(&NodeId::null(), &DateTime::null(), 101),
max_age: 0.0,
timestamps_to_return: TimestampsToReturn::Both,
nodes_to_read: Some(vec![ReadValueId {
node_id: (1, 1).into(),
attribute_id: 1,
..Default::default()
}]),
};
let (mut buffer, channel) = get_buffer_and_channel();
let m: RequestMessage = message.into();
let request_id = buffer.write(1, m, &channel).unwrap();
assert_eq!(request_id, 1);
assert!(buffer.should_encode_chunks());
assert_eq!(buffer.chunks.len(), 1);
buffer.encode_next_chunk(&channel).unwrap();
assert!(buffer.can_read());
let mut cursor = Cursor::new(Vec::new());
buffer.read_into_async(&mut cursor).await.unwrap();
assert!(cursor.get_ref().len() > 50);
}
#[tokio::test]
async fn test_buffer_chunking() {
let message = ReadRequest {
request_header: RequestHeader::new(&NodeId::null(), &DateTime::null(), 101),
max_age: 0.0,
timestamps_to_return: TimestampsToReturn::Both,
nodes_to_read: Some(
(0..1000)
.map(|r| ReadValueId {
node_id: (1, r).into(),
attribute_id: 1,
..Default::default()
})
.collect(),
),
};
let (mut buffer, channel) = get_buffer_and_channel();
let m: RequestMessage = message.into();
let request_id = buffer.write(1, m, &channel).unwrap();
assert_eq!(request_id, 1);
assert_eq!(buffer.chunks.len(), 3);
let mut cursor = Cursor::new(Vec::new());
for _ in 0..3 {
assert!(buffer.should_encode_chunks());
buffer.encode_next_chunk(&channel).unwrap();
assert!(!buffer.should_encode_chunks());
assert!(buffer.can_read());
buffer.read_into_async(&mut cursor).await.unwrap();
}
assert!(!buffer.should_encode_chunks());
assert!(!buffer.can_read());
assert!(cursor.get_ref().len() > 8196 * 2 && cursor.get_ref().len() < 8196 * 3);
}
#[tokio::test]
async fn test_buffer_chunk_storage_is_reused() {
let (mut buffer, channel) = get_buffer_and_channel();
let mut sink = Cursor::new(Vec::new());
let mut warmed_ptr = None;
let mut warmed_capacity = None;
for i in 0..5u32 {
let message = ReadRequest {
request_header: RequestHeader::new(&NodeId::null(), &DateTime::null(), 101),
max_age: 0.0,
timestamps_to_return: TimestampsToReturn::Both,
nodes_to_read: Some(
(0..1000)
.map(|r| ReadValueId {
node_id: (1, r).into(),
attribute_id: 1,
..Default::default()
})
.collect(),
),
};
let m: RequestMessage = message.into();
buffer.write(i + 1, m, &channel).unwrap();
if i >= 1 {
let ptr = buffer.chunk_storage.as_ptr();
let capacity = buffer.chunk_storage.capacity();
if let (Some(warmed_ptr), Some(warmed_capacity)) = (warmed_ptr, warmed_capacity) {
assert_eq!(
ptr, warmed_ptr,
"chunk storage should be reclaimed, not reallocated"
);
assert_eq!(capacity, warmed_capacity, "chunk storage should not grow");
}
warmed_ptr = Some(ptr);
warmed_capacity = Some(capacity);
}
while buffer.should_encode_chunks() {
buffer.encode_next_chunk(&channel).unwrap();
while buffer.can_read() {
buffer.read_into_async(&mut sink).await.unwrap();
}
}
}
}
#[test]
fn test_buffer_too_large_message() {
let message = ReadRequest {
request_header: RequestHeader::new(&NodeId::null(), &DateTime::null(), 101),
max_age: 0.0,
timestamps_to_return: TimestampsToReturn::Both,
nodes_to_read: Some(
(0..10000)
.map(|r| ReadValueId {
node_id: (1, r).into(),
attribute_id: 1,
..Default::default()
})
.collect(),
),
};
let (mut buffer, channel) = get_buffer_and_channel();
let m: RequestMessage = message.into();
let err = buffer.write(1, m, &channel).unwrap_err();
assert_eq!(err.status(), StatusCode::BadRequestTooLarge);
}
#[test]
fn test_buffer_too_many_chunks() {
let message = ReadRequest {
request_header: RequestHeader::new(&NodeId::null(), &DateTime::null(), 101),
max_age: 0.0,
timestamps_to_return: TimestampsToReturn::Both,
nodes_to_read: Some(
(0..4000)
.map(|r| ReadValueId {
node_id: (1, r).into(),
attribute_id: 1,
..Default::default()
})
.collect(),
),
};
let (mut buffer, channel) = get_buffer_and_channel();
let m: RequestMessage = message.into();
let err = buffer.write(1, m, &channel).unwrap_err();
assert_eq!(err.status(), StatusCode::BadCommunicationError);
}
#[tokio::test]
async fn test_buffer_read_partial() {
let message = ReadRequest {
request_header: RequestHeader::new(&NodeId::null(), &DateTime::null(), 101),
max_age: 0.0,
timestamps_to_return: TimestampsToReturn::Both,
nodes_to_read: Some(
(0..1000)
.map(|r| ReadValueId {
node_id: (1, r).into(),
attribute_id: 1,
..Default::default()
})
.collect(),
),
};
let (mut buffer, channel) = get_buffer_and_channel();
let m: RequestMessage = message.into();
let request_id = buffer.write(1, m, &channel).unwrap();
assert_eq!(request_id, 1);
assert_eq!(buffer.chunks.len(), 3);
let mut buf = [0u8; 4098];
let mut cursor = Cursor::new(&mut buf as &mut [u8]);
for _ in 0..2 {
println!("Encode chunks");
assert!(buffer.should_encode_chunks());
buffer.encode_next_chunk(&channel).unwrap();
assert!(!buffer.should_encode_chunks());
assert!(buffer.can_read());
buffer.read_into_async(&mut cursor).await.unwrap();
assert!(buffer.can_read());
assert_eq!(cursor.position(), 4098);
cursor.set_position(0);
buffer.read_into_async(&mut cursor).await.unwrap();
assert!(!buffer.can_read());
assert_eq!(cursor.position(), 4098);
cursor.set_position(0);
}
assert!(buffer.should_encode_chunks());
buffer.encode_next_chunk(&channel).unwrap();
assert!(buffer.can_read());
buffer.read_into_async(&mut cursor).await.unwrap();
assert!(cursor.position() < 4098);
assert!(!buffer.should_encode_chunks());
assert!(!buffer.can_read());
}
}