File
Blob: firmware/vendor/str0m/src/dtls.rs
| 1 | use std::collections::VecDeque; |
| 2 | use std::io; |
| 3 | use std::ops::RangeInclusive; |
| 4 | use std::panic::{RefUnwindSafe, UnwindSafe}; |
| 5 | use std::time::Instant; |
| 6 | |
| 7 | use crate::crypto::Fingerprint; |
| 8 | use crate::crypto::Sha256Provider; |
| 9 | use crate::crypto::dtls::{DtlsCert, DtlsOutput, ProtocolVersion}; |
| 10 | use crate::crypto::dtls::{DtlsInstance, DtlsProvider, DtlsVersion}; |
| 11 | use crate::crypto::{CryptoError, DtlsError}; |
| 12 | use crate::io::DatagramSend; |
| 13 | use crate::util::already_happened; |
| 14 | |
| 15 | /// Encapsulation of DTLS. |
| 16 | /// |
| 17 | /// This is a thin wrapper around `DtlsInstance` that adds fingerprint tracking |
| 18 | /// and active/passive state management. The API follows dimpl's sans-IO pattern. |
| 19 | pub struct Dtls { |
| 20 | /// The underlying DTLS instance. |
| 21 | instance: Box<dyn DtlsInstance>, |
| 22 | |
| 23 | /// The fingerprint of the local certificate. |
| 24 | fingerprint: Fingerprint, |
| 25 | |
| 26 | /// Remote fingerprint (set when received via poll_output). |
| 27 | remote_fingerprint: Option<Fingerprint>, |
| 28 | |
| 29 | /// Whether set_active has been called. |
| 30 | active_state: Option<bool>, |
| 31 | |
| 32 | /// Packets to be sent. |
| 33 | pending_packets: VecDeque<DatagramSend>, |
| 34 | |
| 35 | /// Target MTU (start) and warn threshold (end). The target is forwarded to |
| 36 | /// the backend's fragmenter; outgoing records larger than the warn |
| 37 | /// threshold log a warning. |
| 38 | mtu: RangeInclusive<usize>, |
| 39 | } |
| 40 | |
| 41 | pub(crate) fn is_would_block(error: &DtlsError) -> bool { |
| 42 | match error { |
| 43 | DtlsError::Io(e) => e.kind() == io::ErrorKind::WouldBlock, |
| 44 | DtlsError::CryptoError(crypto_err) => match crypto_err { |
| 45 | CryptoError::Io(e) => e.kind() == io::ErrorKind::WouldBlock, |
| 46 | #[allow(unreachable_patterns)] |
| 47 | _ => false, |
| 48 | }, |
| 49 | } |
| 50 | } |
| 51 | |
| 52 | impl UnwindSafe for Dtls {} |
| 53 | impl RefUnwindSafe for Dtls {} |
| 54 | |
| 55 | impl Dtls { |
| 56 | /// Creates a new DTLS instance. |
| 57 | pub fn new( |
| 58 | cert: &DtlsCert, |
| 59 | dtls_provider: &dyn DtlsProvider, |
| 60 | sha256_provider: &dyn Sha256Provider, |
| 61 | now: Instant, |
| 62 | dtls_version: DtlsVersion, |
| 63 | mtu: RangeInclusive<usize>, |
| 64 | ) -> Result<Self, DtlsError> { |
| 65 | let instance = dtls_provider |
| 66 | .new_dtls(cert, now, dtls_version, Some(*mtu.start())) |
| 67 | .map_err(DtlsError::CryptoError)?; |
| 68 | |
| 69 | // Compute fingerprint from the certificate DER bytes |
| 70 | let fingerprint = Fingerprint { |
| 71 | hash_func: "sha-256".to_string(), |
| 72 | bytes: sha256_provider.sha256(&cert.certificate).to_vec(), |
| 73 | }; |
| 74 | |
| 75 | Ok(Self { |
| 76 | instance, |
| 77 | fingerprint, |
| 78 | remote_fingerprint: None, |
| 79 | active_state: None, |
| 80 | pending_packets: VecDeque::new(), |
| 81 | mtu, |
| 82 | }) |
| 83 | } |
| 84 | |
| 85 | /// Threshold above which an outgoing DTLS record triggers an MTU warning. |
| 86 | pub fn mtu_warn(&self) -> usize { |
| 87 | *self.mtu.end() |
| 88 | } |
| 89 | |
| 90 | /// Tells if this instance has been inited (set_active called). |
| 91 | pub fn is_inited(&self) -> bool { |
| 92 | self.active_state.is_some() |
| 93 | } |
| 94 | |
| 95 | /// Set whether this instance is active (client) or passive (server). |
| 96 | pub fn set_active(&mut self, active: bool) { |
| 97 | self.active_state = Some(active); |
| 98 | self.instance.set_active(active) |
| 99 | } |
| 100 | |
| 101 | /// If set_active was called, returns what was set. |
| 102 | pub fn is_active(&self) -> Option<bool> { |
| 103 | self.active_state |
| 104 | } |
| 105 | |
| 106 | /// The local certificate fingerprint. |
| 107 | pub fn local_fingerprint(&self) -> &Fingerprint { |
| 108 | &self.fingerprint |
| 109 | } |
| 110 | |
| 111 | /// Remote fingerprint, if received. |
| 112 | pub fn remote_fingerprint(&self) -> Option<&Fingerprint> { |
| 113 | self.remote_fingerprint.as_ref() |
| 114 | } |
| 115 | |
| 116 | /// The negotiated DTLS protocol version, or `None` before handshake completion. |
| 117 | pub fn protocol_version(&self) -> Option<ProtocolVersion> { |
| 118 | self.instance.protocol_version() |
| 119 | } |
| 120 | |
| 121 | pub fn is_closed(&self) -> bool { |
| 122 | self.pending_packets.is_empty() && self.instance.is_closed() |
| 123 | } |
| 124 | |
| 125 | /// Set the remote fingerprint (extracted from peer certificate). |
| 126 | pub fn set_remote_fingerprint(&mut self, fingerprint: Fingerprint) { |
| 127 | self.remote_fingerprint = Some(fingerprint); |
| 128 | } |
| 129 | |
| 130 | /// Poll for output from the DTLS instance. |
| 131 | pub fn poll_output<'a>(&mut self, buf: &'a mut [u8]) -> DtlsOutput<'a> { |
| 132 | let next = self.instance.poll_output(buf); |
| 133 | |
| 134 | if let DtlsOutput::Packet(packet) = next { |
| 135 | if packet.len() > self.mtu_warn() { |
| 136 | warn!("DTLS above MTU {}: {}", self.mtu_warn(), packet.len()); |
| 137 | } |
| 138 | self.pending_packets.push_back(packet.to_vec().into()); |
| 139 | |
| 140 | // Return timeout indicating we want another poll straight away |
| 141 | return DtlsOutput::Timeout(already_happened()); |
| 142 | } |
| 143 | |
| 144 | next |
| 145 | } |
| 146 | |
| 147 | pub fn poll_packet(&mut self) -> Option<DatagramSend> { |
| 148 | self.pending_packets.pop_front() |
| 149 | } |
| 150 | |
| 151 | /// Handle an incoming DTLS packet. |
| 152 | pub fn handle_receive(&mut self, packet: &[u8]) -> Result<(), DtlsError> { |
| 153 | if self.active_state.is_none() { |
| 154 | debug!("Ignoring DTLS datagram prior to DTLS start"); |
| 155 | return Ok(()); |
| 156 | } |
| 157 | |
| 158 | self.instance |
| 159 | .handle_packet(packet) |
| 160 | .map_err(|e| DtlsError::CryptoError(CryptoError::Other(format!("DTLS error: {}", e)))) |
| 161 | } |
| 162 | |
| 163 | /// Send application data over DTLS. |
| 164 | pub fn handle_input(&mut self, data: &[u8]) -> Result<(), DtlsError> { |
| 165 | self.instance.send_application_data(data).map_err(|e| { |
| 166 | if matches!(e, dimpl::Error::HandshakePending) { |
| 167 | DtlsError::Io(io::Error::new(io::ErrorKind::WouldBlock, e)) |
| 168 | } else { |
| 169 | DtlsError::CryptoError(CryptoError::Other(format!("DTLS error: {}", e))) |
| 170 | } |
| 171 | }) |
| 172 | } |
| 173 | |
| 174 | /// Handle a timeout event. |
| 175 | pub fn handle_timeout(&mut self, now: Instant) -> Result<(), DtlsError> { |
| 176 | self.instance |
| 177 | .handle_timeout(now) |
| 178 | .map_err(|e| DtlsError::CryptoError(CryptoError::Other(format!("DTLS error: {}", e)))) |
| 179 | } |
| 180 | |
| 181 | /// Initiate graceful shutdown by sending a close_notify alert. |
| 182 | pub fn close(&mut self) -> Result<(), DtlsError> { |
| 183 | self.instance |
| 184 | .close() |
| 185 | .map_err(|e| DtlsError::CryptoError(CryptoError::Other(format!("DTLS error: {}", e)))) |
| 186 | } |
| 187 | } |