///|
fn sctp_writer(capacity : Int) -> @codec.Writer raise SctpError {
@codec.Writer::new(capacity~) catch {
InvalidLength(length) =>
raise InvalidPacket("invalid SCTP writer length \{length}")
Truncated(needed~, remaining~) =>
raise InvalidPacket(
"unexpected SCTP writer truncation \{needed}/\{remaining}",
)
}
}
///|
fn sctp_read_u8(
reader : @codec.Reader,
context : String,
) -> Byte raise SctpError {
reader.read_u8() catch {
InvalidLength(length) =>
raise InvalidPacket("\{context}: invalid length \{length}")
Truncated(needed~, remaining~) =>
raise InvalidPacket("\{context}: need \{needed} bytes, have \{remaining}")
}
}
///|
fn sctp_read_u16(
reader : @codec.Reader,
context : String,
) -> UInt16 raise SctpError {
reader.read_u16_be() catch {
InvalidLength(length) =>
raise InvalidPacket("\{context}: invalid length \{length}")
Truncated(needed~, remaining~) =>
raise InvalidPacket("\{context}: need \{needed} bytes, have \{remaining}")
}
}
///|
fn sctp_read_u32(
reader : @codec.Reader,
context : String,
) -> UInt raise SctpError {
reader.read_u32_be() catch {
InvalidLength(length) =>
raise InvalidPacket("\{context}: invalid length \{length}")
Truncated(needed~, remaining~) =>
raise InvalidPacket("\{context}: need \{needed} bytes, have \{remaining}")
}
}
///|
fn sctp_read_bytes(
reader : @codec.Reader,
length : Int,
context : String,
) -> Bytes raise SctpError {
reader.read_bytes(length) catch {
InvalidLength(invalid_length) =>
raise InvalidPacket("\{context}: invalid length \{invalid_length}")
Truncated(needed~, remaining~) =>
raise InvalidPacket("\{context}: need \{needed} bytes, have \{remaining}")
}
}
///|
fn padding_length(length : Int) -> Int {
(4 - length % 4) % 4
}
///|
fn crc32c(data : Bytes) -> UInt {
let mut checksum = 0xffffffffU
for index = 0; index < data.length(); index = index + 1 {
let byte : Byte = if index >= 8 && index < 12 { 0 } else { data[index] }
checksum = checksum ^ byte.to_uint()
for bit = 0; bit < 8; bit = bit + 1 {
checksum = if (checksum & 1U) != 0U {
(checksum >> 1) ^ 0x82f63b78U
} else {
checksum >> 1
}
}
}
checksum ^ 0xffffffffU
}
///|
fn parameter_type_value(parameter : Parameter) -> (UInt16, Bytes) {
match parameter {
StateCookieParameter(value) => (7, value)
SupportedExtensionsParameter(value) => (0x8008, value)
ForwardTsnSupportedParameter => (0xc000, b"")
UnknownParameter(parameter_type, value) => (parameter_type, value)
}
}
///|
fn encode_parameter(parameter : Parameter) -> Bytes raise SctpError {
let (parameter_type, value) = parameter_type_value(parameter)
if value.length() > 0xfffb {
raise InvalidPacket("SCTP parameter exceeds 65535 bytes")
}
let writer = sctp_writer(4 + value.length())
writer.write_u16_be(parameter_type)
writer.write_u16_be((4 + value.length()).to_uint16())
writer.write_bytes(value)
writer.finish()
}
///|
fn decode_parameter(reader : @codec.Reader) -> Parameter raise SctpError {
let parameter_type = sctp_read_u16(reader, "SCTP parameter type")
let length = sctp_read_u16(reader, "SCTP parameter length").to_int()
if length < 4 {
raise InvalidPacket("SCTP parameter length is smaller than its header")
}
let value = sctp_read_bytes(reader, length - 4, "SCTP parameter value")
match parameter_type {
7 => StateCookieParameter(value)
0x8008 => SupportedExtensionsParameter(value)
0xc000 => {
if !value.is_empty() {
raise InvalidPacket("Forward-TSN-Supported parameter must be empty")
}
ForwardTsnSupportedParameter
}
_ => UnknownParameter(parameter_type, value)
}
}
///|
fn encode_init(init : InitChunk) -> Bytes raise SctpError {
let writer = sctp_writer(64)
writer.write_u32_be(init.initiate_tag)
writer.write_u32_be(init.advertised_receiver_window)
writer.write_u16_be(init.outbound_streams)
writer.write_u16_be(init.inbound_streams)
writer.write_u32_be(init.initial_tsn)
for index = 0; index < init.parameters.length(); index = index + 1 {
let parameter = encode_parameter(init.parameters[index])
writer.write_bytes(parameter)
if index + 1 < init.parameters.length() {
for padding = 0
padding < padding_length(parameter.length())
padding = padding + 1 {
writer.write_u8(0)
}
}
}
writer.finish()
}
///|
fn decode_init(
acknowledgement : Bool,
flags : Byte,
value : Bytes,
) -> InitChunk raise SctpError {
if flags != 0 || value.length() < 16 {
raise InvalidPacket("malformed SCTP INIT chunk")
}
let reader = @codec.Reader::new(value)
let initiate_tag = sctp_read_u32(reader, "SCTP initiate tag")
let advertised_receiver_window = sctp_read_u32(reader, "SCTP receiver window")
let outbound_streams = sctp_read_u16(reader, "SCTP outbound streams")
let inbound_streams = sctp_read_u16(reader, "SCTP inbound streams")
let initial_tsn = sctp_read_u32(reader, "SCTP initial TSN")
let parameters : Array[Parameter] = []
while reader.remaining() > 0 {
if reader.remaining() < 4 {
raise InvalidPacket("truncated SCTP INIT parameter header")
}
let before = reader.position()
parameters.push(decode_parameter(reader))
let consumed = reader.position() - before
let padding = padding_length(consumed)
if padding > reader.remaining() {
if reader.remaining() == 0 {
break
}
raise InvalidPacket("truncated SCTP INIT parameter padding")
}
let padding_bytes = sctp_read_bytes(
reader, padding, "SCTP INIT parameter padding",
)
for byte in padding_bytes {
if byte != 0 {
raise InvalidPacket("SCTP parameter padding must be zero")
}
}
}
InitChunk::new(
acknowledgement~,
initiate_tag~,
advertised_receiver_window~,
outbound_streams~,
inbound_streams~,
initial_tsn~,
parameters~,
)
}
///|
fn encode_data(data : DataChunk) -> (Byte, Bytes) raise SctpError {
if data.user_data.length() > 0xffef {
raise InvalidPacket("SCTP DATA chunk exceeds 65535 bytes")
}
let mut flags : Byte = 0
if data.ending {
flags = flags | 1
}
if data.beginning {
flags = flags | 2
}
if data.unordered {
flags = flags | 4
}
if data.immediate_sack {
flags = flags | 8
}
let writer = sctp_writer(12 + data.user_data.length())
writer.write_u32_be(data.tsn)
let StreamId(stream) = data.stream
writer.write_u16_be(stream)
writer.write_u16_be(data.stream_sequence)
writer.write_u32_be(data.protocol_id)
writer.write_bytes(data.user_data)
(flags, writer.finish())
}
///|
fn decode_data(flags : Byte, value : Bytes) -> DataChunk raise SctpError {
if (flags & 0xf0) != 0 || value.length() < 12 {
raise InvalidPacket("malformed SCTP DATA chunk")
}
let reader = @codec.Reader::new(value)
let tsn = sctp_read_u32(reader, "SCTP DATA TSN")
let stream = StreamId(sctp_read_u16(reader, "SCTP DATA stream"))
let stream_sequence = sctp_read_u16(reader, "SCTP DATA stream sequence")
let protocol_id = sctp_read_u32(reader, "SCTP DATA protocol ID")
let user_data = sctp_read_bytes(
reader,
reader.remaining(),
"SCTP DATA payload",
)
DataChunk::new(
unordered=(flags & 4) != 0,
beginning=(flags & 2) != 0,
ending=(flags & 1) != 0,
immediate_sack=(flags & 8) != 0,
tsn~,
stream~,
stream_sequence~,
protocol_id~,
user_data~,
)
}
///|
fn encode_sack(sack : SackChunk) -> Bytes raise SctpError {
if sack.gap_ack_blocks.length() > 0xffff ||
sack.duplicate_tsns.length() > 0xffff {
raise InvalidPacket("SCTP SACK contains too many entries")
}
let writer = sctp_writer(
12 + sack.gap_ack_blocks.length() * 4 + sack.duplicate_tsns.length() * 4,
)
writer.write_u32_be(sack.cumulative_tsn_ack)
writer.write_u32_be(sack.advertised_receiver_window)
writer.write_u16_be(sack.gap_ack_blocks.length().to_uint16())
writer.write_u16_be(sack.duplicate_tsns.length().to_uint16())
for gap in sack.gap_ack_blocks {
writer.write_u16_be(gap.start)
writer.write_u16_be(gap.end)
}
for duplicate in sack.duplicate_tsns {
writer.write_u32_be(duplicate)
}
writer.finish()
}
///|
fn decode_sack(flags : Byte, value : Bytes) -> SackChunk raise SctpError {
if flags != 0 || value.length() < 12 {
raise InvalidPacket("malformed SCTP SACK chunk")
}
let reader = @codec.Reader::new(value)
let cumulative_tsn_ack = sctp_read_u32(reader, "SCTP cumulative TSN")
let advertised_receiver_window = sctp_read_u32(reader, "SCTP receiver window")
let gap_count = sctp_read_u16(reader, "SCTP gap count").to_int()
let duplicate_count = sctp_read_u16(reader, "SCTP duplicate count").to_int()
if reader.remaining() != gap_count * 4 + duplicate_count * 4 {
raise InvalidPacket("SCTP SACK entry count does not match its length")
}
let gap_ack_blocks : Array[GapAckBlock] = []
for index = 0; index < gap_count; index = index + 1 {
let start = sctp_read_u16(reader, "SCTP gap start")
let end = sctp_read_u16(reader, "SCTP gap end")
if start == 0 || end < start {
raise InvalidPacket("invalid SCTP gap acknowledgement block")
}
gap_ack_blocks.push({ start, end, })
}
let duplicate_tsns : Array[UInt] = []
for index = 0; index < duplicate_count; index = index + 1 {
duplicate_tsns.push(sctp_read_u32(reader, "SCTP duplicate TSN"))
}
SackChunk::new(
cumulative_tsn_ack~,
advertised_receiver_window~,
gap_ack_blocks~,
duplicate_tsns~,
)
}
///|
fn encode_reconfig_parameter(
parameter : ReconfigParameter,
) -> Bytes raise SctpError {
let (parameter_type, value) : (UInt16, Bytes) = match parameter {
OutgoingResetRequest(
request_sequence~,
response_sequence~,
sender_last_tsn~,
streams~
) => {
let writer = sctp_writer(12 + streams.length() * 2)
writer.write_u32_be(request_sequence)
writer.write_u32_be(response_sequence)
writer.write_u32_be(sender_last_tsn)
for stream in streams {
let StreamId(value) = stream
writer.write_u16_be(value)
}
(13, writer.finish())
}
ReconfigResponse(request_sequence~, result~) => {
let writer = sctp_writer(8)
writer.write_u32_be(request_sequence)
writer.write_u32_be(result)
(16, writer.finish())
}
UnknownReconfigParameter(parameter_type, value) => (parameter_type, value)
}
if value.length() > 0xfffb {
raise InvalidPacket("SCTP reconfiguration parameter is too large")
}
let writer = sctp_writer(4 + value.length())
writer.write_u16_be(parameter_type)
writer.write_u16_be((4 + value.length()).to_uint16())
writer.write_bytes(value)
writer.finish()
}
///|
fn decode_reconfig_parameter(
reader : @codec.Reader,
) -> ReconfigParameter raise SctpError {
let parameter_type = sctp_read_u16(
reader, "SCTP reconfiguration parameter type",
)
let length = sctp_read_u16(reader, "SCTP reconfiguration parameter length").to_int()
if length < 4 {
raise InvalidPacket("invalid SCTP reconfiguration parameter length")
}
let value = sctp_read_bytes(
reader,
length - 4,
"SCTP reconfiguration parameter",
)
match parameter_type {
13 => {
if value.length() < 12 || (value.length() - 12) % 2 != 0 {
raise InvalidPacket("malformed outgoing stream reset request")
}
let value_reader = @codec.Reader::new(value)
let request_sequence = sctp_read_u32(
value_reader, "stream reset request sequence",
)
let response_sequence = sctp_read_u32(
value_reader, "stream reset response sequence",
)
let sender_last_tsn = sctp_read_u32(
value_reader, "stream reset sender TSN",
)
let streams : Array[StreamId] = []
while value_reader.remaining() > 0 {
streams.push(
StreamId(sctp_read_u16(value_reader, "stream reset stream ID")),
)
}
OutgoingResetRequest(
request_sequence~,
response_sequence~,
sender_last_tsn~,
streams~,
)
}
16 => {
if value.length() != 8 && value.length() != 16 {
raise InvalidPacket("malformed stream reconfiguration response")
}
let value_reader = @codec.Reader::new(value)
let request_sequence = sctp_read_u32(
value_reader, "stream reset response sequence",
)
let result = sctp_read_u32(value_reader, "stream reset result")
ReconfigResponse(request_sequence~, result~)
}
_ => UnknownReconfigParameter(parameter_type, value)
}
}
///|
fn encode_reconfig(
parameters : Array[ReconfigParameter],
) -> Bytes raise SctpError {
if parameters.is_empty() || parameters.length() > 2 {
raise InvalidPacket("SCTP RE-CONFIG must contain one or two parameters")
}
let writer = sctp_writer(64)
for index = 0; index < parameters.length(); index = index + 1 {
let parameter = encode_reconfig_parameter(parameters[index])
writer.write_bytes(parameter)
if index + 1 < parameters.length() {
for padding = 0
padding < padding_length(parameter.length())
padding = padding + 1 {
writer.write_u8(0)
}
}
}
writer.finish()
}
///|
fn decode_reconfig(
flags : Byte,
value : Bytes,
) -> Array[ReconfigParameter] raise SctpError {
if flags != 0 || value.is_empty() {
raise InvalidPacket("malformed SCTP RE-CONFIG chunk")
}
let reader = @codec.Reader::new(value)
let parameters : Array[ReconfigParameter] = []
while reader.remaining() > 0 {
if parameters.length() >= 2 || reader.remaining() < 4 {
raise InvalidPacket("invalid SCTP RE-CONFIG parameter count")
}
let before = reader.position()
parameters.push(decode_reconfig_parameter(reader))
let consumed = reader.position() - before
let padding = padding_length(consumed)
if padding <= reader.remaining() {
let bytes = sctp_read_bytes(reader, padding, "SCTP RE-CONFIG padding")
for byte in bytes {
if byte != 0 {
raise InvalidPacket("SCTP parameter padding must be zero")
}
}
}
}
parameters
}
///|
fn encode_forward_tsn(forward : ForwardTsn) -> Bytes raise SctpError {
let writer = sctp_writer(4 + forward.streams.length() * 4)
writer.write_u32_be(forward.new_cumulative_tsn)
for stream in forward.streams {
writer.write_u16_be(stream.stream.value())
writer.write_u16_be(stream.sequence)
}
writer.finish()
}
///|
fn decode_forward_tsn(
flags : Byte,
value : Bytes,
) -> ForwardTsn raise SctpError {
if flags != 0 || value.length() < 4 || (value.length() - 4) % 4 != 0 {
raise InvalidPacket("malformed SCTP FORWARD-TSN chunk")
}
let reader = @codec.Reader::new(value)
let new_cumulative_tsn = sctp_read_u32(
reader, "SCTP FORWARD-TSN cumulative TSN",
)
let streams : Array[ForwardTsnStream] = []
while reader.remaining() > 0 {
streams.push(
ForwardTsnStream::new(
stream=StreamId(
sctp_read_u16(reader, "SCTP FORWARD-TSN stream identifier"),
),
sequence=sctp_read_u16(reader, "SCTP FORWARD-TSN stream sequence"),
),
)
}
ForwardTsn::new(new_cumulative_tsn~, streams~)
}
///|
fn chunk_wire_fields(chunk : Chunk) -> (Byte, Byte, Bytes) raise SctpError {
match chunk {
InitChunkValue(init) =>
(if init.acknowledgement { 2 } else { 1 }, 0, encode_init(init))
DataChunkValue(data) => {
let (flags, value) = encode_data(data)
(0, flags, value)
}
SackChunkValue(sack) => (3, 0, encode_sack(sack))
HeartbeatChunk(value) => (4, 0, value)
HeartbeatAckChunk(value) => (5, 0, value)
AbortChunk(value) => (6, 0, value)
ShutdownChunk(cumulative_tsn_ack) => {
let writer = sctp_writer(4)
writer.write_u32_be(cumulative_tsn_ack)
(7, 0, writer.finish())
}
ShutdownAckChunk => (8, 0, b"")
ErrorChunk(value) => (9, 0, value)
CookieEchoChunk(value) => (10, 0, value)
CookieAckChunk => (11, 0, b"")
ShutdownCompleteChunk => (14, 0, b"")
ReconfigChunk(parameters) => (130, 0, encode_reconfig(parameters))
ForwardTsnChunk(forward) => (192, 0, encode_forward_tsn(forward))
UnknownChunk(chunk_type, flags, value) => (chunk_type, flags, value)
}
}
///|
fn encode_chunk(chunk : Chunk) -> Bytes raise SctpError {
let (chunk_type, flags, value) = chunk_wire_fields(chunk)
if value.length() > 0xfffb {
raise InvalidPacket("SCTP chunk exceeds 65535 bytes")
}
let writer = sctp_writer(4 + value.length())
writer.write_u8(chunk_type)
writer.write_u8(flags)
writer.write_u16_be((4 + value.length()).to_uint16())
writer.write_bytes(value)
writer.finish()
}
///|
fn decode_chunk(
chunk_type : Byte,
flags : Byte,
value : Bytes,
) -> Chunk raise SctpError {
match chunk_type {
0 => DataChunkValue(decode_data(flags, value))
1 => InitChunkValue(decode_init(false, flags, value))
2 => InitChunkValue(decode_init(true, flags, value))
3 => SackChunkValue(decode_sack(flags, value))
4 => HeartbeatChunk(value)
5 => HeartbeatAckChunk(value)
6 => AbortChunk(value)
7 => {
if flags != 0 || value.length() != 4 {
raise InvalidPacket("malformed SCTP SHUTDOWN chunk")
}
let reader = @codec.Reader::new(value)
ShutdownChunk(sctp_read_u32(reader, "SCTP shutdown TSN"))
}
8 => {
if flags != 0 || !value.is_empty() {
raise InvalidPacket("malformed SCTP SHUTDOWN-ACK chunk")
}
ShutdownAckChunk
}
9 => ErrorChunk(value)
10 => CookieEchoChunk(value)
11 => {
if flags != 0 || !value.is_empty() {
raise InvalidPacket("malformed SCTP COOKIE-ACK chunk")
}
CookieAckChunk
}
14 => {
if !value.is_empty() {
raise InvalidPacket("malformed SCTP SHUTDOWN-COMPLETE chunk")
}
ShutdownCompleteChunk
}
130 => ReconfigChunk(decode_reconfig(flags, value))
192 => ForwardTsnChunk(decode_forward_tsn(flags, value))
_ => UnknownChunk(chunk_type, flags, value)
}
}
///|
pub fn Packet::encode(self : Packet) -> Bytes raise SctpError {
if self.source_port == 0 || self.destination_port == 0 {
raise InvalidPacket("SCTP ports must be nonzero")
}
if self.chunks.any(chunk => {
match chunk {
InitChunkValue(init) => !init.acknowledgement
_ => false
}
}) &&
(self.chunks.length() != 1 || self.verification_tag != 0U) {
raise InvalidPacket(
"SCTP INIT must be unbundled with a zero verification tag",
)
}
let writer = sctp_writer(256)
writer.write_u16_be(self.source_port)
writer.write_u16_be(self.destination_port)
writer.write_u32_be(self.verification_tag)
writer.write_u32_be(0)
for chunk in self.chunks {
let encoded = encode_chunk(chunk)
writer.write_bytes(encoded)
for padding = 0
padding < padding_length(encoded.length())
padding = padding + 1 {
writer.write_u8(0)
}
}
let bytes = writer.finish().to_array()
let checksum = crc32c(Bytes::from_array(bytes))
bytes[8] = checksum.to_byte()
bytes[9] = (checksum >> 8).to_byte()
bytes[10] = (checksum >> 16).to_byte()
bytes[11] = (checksum >> 24).to_byte()
Bytes::from_array(bytes)
}
///|
pub fn Packet::decode(raw : Bytes) -> Packet raise SctpError {
if raw.length() < 12 {
raise InvalidPacket("SCTP packet is shorter than its common header")
}
let their_checksum = raw[8].to_uint() |
(raw[9].to_uint() << 8) |
(raw[10].to_uint() << 16) |
(raw[11].to_uint() << 24)
if their_checksum != crc32c(raw) {
raise ChecksumMismatch
}
let reader = @codec.Reader::new(raw)
let source_port = sctp_read_u16(reader, "SCTP source port")
let destination_port = sctp_read_u16(reader, "SCTP destination port")
let verification_tag = sctp_read_u32(reader, "SCTP verification tag")
ignore(sctp_read_u32(reader, "SCTP checksum"))
if source_port == 0 || destination_port == 0 {
raise InvalidPacket("SCTP ports must be nonzero")
}
let chunks : Array[Chunk] = []
while reader.remaining() > 0 {
if reader.remaining() < 4 {
raise InvalidPacket("truncated SCTP chunk header")
}
let chunk_type = sctp_read_u8(reader, "SCTP chunk type")
let flags = sctp_read_u8(reader, "SCTP chunk flags")
let length = sctp_read_u16(reader, "SCTP chunk length").to_int()
if length < 4 {
raise InvalidPacket("SCTP chunk length is smaller than its header")
}
let value = sctp_read_bytes(reader, length - 4, "SCTP chunk value")
chunks.push(decode_chunk(chunk_type, flags, value))
let padding = padding_length(length)
let padding_bytes = sctp_read_bytes(reader, padding, "SCTP chunk padding")
for byte in padding_bytes {
if byte != 0 {
raise InvalidPacket("SCTP chunk padding must be zero")
}
}
}
if chunks.any(chunk => {
match chunk {
InitChunkValue(init) => !init.acknowledgement
_ => false
}
}) &&
(chunks.length() != 1 || verification_tag != 0U) {
raise InvalidPacket(
"SCTP INIT must be unbundled with a zero verification tag",
)
}
{ source_port, destination_port, verification_tag, chunks, }
}