use virtio_accel_core::{ByteSink, ByteSource};
use virtio_accel_proto::StatusCode;
use crate::{
ChainLayoutError, ChainRegion, DecodedRequest, FrameDecodeError, FrameDecoder,
ResponseWriteError, ResponseWriter, UnrecoverableDecodeError, validate_chain_layout,
};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum UnusableFrame {
ChainLayout(ChainLayoutError),
Request(UnrecoverableDecodeError),
InsufficientResponse {
request_id: u64,
required: u64,
available: u64,
},
}
#[derive(Debug)]
pub enum FramePreflight<'a> {
Ready(DecodedRequest<'a>),
Rejected {
request_id: u64,
status: StatusCode,
used: u32,
},
Unusable(UnusableFrame),
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum FramePreflightError {
ResponseWrite(ResponseWriteError),
}
pub fn preflight_command_frame<'a>(
decoder: &FrameDecoder,
regions: &[ChainRegion],
request: &'a dyn ByteSource,
response: &mut dyn ByteSink,
) -> Result<FramePreflight<'a>, FramePreflightError> {
let layout = match validate_chain_layout(regions, decoder.limits().max_chain_descriptors()) {
Ok(layout) => layout,
Err(error) => {
return Ok(FramePreflight::Unusable(UnusableFrame::ChainLayout(error)));
}
};
if let Err(error) = layout.validate_port_lengths(request.len(), response.len()) {
return Ok(FramePreflight::Unusable(UnusableFrame::ChainLayout(error)));
}
match decoder.decode(request, response.len()) {
Ok(request) => Ok(FramePreflight::Ready(request)),
Err(FrameDecodeError::Protocol { request_id, status }) => {
let used = ResponseWriter::new(response, decoder.limits().max_response_bytes())
.write_empty(status, request_id)
.map_err(FramePreflightError::ResponseWrite)?;
Ok(FramePreflight::Rejected {
request_id,
status,
used,
})
}
Err(FrameDecodeError::Unrecoverable(error)) => {
Ok(FramePreflight::Unusable(UnusableFrame::Request(error)))
}
Err(FrameDecodeError::InsufficientResponse {
request_id,
required,
available,
}) => Ok(FramePreflight::Unusable(
UnusableFrame::InsufficientResponse {
request_id,
required,
available,
},
)),
}
}