///|
fn write_opaque8(
writer : @codec.Writer,
value : Bytes,
context : String,
) -> Unit raise DtlsError {
if value.length() > 0xff {
raise InvalidHandshake("\{context} exceeds 255 bytes")
}
writer.write_u8(value.length().to_byte())
writer.write_bytes(value)
}
///|
fn write_opaque16(
writer : @codec.Writer,
value : Bytes,
context : String,
) -> Unit raise DtlsError {
if value.length() > 0xffff {
raise InvalidHandshake("\{context} exceeds 65535 bytes")
}
writer.write_u16_be(value.length().to_uint16())
writer.write_bytes(value)
}
///|
fn write_opaque24(
writer : @codec.Writer,
value : Bytes,
context : String,
) -> Unit raise DtlsError {
if value.length() > 0xffffff {
raise InvalidHandshake("\{context} exceeds 24-bit framing")
}
writer.write_u24_be(value.length().reinterpret_as_uint())
writer.write_bytes(value)
}
///|
fn read_opaque8(
reader : @codec.Reader,
context : String,
) -> Bytes raise DtlsError {
let length = hs_read_u8(reader, context).to_int()
hs_read_bytes(reader, length, context)
}
///|
fn read_opaque16(
reader : @codec.Reader,
context : String,
) -> Bytes raise DtlsError {
let length = hs_read_u16(reader, context).to_int()
hs_read_bytes(reader, length, context)
}
///|
fn read_opaque24(
reader : @codec.Reader,
context : String,
) -> Bytes raise DtlsError {
let length = hs_read_u24(reader, context).reinterpret_as_int()
hs_read_bytes(reader, length, context)
}
///|
pub struct ClientHelloMessage {
random : Bytes
session_id : Bytes
cookie : Bytes
cipher_suites : Array[UInt16]
compression_methods : Array[Byte]
extensions : HelloExtensions
} derive(Debug, Eq)
///|
pub fn ClientHelloMessage::new(
random~ : Bytes,
cookie? : Bytes = b"",
cipher_suites? : Array[CipherSuite] = [EcdheEcdsaAes128GcmSha256],
extensions~ : HelloExtensions,
) -> ClientHelloMessage raise DtlsError {
if random.length() != 32 {
raise InvalidHandshake("ClientHello random must contain 32 bytes")
}
if cookie.length() > 255 {
raise InvalidHandshake("ClientHello cookie exceeds 255 bytes")
}
if cipher_suites.is_empty() {
raise InvalidHandshake("ClientHello must offer a cipher suite")
}
let cipher_suite_codes : Array[UInt16] = []
for suite in cipher_suites {
if cipher_suite_codes.contains(suite.code()) {
raise InvalidHandshake("ClientHello contains a duplicate cipher suite")
}
cipher_suite_codes.push(suite.code())
}
{
random,
session_id: b"",
cookie,
cipher_suites: cipher_suite_codes,
compression_methods: [0],
extensions,
}
}
///|
pub fn ClientHelloMessage::random(self : ClientHelloMessage) -> Bytes {
self.random
}
///|
pub fn ClientHelloMessage::cookie(self : ClientHelloMessage) -> Bytes {
self.cookie
}
///|
pub fn ClientHelloMessage::cipher_suites(
self : ClientHelloMessage,
) -> Array[UInt16] {
self.cipher_suites.copy()
}
///|
pub fn ClientHelloMessage::extensions(
self : ClientHelloMessage,
) -> HelloExtensions {
self.extensions
}
///|
fn ClientHelloMessage::encode_body(
self : ClientHelloMessage,
) -> Bytes raise DtlsError {
if self.session_id.length() > 32 ||
self.cipher_suites.is_empty() ||
self.cipher_suites.length() > 0x7fff ||
self.compression_methods.is_empty() ||
self.compression_methods.length() > 0xff {
raise InvalidHandshake("invalid ClientHello vector length")
}
let writer = dtls_writer(256, "ClientHello")
writer.write_u8(Dtls12.major())
writer.write_u8(Dtls12.minor())
writer.write_bytes(self.random)
write_opaque8(writer, self.session_id, "ClientHello session ID")
write_opaque8(writer, self.cookie, "ClientHello cookie")
writer.write_u16_be((self.cipher_suites.length() * 2).to_uint16())
for cipher_suite in self.cipher_suites {
writer.write_u16_be(cipher_suite)
}
writer.write_u8(self.compression_methods.length().to_byte())
for compression_method in self.compression_methods {
writer.write_u8(compression_method)
}
writer.write_bytes(self.extensions.encode())
writer.finish()
}
///|
fn ClientHelloMessage::decode_body(
body : Bytes,
) -> ClientHelloMessage raise DtlsError {
let reader = @codec.Reader::new(body)
let version = ProtocolVersion::from_bytes(
hs_read_u8(reader, "ClientHello version"),
hs_read_u8(reader, "ClientHello version"),
)
if version != Dtls12 {
raise UnsupportedVersion
}
let random = hs_read_bytes(reader, 32, "ClientHello random")
let session_id = read_opaque8(reader, "ClientHello session ID")
if session_id.length() > 32 {
raise InvalidHandshake("ClientHello session ID exceeds 32 bytes")
}
let cookie = read_opaque8(reader, "ClientHello cookie")
let cipher_suite_bytes = hs_read_u16(reader, "ClientHello cipher suites").to_int()
if cipher_suite_bytes == 0 ||
cipher_suite_bytes % 2 != 0 ||
reader.remaining() < cipher_suite_bytes {
raise InvalidHandshake("malformed ClientHello cipher-suite vector")
}
let cipher_suites : Array[UInt16] = []
for offset = 0; offset < cipher_suite_bytes; offset = offset + 2 {
cipher_suites.push(hs_read_u16(reader, "ClientHello cipher suite"))
}
let compression_count = hs_read_u8(reader, "ClientHello compression methods").to_int()
if compression_count == 0 || reader.remaining() < compression_count {
raise InvalidHandshake("malformed ClientHello compression vector")
}
let compression_methods = hs_read_bytes(
reader, compression_count, "ClientHello compression methods",
).to_array()
if !compression_methods.contains(0) {
raise InvalidHandshake("ClientHello does not offer null compression")
}
if reader.remaining() < 2 {
raise InvalidHandshake("ClientHello extensions are missing")
}
let extensions = HelloExtensions::decode(
hs_read_bytes(reader, reader.remaining(), "ClientHello extensions"),
)
{
random,
session_id,
cookie,
cipher_suites,
compression_methods,
extensions,
}
}
///|
pub struct ServerHelloMessage {
random : Bytes
session_id : Bytes
cipher_suite : UInt16
compression_method : Byte
extensions : HelloExtensions
} derive(Debug, Eq)
///|
pub fn ServerHelloMessage::new(
random~ : Bytes,
cipher_suite? : CipherSuite = EcdheEcdsaAes128GcmSha256,
selected_srtp_profile? : UInt16,
) -> ServerHelloMessage raise DtlsError {
if random.length() != 32 {
raise InvalidHandshake("ServerHello random must contain 32 bytes")
}
{
random,
session_id: b"",
cipher_suite: cipher_suite.code(),
compression_method: 0,
extensions: HelloExtensions::server(selected_srtp_profile),
}
}
///|
pub fn ServerHelloMessage::random(self : ServerHelloMessage) -> Bytes {
self.random
}
///|
pub fn ServerHelloMessage::cipher_suite(self : ServerHelloMessage) -> UInt16 {
self.cipher_suite
}
///|
pub fn ServerHelloMessage::extensions(
self : ServerHelloMessage,
) -> HelloExtensions {
self.extensions
}
///|
fn ServerHelloMessage::encode_body(
self : ServerHelloMessage,
) -> Bytes raise DtlsError {
if self.random.length() != 32 || self.session_id.length() > 32 {
raise InvalidHandshake("invalid ServerHello field length")
}
let writer = dtls_writer(160, "ServerHello")
writer.write_u8(Dtls12.major())
writer.write_u8(Dtls12.minor())
writer.write_bytes(self.random)
write_opaque8(writer, self.session_id, "ServerHello session ID")
writer.write_u16_be(self.cipher_suite)
writer.write_u8(self.compression_method)
writer.write_bytes(self.extensions.encode())
writer.finish()
}
///|
fn ServerHelloMessage::decode_body(
body : Bytes,
) -> ServerHelloMessage raise DtlsError {
let reader = @codec.Reader::new(body)
let version = ProtocolVersion::from_bytes(
hs_read_u8(reader, "ServerHello version"),
hs_read_u8(reader, "ServerHello version"),
)
if version != Dtls12 {
raise UnsupportedVersion
}
let random = hs_read_bytes(reader, 32, "ServerHello random")
let session_id = read_opaque8(reader, "ServerHello session ID")
if session_id.length() > 32 {
raise InvalidHandshake("ServerHello session ID exceeds 32 bytes")
}
let cipher_suite = hs_read_u16(reader, "ServerHello cipher suite")
let compression_method = hs_read_u8(reader, "ServerHello compression method")
if compression_method != 0 {
raise InvalidHandshake("ServerHello selected non-null compression")
}
let extensions = HelloExtensions::decode(
hs_read_bytes(reader, reader.remaining(), "ServerHello extensions"),
)
{ random, session_id, cipher_suite, compression_method, extensions, }
}
///|
pub struct HelloVerifyRequestMessage {
cookie : Bytes
} derive(Debug, Eq)
///|
fn HelloVerifyRequestMessage::encode_body(
self : HelloVerifyRequestMessage,
) -> Bytes raise DtlsError {
let writer = dtls_writer(3 + self.cookie.length(), "HelloVerifyRequest")
writer.write_u8(Dtls12.major())
writer.write_u8(Dtls12.minor())
write_opaque8(writer, self.cookie, "HelloVerifyRequest cookie")
writer.finish()
}
///|
fn HelloVerifyRequestMessage::decode_body(
body : Bytes,
) -> HelloVerifyRequestMessage raise DtlsError {
let reader = @codec.Reader::new(body)
ignore(
ProtocolVersion::from_bytes(
hs_read_u8(reader, "HelloVerifyRequest version"),
hs_read_u8(reader, "HelloVerifyRequest version"),
),
)
let cookie = read_opaque8(reader, "HelloVerifyRequest cookie")
if reader.remaining() != 0 {
raise InvalidHandshake("trailing HelloVerifyRequest bytes")
}
{ cookie, }
}
///|
pub struct CertificateMessage {
certificates : Array[Bytes]
} derive(Debug, Eq)
///|
fn CertificateMessage::encode_body(
self : CertificateMessage,
) -> Bytes raise DtlsError {
let certificates = dtls_writer(1024, "Certificate list")
for certificate in self.certificates {
write_opaque24(certificates, certificate, "certificate")
}
let certificates = certificates.finish()
let writer = dtls_writer(certificates.length() + 3, "Certificate")
write_opaque24(writer, certificates, "certificate list")
writer.finish()
}
///|
fn CertificateMessage::decode_body(
body : Bytes,
) -> CertificateMessage raise DtlsError {
let reader = @codec.Reader::new(body)
let encoded_certificates = read_opaque24(reader, "certificate list")
if reader.remaining() != 0 {
raise InvalidHandshake("trailing Certificate message bytes")
}
let certificate_reader = @codec.Reader::new(encoded_certificates)
let certificates : Array[Bytes] = []
while certificate_reader.remaining() > 0 {
certificates.push(read_opaque24(certificate_reader, "certificate"))
}
{ certificates, }
}
///|
pub struct ServerKeyExchangeMessage {
named_curve : UInt16
public_key : Bytes
signature_algorithm : UInt16
signature : Bytes
identity_hint : Bytes?
} derive(Debug, Eq)
///|
fn ServerKeyExchangeMessage::parameters(
self : ServerKeyExchangeMessage,
) -> Bytes raise DtlsError {
if self.identity_hint is Some(_) {
raise InvalidHandshake("PSK ServerKeyExchange has no ECDHE parameters")
}
let writer = dtls_writer(4 + self.public_key.length(), "ECDHE parameters")
writer.write_u8(3)
writer.write_u16_be(self.named_curve)
write_opaque8(writer, self.public_key, "ECDHE public key")
writer.finish()
}
///|
fn ServerKeyExchangeMessage::encode_body(
self : ServerKeyExchangeMessage,
) -> Bytes raise DtlsError {
match self.identity_hint {
Some(identity_hint) => {
let writer = dtls_writer(
2 + identity_hint.length(),
"PSK ServerKeyExchange",
)
write_opaque16(writer, identity_hint, "PSK identity hint")
return writer.finish()
}
None => ()
}
let writer = dtls_writer(160, "ServerKeyExchange")
writer.write_bytes(self.parameters())
writer.write_u16_be(self.signature_algorithm)
write_opaque16(writer, self.signature, "ServerKeyExchange signature")
writer.finish()
}
///|
fn ServerKeyExchangeMessage::decode_body(
body : Bytes,
) -> ServerKeyExchangeMessage raise DtlsError {
if body.length() >= 2 {
let identity_length = ((body[0].to_uint() << 8) | body[1].to_uint()).reinterpret_as_int()
if identity_length + 2 == body.length() {
return {
named_curve: 0,
public_key: b"",
signature_algorithm: 0,
signature: b"",
identity_hint: Some(body[2:].to_owned()),
}
}
}
let reader = @codec.Reader::new(body)
if hs_read_u8(reader, "elliptic curve type") != 3 {
raise InvalidHandshake("only named ECDHE curves are supported")
}
let named_curve = hs_read_u16(reader, "named curve")
let public_key = read_opaque8(reader, "ECDHE public key")
let signature_algorithm = hs_read_u16(
reader, "ServerKeyExchange signature algorithm",
)
let signature = read_opaque16(reader, "ServerKeyExchange signature")
if reader.remaining() != 0 {
raise InvalidHandshake("trailing ServerKeyExchange bytes")
}
{
named_curve,
public_key,
signature_algorithm,
signature,
identity_hint: None,
}
}
///|
pub struct CertificateRequestMessage {
certificate_types : Array[Byte]
signature_algorithms : Array[UInt16]
} derive(Debug, Eq)
///|
fn CertificateRequestMessage::webrtc(
certificate_type : Byte,
signature_algorithm : UInt16,
) -> CertificateRequestMessage {
{
certificate_types: [certificate_type],
signature_algorithms: [signature_algorithm],
}
}
///|
fn CertificateRequestMessage::encode_body(
self : CertificateRequestMessage,
) -> Bytes raise DtlsError {
if self.certificate_types.is_empty() || self.certificate_types.length() > 0xff {
raise InvalidHandshake("invalid certificate-type vector")
}
let writer = dtls_writer(32, "CertificateRequest")
writer.write_u8(self.certificate_types.length().to_byte())
for certificate_type in self.certificate_types {
writer.write_u8(certificate_type)
}
writer.write_bytes(u16_vector(self.signature_algorithms))
writer.write_u16_be(0)
writer.finish()
}
///|
fn CertificateRequestMessage::decode_body(
body : Bytes,
) -> CertificateRequestMessage raise DtlsError {
let reader = @codec.Reader::new(body)
let certificate_type_count = hs_read_u8(
reader, "CertificateRequest certificate types",
).to_int()
if certificate_type_count == 0 {
raise InvalidHandshake("CertificateRequest has no certificate type")
}
let certificate_types = hs_read_bytes(
reader, certificate_type_count, "CertificateRequest certificate types",
).to_array()
let algorithm_length = hs_read_u16(
reader, "CertificateRequest signature algorithms",
).to_int()
if algorithm_length == 0 ||
algorithm_length % 2 != 0 ||
reader.remaining() < algorithm_length + 2 {
raise InvalidHandshake("malformed CertificateRequest signature vector")
}
let signature_algorithms : Array[UInt16] = []
for offset = 0; offset < algorithm_length; offset = offset + 2 {
signature_algorithms.push(
hs_read_u16(reader, "CertificateRequest signature algorithm"),
)
}
let authorities_length = hs_read_u16(reader, "CertificateRequest authorities").to_int()
ignore(
hs_read_bytes(reader, authorities_length, "CertificateRequest authorities"),
)
if reader.remaining() != 0 {
raise InvalidHandshake("trailing CertificateRequest bytes")
}
{ certificate_types, signature_algorithms, }
}
///|
pub struct ClientKeyExchangeMessage {
public_key : Bytes
identity_hint : Bytes?
} derive(Debug, Eq)
///|
fn ClientKeyExchangeMessage::encode_body(
self : ClientKeyExchangeMessage,
) -> Bytes raise DtlsError {
match self.identity_hint {
Some(identity_hint) => {
let writer = dtls_writer(
2 + identity_hint.length(),
"PSK ClientKeyExchange",
)
write_opaque16(writer, identity_hint, "PSK identity")
return writer.finish()
}
None => ()
}
let writer = dtls_writer(66, "ClientKeyExchange")
write_opaque8(writer, self.public_key, "ECDHE public key")
writer.finish()
}
///|
fn ClientKeyExchangeMessage::decode_body(
body : Bytes,
) -> ClientKeyExchangeMessage raise DtlsError {
if body.length() >= 2 {
let identity_length = ((body[0].to_uint() << 8) | body[1].to_uint()).reinterpret_as_int()
if identity_length + 2 == body.length() {
return { public_key: b"", identity_hint: Some(body[2:].to_owned()), }
}
}
let reader = @codec.Reader::new(body)
let public_key = read_opaque8(reader, "ECDHE public key")
if reader.remaining() != 0 {
raise InvalidHandshake("trailing ClientKeyExchange bytes")
}
{ public_key, identity_hint: None, }
}
///|
pub struct CertificateVerifyMessage {
signature_algorithm : UInt16
signature : Bytes
} derive(Debug, Eq)
///|
fn CertificateVerifyMessage::encode_body(
self : CertificateVerifyMessage,
) -> Bytes raise DtlsError {
let writer = dtls_writer(4 + self.signature.length(), "CertificateVerify")
writer.write_u16_be(self.signature_algorithm)
write_opaque16(writer, self.signature, "CertificateVerify signature")
writer.finish()
}
///|
fn CertificateVerifyMessage::decode_body(
body : Bytes,
) -> CertificateVerifyMessage raise DtlsError {
let reader = @codec.Reader::new(body)
let signature_algorithm = hs_read_u16(
reader, "CertificateVerify signature algorithm",
)
let signature = read_opaque16(reader, "CertificateVerify signature")
if reader.remaining() != 0 {
raise InvalidHandshake("trailing CertificateVerify bytes")
}
{ signature_algorithm, signature, }
}
///|
pub struct FinishedMessage {
verify_data : Bytes
} derive(Debug, Eq)
///|
fn FinishedMessage::encode_body(
self : FinishedMessage,
) -> Bytes raise DtlsError {
if self.verify_data.length() != 12 {
raise InvalidHandshake("Finished verify_data must contain 12 bytes")
}
self.verify_data
}
///|
fn FinishedMessage::decode_body(
body : Bytes,
) -> FinishedMessage raise DtlsError {
if body.length() != 12 {
raise InvalidHandshake("Finished verify_data must contain 12 bytes")
}
{ verify_data: body, }
}
///|
pub(all) enum HandshakeMessage {
ClientHelloHandshake(ClientHelloMessage)
ServerHelloHandshake(ServerHelloMessage)
HelloVerifyRequestHandshake(HelloVerifyRequestMessage)
CertificateHandshake(CertificateMessage)
ServerKeyExchangeHandshake(ServerKeyExchangeMessage)
CertificateRequestHandshake(CertificateRequestMessage)
ServerHelloDoneHandshake
CertificateVerifyHandshake(CertificateVerifyMessage)
ClientKeyExchangeHandshake(ClientKeyExchangeMessage)
FinishedHandshake(FinishedMessage)
} derive(Debug, Eq)
///|
pub fn HandshakeMessage::handshake_type(
self : HandshakeMessage,
) -> HandshakeType {
match self {
ClientHelloHandshake(_) => ClientHello
ServerHelloHandshake(_) => ServerHello
HelloVerifyRequestHandshake(_) => HelloVerifyRequest
CertificateHandshake(_) => Certificate
ServerKeyExchangeHandshake(_) => ServerKeyExchange
CertificateRequestHandshake(_) => CertificateRequest
ServerHelloDoneHandshake => ServerHelloDone
CertificateVerifyHandshake(_) => CertificateVerify
ClientKeyExchangeHandshake(_) => ClientKeyExchange
FinishedHandshake(_) => Finished
}
}
///|
pub fn HandshakeMessage::encode_body(
self : HandshakeMessage,
) -> Bytes raise DtlsError {
match self {
ClientHelloHandshake(message) => message.encode_body()
ServerHelloHandshake(message) => message.encode_body()
HelloVerifyRequestHandshake(message) => message.encode_body()
CertificateHandshake(message) => message.encode_body()
ServerKeyExchangeHandshake(message) => message.encode_body()
CertificateRequestHandshake(message) => message.encode_body()
ServerHelloDoneHandshake => b""
CertificateVerifyHandshake(message) => message.encode_body()
ClientKeyExchangeHandshake(message) => message.encode_body()
FinishedHandshake(message) => message.encode_body()
}
}
///|
pub fn HandshakeMessage::to_fragment(
self : HandshakeMessage,
message_sequence : UInt16,
) -> HandshakeFragment raise DtlsError {
let body = self.encode_body()
HandshakeFragment::new(
handshake_type=self.handshake_type(),
total_length=body.length().reinterpret_as_uint(),
message_sequence~,
fragment_offset=0,
body~,
)
}
///|
pub fn HandshakeMessage::decode(
fragment : HandshakeFragment,
) -> HandshakeMessage raise DtlsError {
if !fragment.is_complete() {
raise InvalidHandshake(
"fragmented handshake must be reassembled before decoding",
)
}
match fragment.handshake_type {
ClientHello =>
ClientHelloHandshake(ClientHelloMessage::decode_body(fragment.body))
ServerHello =>
ServerHelloHandshake(ServerHelloMessage::decode_body(fragment.body))
HelloVerifyRequest =>
HelloVerifyRequestHandshake(
HelloVerifyRequestMessage::decode_body(fragment.body),
)
Certificate =>
CertificateHandshake(CertificateMessage::decode_body(fragment.body))
ServerKeyExchange =>
ServerKeyExchangeHandshake(
ServerKeyExchangeMessage::decode_body(fragment.body),
)
CertificateRequest =>
CertificateRequestHandshake(
CertificateRequestMessage::decode_body(fragment.body),
)
ServerHelloDone => {
if !fragment.body.is_empty() {
raise InvalidHandshake("ServerHelloDone must be empty")
}
ServerHelloDoneHandshake
}
CertificateVerify =>
CertificateVerifyHandshake(
CertificateVerifyMessage::decode_body(fragment.body),
)
ClientKeyExchange =>
ClientKeyExchangeHandshake(
ClientKeyExchangeMessage::decode_body(fragment.body),
)
Finished => FinishedHandshake(FinishedMessage::decode_body(fragment.body))
_ =>
raise InvalidHandshake(
"unsupported DTLS handshake type \{fragment.handshake_type.code()}",
)
}
}