use bytes::Bytes;
use crate::codec::write_buffer::WriteBuffer;
use crate::solicit::error_code::ErrorCodeOrUnknown;
use crate::solicit::frame::flags::*;
use crate::solicit::frame::parse_stream_id;
use crate::solicit::frame::Frame;
use crate::solicit::frame::FrameBuilder;
use crate::solicit::frame::FrameHeader;
use crate::solicit::frame::FrameIR;
use crate::solicit::frame::ParseFrameError;
use crate::solicit::frame::ParseFrameResult;
use crate::solicit::frame::RawFrame;
use crate::solicit::stream_id::StreamId;
use crate::ErrorCode;
pub const GOAWAY_MIN_FRAME_LEN: u32 = 8;
pub const GOAWAY_FRAME_TYPE: u8 = 0x7;
#[derive(Clone, Debug, PartialEq)]
pub struct GoawayFrame {
pub last_stream_id: StreamId,
pub(crate) error_code: ErrorCodeOrUnknown,
pub debug_data: Bytes,
flags: Flags<NoFlag>,
}
impl GoawayFrame {
pub fn new(last_stream_id: StreamId, error_code: ErrorCode) -> Self {
GoawayFrame::with_debug_data(last_stream_id, error_code, Bytes::new())
}
pub fn with_debug_data(
last_stream_id: StreamId,
error_code: ErrorCode,
debug_data: Bytes,
) -> Self {
GoawayFrame {
last_stream_id: last_stream_id,
error_code: error_code.into(),
debug_data: debug_data,
flags: Flags::default(),
}
}
pub fn error_code(&self) -> ErrorCode {
self.error_code.into()
}
pub fn raw_error_code(&self) -> u32 {
self.error_code.0
}
pub fn last_stream_id(&self) -> StreamId {
self.last_stream_id
}
pub fn debug_data(&self) -> &Bytes {
&self.debug_data
}
pub fn payload_len(&self) -> u32 {
GOAWAY_MIN_FRAME_LEN + self.debug_data.len() as u32
}
}
impl Frame for GoawayFrame {
type FlagType = NoFlag;
fn from_raw(raw_frame: &RawFrame) -> ParseFrameResult<Self> {
let FrameHeader {
payload_len,
frame_type,
flags,
stream_id,
} = raw_frame.header();
if payload_len < GOAWAY_MIN_FRAME_LEN {
return Err(ParseFrameError::IncorrectPayloadLen);
}
if frame_type != GOAWAY_FRAME_TYPE {
return Err(ParseFrameError::InternalError);
}
if stream_id != 0x0 {
return Err(ParseFrameError::StreamIdMustBeNonZero);
}
let last_stream_id = parse_stream_id(&raw_frame.payload());
let error = unpack_octets_4!(raw_frame.payload(), 4, u32);
let debug_data = raw_frame.payload().slice(GOAWAY_MIN_FRAME_LEN as usize..);
Ok(GoawayFrame {
last_stream_id,
error_code: ErrorCodeOrUnknown(error),
debug_data,
flags: Flags::new(flags),
})
}
fn flags(&self) -> Flags<NoFlag> {
self.flags
}
fn get_stream_id(&self) -> StreamId {
0
}
fn get_header(&self) -> FrameHeader {
FrameHeader {
payload_len: self.payload_len(),
frame_type: GOAWAY_FRAME_TYPE,
flags: self.flags.0,
stream_id: 0,
}
}
}
impl FrameIR for GoawayFrame {
fn serialize_into(self, builder: &mut WriteBuffer) {
builder.write_header(self.get_header());
builder.write_u32(self.last_stream_id);
builder.write_u32(self.error_code.0);
builder.extend_from_bytes(self.debug_data);
}
}
#[cfg(test)]
mod tests {
use super::GoawayFrame;
use crate::solicit::frame::Frame;
use crate::solicit::frame::FrameHeader;
use crate::solicit::frame::FrameIR;
use crate::solicit::tests::common::raw_frame_from_parts;
use crate::ErrorCode;
use bytes::Bytes;
#[test]
fn test_parse_valid_no_debug_data() {
let raw =
raw_frame_from_parts(FrameHeader::new(8, 0x7, 0, 0), vec![0, 0, 0, 0, 0, 0, 0, 1]);
let frame = GoawayFrame::from_raw(&raw).expect("Expected successful parse");
assert_eq!(frame.error_code(), ErrorCode::ProtocolError);
assert_eq!(frame.last_stream_id(), 0);
assert_eq!(frame.debug_data(), &Bytes::new());
}
#[test]
fn test_parse_valid_no_debug_data_2() {
let raw =
raw_frame_from_parts(FrameHeader::new(8, 0x7, 0, 0), vec![0, 0, 1, 0, 0, 0, 0, 1]);
let frame = GoawayFrame::from_raw(&raw).expect("Expected successful parse");
assert_eq!(frame.error_code(), ErrorCode::ProtocolError);
assert_eq!(frame.last_stream_id(), 0x00000100);
assert_eq!(frame.debug_data(), &Bytes::new());
}
#[test]
fn test_parse_valid_with_debug_data() {
let raw = raw_frame_from_parts(
FrameHeader::new(12, 0x7, 0, 0),
vec![0, 0, 0, 0, 0, 0, 0, 1, 1, 2, 3, 4],
);
let frame = GoawayFrame::from_raw(&raw).expect("Expected successful parse");
assert_eq!(frame.error_code(), ErrorCode::ProtocolError);
assert_eq!(frame.last_stream_id(), 0);
assert_eq!(frame.debug_data(), &Bytes::from(&[1, 2, 3, 4][..]));
}
#[test]
fn test_parse_ignores_reserved_bit() {
let raw = raw_frame_from_parts(
FrameHeader::new(8, 0x7, 0, 0),
vec![0x80, 0, 0, 0, 0, 0, 0, 1],
);
let frame = GoawayFrame::from_raw(&raw).expect("Expected successful parse");
assert_eq!(frame.error_code(), ErrorCode::ProtocolError);
assert_eq!(frame.last_stream_id(), 0);
assert_eq!(frame.debug_data(), &Bytes::new());
}
#[test]
fn test_parse_invalid_id() {
let raw = raw_frame_from_parts(
FrameHeader::new(12, 0x1, 0, 0),
vec![0, 0, 0, 0, 0, 0, 0, 1, 1, 2, 3, 4],
);
assert!(GoawayFrame::from_raw(&raw).is_err(), "expected invalid id");
}
#[test]
fn test_parse_invalid_stream_id() {
let raw =
raw_frame_from_parts(FrameHeader::new(8, 0x7, 0, 3), vec![0, 0, 0, 0, 0, 0, 0, 1]);
assert!(
GoawayFrame::from_raw(&raw).is_err(),
"expected invalid stream id"
);
}
#[test]
fn test_parse_invalid_length() {
let raw = raw_frame_from_parts(FrameHeader::new(7, 0x1, 0, 0), vec![0, 0, 0, 0, 0, 0, 1]);
assert!(GoawayFrame::from_raw(&raw).is_err(), "expected too short");
}
#[test]
fn test_serialize_no_debug_data() {
let frame = GoawayFrame::new(0, ErrorCode::ProtocolError);
let expected: Vec<u8> =
raw_frame_from_parts(FrameHeader::new(8, 0x7, 0, 0), vec![0, 0, 0, 0, 0, 0, 0, 1])
.as_ref()
.to_owned();
let raw = frame.serialize_into_vec();
assert_eq!(expected, raw);
}
#[test]
fn test_serialize_with_debug_data() {
let frame = GoawayFrame::with_debug_data(
0,
ErrorCode::ProtocolError.into(),
Bytes::from_static(b"Hi!"),
);
let expected: Vec<u8> = raw_frame_from_parts(
FrameHeader::new(11, 0x7, 0, 0),
vec![0, 0, 0, 0, 0, 0, 0, 1, b'H', b'i', b'!'],
)
.as_ref()
.to_owned();
let raw = frame.serialize_into_vec();
assert_eq!(expected, raw);
}
}