///|
const RECORD_HEADER_LENGTH : Int = 13
///|
const HANDSHAKE_HEADER_LENGTH : Int = 12
///|
fn dtls_writer(
capacity : Int,
context : String,
) -> @codec.Writer raise DtlsError {
@codec.Writer::new(capacity~) catch {
InvalidLength(length) =>
raise InvalidRecord("\{context}: invalid length \{length}")
Truncated(needed~, remaining~) =>
raise InvalidRecord(
"\{context}: unexpected truncation \{needed}/\{remaining}",
)
}
}
///|
fn read_u8(reader : @codec.Reader, context : String) -> Byte raise DtlsError {
reader.read_u8() catch {
InvalidLength(length) =>
raise InvalidRecord("\{context}: invalid length \{length}")
Truncated(needed~, remaining~) =>
raise InvalidRecord("\{context}: need \{needed} bytes, have \{remaining}")
}
}
///|
fn read_u16(reader : @codec.Reader, context : String) -> UInt16 raise DtlsError {
reader.read_u16_be() catch {
InvalidLength(length) =>
raise InvalidRecord("\{context}: invalid length \{length}")
Truncated(needed~, remaining~) =>
raise InvalidRecord("\{context}: need \{needed} bytes, have \{remaining}")
}
}
///|
fn read_u24(reader : @codec.Reader, context : String) -> UInt raise DtlsError {
reader.read_u24_be() catch {
InvalidLength(length) =>
raise InvalidHandshake("\{context}: invalid length \{length}")
Truncated(needed~, remaining~) =>
raise InvalidHandshake(
"\{context}: need \{needed} bytes, have \{remaining}",
)
}
}
///|
fn read_u32(reader : @codec.Reader, context : String) -> UInt raise DtlsError {
reader.read_u32_be() catch {
InvalidLength(length) =>
raise InvalidRecord("\{context}: invalid length \{length}")
Truncated(needed~, remaining~) =>
raise InvalidRecord("\{context}: need \{needed} bytes, have \{remaining}")
}
}
///|
fn read_bytes(
reader : @codec.Reader,
length : Int,
context : String,
) -> Bytes raise DtlsError {
reader.read_bytes(length) catch {
InvalidLength(invalid_length) =>
raise InvalidRecord("\{context}: invalid length \{invalid_length}")
Truncated(needed~, remaining~) =>
raise InvalidRecord("\{context}: need \{needed} bytes, have \{remaining}")
}
}
///|
fn ContentType::from_code(code : Byte) -> ContentType raise DtlsError {
match code {
20 => ChangeCipherSpec
21 => Alert
22 => Handshake
23 => ApplicationData
_ => raise InvalidRecord("unknown DTLS content type \{code}")
}
}
///|
fn ProtocolVersion::from_bytes(
major : Byte,
minor : Byte,
) -> ProtocolVersion raise DtlsError {
match (major, minor) {
(0xfe, 0xff) => Dtls10
(0xfe, 0xfd) => Dtls12
_ => raise UnsupportedVersion
}
}
///|
fn HandshakeType::from_code(code : Byte) -> HandshakeType {
match code {
0 => HelloRequest
1 => ClientHello
2 => ServerHello
3 => HelloVerifyRequest
11 => Certificate
12 => ServerKeyExchange
13 => CertificateRequest
14 => ServerHelloDone
15 => CertificateVerify
16 => ClientKeyExchange
20 => Finished
_ => UnknownHandshake(code)
}
}
///|
fn RecordHeader::encode_into(
self : RecordHeader,
writer : @codec.Writer,
) -> Unit {
writer.write_u8(self.content_type.code())
writer.write_u8(self.version.major())
writer.write_u8(self.version.minor())
writer.write_u16_be(self.epoch)
writer.write_u16_be((self.sequence_number >> 32).to_uint16())
writer.write_u32_be(self.sequence_number.to_uint())
writer.write_u16_be(self.content_length)
}
///|
pub fn Record::encode(self : Record) -> Bytes raise DtlsError {
if self.payload.length() != self.header.content_length.to_int() {
raise InvalidRecord("DTLS record header length does not match payload")
}
let writer = dtls_writer(
RECORD_HEADER_LENGTH + self.payload.length(),
"DTLS record",
)
self.header.encode_into(writer)
writer.write_bytes(self.payload)
writer.finish()
}
///|
fn decode_record(reader : @codec.Reader) -> Record raise DtlsError {
let content_type = ContentType::from_code(
read_u8(reader, "DTLS record content type"),
)
let major = read_u8(reader, "DTLS record version")
let minor = read_u8(reader, "DTLS record version")
let version = ProtocolVersion::from_bytes(major, minor)
let epoch = read_u16(reader, "DTLS record epoch")
let sequence_number = (
read_u16(reader, "DTLS record sequence").to_uint64() << 32
) |
read_u32(reader, "DTLS record sequence").to_uint64()
let content_length = read_u16(reader, "DTLS record length")
let payload = read_bytes(
reader,
content_length.to_int(),
"DTLS record payload",
)
{
header: RecordHeader::new(
content_type~,
version~,
epoch~,
sequence_number~,
content_length~,
),
payload,
}
}
///|
pub fn decode_records(datagram : Bytes) -> Array[Record] raise DtlsError {
if datagram.is_empty() {
raise InvalidRecord("empty DTLS datagram")
}
let reader = @codec.Reader::new(datagram)
let records : Array[Record] = []
while reader.remaining() > 0 {
if reader.remaining() < RECORD_HEADER_LENGTH {
raise InvalidRecord("truncated DTLS record header")
}
records.push(decode_record(reader))
}
records
}
///|
pub fn encode_records(records : Array[Record]) -> Bytes raise DtlsError {
if records.is_empty() {
raise InvalidRecord("cannot encode an empty DTLS datagram")
}
let encoded : Array[Bytes] = []
let mut total_length = 0
for record in records {
let bytes = record.encode()
total_length += bytes.length()
encoded.push(bytes)
}
let writer = dtls_writer(total_length, "DTLS datagram")
for bytes in encoded {
writer.write_bytes(bytes)
}
writer.finish()
}
///|
pub fn HandshakeFragment::encode(
self : HandshakeFragment,
) -> Bytes raise DtlsError {
let writer = dtls_writer(
HANDSHAKE_HEADER_LENGTH + self.body.length(),
"DTLS handshake",
)
writer.write_u8(self.handshake_type.code())
writer.write_u24_be(self.total_length)
writer.write_u16_be(self.message_sequence)
writer.write_u24_be(self.fragment_offset)
writer.write_u24_be(self.body.length().reinterpret_as_uint())
writer.write_bytes(self.body)
writer.finish()
}
///|
fn decode_handshake(
reader : @codec.Reader,
) -> HandshakeFragment raise DtlsError {
let handshake_type = HandshakeType::from_code(
read_u8(reader, "DTLS handshake type"),
)
let total_length = read_u24(reader, "DTLS handshake length")
let message_sequence = read_u16(reader, "DTLS handshake sequence")
let fragment_offset = read_u24(reader, "DTLS handshake fragment offset")
let fragment_length = read_u24(reader, "DTLS handshake fragment length")
let body = read_bytes(
reader,
fragment_length.reinterpret_as_int(),
"DTLS handshake fragment",
) catch {
InvalidRecord(message) => raise InvalidHandshake(message)
error => raise error
}
HandshakeFragment::new(
handshake_type~,
total_length~,
message_sequence~,
fragment_offset~,
body~,
)
}
///|
pub fn decode_handshakes(
payload : Bytes,
) -> Array[HandshakeFragment] raise DtlsError {
if payload.is_empty() {
raise InvalidHandshake("empty DTLS handshake record")
}
let reader = @codec.Reader::new(payload)
let handshakes : Array[HandshakeFragment] = []
while reader.remaining() > 0 {
if reader.remaining() < HANDSHAKE_HEADER_LENGTH {
raise InvalidHandshake("truncated DTLS handshake header")
}
handshakes.push(decode_handshake(reader))
}
handshakes
}