Skip to main content

stun_proto/
agent.rs

1// Copyright (C) 2020 Matthew Waters <matthew@centricular.com>
2//
3// Licensed under the Apache License, Version 2.0 <LICENSE-APACHE or
4// http://www.apache.org/licenses/LICENSE-2.0> or the MIT license
5// <LICENSE-MIT or http://opensource.org/licenses/MIT>, at your
6// option. This file may not be copied, modified, or distributed
7// except according to those terms.
8//
9// SPDX-License-Identifier: MIT OR Apache-2.0
10
11//! # STUN agent
12//!
13//! A STUN Agent that follows the procedures of [RFC5389] and [RFC8489] and is implemented with the
14//! sans-IO pattern. This agent does no IO processing and operates solely on inputs it is
15//! provided.
16//!
17//! [RFC8489]: https://tools.ietf.org/html/rfc8489
18//! [RFC5389]: https://tools.ietf.org/html/rfc5389
19
20use core::net::SocketAddr;
21use core::sync::atomic::{AtomicUsize, Ordering};
22
23use alloc::collections::{BTreeMap, BTreeSet};
24use alloc::vec;
25use alloc::vec::Vec;
26use core::time::Duration;
27
28use crate::stats::StunAgentStats;
29use crate::Instant;
30
31use stun_types::attribute::*;
32use stun_types::data::Data;
33use stun_types::message::*;
34
35use stun_types::TransportType;
36
37use tracing::{debug, trace};
38
39static STUN_AGENT_COUNT: AtomicUsize = AtomicUsize::new(0);
40
41/// Implementation of a STUN agent.
42#[derive(Debug)]
43pub struct StunAgent {
44    id: usize,
45    transport: TransportType,
46    local_addr: SocketAddr,
47    remote_addr: Option<SocketAddr>,
48    validated_peers: BTreeSet<SocketAddr>,
49    outstanding_requests: BTreeMap<TransactionId, StunRequestState>,
50    request_timeouts: Vec<Duration>,
51    last_retransmit_timeout: Duration,
52    stats: Option<StunAgentStats>,
53}
54
55/// Builder struct for a [`StunAgent`].
56#[derive(Debug)]
57pub struct StunAgentBuilder {
58    transport: TransportType,
59    local_addr: SocketAddr,
60    remote_addr: Option<SocketAddr>,
61    rto: RequestRto,
62    stats: bool,
63}
64
65impl StunAgentBuilder {
66    /// Set the remote address the [`StunAgent`] will be configured to only send data to.
67    pub fn remote_addr(mut self, addr: SocketAddr) -> Self {
68        self.remote_addr = Some(addr);
69        self
70    }
71
72    /// Configure the default timeouts and retransmissions for each STUN request.
73    ///
74    /// - `initial` - the initial time between consecutive transmissions. If 0, or 1, then only a
75    ///   single request will be performed.
76    /// - `max` - the maximum amount of time between consecutive retransmits.
77    /// - `retransmits` - the total number of transmissions of the request.
78    /// - `final_retransmit_timeout` - the amount of time after the final transmission to wait
79    ///   for a response before considering the request as having timed out.
80    ///
81    /// As specified in RFC 8489, `initial_rto` should be >= 500ms (unless specific information is
82    /// available on the RTT, `max` is `Duration::MAX`, `retransmits` has a default value of 7,
83    /// and `last_retransmit_timeout` should be `16 * initial_rto`.
84    ///
85    /// STUN transactions over TCP will only send a single request and have a timeout of the sum of
86    /// the timeouts of a UDP transaction.
87    pub fn request_retransmits(
88        mut self,
89        initial: Duration,
90        max: Duration,
91        retransmits: u32,
92        final_retransmit_timeout: Duration,
93    ) -> Self {
94        self.rto.initial = initial;
95        self.rto.max = max;
96        self.rto.retransmits = retransmits;
97        self.rto.last_retransmit = final_retransmit_timeout;
98        self
99    }
100
101    /// Enable statistics tracking on the built [`StunAgent`].
102    ///
103    /// When enabled, the agent tracks round-trip times, bytes sent and received,
104    /// transaction counts, and timeout/cancellation counts. Access the collected
105    /// statistics via [`StunAgent::stats`] or [`StunAgent::stats_mut`].
106    pub fn stats(mut self, stats: bool) -> Self {
107        self.stats = stats;
108        self
109    }
110
111    /// Build the [`StunAgent`].
112    pub fn build(self) -> StunAgent {
113        let id = STUN_AGENT_COUNT.fetch_add(1, Ordering::SeqCst);
114        let (request_timeouts, last_retransmit_timeout) =
115            self.rto.calculate_timeouts(self.transport);
116        let stats = if self.stats {
117            Some(StunAgentStats::default())
118        } else {
119            None
120        };
121        StunAgent {
122            id,
123            transport: self.transport,
124            local_addr: self.local_addr,
125            remote_addr: self.remote_addr,
126            validated_peers: Default::default(),
127            outstanding_requests: Default::default(),
128            request_timeouts,
129            last_retransmit_timeout,
130            stats,
131        }
132    }
133}
134
135impl StunAgent {
136    /// Create a new [`StunAgentBuilder`].
137    pub fn builder(transport: TransportType, local_addr: SocketAddr) -> StunAgentBuilder {
138        StunAgentBuilder {
139            transport,
140            local_addr,
141            remote_addr: None,
142            rto: Default::default(),
143            stats: false,
144        }
145    }
146
147    /// The [`TransportType`] of this [`StunAgent`].
148    pub fn transport(&self) -> TransportType {
149        self.transport
150    }
151
152    /// The local address of this [`StunAgent`].
153    pub fn local_addr(&self) -> SocketAddr {
154        self.local_addr
155    }
156
157    /// The remote address of this [`StunAgent`].
158    pub fn remote_addr(&self) -> Option<SocketAddr> {
159        self.remote_addr
160    }
161
162    /// Returns the statistics for this agent if stats collection is enabled.
163    pub fn stats(&self) -> Option<&StunAgentStats> {
164        self.stats.as_ref()
165    }
166
167    /// Returns the mutable statistics for this agent if stats collection is enabled.
168    pub fn stats_mut(&mut self) -> Option<&mut StunAgentStats> {
169        self.stats.as_mut()
170    }
171
172    /// Perform any operations needed to be able to send data to a peer.
173    pub fn send_data<T: AsRef<[u8]>>(&self, bytes: T, to: SocketAddr) -> Transmit<T> {
174        send_data(self.transport, bytes, self.local_addr, to)
175    }
176
177    /// Perform any operations needed to be able to send a [`Message`] to a peer.
178    ///
179    /// The returned [`Transmit`] must be sent to the respective peer after this call.
180    ///
181    /// # Panics
182    ///
183    /// - If the STUN Message is a request. Use [`send_request()`](StunAgent::send_request) instead.
184    #[tracing::instrument(name = "stun_agent_send",
185        skip(self, msg),
186        fields(
187            transport = %self.transport,
188            from = %self.local_addr,
189            transaction_id,
190        )
191    )]
192    pub fn send<T: AsRef<[u8]>>(
193        &mut self,
194        msg: T,
195        to: SocketAddr,
196        now: Instant,
197    ) -> Result<Transmit<T>, StunError> {
198        let data = msg.as_ref();
199        let hdr = MessageHeader::from_bytes(data)?;
200        tracing::Span::current().record(
201            "transaction_id",
202            tracing::field::display(hdr.transaction_id()),
203        );
204        let cls = hdr.get_type().class();
205        assert_ne!(cls, MessageClass::Request);
206        trace!("Sending {} to {to}", hdr.get_type());
207        if let Some(s) = &mut self.stats {
208            if cls == MessageClass::Indication {
209                s.record_indication_sent(data.len() as u64);
210            } else {
211                s.record_response_sent(data.len() as u64);
212            }
213        }
214        Ok(Transmit::new(msg, self.transport, self.local_addr, to))
215    }
216
217    /// Perform any operations needed to be able to send a request [`Message`] to a peer.
218    ///
219    /// The returned [`Transmit`] must be sent to the respective peer after this call.
220    ///
221    /// # Panics
222    ///
223    /// - If the STUN Message is not a request. Use [`send()`](StunAgent::send) instead.
224    #[tracing::instrument(name = "stun_agent_send_request",
225        skip(self, msg),
226        fields(
227            transport = %self.transport,
228            from = %self.local_addr,
229            transaction_id,
230        )
231    )]
232    pub fn send_request<'a, T: AsRef<[u8]>>(
233        &'a mut self,
234        msg: T,
235        to: SocketAddr,
236        now: Instant,
237    ) -> Result<Transmit<Data<'a>>, StunError> {
238        let data = msg.as_ref();
239        let hdr = MessageHeader::from_bytes(data)?;
240        assert!(hdr.get_type().has_class(MessageClass::Request));
241        let transaction_id = hdr.transaction_id();
242        tracing::Span::current().record("transaction_id", tracing::field::display(transaction_id));
243        let state = match self.outstanding_requests.entry(transaction_id) {
244            alloc::collections::btree_map::Entry::Vacant(entry) => {
245                let integrity_algorithm = MessageAttributesIter::new(data)
246                    .filter_map(|(_offset, attr)| match attr.get_type() {
247                        MessageIntegrity::TYPE => Some(IntegrityAlgorithm::Sha1),
248                        MessageIntegritySha256::TYPE => Some(IntegrityAlgorithm::Sha256),
249                        _ => None,
250                    })
251                    .last();
252                trace!("Adding request to {to} with integrity algorithm: {integrity_algorithm:?}");
253                if let Some(s) = &mut self.stats {
254                    s.record_request_sent(data.len() as u64);
255                }
256                entry.insert(StunRequestState::new(
257                    msg,
258                    self.transport,
259                    self.local_addr,
260                    to,
261                    transaction_id,
262                    integrity_algorithm,
263                    self.request_timeouts.clone(),
264                    self.last_retransmit_timeout,
265                ))
266            }
267            alloc::collections::btree_map::Entry::Occupied(_entry) => {
268                return Err(StunError::AlreadyInProgress);
269            }
270        };
271        let Some(transmit) = state.poll_transmit(now) else {
272            unreachable!();
273        };
274        Ok(Transmit::new(
275            Data::from(transmit.data),
276            transmit.transport,
277            transmit.from,
278            transmit.to,
279        ))
280    }
281
282    /// Returns whether this agent has received or sent a STUN message with this peer. Failure may
283    /// be the result of an attacker and the caller must drop any non-STUN data received before this
284    /// functions returns `true`.
285    ///
286    /// If non-STUN data is received over a TCP connection from an unvalidated peer, the caller
287    /// must immediately close the TCP connection.
288    pub fn is_validated_peer(&self, remote_addr: SocketAddr) -> bool {
289        self.validated_peers.contains(&remote_addr)
290    }
291
292    /// Indicate to the STUN agent that STUN messages have been sent/received to/from a peer.
293    #[tracing::instrument(
294        name = "stun_validated_peer"
295        skip(self),
296        fields(stun_id = self.id)
297    )]
298    pub fn validated_peer(&mut self, addr: SocketAddr) {
299        if !self.validated_peers.contains(&addr) {
300            debug!("validated peer {:?}", addr);
301            self.validated_peers.insert(addr);
302        }
303    }
304
305    /// Provide data received on a socket from a peer for handling by the [`StunAgent`] after it
306    /// has successfully passed authentication.
307    ///
308    /// For responses, this will cause the associated request to be removed from the agent if it
309    /// exists.
310    ///
311    /// The return value indicates whether the message passes internal checks and should be acted
312    /// upon.
313    ///
314    /// This function is provided for backwards compatibility and new code should use
315    /// [`StunAgent::handle_stun_message_with_time`] in order to compute round trip statistics.
316    #[tracing::instrument(
317        name = "stun_handle_message"
318        skip(self, msg, from),
319        fields(
320            transaction_id = %msg.transaction_id(),
321        )
322    )]
323    #[deprecated = "Use handle_stun_message_with_time() to be able to retrieve round trip statistics"]
324    // FIXME: 3.0: remove
325    pub fn handle_stun_message(&mut self, msg: &Message<'_>, from: SocketAddr) -> bool {
326        self.handle_stun_message_internal(msg, from, None)
327    }
328
329    /// Provide data received on a socket from a peer for handling by the [`StunAgent`] after it
330    /// has successfully passed authentication.
331    ///
332    /// For responses, this will cause the associated request to be removed from the agent if it
333    /// exists. If statistics are enabled, then `time` is used to compute round trip statistics.
334    ///
335    /// The return value indicates whether the message passes internal checks and should be acted
336    /// upon.
337    pub fn handle_stun_message_with_time(
338        &mut self,
339        msg: &Message<'_>,
340        from: SocketAddr,
341        now: Instant,
342    ) -> bool {
343        self.handle_stun_message_internal(msg, from, Some(now))
344    }
345
346    fn handle_stun_message_internal(
347        &mut self,
348        msg: &Message<'_>,
349        from: SocketAddr,
350        now: Option<Instant>,
351    ) -> bool {
352        let outstanding = if msg.is_response() {
353            self.take_outstanding_request(&msg.transaction_id())
354        } else {
355            None
356        };
357        if msg.is_response() && outstanding.is_none() {
358            trace!("original request disappeared");
359            return false;
360        }
361        if let Some(s) = &mut self.stats {
362            let msg_len = msg.as_bytes().len() as u64;
363            match msg.class() {
364                MessageClass::Request => s.record_request_received(msg_len),
365                MessageClass::Indication => s.record_indication_received(msg_len),
366                MessageClass::Error | MessageClass::Success => {
367                    if let Some(state) = &outstanding {
368                        if let Some(rtt) = now
369                            .zip(
370                                state
371                                    .last_send_time
372                                    // Only requests that did not require a retransmit count towards rtt.
373                                    .filter(|_| state.timeout_i == 0)
374                                    .zip(state.first_send_time),
375                            )
376                            .map(|(now, (last, _first))| now - last)
377                        {
378                            s.record_rtt(rtt);
379                        }
380                    }
381                    s.record_response_received(msg_len);
382                }
383            }
384        }
385        self.validated_peer(from);
386        true
387    }
388
389    #[tracing::instrument(
390        skip(self, transaction_id),
391        fields(transaction_id = %transaction_id)
392    )]
393    fn take_outstanding_request(
394        &mut self,
395        transaction_id: &TransactionId,
396    ) -> Option<StunRequestState> {
397        if let Some(request) = self.outstanding_requests.remove(transaction_id) {
398            trace!("removing request");
399            Some(request)
400        } else {
401            trace!("no outstanding request");
402            None
403        }
404    }
405
406    /// Retrieve a reference to an outstanding STUN request. Outstanding requests are kept until
407    /// either:
408    /// - [`handle_stun_message()`](StunAgent::handle_stun_message) is called, or
409    /// - [`poll()`](StunAgent::poll) returns [`StunAgentPollRet::TransactionCancelled`] or
410    ///   [`StunAgentPollRet::TransactionTimedOut`] for the request.
411    pub fn request_transaction(&self, transaction_id: TransactionId) -> Option<StunRequest<'_>> {
412        self.request_state(transaction_id)
413            .map(|request| StunRequest {
414                agent: self,
415                peer_address: request.to,
416                request_integrity: request.request_integrity,
417            })
418    }
419
420    /// Retrieve a mutable reference to an outstanding STUN request. Outstanding requests are kept
421    /// until either:
422    /// - [`handle_stun_message()`](StunAgent::handle_stun_message) is called, or
423    /// - [`poll()`](StunAgent::poll) returns [`StunAgentPollRet::TransactionCancelled`] or
424    ///   [`StunAgentPollRet::TransactionTimedOut`] for the request.
425    pub fn mut_request_transaction(
426        &mut self,
427        transaction_id: TransactionId,
428    ) -> Option<StunRequestMut<'_>> {
429        if let Some(request) = self.mut_request_state(transaction_id) {
430            let peer_address = request.to;
431            let request_integrity = request.request_integrity;
432            Some(StunRequestMut {
433                agent: self,
434                transaction_id,
435                peer_address,
436                request_integrity,
437            })
438        } else {
439            None
440        }
441    }
442
443    fn mut_request_state(
444        &mut self,
445        transaction_id: TransactionId,
446    ) -> Option<&mut StunRequestState> {
447        self.outstanding_requests.get_mut(&transaction_id)
448    }
449
450    fn request_state(&self, transaction_id: TransactionId) -> Option<&StunRequestState> {
451        self.outstanding_requests.get(&transaction_id)
452    }
453
454    /// Poll the agent for making further progress on any outstanding requests. The returned value
455    /// indicates the current state and anything the caller needs to perform.
456    ///
457    /// Upon expiry of the timer from [`StunAgentPollRet::WaitUntil`],
458    /// [`poll_transmit()`](StunAgent::poll_transmit) must be called.
459    #[tracing::instrument(
460        name = "stun_agent_poll"
461        level = "debug",
462        skip(self),
463    )]
464    pub fn poll(&mut self, now: Instant) -> StunAgentPollRet {
465        let mut lowest_wait = now + Duration::from_secs(3600);
466        let mut timeout = None;
467        let mut cancelled = None;
468        for (transaction_id, request) in self.outstanding_requests.iter_mut() {
469            debug_assert_eq!(transaction_id, &request.transaction_id);
470            match request.poll(now) {
471                StunRequestPollRet::Cancelled => {
472                    cancelled = Some(*transaction_id);
473                    break;
474                }
475                StunRequestPollRet::WaitUntil(wait_until) => {
476                    if wait_until < lowest_wait {
477                        lowest_wait = wait_until;
478                    }
479                }
480                StunRequestPollRet::TimedOut => {
481                    timeout = Some(*transaction_id);
482                    break;
483                }
484            }
485        }
486        if let Some(transaction) = timeout {
487            if let Some(_state) = self.outstanding_requests.remove(&transaction) {
488                if let Some(s) = &mut self.stats {
489                    s.record_timeout();
490                }
491                return StunAgentPollRet::TransactionTimedOut(transaction);
492            }
493        }
494        if let Some(transaction) = cancelled {
495            if let Some(_state) = self.outstanding_requests.remove(&transaction) {
496                if let Some(s) = &mut self.stats {
497                    s.record_cancelled();
498                }
499                return StunAgentPollRet::TransactionCancelled(transaction);
500            }
501        }
502        StunAgentPollRet::WaitUntil(lowest_wait)
503    }
504
505    /// Poll for any transmissions that may need to be performed.
506    #[tracing::instrument(
507        name = "stun_agent_poll_transmit"
508        level = "debug",
509        skip(self),
510    )]
511    pub fn poll_transmit(&mut self, now: Instant) -> Option<Transmit<&[u8]>> {
512        let transmit = self
513            .outstanding_requests
514            .values_mut()
515            .filter_map(|request| request.poll_transmit(now))
516            .next();
517        if let Some(t) = &transmit {
518            if let Some(s) = &mut self.stats {
519                s.record_retransmit_bytes(t.data.as_ref().len() as u64);
520            }
521        }
522        transmit
523    }
524}
525
526/// Return value for [`StunAgent::poll`].
527#[derive(Debug)]
528pub enum StunAgentPollRet {
529    /// An oustanding transaction timed out and has been removed from the agent.
530    TransactionTimedOut(TransactionId),
531    /// An oustanding transaction was cancelled and has been removed from the agent.
532    TransactionCancelled(TransactionId),
533    /// Wait until the specified time has passed.
534    WaitUntil(Instant),
535}
536
537fn send_data<T: AsRef<[u8]>>(
538    transport: TransportType,
539    bytes: T,
540    from: SocketAddr,
541    to: SocketAddr,
542) -> Transmit<T> {
543    Transmit::new(bytes, transport, from, to)
544}
545
546/// A piece of data that needs to, or has been transmitted.
547#[derive(Debug)]
548pub struct Transmit<T: AsRef<[u8]>> {
549    /// The data blob.
550    pub data: T,
551    /// The transport for the transmission.
552    pub transport: TransportType,
553    /// The source address of the transmission.
554    pub from: SocketAddr,
555    /// The destination address of the transmission.
556    pub to: SocketAddr,
557}
558
559impl<T: AsRef<[u8]>> core::fmt::Display for Transmit<T> {
560    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
561        write!(
562            f,
563            "Transmit({}: {} -> {} of {} bytes)",
564            self.transport,
565            self.from,
566            self.to,
567            self.data.as_ref().len()
568        )
569    }
570}
571
572impl<T: AsRef<[u8]>> Transmit<T> {
573    /// Construct a new [`Transmit`] with the specifid data and 5-tuple.
574    pub fn new(data: T, transport: TransportType, from: SocketAddr, to: SocketAddr) -> Self {
575        Self {
576            data,
577            transport,
578            from,
579            to,
580        }
581    }
582
583    /// Reinterpret the data of a [`Transmit`] into a different type.
584    ///
585    /// # Examples
586    ///
587    /// ```
588    /// # use stun_proto::agent::Transmit;
589    /// # use stun_proto::types::TransportType;
590    /// # use core::net::SocketAddr;
591    /// let local_addr = "10.0.0.1:1000".parse().unwrap();
592    /// let remote_addr = "10.0.0.2:2000".parse().unwrap();
593    /// let slice = [42; 8];
594    /// let transmit = Transmit::new(slice.clone(), TransportType::Udp, local_addr, remote_addr);
595    /// // change the data type of the `Transmit` into a `Vec<u8>`.
596    /// let transmit = transmit.reinterpret_data(|data| data.to_vec());
597    /// # assert_eq!(transmit.transport, TransportType::Udp);
598    /// # assert_eq!(transmit.from, local_addr);
599    /// # assert_eq!(transmit.to, remote_addr);
600    /// assert_eq!(transmit.data, slice.as_slice());
601    /// ```
602    pub fn reinterpret_data<O: AsRef<[u8]>, F: FnOnce(T) -> O>(self, f: F) -> Transmit<O> {
603        Transmit {
604            data: f(self.data),
605            transport: self.transport,
606            from: self.from,
607            to: self.to,
608        }
609    }
610}
611
612impl Transmit<Data<'_>> {
613    /// Construct a new owned [`Transmit`] from a provided [`Transmit`].
614    pub fn into_owned<'b>(self) -> Transmit<Data<'b>> {
615        self.reinterpret_data(|data| data.into_owned())
616    }
617}
618
619/// Return value for [`StunRequest::poll`].
620#[derive(Debug)]
621enum StunRequestPollRet {
622    /// Wait until the specified time has passed.
623    WaitUntil(Instant),
624    /// The request has been cancelled and will not make further progress.
625    Cancelled,
626    /// The request timed out.
627    TimedOut,
628}
629
630#[derive(Debug)]
631struct RequestRto {
632    initial: Duration,
633    max: Duration,
634    retransmits: u32,
635    last_retransmit: Duration,
636}
637
638impl Default for RequestRto {
639    fn default() -> Self {
640        Self {
641            initial: Duration::from_millis(500),
642            max: Duration::MAX,
643            retransmits: 7,
644            last_retransmit: Duration::from_millis(8),
645        }
646    }
647}
648
649impl RequestRto {
650    fn calculate_timeouts(&self, transport: TransportType) -> (Vec<Duration>, Duration) {
651        match transport {
652            TransportType::Udp => {
653                let timeouts = (0..self.retransmits.max(1) - 1)
654                    .map(|i| (self.initial * 2u32.pow(i)).min(self.max))
655                    .collect::<Vec<_>>();
656                (timeouts, self.last_retransmit)
657            }
658            TransportType::Tcp => {
659                let timeouts = vec![];
660                let last_retransmit_timeout = self.last_retransmit
661                    + (0..self.retransmits.max(1) - 1).fold(Duration::ZERO, |acc, i| {
662                        acc + (self.initial * 2u32.pow(i)).min(self.max)
663                    });
664                (timeouts, last_retransmit_timeout)
665            }
666        }
667    }
668}
669
670#[derive(Debug)]
671struct StunRequestState {
672    transaction_id: TransactionId,
673    request_integrity: Option<IntegrityAlgorithm>,
674    bytes: Vec<u8>,
675    transport: TransportType,
676    from: SocketAddr,
677    to: SocketAddr,
678    timeouts: Vec<Duration>,
679    last_retransmit_timeout: Duration,
680    recv_cancelled: bool,
681    send_cancelled: bool,
682    timeout_i: usize,
683    first_send_time: Option<Instant>,
684    last_send_time: Option<Instant>,
685}
686
687impl StunRequestState {
688    #[allow(clippy::too_many_arguments)]
689    fn new<T: AsRef<[u8]>>(
690        request: T,
691        transport: TransportType,
692        from: SocketAddr,
693        to: SocketAddr,
694        transaction_id: TransactionId,
695        integrity_algorithm: Option<IntegrityAlgorithm>,
696        timeouts: Vec<Duration>,
697        last_retransmit_timeout: Duration,
698    ) -> Self {
699        let data = request.as_ref();
700        Self {
701            transaction_id,
702            bytes: data.to_vec(),
703            transport,
704            from,
705            to,
706            request_integrity: integrity_algorithm,
707            timeouts,
708            timeout_i: 0,
709            last_retransmit_timeout,
710            recv_cancelled: false,
711            send_cancelled: false,
712            first_send_time: None,
713            last_send_time: None,
714        }
715    }
716
717    #[tracing::instrument(skip(self, now), level = "trace")]
718    fn next_send_time(&self, now: Instant) -> Option<Instant> {
719        let Some(last_send) = self.last_send_time else {
720            trace!("not sent yet -> send immediately");
721            return Some(now);
722        };
723        if self.timeout_i >= self.timeouts.len() {
724            let next_send = last_send + self.last_retransmit_timeout;
725            trace!("final retransmission, final timeout ends at {next_send:?}");
726            if next_send > now {
727                return Some(next_send);
728            }
729            return None;
730        }
731        let next_send = last_send + self.timeouts[self.timeout_i];
732        Some(next_send)
733    }
734
735    #[tracing::instrument(
736        name = "stun_request_poll"
737        level = "debug",
738        ret,
739        skip(self, now),
740        fields(transaction_id = %self.transaction_id),
741    )]
742    fn poll(&mut self, now: Instant) -> StunRequestPollRet {
743        if self.recv_cancelled {
744            return StunRequestPollRet::Cancelled;
745        }
746        // TODO: account for TCP connect in timeout
747        let Some(next_send) = self.next_send_time(now) else {
748            return StunRequestPollRet::TimedOut;
749        };
750        if next_send >= now {
751            if self.send_cancelled && self.timeout_i >= self.timeouts.len() {
752                // this cancellation may need a different value
753                return StunRequestPollRet::Cancelled;
754            }
755            return StunRequestPollRet::WaitUntil(next_send);
756        }
757        StunRequestPollRet::WaitUntil(now)
758    }
759
760    #[tracing::instrument(
761        name = "stun_request_poll_transmit",
762        skip(self, now),
763        fields(transaction_id = %self.transaction_id)
764    )]
765    fn poll_transmit(&mut self, now: Instant) -> Option<Transmit<&[u8]>> {
766        if self.recv_cancelled {
767            return None;
768        };
769        let next_send = self.next_send_time(now)?;
770
771        if next_send > now {
772            return None;
773        }
774        if self.last_send_time.is_some() {
775            self.timeout_i += 1;
776        }
777        if self.first_send_time.is_none() {
778            self.first_send_time = Some(now);
779        }
780        self.last_send_time = Some(now);
781        if self.send_cancelled {
782            return None;
783        };
784        trace!(
785            "sending {} bytes over {:?} from {:?} to {:?}",
786            self.bytes.len(),
787            self.transport,
788            self.from,
789            self.to
790        );
791        Some(send_data(
792            self.transport,
793            self.bytes.as_slice(),
794            self.from,
795            self.to,
796        ))
797    }
798}
799
800/// A STUN Request.
801#[derive(Debug, Clone)]
802pub struct StunRequest<'a> {
803    agent: &'a StunAgent,
804    peer_address: SocketAddr,
805    request_integrity: Option<IntegrityAlgorithm>,
806}
807
808impl StunRequest<'_> {
809    /// The remote address the request is sent to.
810    pub fn peer_address(&self) -> SocketAddr {
811        self.peer_address
812    }
813
814    /// The integrity algorithm present on the request.
815    pub fn integrity(&self) -> Option<IntegrityAlgorithm> {
816        self.request_integrity
817    }
818
819    /// The [`StunAgent`] this request is being sent with.
820    pub fn agent(&self) -> &StunAgent {
821        self.agent
822    }
823}
824
825/// A STUN Request.
826#[derive(Debug)]
827pub struct StunRequestMut<'a> {
828    agent: &'a mut StunAgent,
829    transaction_id: TransactionId,
830    peer_address: SocketAddr,
831    request_integrity: Option<IntegrityAlgorithm>,
832}
833
834impl StunRequestMut<'_> {
835    /// The remote address the request is sent to.
836    pub fn peer_address(&self) -> SocketAddr {
837        self.peer_address
838    }
839
840    /// The integrity algorithm present on the request.
841    pub fn integrity(&self) -> Option<IntegrityAlgorithm> {
842        self.request_integrity
843    }
844
845    /// Do not retransmit further. This will still allow for a reply to occur within the configured
846    /// timeouts, but will never send a retransmission. If no response is received, this will cause
847    /// [`StunAgent::poll()`] to return [`StunAgentPollRet::TransactionCancelled`] for this request.
848    pub fn cancel_retransmissions(&mut self) {
849        if let Some(state) = self.agent.mut_request_state(self.transaction_id) {
850            state.send_cancelled = true;
851        }
852    }
853
854    /// Do not wait for any kind of response. This will cause [`StunAgent::poll()`] to return
855    /// [`StunAgentPollRet::TransactionCancelled`] for this request.
856    pub fn cancel(&mut self) {
857        if let Some(state) = self.agent.mut_request_state(self.transaction_id) {
858            state.send_cancelled = true;
859            state.recv_cancelled = true;
860        }
861    }
862
863    /// The [`StunAgent`] this request is being sent with.
864    pub fn agent(&self) -> &StunAgent {
865        self.agent
866    }
867
868    /// The mutable [`StunAgent`] this request is being sent with.
869    pub fn mut_agent(&mut self) -> &mut StunAgent {
870        self.agent
871    }
872
873    /// Configure the timeouts and retransmissions for the STUN request.
874    ///
875    /// This the same as calling `[configure_timeout_with_max`] with a `max` of `Duration::MAX`.
876    pub fn configure_timeout(
877        &mut self,
878        initial_rto: Duration,
879        retransmits: u32,
880        last_retransmit_timeout: Duration,
881    ) {
882        self.configure_timeout_with_max(
883            initial_rto,
884            retransmits,
885            last_retransmit_timeout,
886            Duration::MAX,
887        );
888    }
889
890    /// Configure the timeouts and retransmissions for the STUN request.
891    ///
892    /// - `initial` - the initial time between consecutive transmissions. If 0, or 1, then only a
893    ///   single request will be performed.
894    /// - `max` - the maximum amount of time between consecutive retransmits.
895    /// - `retransmits` - the total number of transmissions of the request.
896    /// - `final_retransmit_timeout` - the amount of time after the final transmission to wait
897    ///   for a response before considering the request as having timed out.
898    ///
899    /// As specified in RFC 8489, `initial_rto` should be >= 500ms (unless specific information is
900    /// available on the RTT, `max` is `Duration::MAX`, `retransmits` has a default value of 7,
901    /// and `last_retransmit_timeout` should be `16 * initial_rto`.
902    ///
903    /// STUN transactions over TCP will only send a single request and have a timeout of the sum of
904    /// the timeouts of a UDP transaction.
905    pub fn configure_timeout_with_max(
906        &mut self,
907        initial_rto: Duration,
908        retransmits: u32,
909        last_retransmit_timeout: Duration,
910        max_rto: Duration,
911    ) {
912        if let Some(state) = self.agent.mut_request_state(self.transaction_id) {
913            let (timeouts, final_wait) = RequestRto {
914                initial: initial_rto,
915                max: max_rto,
916                retransmits,
917                last_retransmit: last_retransmit_timeout,
918            }
919            .calculate_timeouts(state.transport);
920            state.timeouts = timeouts;
921            state.last_retransmit_timeout = final_wait;
922        }
923    }
924}
925
926/// STUN errors.
927#[derive(Debug, thiserror::Error)]
928#[non_exhaustive]
929pub enum StunError {
930    /// The operation is already in progress.
931    #[error("The operation is already in progress")]
932    AlreadyInProgress,
933    /// A resource was not found.
934    #[error("A required resource could not be found")]
935    ResourceNotFound,
936    /// An operation timed out without a response.
937    #[error("An operation timed out")]
938    TimedOut,
939    /// Unexpected data was received or an operation is not allowed at this time.
940    #[error("Unexpected data was received")]
941    ProtocolViolation,
942    /// An operation was cancelled.
943    #[error("Operation was aborted")]
944    Aborted,
945    /// A parsing error. The contained error contains more details.
946    #[error("{}", .0)]
947    ParseError(StunParseError),
948    /// A writing error. The contained error contains more details.
949    #[error("{}", .0)]
950    WriteError(StunWriteError),
951}
952
953impl From<StunParseError> for StunError {
954    fn from(e: StunParseError) -> Self {
955        StunError::ParseError(e)
956    }
957}
958
959impl From<StunWriteError> for StunError {
960    fn from(e: StunWriteError) -> Self {
961        StunError::WriteError(e)
962    }
963}
964
965#[cfg(test)]
966pub(crate) mod tests {
967    use alloc::string::String;
968    use tracing::error;
969
970    use crate::auth::ShortTermAuth;
971
972    use super::*;
973
974    #[test]
975    fn agent_getters_setters() {
976        let _log = crate::tests::test_init_log();
977        let local_addr = "10.0.0.1:12345".parse().unwrap();
978        let remote_addr = "10.0.0.2:3478".parse().unwrap();
979        let agent = StunAgent::builder(TransportType::Udp, local_addr)
980            .remote_addr(remote_addr)
981            .build();
982
983        assert_eq!(agent.transport(), TransportType::Udp);
984        assert_eq!(agent.local_addr(), local_addr);
985        assert_eq!(agent.remote_addr(), Some(remote_addr));
986    }
987
988    #[test]
989    fn request() {
990        let _log = crate::tests::test_init_log();
991        let local_addr = "127.0.0.1:2000".parse().unwrap();
992        let remote_addr = "127.0.0.1:1000".parse().unwrap();
993        let mut agent = StunAgent::builder(TransportType::Udp, local_addr)
994            .remote_addr(remote_addr)
995            .build();
996        let now = Instant::ZERO;
997
998        let msg = Message::builder_request(BINDING, MessageWriteVec::new());
999        let transaction_id = msg.transaction_id();
1000        let transmit = agent
1001            .send_request(msg.finish(), remote_addr, now)
1002            .unwrap()
1003            .into_owned();
1004        let request = agent.request_transaction(transaction_id).unwrap();
1005        assert!(request.integrity().is_none());
1006        assert_eq!(transmit.transport, TransportType::Udp);
1007        assert_eq!(transmit.from, local_addr);
1008        assert_eq!(transmit.to, remote_addr);
1009        let request = Message::from_bytes(&transmit.data).unwrap();
1010        let response = Message::builder_error(&request, MessageWriteVec::new());
1011        let resp_data = response.finish();
1012        let response = Message::from_bytes(&resp_data).unwrap();
1013        assert!(agent.handle_stun_message_with_time(&response, remote_addr, now));
1014        assert!(agent.request_transaction(transaction_id).is_none());
1015        assert!(agent.mut_request_transaction(transaction_id).is_none());
1016
1017        let ret = agent.poll(now);
1018        assert!(matches!(ret, StunAgentPollRet::WaitUntil(_)));
1019    }
1020
1021    #[test]
1022    fn indication_with_invalid_response() {
1023        let _log = crate::tests::test_init_log();
1024        let local_addr = "127.0.0.1:2000".parse().unwrap();
1025        let remote_addr = "127.0.0.1:1000".parse().unwrap();
1026        let now = Instant::ZERO;
1027
1028        let mut agent = StunAgent::builder(TransportType::Udp, local_addr)
1029            .remote_addr(remote_addr)
1030            .build();
1031        let transaction_id = TransactionId::generate();
1032        let msg = Message::builder(
1033            MessageType::from_class_method(MessageClass::Indication, BINDING),
1034            transaction_id,
1035            MessageWriteVec::new(),
1036        );
1037        let transmit = agent
1038            .send(msg.finish(), remote_addr, Instant::ZERO)
1039            .unwrap();
1040        assert_eq!(transmit.transport, TransportType::Udp);
1041        assert_eq!(transmit.from, local_addr);
1042        assert_eq!(transmit.to, remote_addr);
1043        let _indication = Message::from_bytes(&transmit.data).unwrap();
1044        assert!(agent.request_transaction(transaction_id).is_none());
1045        assert!(agent.mut_request_transaction(transaction_id).is_none());
1046        // you should definitely never do this ;). Indications should never get replies.
1047        let response = Message::builder(
1048            MessageType::from_class_method(MessageClass::Error, BINDING),
1049            transaction_id,
1050            MessageWriteVec::new(),
1051        );
1052        let resp_data = response.finish();
1053        let response = Message::from_bytes(&resp_data).unwrap();
1054        // response without a request is dropped.
1055        assert!(!agent.handle_stun_message_with_time(&response, remote_addr, now))
1056    }
1057
1058    #[test]
1059    fn request_with_credentials() {
1060        let _log = crate::tests::test_init_log();
1061        let local_addr = "10.0.0.1:12345".parse().unwrap();
1062        let remote_addr = "10.0.0.2:3478".parse().unwrap();
1063        let now = Instant::ZERO;
1064
1065        let mut auth = ShortTermAuth::new();
1066        let mut agent = StunAgent::builder(TransportType::Udp, local_addr).build();
1067        let credentials = ShortTermCredentials::new(String::from("local_password"));
1068        auth.set_credentials(credentials.clone(), IntegrityAlgorithm::Sha1);
1069
1070        // unvalidated peer data should be dropped
1071        assert!(!agent.is_validated_peer(remote_addr));
1072
1073        let mut msg = Message::builder_request(BINDING, MessageWriteVec::new());
1074        let transaction_id = msg.transaction_id();
1075        msg.add_message_integrity(&credentials.clone().into(), IntegrityAlgorithm::Sha1)
1076            .unwrap();
1077        error!("send");
1078        let transmit = agent
1079            .send_request(msg.finish(), remote_addr, Instant::ZERO)
1080            .unwrap();
1081        error!("sent");
1082
1083        let request = Message::from_bytes(&transmit.data).unwrap();
1084
1085        error!("generate response");
1086        let mut response = Message::builder_success(&request, MessageWriteVec::new());
1087        let xor_addr = XorMappedAddress::new(transmit.from, request.transaction_id());
1088        response.add_attribute(&xor_addr).unwrap();
1089        response
1090            .add_message_integrity(&credentials.into(), IntegrityAlgorithm::Sha1)
1091            .unwrap();
1092        error!("{response:?}");
1093
1094        let data = response.finish();
1095        error!("{data:?}");
1096        let response = Message::from_bytes(&data).unwrap();
1097        error!("{response}");
1098        assert_eq!(
1099            auth.validate_incoming_message(&response).unwrap(),
1100            Some(IntegrityAlgorithm::Sha1)
1101        );
1102        let request = agent
1103            .request_transaction(response.transaction_id())
1104            .unwrap();
1105        assert_eq!(request.integrity(), Some(IntegrityAlgorithm::Sha1));
1106        assert!(agent.handle_stun_message_with_time(&response, remote_addr, now));
1107
1108        assert_eq!(response.transaction_id(), transaction_id);
1109        assert!(agent.request_transaction(transaction_id).is_none());
1110        assert!(agent.mut_request_transaction(transaction_id).is_none());
1111        assert!(agent.is_validated_peer(remote_addr));
1112    }
1113
1114    #[test]
1115    fn request_unanswered() {
1116        let _log = crate::tests::test_init_log();
1117        let local_addr = "127.0.0.1:2000".parse().unwrap();
1118        let remote_addr = "127.0.0.1:1000".parse().unwrap();
1119        let mut agent = StunAgent::builder(TransportType::Udp, local_addr)
1120            .remote_addr(remote_addr)
1121            .build();
1122        let msg = Message::builder_request(BINDING, MessageWriteVec::new());
1123        let transaction_id = msg.transaction_id();
1124        agent
1125            .send_request(msg.finish(), remote_addr, Instant::ZERO)
1126            .unwrap();
1127        let mut now = Instant::ZERO;
1128        loop {
1129            let _ = agent.poll_transmit(now);
1130            match agent.poll(now) {
1131                StunAgentPollRet::WaitUntil(new_now) => {
1132                    now = new_now;
1133                }
1134                StunAgentPollRet::TransactionTimedOut(_) => break,
1135                _ => unreachable!(),
1136            }
1137        }
1138        assert!(agent.request_transaction(transaction_id).is_none());
1139        assert!(agent.mut_request_transaction(transaction_id).is_none());
1140
1141        // unvalidated peer data should be dropped
1142        assert!(!agent.is_validated_peer(remote_addr));
1143    }
1144
1145    #[test]
1146    fn request_custom_timeout() {
1147        let _log = crate::tests::test_init_log();
1148        let local_addr = "127.0.0.1:2000".parse().unwrap();
1149        let remote_addr = "127.0.0.1:1000".parse().unwrap();
1150        let mut agent = StunAgent::builder(TransportType::Udp, local_addr)
1151            .remote_addr(remote_addr)
1152            .build();
1153        let msg = Message::builder_request(BINDING, MessageWriteVec::new());
1154        let transaction_id = msg.transaction_id();
1155        let mut now = Instant::ZERO;
1156        agent.send_request(msg.finish(), remote_addr, now).unwrap();
1157        let mut transaction = agent.mut_request_transaction(transaction_id).unwrap();
1158        transaction.configure_timeout_with_max(
1159            Duration::from_secs(1),
1160            4,
1161            Duration::from_secs(10),
1162            Duration::from_secs(2),
1163        );
1164        let StunAgentPollRet::WaitUntil(wait) = agent.poll(now) else {
1165            unreachable!();
1166        };
1167        assert_eq!(wait - now, Duration::from_secs(1));
1168        now = wait;
1169        // a poll with the same instant should not busy loop
1170        let StunAgentPollRet::WaitUntil(wait) = agent.poll(now) else {
1171            unreachable!();
1172        };
1173        assert_eq!(wait, now);
1174        let Some(_) = agent.poll_transmit(now) else {
1175            unreachable!();
1176        };
1177        let StunAgentPollRet::WaitUntil(wait) = agent.poll(now) else {
1178            unreachable!();
1179        };
1180        assert_eq!(wait - now, Duration::from_secs(2));
1181        now = wait;
1182        let Some(_) = agent.poll_transmit(now) else {
1183            unreachable!();
1184        };
1185        let StunAgentPollRet::WaitUntil(wait) = agent.poll(now) else {
1186            unreachable!();
1187        };
1188        assert_eq!(wait - now, Duration::from_secs(2));
1189        now = wait;
1190        let Some(_) = agent.poll_transmit(now) else {
1191            unreachable!();
1192        };
1193        let StunAgentPollRet::WaitUntil(wait) = agent.poll(now) else {
1194            unreachable!();
1195        };
1196        assert_eq!(wait - now, Duration::from_secs(10));
1197        now = wait;
1198        let StunAgentPollRet::TransactionTimedOut(timed_out) = agent.poll(now) else {
1199            unreachable!();
1200        };
1201        assert_eq!(timed_out, transaction_id);
1202
1203        assert!(agent.request_transaction(transaction_id).is_none());
1204        assert!(agent.mut_request_transaction(transaction_id).is_none());
1205
1206        // unvalidated peer data should be dropped
1207        assert!(!agent.is_validated_peer(remote_addr));
1208    }
1209
1210    #[test]
1211    fn request_no_retransmit() {
1212        let _log = crate::tests::test_init_log();
1213        let local_addr = "127.0.0.1:2000".parse().unwrap();
1214        let remote_addr = "127.0.0.1:1000".parse().unwrap();
1215        let mut agent = StunAgent::builder(TransportType::Udp, local_addr)
1216            .remote_addr(remote_addr)
1217            .build();
1218        let msg = Message::builder_request(BINDING, MessageWriteVec::new());
1219        let transaction_id = msg.transaction_id();
1220        let mut now = Instant::ZERO;
1221        agent.send_request(msg.finish(), remote_addr, now).unwrap();
1222        let mut transaction = agent.mut_request_transaction(transaction_id).unwrap();
1223        transaction.configure_timeout(Duration::from_secs(1), 0, Duration::from_secs(10));
1224        let StunAgentPollRet::WaitUntil(wait) = agent.poll(now) else {
1225            unreachable!();
1226        };
1227        assert_eq!(wait - now, Duration::from_secs(10));
1228        now = wait;
1229        let StunAgentPollRet::TransactionTimedOut(timed_out) = agent.poll(now) else {
1230            unreachable!();
1231        };
1232        assert_eq!(timed_out, transaction_id);
1233
1234        assert!(agent.request_transaction(transaction_id).is_none());
1235        assert!(agent.mut_request_transaction(transaction_id).is_none());
1236
1237        // unvalidated peer data should be dropped
1238        assert!(!agent.is_validated_peer(remote_addr));
1239    }
1240
1241    #[test]
1242    fn request_tcp_custom_timeout() {
1243        let _log = crate::tests::test_init_log();
1244        let local_addr = "127.0.0.1:2000".parse().unwrap();
1245        let remote_addr = "127.0.0.1:1000".parse().unwrap();
1246        let mut agent = StunAgent::builder(TransportType::Tcp, local_addr)
1247            .remote_addr(remote_addr)
1248            .request_retransmits(
1249                Duration::from_secs(1),
1250                Duration::from_secs(2),
1251                4,
1252                Duration::from_secs(3),
1253            )
1254            .build();
1255        let msg = Message::builder_request(BINDING, MessageWriteVec::new());
1256        let transaction_id = msg.transaction_id();
1257        let mut now = Instant::ZERO;
1258        agent.send_request(msg.finish(), remote_addr, now).unwrap();
1259        let StunAgentPollRet::WaitUntil(wait) = agent.poll(now) else {
1260            unreachable!();
1261        };
1262        assert_eq!(wait - now, Duration::from_secs(1 + 2 + 2 + 3));
1263        now = wait;
1264        let StunAgentPollRet::TransactionTimedOut(timed_out) = agent.poll(now) else {
1265            unreachable!();
1266        };
1267        assert_eq!(timed_out, transaction_id);
1268
1269        assert!(agent.request_transaction(transaction_id).is_none());
1270        assert!(agent.mut_request_transaction(transaction_id).is_none());
1271
1272        // unvalidated peer data should be dropped
1273        assert!(!agent.is_validated_peer(remote_addr));
1274    }
1275
1276    #[test]
1277    fn request_without_credentials() {
1278        let _log = crate::tests::test_init_log();
1279        let local_addr = "10.0.0.1:12345".parse().unwrap();
1280        let remote_addr = "10.0.0.2:3478".parse().unwrap();
1281        let now = Instant::ZERO;
1282
1283        let mut agent = StunAgent::builder(TransportType::Udp, local_addr).build();
1284
1285        // unvalidated peer data should be dropped
1286        assert!(!agent.is_validated_peer(remote_addr));
1287
1288        let msg = Message::builder_request(BINDING, MessageWriteVec::new());
1289        let transaction_id = msg.transaction_id();
1290        let transmit = agent
1291            .send_request(msg.finish(), remote_addr, Instant::ZERO)
1292            .unwrap();
1293
1294        let request = Message::from_bytes(&transmit.data).unwrap();
1295
1296        let mut response = Message::builder_success(&request, MessageWriteVec::new());
1297        let xor_addr = XorMappedAddress::new(transmit.from, request.transaction_id());
1298        response.add_attribute(&xor_addr).unwrap();
1299
1300        let data = response.finish();
1301        let to = transmit.to;
1302        trace!("data: {data:?}");
1303        let response = Message::from_bytes(&data).unwrap();
1304        let request = agent
1305            .request_transaction(response.transaction_id())
1306            .unwrap();
1307        assert_eq!(request.integrity(), None);
1308        assert!(agent.handle_stun_message_with_time(&response, to, now));
1309        assert_eq!(response.transaction_id(), transaction_id);
1310        assert!(agent.request_transaction(transaction_id).is_none());
1311        assert!(agent.mut_request_transaction(transaction_id).is_none());
1312        assert!(agent.is_validated_peer(remote_addr));
1313    }
1314
1315    #[test]
1316    fn response_with_incorrect_credentials() {
1317        let _log = crate::tests::test_init_log();
1318        let local_addr = "10.0.0.1:12345".parse().unwrap();
1319        let remote_addr = "10.0.0.2:3478".parse().unwrap();
1320        let now = Instant::ZERO;
1321
1322        let mut auth = ShortTermAuth::new();
1323        let mut agent = StunAgent::builder(TransportType::Udp, local_addr).build();
1324        let credentials = ShortTermCredentials::new(String::from("local_password"));
1325        let wrong_credentials = ShortTermCredentials::new(String::from("wrong_password"));
1326        auth.set_credentials(credentials.clone(), IntegrityAlgorithm::Sha1);
1327
1328        let mut msg = Message::builder_request(BINDING, MessageWriteVec::new());
1329        msg.add_message_integrity(&credentials.clone().into(), IntegrityAlgorithm::Sha1)
1330            .unwrap();
1331        let transmit = agent
1332            .send_request(msg.finish(), remote_addr, Instant::ZERO)
1333            .unwrap();
1334        let data = transmit.data;
1335
1336        let request = Message::from_bytes(&data).unwrap();
1337
1338        let mut response = Message::builder_success(&request, MessageWriteVec::new());
1339        let xor_addr = XorMappedAddress::new(transmit.from, request.transaction_id());
1340        response.add_attribute(&xor_addr).unwrap();
1341        // wrong credentials, should be `remote_credentials`
1342        response
1343            .add_message_integrity(&wrong_credentials.into(), IntegrityAlgorithm::Sha1)
1344            .unwrap();
1345
1346        let data = response.finish();
1347        let response = Message::from_bytes(&data).unwrap();
1348        // reply is ignored as it does not have credentials
1349        let request = agent
1350            .request_transaction(response.transaction_id())
1351            .unwrap();
1352        assert_eq!(request.integrity(), Some(IntegrityAlgorithm::Sha1));
1353        assert!(matches!(
1354            auth.validate_incoming_message(&response),
1355            Err(ValidateError::IntegrityFailed)
1356        ));
1357
1358        // unvalidated peer data should be dropped
1359        assert!(!agent.is_validated_peer(remote_addr));
1360
1361        // however signifying success will cause peer validation to succeed
1362        assert!(agent.handle_stun_message_with_time(&response, remote_addr, now));
1363        assert!(!agent.handle_stun_message_with_time(&response, remote_addr, now));
1364        assert!(agent.is_validated_peer(remote_addr));
1365    }
1366
1367    #[test]
1368    fn duplicate_response_ignored() {
1369        let _log = crate::tests::test_init_log();
1370        let local_addr = "10.0.0.1:12345".parse().unwrap();
1371        let remote_addr = "10.0.0.2:3478".parse().unwrap();
1372        let now = Instant::ZERO;
1373
1374        let mut agent = StunAgent::builder(TransportType::Udp, local_addr).build();
1375        assert!(!agent.is_validated_peer(remote_addr));
1376
1377        let msg = Message::builder_request(BINDING, MessageWriteVec::new());
1378        let transmit = agent
1379            .send_request(msg.finish(), remote_addr, Instant::ZERO)
1380            .unwrap();
1381        let data = transmit.data;
1382
1383        let request = Message::from_bytes(&data).unwrap();
1384
1385        let mut response = Message::builder_success(&request, MessageWriteVec::new());
1386        let xor_addr = XorMappedAddress::new(transmit.from, request.transaction_id());
1387        response.add_attribute(&xor_addr).unwrap();
1388
1389        let data = response.finish();
1390        let to = transmit.to;
1391        let response = Message::from_bytes(&data).unwrap();
1392        assert!(agent.handle_stun_message_with_time(&response, to, now));
1393
1394        let response = Message::from_bytes(&data).unwrap();
1395        assert!(!agent.handle_stun_message_with_time(&response, to, now));
1396    }
1397
1398    #[test]
1399    fn request_cancel() {
1400        let _log = crate::tests::test_init_log();
1401        let local_addr = "10.0.0.1:12345".parse().unwrap();
1402        let remote_addr = "10.0.0.2:3478".parse().unwrap();
1403
1404        let mut agent = StunAgent::builder(TransportType::Udp, local_addr).build();
1405
1406        let msg = Message::builder_request(BINDING, MessageWriteVec::new());
1407        let transaction_id = msg.transaction_id();
1408        let _transmit = agent
1409            .send_request(msg.finish(), remote_addr, Instant::ZERO)
1410            .unwrap();
1411
1412        let mut request = agent.mut_request_transaction(transaction_id).unwrap();
1413        assert_eq!(request.integrity(), None);
1414        assert_eq!(request.agent().local_addr(), local_addr);
1415        assert_eq!(request.mut_agent().local_addr(), local_addr);
1416        assert_eq!(request.peer_address(), remote_addr);
1417        request.cancel();
1418
1419        let ret = agent.poll(Instant::ZERO);
1420        let StunAgentPollRet::TransactionCancelled(_request) = ret else {
1421            unreachable!();
1422        };
1423        assert_eq!(transaction_id, transaction_id);
1424        assert!(agent.request_transaction(transaction_id).is_none());
1425        assert!(agent.mut_request_transaction(transaction_id).is_none());
1426        assert!(!agent.is_validated_peer(remote_addr));
1427    }
1428
1429    #[test]
1430    fn request_cancel_send() {
1431        let _log = crate::tests::test_init_log();
1432        let local_addr = "10.0.0.1:12345".parse().unwrap();
1433        let remote_addr = "10.0.0.2:3478".parse().unwrap();
1434
1435        let mut agent = StunAgent::builder(TransportType::Udp, local_addr).build();
1436
1437        let msg = Message::builder_request(BINDING, MessageWriteVec::new());
1438        let transaction_id = msg.transaction_id();
1439        let _transmit = agent
1440            .send_request(msg.finish(), remote_addr, Instant::ZERO)
1441            .unwrap();
1442
1443        let mut request = agent.mut_request_transaction(transaction_id).unwrap();
1444        assert_eq!(request.integrity(), None);
1445        assert_eq!(request.agent().local_addr(), local_addr);
1446        assert_eq!(request.mut_agent().local_addr(), local_addr);
1447        assert_eq!(request.peer_address(), remote_addr);
1448        request.cancel_retransmissions();
1449
1450        let mut now = Instant::ZERO;
1451        let start = now;
1452        loop {
1453            match agent.poll(now) {
1454                StunAgentPollRet::WaitUntil(new_now) => {
1455                    assert_ne!(new_now, now);
1456                    now = new_now;
1457                }
1458                StunAgentPollRet::TransactionCancelled(_) => break,
1459                _ => unreachable!(),
1460            }
1461            let _ = agent.poll_transmit(now);
1462        }
1463        assert!(now - start > Duration::from_secs(20));
1464        assert!(agent.request_transaction(transaction_id).is_none());
1465        assert!(agent.mut_request_transaction(transaction_id).is_none());
1466        assert!(!agent.is_validated_peer(remote_addr));
1467    }
1468
1469    #[test]
1470    fn request_duplicate() {
1471        let _log = crate::tests::test_init_log();
1472        let local_addr = "10.0.0.1:12345".parse().unwrap();
1473        let remote_addr = "10.0.0.2:3478".parse().unwrap();
1474        let now = Instant::ZERO;
1475
1476        let mut agent = StunAgent::builder(TransportType::Udp, local_addr).build();
1477
1478        let msg = Message::builder_request(BINDING, MessageWriteVec::new());
1479        let transaction_id = msg.transaction_id();
1480        let msg = msg.finish();
1481        let transmit = agent
1482            .send_request(msg.clone(), remote_addr, Instant::ZERO)
1483            .unwrap();
1484        let to = transmit.to;
1485        let request = Message::from_bytes(&transmit.data).unwrap();
1486
1487        let mut response = Message::builder_success(&request, MessageWriteVec::new());
1488        let xor_addr = XorMappedAddress::new(transmit.from, transaction_id);
1489        response.add_attribute(&xor_addr).unwrap();
1490
1491        assert!(matches!(
1492            agent.send_request(msg, remote_addr, Instant::ZERO),
1493            Err(StunError::AlreadyInProgress)
1494        ));
1495
1496        // the original transaction should still exist
1497        let request = agent.request_transaction(transaction_id).unwrap();
1498        assert_eq!(request.peer_address(), remote_addr);
1499
1500        let data = response.finish();
1501        let response = Message::from_bytes(&data).unwrap();
1502        assert!(agent.handle_stun_message_with_time(&response, to, now));
1503
1504        assert!(agent.is_validated_peer(to));
1505    }
1506
1507    #[test]
1508    fn incoming_request() {
1509        let _log = crate::tests::test_init_log();
1510        let local_addr = "10.0.0.1:12345".parse().unwrap();
1511        let remote_addr = "10.0.0.2:3478".parse().unwrap();
1512        let now = Instant::ZERO;
1513
1514        let mut agent = StunAgent::builder(TransportType::Udp, local_addr).build();
1515
1516        let msg = Message::builder_request(BINDING, MessageWriteVec::new());
1517        let data = msg.finish();
1518        let stun = Message::from_bytes(&data).unwrap();
1519        error!("{stun:?}");
1520        assert!(agent.handle_stun_message_with_time(&stun, remote_addr, now));
1521        agent.validated_peer(remote_addr);
1522        assert!(agent.is_validated_peer(remote_addr));
1523    }
1524
1525    #[test]
1526    fn tcp_request() {
1527        let _log = crate::tests::test_init_log();
1528        let local_addr = "127.0.0.1:2000".parse().unwrap();
1529        let remote_addr = "127.0.0.1:1000".parse().unwrap();
1530        let mut agent = StunAgent::builder(TransportType::Tcp, local_addr)
1531            .remote_addr(remote_addr)
1532            .build();
1533
1534        let msg = Message::builder_request(BINDING, MessageWriteVec::new());
1535        let transaction_id = msg.transaction_id();
1536        let transmit = agent
1537            .send_request(msg.finish(), remote_addr, Instant::ZERO)
1538            .unwrap();
1539        assert_eq!(transmit.transport, TransportType::Tcp);
1540        assert_eq!(transmit.from, local_addr);
1541        assert_eq!(transmit.to, remote_addr);
1542
1543        let request = Message::from_bytes(&transmit.data).unwrap();
1544        assert_eq!(request.transaction_id(), transaction_id);
1545    }
1546
1547    #[test]
1548    fn transmit_into_owned() {
1549        let data = [0x10, 0x20];
1550        let transport = TransportType::Udp;
1551        let from = "127.0.0.1:1000".parse().unwrap();
1552        let to = "127.0.0.1:2000".parse().unwrap();
1553        let transmit = Transmit::new(Data::from(data.as_ref()), TransportType::Udp, from, to);
1554        let owned = transmit.into_owned();
1555        assert_eq!(owned.data.as_ref(), data.as_ref());
1556        assert_eq!(owned.transport, transport);
1557        assert_eq!(owned.from, from);
1558        assert_eq!(owned.to, to);
1559        error!("{owned}");
1560    }
1561
1562    #[test]
1563    fn transmit_display() {
1564        let data = [0x10, 0x20];
1565        let from = "127.0.0.1:1000".parse().unwrap();
1566        let to = "127.0.0.1:2000".parse().unwrap();
1567        assert_eq!(
1568            alloc::format!(
1569                "{}",
1570                Transmit::new(Data::from(data.as_ref()), TransportType::Udp, from, to)
1571            ),
1572            String::from("Transmit(UDP: 127.0.0.1:1000 -> 127.0.0.1:2000 of 2 bytes)")
1573        );
1574    }
1575
1576    #[test]
1577    fn request_retransmits() {
1578        let _log = crate::tests::test_init_log();
1579        let rto = RequestRto {
1580            initial: Duration::from_millis(1),
1581            max: Duration::MAX,
1582            retransmits: 0,
1583            last_retransmit: Duration::from_secs(1),
1584        };
1585        let (timeouts, last_transmit_timeout) = rto.calculate_timeouts(TransportType::Udp);
1586        assert_eq!(timeouts, vec![]);
1587        assert_eq!(last_transmit_timeout, Duration::from_secs(1));
1588        let (timeouts, last_transmit_timeout) = rto.calculate_timeouts(TransportType::Tcp);
1589        assert_eq!(timeouts, vec![]);
1590        assert_eq!(last_transmit_timeout, Duration::from_secs(1));
1591
1592        let rto = RequestRto {
1593            initial: Duration::from_millis(1),
1594            max: Duration::MAX,
1595            retransmits: 1,
1596            last_retransmit: Duration::from_secs(1),
1597        };
1598        let (timeouts, last_transmit_timeout) = rto.calculate_timeouts(TransportType::Udp);
1599        assert_eq!(timeouts, vec![]);
1600        assert_eq!(last_transmit_timeout, Duration::from_secs(1));
1601        let (timeouts, last_transmit_timeout) = rto.calculate_timeouts(TransportType::Tcp);
1602        assert_eq!(timeouts, vec![]);
1603        assert_eq!(last_transmit_timeout, Duration::from_secs(1));
1604
1605        let rto = RequestRto {
1606            initial: Duration::from_millis(1),
1607            max: Duration::MAX,
1608            retransmits: 2,
1609            last_retransmit: Duration::from_secs(1),
1610        };
1611        let (timeouts, last_transmit_timeout) = rto.calculate_timeouts(TransportType::Udp);
1612        assert_eq!(timeouts, vec![Duration::from_millis(1)]);
1613        assert_eq!(last_transmit_timeout, Duration::from_secs(1));
1614        let (timeouts, last_transmit_timeout) = rto.calculate_timeouts(TransportType::Tcp);
1615        assert_eq!(timeouts, vec![]);
1616        assert_eq!(
1617            last_transmit_timeout,
1618            Duration::from_secs(1) + Duration::from_millis(1)
1619        );
1620    }
1621
1622    #[test]
1623    fn stats_send_receive_request() {
1624        let _log = crate::tests::test_init_log();
1625        let local_addr = "127.0.0.1:2000".parse().unwrap();
1626        let remote_addr = "127.0.0.1:1000".parse().unwrap();
1627        let now = Instant::ZERO;
1628
1629        let mut agent = StunAgent::builder(TransportType::Udp, local_addr)
1630            .stats(true)
1631            .build();
1632
1633        assert!(agent.stats().is_some());
1634        let stats = agent.stats().unwrap();
1635        assert_eq!(stats.requests_sent(), 0);
1636        assert_eq!(stats.responses_received(), 0);
1637        assert_eq!(stats.bytes_sent(), 0);
1638        assert_eq!(stats.rtt_count(), 0);
1639
1640        let msg = Message::builder_request(BINDING, MessageWriteVec::new());
1641        let transmit = agent.send_request(msg.finish(), remote_addr, now).unwrap();
1642        let request_data = transmit.data.as_ref().to_vec();
1643
1644        let stats = agent.stats().unwrap();
1645        assert_eq!(stats.requests_sent(), 1);
1646        assert_eq!(stats.bytes_sent(), request_data.len() as u64);
1647
1648        // Build and handle response with RTT tracking
1649        let request = Message::from_bytes(&request_data).unwrap();
1650        let mut response = Message::builder_success(&request, MessageWriteVec::new());
1651        let xor_addr =
1652            XorMappedAddress::new("10.0.0.1:12345".parse().unwrap(), request.transaction_id());
1653        response.add_attribute(&xor_addr).unwrap();
1654        let response_data = response.finish();
1655        let response = Message::from_bytes(&response_data).unwrap();
1656
1657        let receive_time = Duration::from_millis(10);
1658        assert!(agent.poll_transmit(now + receive_time).is_none());
1659        assert!(agent.handle_stun_message_with_time(&response, remote_addr, now + receive_time));
1660
1661        let stats = agent.stats().unwrap();
1662        assert_eq!(stats.responses_received(), 1);
1663        assert_eq!(stats.bytes_received(), response_data.len() as u64);
1664        assert_eq!(stats.rtt_count(), 1);
1665        assert_eq!(stats.rtt_min(), Some(receive_time));
1666        assert_eq!(stats.rtt_max(), Some(receive_time));
1667        assert_eq!(stats.rtt_average(), Some(receive_time));
1668
1669        assert!(agent.handle_stun_message_with_time(&request, remote_addr, now));
1670        let stats = agent.stats().unwrap();
1671        assert_eq!(stats.requests_sent(), 1);
1672        assert_eq!(
1673            stats.bytes_received(),
1674            (request_data.len() + response_data.len()) as u64
1675        );
1676
1677        agent.send(&response, remote_addr, now).unwrap();
1678        let stats = agent.stats().unwrap();
1679        assert_eq!(stats.responses_sent(), 1);
1680        assert_eq!(
1681            stats.bytes_sent(),
1682            (request_data.len() + response_data.len()) as u64
1683        );
1684    }
1685
1686    #[test]
1687    fn stats_disabled_by_default() {
1688        let _log = crate::tests::test_init_log();
1689        let local_addr = "127.0.0.1:2000".parse().unwrap();
1690        let agent = StunAgent::builder(TransportType::Udp, local_addr).build();
1691        assert!(agent.stats().is_none());
1692    }
1693
1694    #[test]
1695    fn stats_timeout_and_cancel() {
1696        let _log = crate::tests::test_init_log();
1697        let local_addr = "127.0.0.1:2000".parse().unwrap();
1698        let remote_addr = "127.0.0.1:1000".parse().unwrap();
1699        let mut agent = StunAgent::builder(TransportType::Udp, local_addr)
1700            .stats(true)
1701            .build();
1702
1703        // Send a request that will timeout
1704        let msg = Message::builder_request(BINDING, MessageWriteVec::new());
1705        agent
1706            .send_request(msg.finish(), remote_addr, Instant::ZERO)
1707            .unwrap();
1708
1709        let mut now = Instant::ZERO;
1710        loop {
1711            let _ = agent.poll_transmit(now);
1712            match agent.poll(now) {
1713                StunAgentPollRet::WaitUntil(new_now) => now = new_now,
1714                StunAgentPollRet::TransactionTimedOut(_) => break,
1715                _ => unreachable!(),
1716            }
1717        }
1718
1719        let stats = agent.stats().unwrap();
1720        assert_eq!(stats.transactions_timed_out(), 1);
1721
1722        // Send another request that we'll cancel
1723        let msg = Message::builder_request(BINDING, MessageWriteVec::new());
1724        let cancel_id = msg.transaction_id();
1725        agent.send_request(msg.finish(), remote_addr, now).unwrap();
1726
1727        let mut req = agent.mut_request_transaction(cancel_id).unwrap();
1728        req.cancel();
1729
1730        let ret = agent.poll(now);
1731        assert!(matches!(ret, StunAgentPollRet::TransactionCancelled(_)));
1732
1733        let stats = agent.stats().unwrap();
1734        assert_eq!(stats.transactions_cancelled(), 1);
1735    }
1736
1737    #[test]
1738    fn stats_rtt_only_no_retransmit() {
1739        let _log = crate::tests::test_init_log();
1740        let local_addr = "127.0.0.1:2000".parse().unwrap();
1741        let remote_addr = "127.0.0.1:1000".parse().unwrap();
1742
1743        let mut agent = StunAgent::builder(TransportType::Udp, local_addr)
1744            .request_retransmits(
1745                Duration::from_millis(1),
1746                Duration::MAX,
1747                2,
1748                Duration::from_millis(1),
1749            )
1750            .stats(true)
1751            .build();
1752
1753        let msg = Message::builder_request(BINDING, MessageWriteVec::new());
1754        let transmit = agent
1755            .send_request(msg.finish(), remote_addr, Instant::ZERO)
1756            .unwrap();
1757        let request_data = transmit.data.as_ref().to_vec();
1758        let from_addr = transmit.from;
1759
1760        // Retransmit once by polling
1761        let retransmit_time = Instant::ZERO + Duration::from_millis(1);
1762        let _retransmit = agent.poll_transmit(retransmit_time).unwrap();
1763
1764        // Now build and handle response
1765        let request = Message::from_bytes(&request_data).unwrap();
1766        let mut response = Message::builder_success(&request, MessageWriteVec::new());
1767        let xor_addr = XorMappedAddress::new(from_addr, request.transaction_id());
1768        response.add_attribute(&xor_addr).unwrap();
1769        let response_data = response.finish();
1770        let response = Message::from_bytes(&response_data).unwrap();
1771
1772        let receive_time = retransmit_time + Duration::from_millis(10);
1773        assert!(agent.handle_stun_message_with_time(&response, remote_addr, receive_time));
1774
1775        let stats = agent.stats().unwrap();
1776        assert_eq!(stats.responses_received(), 1);
1777        assert_eq!(stats.rtt_count(), 0);
1778    }
1779
1780    #[test]
1781    fn stats_error_response() {
1782        let _log = crate::tests::test_init_log();
1783        let local_addr = "127.0.0.1:2000".parse().unwrap();
1784        let remote_addr = "127.0.0.1:1000".parse().unwrap();
1785        let now = Instant::ZERO;
1786
1787        let mut agent = StunAgent::builder(TransportType::Udp, local_addr)
1788            .stats(true)
1789            .build();
1790
1791        let msg = Message::builder_request(BINDING, MessageWriteVec::new());
1792        let transmit = agent
1793            .send_request(msg.finish(), remote_addr, Instant::ZERO)
1794            .unwrap();
1795        let request_data = transmit.data.as_ref().to_vec();
1796
1797        let request = Message::from_bytes(&request_data).unwrap();
1798        let mut error_response = Message::builder_error(&request, MessageWriteVec::new());
1799        let error_code = ErrorCode::builder(ErrorCode::BAD_REQUEST).build().unwrap();
1800        error_response.add_attribute(&error_code).unwrap();
1801        let error_data = error_response.finish();
1802        let error_msg = Message::from_bytes(&error_data).unwrap();
1803
1804        assert!(agent.handle_stun_message_with_time(&error_msg, remote_addr, now));
1805
1806        let stats = agent.stats().unwrap();
1807        assert_eq!(stats.responses_received(), 1);
1808        assert_eq!(stats.bytes_received(), error_data.len() as u64);
1809    }
1810
1811    #[test]
1812    fn stats_send_receive_indication() {
1813        let _log = crate::tests::test_init_log();
1814        let local_addr = "127.0.0.1:2000".parse().unwrap();
1815        let remote_addr = "127.0.0.1:1000".parse().unwrap();
1816        let mut agent = StunAgent::builder(TransportType::Udp, local_addr)
1817            .stats(true)
1818            .build();
1819        let now = Instant::ZERO;
1820
1821        // Send a request that will timeout
1822        let msg = Message::builder_indication(BINDING, MessageWriteVec::new());
1823        let msg_data = msg.finish();
1824        agent
1825            .send(msg_data.clone(), remote_addr, Instant::ZERO)
1826            .unwrap();
1827
1828        let stats = agent.stats().unwrap();
1829        assert_eq!(stats.indications_sent(), 1);
1830        assert_eq!(stats.bytes_sent(), msg_data.len() as u64);
1831
1832        let msg = Message::from_bytes(&msg_data).unwrap();
1833        assert!(agent.handle_stun_message_with_time(&msg, remote_addr, now));
1834
1835        let stats = agent.stats().unwrap();
1836        assert_eq!(stats.indications_received(), 1);
1837        assert_eq!(stats.bytes_received(), msg_data.len() as u64);
1838    }
1839
1840    #[test]
1841    fn request_removal_from_mut_agent() {
1842        let _log = crate::tests::test_init_log();
1843        let local_addr = "10.0.0.1:12345".parse().unwrap();
1844        let remote_addr = "10.0.0.2:3478".parse().unwrap();
1845        let now = Instant::ZERO;
1846
1847        let mut agent = StunAgent::builder(TransportType::Udp, local_addr).build();
1848
1849        let msg = Message::builder_request(BINDING, MessageWriteVec::new());
1850        let transaction_id = msg.transaction_id();
1851        let transmit = agent
1852            .send_request(msg.finish(), remote_addr, Instant::ZERO)
1853            .unwrap();
1854
1855        let request = Message::from_bytes(&transmit.data).unwrap();
1856
1857        let mut response = Message::builder_success(&request, MessageWriteVec::new());
1858        let xor_addr = XorMappedAddress::new(transmit.from, request.transaction_id());
1859        response.add_attribute(&xor_addr).unwrap();
1860
1861        let data = response.finish();
1862        let to = transmit.to;
1863        trace!("data: {data:?}");
1864        let response = Message::from_bytes(&data).unwrap();
1865        assert_eq!(response.transaction_id(), transaction_id);
1866
1867        // handling the success response while holding the outstanding request should not panic and
1868        // still result in valid retrieval of configuration.
1869        let mut request = agent
1870            .mut_request_transaction(response.transaction_id())
1871            .unwrap();
1872        assert_eq!(request.integrity(), None);
1873        assert_eq!(request.peer_address(), to);
1874        assert!(request
1875            .mut_agent()
1876            .handle_stun_message_with_time(&response, to, now));
1877
1878        assert_eq!(request.integrity(), None);
1879        assert_eq!(request.peer_address(), to);
1880        // attempting to modify the request is ignored.
1881        request.cancel_retransmissions();
1882        request.cancel();
1883        request.configure_timeout(Duration::from_secs(1), 3, Duration::from_secs(2));
1884
1885        // the request is no longer stored by the agent.
1886        assert!(request
1887            .agent()
1888            .request_transaction(transaction_id)
1889            .is_none());
1890        assert!(request
1891            .mut_agent()
1892            .mut_request_transaction(transaction_id)
1893            .is_none());
1894    }
1895}