1use 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#[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#[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 pub fn remote_addr(mut self, addr: SocketAddr) -> Self {
68 self.remote_addr = Some(addr);
69 self
70 }
71
72 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 pub fn stats(mut self, stats: bool) -> Self {
107 self.stats = stats;
108 self
109 }
110
111 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 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 pub fn transport(&self) -> TransportType {
149 self.transport
150 }
151
152 pub fn local_addr(&self) -> SocketAddr {
154 self.local_addr
155 }
156
157 pub fn remote_addr(&self) -> Option<SocketAddr> {
159 self.remote_addr
160 }
161
162 pub fn stats(&self) -> Option<&StunAgentStats> {
164 self.stats.as_ref()
165 }
166
167 pub fn stats_mut(&mut self) -> Option<&mut StunAgentStats> {
169 self.stats.as_mut()
170 }
171
172 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 #[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 #[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 pub fn is_validated_peer(&self, remote_addr: SocketAddr) -> bool {
289 self.validated_peers.contains(&remote_addr)
290 }
291
292 #[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 #[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 pub fn handle_stun_message(&mut self, msg: &Message<'_>, from: SocketAddr) -> bool {
326 self.handle_stun_message_internal(msg, from, None)
327 }
328
329 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 .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 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 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 #[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 #[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#[derive(Debug)]
528pub enum StunAgentPollRet {
529 TransactionTimedOut(TransactionId),
531 TransactionCancelled(TransactionId),
533 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#[derive(Debug)]
548pub struct Transmit<T: AsRef<[u8]>> {
549 pub data: T,
551 pub transport: TransportType,
553 pub from: SocketAddr,
555 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 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 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 pub fn into_owned<'b>(self) -> Transmit<Data<'b>> {
615 self.reinterpret_data(|data| data.into_owned())
616 }
617}
618
619#[derive(Debug)]
621enum StunRequestPollRet {
622 WaitUntil(Instant),
624 Cancelled,
626 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 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 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#[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 pub fn peer_address(&self) -> SocketAddr {
811 self.peer_address
812 }
813
814 pub fn integrity(&self) -> Option<IntegrityAlgorithm> {
816 self.request_integrity
817 }
818
819 pub fn agent(&self) -> &StunAgent {
821 self.agent
822 }
823}
824
825#[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 pub fn peer_address(&self) -> SocketAddr {
837 self.peer_address
838 }
839
840 pub fn integrity(&self) -> Option<IntegrityAlgorithm> {
842 self.request_integrity
843 }
844
845 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 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 pub fn agent(&self) -> &StunAgent {
865 self.agent
866 }
867
868 pub fn mut_agent(&mut self) -> &mut StunAgent {
870 self.agent
871 }
872
873 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 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#[derive(Debug, thiserror::Error)]
928#[non_exhaustive]
929pub enum StunError {
930 #[error("The operation is already in progress")]
932 AlreadyInProgress,
933 #[error("A required resource could not be found")]
935 ResourceNotFound,
936 #[error("An operation timed out")]
938 TimedOut,
939 #[error("Unexpected data was received")]
941 ProtocolViolation,
942 #[error("Operation was aborted")]
944 Aborted,
945 #[error("{}", .0)]
947 ParseError(StunParseError),
948 #[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 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 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 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 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 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 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 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 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 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 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 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 assert!(!agent.is_validated_peer(remote_addr));
1360
1361 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 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 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 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 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 let retransmit_time = Instant::ZERO + Duration::from_millis(1);
1762 let _retransmit = agent.poll_transmit(retransmit_time).unwrap();
1763
1764 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 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 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 request.cancel_retransmissions();
1882 request.cancel();
1883 request.configure_timeout(Duration::from_secs(1), 3, Duration::from_secs(2));
1884
1885 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}