Skip to content
File

Blob: firmware/vendor/str0m/src/dtls.rs

rust188 lines
1use std::collections::VecDeque;
2use std::io;
3use std::ops::RangeInclusive;
4use std::panic::{RefUnwindSafe, UnwindSafe};
5use std::time::Instant;
6 
7use crate::crypto::Fingerprint;
8use crate::crypto::Sha256Provider;
9use crate::crypto::dtls::{DtlsCert, DtlsOutput, ProtocolVersion};
10use crate::crypto::dtls::{DtlsInstance, DtlsProvider, DtlsVersion};
11use crate::crypto::{CryptoError, DtlsError};
12use crate::io::DatagramSend;
13use 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.
19pub 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 
41pub(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 
52impl UnwindSafe for Dtls {}
53impl RefUnwindSafe for Dtls {}
54 
55impl 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}