1 #[cfg(test)]
2 mod relay_conn_test;
3
4 // client implements the API for a TURN client
5 use super::binding::*;
6 use super::periodic_timer::*;
7 use super::permission::*;
8 use super::transaction::*;
9 use crate::proto;
10 use crate::Error;
11
12 use stun::agent::*;
13 use stun::attributes::*;
14 use stun::error_code::*;
15 use stun::fingerprint::*;
16 use stun::integrity::*;
17 use stun::message::*;
18 use stun::textattrs::*;
19
20 use util::Conn;
21
22 use std::io;
23 use std::net::SocketAddr;
24 use std::sync::Arc;
25 use tokio::sync::{mpsc, Mutex};
26 use tokio::time::{Duration, Instant};
27
28 use async_trait::async_trait;
29
30 const PERM_REFRESH_INTERVAL: Duration = Duration::from_secs(120);
31 const MAX_RETRY_ATTEMPTS: u16 = 3;
32
33 pub(crate) struct InboundData {
34 pub(crate) data: Vec<u8>,
35 pub(crate) from: SocketAddr,
36 }
37
38 // UDPConnObserver is an interface to UDPConn observer
39 #[async_trait]
40 pub trait RelayConnObserver {
turn_server_addr(&self) -> String41 fn turn_server_addr(&self) -> String;
username(&self) -> Username42 fn username(&self) -> Username;
realm(&self) -> Realm43 fn realm(&self) -> Realm;
write_to(&self, data: &[u8], to: &str) -> Result<usize, util::Error>44 async fn write_to(&self, data: &[u8], to: &str) -> Result<usize, util::Error>;
perform_transaction( &mut self, msg: &Message, to: &str, ignore_result: bool, ) -> Result<TransactionResult, Error>45 async fn perform_transaction(
46 &mut self,
47 msg: &Message,
48 to: &str,
49 ignore_result: bool,
50 ) -> Result<TransactionResult, Error>;
51 }
52
53 // RelayConnConfig is a set of configuration params use by NewUDPConn
54 pub(crate) struct RelayConnConfig {
55 pub(crate) relayed_addr: SocketAddr,
56 pub(crate) integrity: MessageIntegrity,
57 pub(crate) nonce: Nonce,
58 pub(crate) lifetime: Duration,
59 pub(crate) binding_mgr: Arc<Mutex<BindingManager>>,
60 pub(crate) read_ch_rx: Arc<Mutex<mpsc::Receiver<InboundData>>>,
61 }
62
63 pub struct RelayConnInternal<T: 'static + RelayConnObserver + Send + Sync> {
64 obs: Arc<Mutex<T>>,
65 relayed_addr: SocketAddr,
66 perm_map: PermissionMap,
67 binding_mgr: Arc<Mutex<BindingManager>>,
68 integrity: MessageIntegrity,
69 nonce: Nonce,
70 lifetime: Duration,
71 }
72
73 // RelayConn is the implementation of the Conn interfaces for UDP Relayed network connections.
74 pub struct RelayConn<T: 'static + RelayConnObserver + Send + Sync> {
75 relayed_addr: SocketAddr,
76 read_ch_rx: Arc<Mutex<mpsc::Receiver<InboundData>>>,
77 relay_conn: Arc<Mutex<RelayConnInternal<T>>>,
78 refresh_alloc_timer: PeriodicTimer,
79 refresh_perms_timer: PeriodicTimer,
80 }
81
82 impl<T: 'static + RelayConnObserver + Send + Sync> RelayConn<T> {
83 // new creates a new instance of UDPConn
new(obs: Arc<Mutex<T>>, config: RelayConnConfig) -> Self84 pub(crate) async fn new(obs: Arc<Mutex<T>>, config: RelayConnConfig) -> Self {
85 log::debug!("initial lifetime: {} seconds", config.lifetime.as_secs());
86
87 let c = RelayConn {
88 refresh_alloc_timer: PeriodicTimer::new(TimerIdRefresh::Alloc, config.lifetime / 2),
89 refresh_perms_timer: PeriodicTimer::new(TimerIdRefresh::Perms, PERM_REFRESH_INTERVAL),
90 relayed_addr: config.relayed_addr,
91 read_ch_rx: Arc::clone(&config.read_ch_rx),
92 relay_conn: Arc::new(Mutex::new(RelayConnInternal::new(obs, config))),
93 };
94
95 let rci1 = Arc::clone(&c.relay_conn);
96 let rci2 = Arc::clone(&c.relay_conn);
97
98 if c.refresh_alloc_timer.start(rci1).await {
99 log::debug!("refresh_alloc_timer started");
100 }
101 if c.refresh_perms_timer.start(rci2).await {
102 log::debug!("refresh_perms_timer started");
103 }
104
105 c
106 }
107 }
108
109 #[async_trait]
110 impl<T: RelayConnObserver + Send + Sync> Conn for RelayConn<T> {
connect(&self, _addr: SocketAddr) -> Result<(), util::Error>111 async fn connect(&self, _addr: SocketAddr) -> Result<(), util::Error> {
112 Err(io::Error::new(io::ErrorKind::Other, "Not applicable").into())
113 }
114
recv(&self, _buf: &mut [u8]) -> Result<usize, util::Error>115 async fn recv(&self, _buf: &mut [u8]) -> Result<usize, util::Error> {
116 Err(io::Error::new(io::ErrorKind::Other, "Not applicable").into())
117 }
118
119 // ReadFrom reads a packet from the connection,
120 // copying the payload into p. It returns the number of
121 // bytes copied into p and the return address that
122 // was on the packet.
123 // It returns the number of bytes read (0 <= n <= len(p))
124 // and any error encountered. Callers should always process
125 // the n > 0 bytes returned before considering the error err.
126 // ReadFrom can be made to time out and return
127 // an Error with Timeout() == true after a fixed time limit;
128 // see SetDeadline and SetReadDeadline.
recv_from(&self, p: &mut [u8]) -> Result<(usize, SocketAddr), util::Error>129 async fn recv_from(&self, p: &mut [u8]) -> Result<(usize, SocketAddr), util::Error> {
130 let mut read_ch_rx = self.read_ch_rx.lock().await;
131
132 if let Some(ib_data) = read_ch_rx.recv().await {
133 let n = ib_data.data.len();
134 if p.len() < n {
135 return Err(io::Error::new(
136 io::ErrorKind::InvalidInput,
137 Error::ErrShortBuffer.to_string(),
138 )
139 .into());
140 }
141 p[..n].copy_from_slice(&ib_data.data);
142 Ok((n, ib_data.from))
143 } else {
144 Err(io::Error::new(
145 io::ErrorKind::ConnectionAborted,
146 Error::ErrAlreadyClosed.to_string(),
147 )
148 .into())
149 }
150 }
151
send(&self, _buf: &[u8]) -> Result<usize, util::Error>152 async fn send(&self, _buf: &[u8]) -> Result<usize, util::Error> {
153 Err(io::Error::new(io::ErrorKind::Other, "Not applicable").into())
154 }
155
156 // write_to writes a packet with payload p to addr.
157 // write_to can be made to time out and return
158 // an Error with Timeout() == true after a fixed time limit;
159 // see SetDeadline and SetWriteDeadline.
160 // On packet-oriented connections, write timeouts are rare.
send_to(&self, p: &[u8], addr: SocketAddr) -> Result<usize, util::Error>161 async fn send_to(&self, p: &[u8], addr: SocketAddr) -> Result<usize, util::Error> {
162 let mut relay_conn = self.relay_conn.lock().await;
163 match relay_conn.send_to(p, addr).await {
164 Ok(n) => Ok(n),
165 Err(err) => Err(io::Error::new(io::ErrorKind::Other, err.to_string()).into()),
166 }
167 }
168
169 // LocalAddr returns the local network address.
local_addr(&self) -> Result<SocketAddr, util::Error>170 fn local_addr(&self) -> Result<SocketAddr, util::Error> {
171 Ok(self.relayed_addr)
172 }
173
remote_addr(&self) -> Option<SocketAddr>174 fn remote_addr(&self) -> Option<SocketAddr> {
175 None
176 }
177
178 // Close closes the connection.
179 // Any blocked ReadFrom or write_to operations will be unblocked and return errors.
close(&self) -> Result<(), util::Error>180 async fn close(&self) -> Result<(), util::Error> {
181 self.refresh_alloc_timer.stop().await;
182 self.refresh_perms_timer.stop().await;
183
184 let mut relay_conn = self.relay_conn.lock().await;
185 let _ = relay_conn
186 .close()
187 .await
188 .map_err(|err| util::Error::Other(format!("{err}")));
189 Ok(())
190 }
191 }
192
193 impl<T: RelayConnObserver + Send + Sync> RelayConnInternal<T> {
194 // new creates a new instance of UDPConn
new(obs: Arc<Mutex<T>>, config: RelayConnConfig) -> Self195 fn new(obs: Arc<Mutex<T>>, config: RelayConnConfig) -> Self {
196 RelayConnInternal {
197 obs,
198 relayed_addr: config.relayed_addr,
199 perm_map: PermissionMap::new(),
200 binding_mgr: config.binding_mgr,
201 integrity: config.integrity,
202 nonce: config.nonce,
203 lifetime: config.lifetime,
204 }
205 }
206
207 // write_to writes a packet with payload p to addr.
208 // write_to can be made to time out and return
209 // an Error with Timeout() == true after a fixed time limit;
210 // see SetDeadline and SetWriteDeadline.
211 // On packet-oriented connections, write timeouts are rare.
send_to(&mut self, p: &[u8], addr: SocketAddr) -> Result<usize, Error>212 async fn send_to(&mut self, p: &[u8], addr: SocketAddr) -> Result<usize, Error> {
213 // check if we have a permission for the destination IP addr
214 let perm = if let Some(perm) = self.perm_map.find(&addr) {
215 Arc::clone(perm)
216 } else {
217 let perm = Arc::new(Permission::default());
218 self.perm_map.insert(&addr, Arc::clone(&perm));
219 perm
220 };
221
222 let mut result = Ok(());
223 for _ in 0..MAX_RETRY_ATTEMPTS {
224 result = self.create_perm(&perm, addr).await;
225 if let Err(err) = &result {
226 if Error::ErrTryAgain != *err {
227 break;
228 }
229 }
230 }
231 result?;
232
233 let number = {
234 let (bind_st, bind_at, bind_number, bind_addr) = {
235 let mut binding_mgr = self.binding_mgr.lock().await;
236 let b = if let Some(b) = binding_mgr.find_by_addr(&addr) {
237 b
238 } else {
239 binding_mgr
240 .create(addr)
241 .ok_or_else(|| Error::Other("Addr not found".to_owned()))?
242 };
243 (b.state(), b.refreshed_at(), b.number, b.addr)
244 };
245
246 if bind_st == BindingState::Idle
247 || bind_st == BindingState::Request
248 || bind_st == BindingState::Failed
249 {
250 // block only callers with the same binding until
251 // the binding transaction has been complete
252 // binding state may have been changed while waiting. check again.
253 if bind_st == BindingState::Idle {
254 let binding_mgr = Arc::clone(&self.binding_mgr);
255 let rc_obs = Arc::clone(&self.obs);
256 let nonce = self.nonce.clone();
257 let integrity = self.integrity.clone();
258 {
259 let mut bm = binding_mgr.lock().await;
260 if let Some(b) = bm.get_by_addr(&bind_addr) {
261 b.set_state(BindingState::Request);
262 }
263 }
264 tokio::spawn(async move {
265 let result = RelayConnInternal::bind(
266 rc_obs,
267 bind_addr,
268 bind_number,
269 nonce,
270 integrity,
271 )
272 .await;
273
274 {
275 let mut bm = binding_mgr.lock().await;
276 if let Err(err) = result {
277 if Error::ErrUnexpectedResponse != err {
278 bm.delete_by_addr(&bind_addr);
279 } else if let Some(b) = bm.get_by_addr(&bind_addr) {
280 b.set_state(BindingState::Failed);
281 }
282
283 // keep going...
284 log::warn!("bind() failed: {}", err);
285 } else if let Some(b) = bm.get_by_addr(&bind_addr) {
286 b.set_state(BindingState::Ready);
287 }
288 }
289 });
290 }
291
292 // send data using SendIndication
293 let peer_addr = socket_addr2peer_address(&addr);
294 let mut msg = Message::new();
295 msg.build(&[
296 Box::new(TransactionId::new()),
297 Box::new(MessageType::new(METHOD_SEND, CLASS_INDICATION)),
298 Box::new(proto::data::Data(p.to_vec())),
299 Box::new(peer_addr),
300 Box::new(FINGERPRINT),
301 ])?;
302
303 // indication has no transaction (fire-and-forget)
304 let obs = self.obs.lock().await;
305 let turn_server_addr = obs.turn_server_addr();
306 return Ok(obs.write_to(&msg.raw, &turn_server_addr).await?);
307 }
308
309 // binding is either ready
310
311 // check if the binding needs a refresh
312 if bind_st == BindingState::Ready
313 && Instant::now()
314 .checked_duration_since(bind_at)
315 .unwrap_or_else(|| Duration::from_secs(0))
316 > Duration::from_secs(5 * 60)
317 {
318 let binding_mgr = Arc::clone(&self.binding_mgr);
319 let rc_obs = Arc::clone(&self.obs);
320 let nonce = self.nonce.clone();
321 let integrity = self.integrity.clone();
322 {
323 let mut bm = binding_mgr.lock().await;
324 if let Some(b) = bm.get_by_addr(&bind_addr) {
325 b.set_state(BindingState::Refresh);
326 }
327 }
328 tokio::spawn(async move {
329 let result =
330 RelayConnInternal::bind(rc_obs, bind_addr, bind_number, nonce, integrity)
331 .await;
332
333 {
334 let mut bm = binding_mgr.lock().await;
335 if let Err(err) = result {
336 if Error::ErrUnexpectedResponse != err {
337 bm.delete_by_addr(&bind_addr);
338 } else if let Some(b) = bm.get_by_addr(&bind_addr) {
339 b.set_state(BindingState::Failed);
340 }
341
342 // keep going...
343 log::warn!("bind() for refresh failed: {}", err);
344 } else if let Some(b) = bm.get_by_addr(&bind_addr) {
345 b.set_refreshed_at(Instant::now());
346 b.set_state(BindingState::Ready);
347 }
348 }
349 });
350 }
351
352 bind_number
353 };
354
355 // send via ChannelData
356 self.send_channel_data(p, number).await
357 }
358
359 // This func-block would block, per destination IP (, or perm), until
360 // the perm state becomes "requested". Purpose of this is to guarantee
361 // the order of packets (within the same perm).
362 // Note that CreatePermission transaction may not be complete before
363 // all the data transmission. This is done assuming that the request
364 // will be mostly likely successful and we can tolerate some loss of
365 // UDP packet (or reorder), inorder to minimize the latency in most cases.
create_perm(&mut self, perm: &Arc<Permission>, addr: SocketAddr) -> Result<(), Error>366 async fn create_perm(&mut self, perm: &Arc<Permission>, addr: SocketAddr) -> Result<(), Error> {
367 if perm.state() == PermState::Idle {
368 // punch a hole! (this would block a bit..)
369 if let Err(err) = self.create_permissions(&[addr]).await {
370 self.perm_map.delete(&addr);
371 return Err(err);
372 }
373 perm.set_state(PermState::Permitted);
374 }
375 Ok(())
376 }
377
send_channel_data(&self, data: &[u8], ch_num: u16) -> Result<usize, Error>378 async fn send_channel_data(&self, data: &[u8], ch_num: u16) -> Result<usize, Error> {
379 let mut ch_data = proto::chandata::ChannelData {
380 data: data.to_vec(),
381 number: proto::channum::ChannelNumber(ch_num),
382 ..Default::default()
383 };
384 ch_data.encode();
385
386 let obs = self.obs.lock().await;
387 Ok(obs.write_to(&ch_data.raw, &obs.turn_server_addr()).await?)
388 }
389
create_permissions(&mut self, addrs: &[SocketAddr]) -> Result<(), Error>390 async fn create_permissions(&mut self, addrs: &[SocketAddr]) -> Result<(), Error> {
391 let res = {
392 let msg = {
393 let obs = self.obs.lock().await;
394 let mut setters: Vec<Box<dyn Setter>> = vec![
395 Box::new(TransactionId::new()),
396 Box::new(MessageType::new(METHOD_CREATE_PERMISSION, CLASS_REQUEST)),
397 ];
398
399 for addr in addrs {
400 setters.push(Box::new(socket_addr2peer_address(addr)));
401 }
402
403 setters.push(Box::new(obs.username()));
404 setters.push(Box::new(obs.realm()));
405 setters.push(Box::new(self.nonce.clone()));
406 setters.push(Box::new(self.integrity.clone()));
407 setters.push(Box::new(FINGERPRINT));
408
409 let mut msg = Message::new();
410 msg.build(&setters)?;
411 msg
412 };
413
414 let mut obs = self.obs.lock().await;
415 let turn_server_addr = obs.turn_server_addr();
416
417 log::debug!("UDPConn.createPermissions call PerformTransaction 1");
418 let tr_res = obs
419 .perform_transaction(&msg, &turn_server_addr, false)
420 .await?;
421
422 tr_res.msg
423 };
424
425 if res.typ.class == CLASS_ERROR_RESPONSE {
426 let mut code = ErrorCodeAttribute::default();
427 let result = code.get_from(&res);
428 if result.is_err() {
429 return Err(Error::Other(format!("{}", res.typ)));
430 } else if code.code == CODE_STALE_NONCE {
431 self.set_nonce_from_msg(&res);
432 return Err(Error::ErrTryAgain);
433 } else {
434 return Err(Error::Other(format!("{} (error {})", res.typ, code)));
435 }
436 }
437
438 Ok(())
439 }
440
set_nonce_from_msg(&mut self, msg: &Message)441 pub fn set_nonce_from_msg(&mut self, msg: &Message) {
442 // Update nonce
443 match Nonce::get_from_as(msg, ATTR_NONCE) {
444 Ok(nonce) => {
445 self.nonce = nonce;
446 log::debug!("refresh allocation: 438, got new nonce.");
447 }
448 Err(_) => log::warn!("refresh allocation: 438 but no nonce."),
449 }
450 }
451
452 // Close closes the connection.
453 // Any blocked ReadFrom or write_to operations will be unblocked and return errors.
close(&mut self) -> Result<(), Error>454 pub async fn close(&mut self) -> Result<(), Error> {
455 self.refresh_allocation(Duration::from_secs(0), true /* dontWait=true */)
456 .await
457 }
458
refresh_allocation( &mut self, lifetime: Duration, dont_wait: bool, ) -> Result<(), Error>459 async fn refresh_allocation(
460 &mut self,
461 lifetime: Duration,
462 dont_wait: bool,
463 ) -> Result<(), Error> {
464 let res = {
465 let mut obs = self.obs.lock().await;
466
467 let mut msg = Message::new();
468 msg.build(&[
469 Box::new(TransactionId::new()),
470 Box::new(MessageType::new(METHOD_REFRESH, CLASS_REQUEST)),
471 Box::new(proto::lifetime::Lifetime(lifetime)),
472 Box::new(obs.username()),
473 Box::new(obs.realm()),
474 Box::new(self.nonce.clone()),
475 Box::new(self.integrity.clone()),
476 Box::new(FINGERPRINT),
477 ])?;
478
479 log::debug!("send refresh request (dont_wait={})", dont_wait);
480 let turn_server_addr = obs.turn_server_addr();
481 let tr_res = obs
482 .perform_transaction(&msg, &turn_server_addr, dont_wait)
483 .await?;
484
485 if dont_wait {
486 log::debug!("refresh request sent");
487 return Ok(());
488 }
489
490 log::debug!("refresh request sent, and waiting response");
491
492 tr_res.msg
493 };
494
495 if res.typ.class == CLASS_ERROR_RESPONSE {
496 let mut code = ErrorCodeAttribute::default();
497 let result = code.get_from(&res);
498 if result.is_err() {
499 return Err(Error::Other(format!("{}", res.typ)));
500 } else if code.code == CODE_STALE_NONCE {
501 self.set_nonce_from_msg(&res);
502 return Err(Error::ErrTryAgain);
503 } else {
504 return Ok(());
505 }
506 }
507
508 // Getting lifetime from response
509 let mut updated_lifetime = proto::lifetime::Lifetime::default();
510 updated_lifetime.get_from(&res)?;
511
512 self.lifetime = updated_lifetime.0;
513 log::debug!("updated lifetime: {} seconds", self.lifetime.as_secs());
514 Ok(())
515 }
516
refresh_permissions(&mut self) -> Result<(), Error>517 async fn refresh_permissions(&mut self) -> Result<(), Error> {
518 let addrs = self.perm_map.addrs();
519 if addrs.is_empty() {
520 log::debug!("no permission to refresh");
521 return Ok(());
522 }
523
524 if let Err(err) = self.create_permissions(&addrs).await {
525 if Error::ErrTryAgain != err {
526 log::error!("fail to refresh permissions: {}", err);
527 }
528 return Err(err);
529 }
530
531 log::debug!("refresh permissions successful");
532 Ok(())
533 }
534
bind( rc_obs: Arc<Mutex<T>>, bind_addr: SocketAddr, bind_number: u16, nonce: Nonce, integrity: MessageIntegrity, ) -> Result<(), Error>535 async fn bind(
536 rc_obs: Arc<Mutex<T>>,
537 bind_addr: SocketAddr,
538 bind_number: u16,
539 nonce: Nonce,
540 integrity: MessageIntegrity,
541 ) -> Result<(), Error> {
542 let (msg, turn_server_addr) = {
543 let obs = rc_obs.lock().await;
544
545 let setters: Vec<Box<dyn Setter>> = vec![
546 Box::new(TransactionId::new()),
547 Box::new(MessageType::new(METHOD_CHANNEL_BIND, CLASS_REQUEST)),
548 Box::new(socket_addr2peer_address(&bind_addr)),
549 Box::new(proto::channum::ChannelNumber(bind_number)),
550 Box::new(obs.username()),
551 Box::new(obs.realm()),
552 Box::new(nonce),
553 Box::new(integrity),
554 Box::new(FINGERPRINT),
555 ];
556
557 let mut msg = Message::new();
558 msg.build(&setters)?;
559
560 (msg, obs.turn_server_addr())
561 };
562
563 log::debug!("UDPConn.bind call PerformTransaction 1");
564 let tr_res = {
565 let mut obs = rc_obs.lock().await;
566 obs.perform_transaction(&msg, &turn_server_addr, false)
567 .await?
568 };
569
570 let res = tr_res.msg;
571
572 if res.typ != MessageType::new(METHOD_CHANNEL_BIND, CLASS_SUCCESS_RESPONSE) {
573 return Err(Error::ErrUnexpectedResponse);
574 }
575
576 log::debug!("channel binding successful: {} {}", bind_addr, bind_number);
577
578 // Success.
579 Ok(())
580 }
581 }
582
583 #[async_trait]
584 impl<T: RelayConnObserver + Send + Sync> PeriodicTimerTimeoutHandler for RelayConnInternal<T> {
on_timeout(&mut self, id: TimerIdRefresh)585 async fn on_timeout(&mut self, id: TimerIdRefresh) {
586 log::debug!("refresh timer {:?} expired", id);
587 match id {
588 TimerIdRefresh::Alloc => {
589 let lifetime = self.lifetime;
590 // limit the max retries on errTryAgain to 3
591 // when stale nonce returns, sencond retry should succeed
592 let mut result = Ok(());
593 for _ in 0..MAX_RETRY_ATTEMPTS {
594 result = self.refresh_allocation(lifetime, false).await;
595 if let Err(err) = &result {
596 if Error::ErrTryAgain != *err {
597 break;
598 }
599 }
600 }
601 if result.is_err() {
602 log::warn!("refresh allocation failed");
603 }
604 }
605 TimerIdRefresh::Perms => {
606 let mut result = Ok(());
607 for _ in 0..MAX_RETRY_ATTEMPTS {
608 result = self.refresh_permissions().await;
609 if let Err(err) = &result {
610 if Error::ErrTryAgain != *err {
611 break;
612 }
613 }
614 }
615 if result.is_err() {
616 log::warn!("refresh permissions failed");
617 }
618 }
619 }
620 }
621 }
622
socket_addr2peer_address(addr: &SocketAddr) -> proto::peeraddr::PeerAddress623 fn socket_addr2peer_address(addr: &SocketAddr) -> proto::peeraddr::PeerAddress {
624 proto::peeraddr::PeerAddress {
625 ip: addr.ip(),
626 port: addr.port(),
627 }
628 }
629