1use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
41use std::sync::Arc;
42use std::time::Duration;
43
44use futures_util::stream::SplitSink;
45use futures_util::stream::SplitStream;
46use futures_util::{SinkExt, StreamExt};
47use tokio::net::TcpStream;
48use tokio::sync::Mutex;
49use tokio_tungstenite::tungstenite::protocol::frame::coding::CloseCode;
50use tokio_tungstenite::tungstenite::protocol::frame::CloseFrame;
51use tokio_tungstenite::tungstenite::protocol::WebSocketConfig;
52use tokio_tungstenite::tungstenite::Message;
53use tokio_tungstenite::WebSocketStream;
54
55use crate::network::{NetworkProvider, PeerAddr};
56use crate::session::hello::{
57 HelloEnvelope, HelloKind, PeerIdentity, CLOSE_APP_MISMATCH, CLOSE_HELLO_PROTOCOL,
58 CLOSE_IDENTITY_MISMATCH, HELLO_TIMEOUT,
59};
60
61use super::{
62 resolve_dial_addr, FramedStream, StreamListener, StreamTransport, TransportError, WsConfig,
63};
64
65pub struct WebSocketTransport<N: NetworkProvider> {
85 network: Arc<N>,
87 config: WsConfig,
89}
90
91impl<N: NetworkProvider + 'static> WebSocketTransport<N> {
92 pub fn new(network: Arc<N>, config: WsConfig) -> Self {
97 Self { network, config }
98 }
99
100 fn local_hello(&self) -> HelloEnvelope {
107 let identity = self.network.local_identity();
108 HelloEnvelope::new(PeerIdentity {
109 app_id: identity.app_id,
110 device_id: identity.device_id,
111 device_name: identity.device_name,
112 os: std::env::consts::OS.to_string(),
113 tailscale_id: identity.tailscale_id,
114 })
115 }
116
117 fn ws_protocol_config(&self) -> WebSocketConfig {
119 let mut config = WebSocketConfig::default();
120 config.max_message_size = Some(self.config.max_message_size);
121 config.max_frame_size = Some(self.config.max_message_size);
122 config
123 }
124}
125
126async fn close_ws_with_code(ws: &mut WebSocketStream<TcpStream>, code: u16, reason: &str) {
133 let close_frame = CloseFrame {
134 code: CloseCode::from(code),
135 reason: reason.to_string().into(),
136 };
137 let _ = ws.send(Message::Close(Some(close_frame))).await;
138 let _ = ws.close(None).await;
139}
140
141const MAX_CONTROL_FRAMES_BEFORE_HELLO: usize = 16;
146
147async fn receive_hello(
157 ws: &mut WebSocketStream<TcpStream>,
158) -> Result<HelloEnvelope, TransportError> {
159 let fut = async {
160 let mut control_frame_count: usize = 0;
161 loop {
162 match ws.next().await {
163 Some(Ok(Message::Text(text))) => {
164 return serde_json::from_str::<HelloEnvelope>(&text)
165 .map_err(|e| TransportError::HelloMalformed(format!("parse hello: {e}")));
166 }
167 Some(Ok(Message::Binary(data))) => {
168 return serde_json::from_slice::<HelloEnvelope>(&data)
169 .map_err(|e| TransportError::HelloMalformed(format!("parse hello: {e}")));
170 }
171 Some(Ok(Message::Ping(_) | Message::Pong(_) | Message::Frame(_))) => {
172 control_frame_count += 1;
173 if control_frame_count > MAX_CONTROL_FRAMES_BEFORE_HELLO {
174 return Err(TransportError::HelloMalformed(
175 "too many control frames before hello".to_string(),
176 ));
177 }
178 continue;
179 }
180 Some(Ok(Message::Close(_))) => {
181 return Err(TransportError::HelloMalformed(
182 "peer closed connection before hello".to_string(),
183 ));
184 }
185 Some(Err(e)) => {
186 return Err(TransportError::HelloMalformed(format!(
187 "receive hello: {e}"
188 )));
189 }
190 None => {
191 return Err(TransportError::HelloMalformed(
192 "connection closed before hello".to_string(),
193 ));
194 }
195 }
196 }
197 };
198
199 match tokio::time::timeout(HELLO_TIMEOUT, fut).await {
200 Ok(inner) => inner,
201 Err(_) => Err(TransportError::HelloTimeout),
202 }
203}
204
205const MAX_APP_ID_LEN: usize = 32;
209const MAX_DEVICE_NAME_LEN: usize = 512;
210const MAX_DEVICE_ID_LEN: usize = 64;
211const MAX_TAILSCALE_ID_LEN: usize = 256;
212const MAX_OS_LEN: usize = 32;
213
214fn validate_hello(
217 remote: HelloEnvelope,
218 local_app_id: &str,
219) -> Result<PeerIdentity, TransportError> {
220 if remote.kind != HelloKind::Hello {
221 return Err(TransportError::HelloMalformed(format!(
222 "unexpected hello kind: {:?}",
223 remote.kind
224 )));
225 }
226 if remote.version < HelloEnvelope::MIN_SUPPORTED_VERSION {
227 return Err(TransportError::HelloMalformed(format!(
228 "unsupported hello version {} (minimum supported: {})",
229 remote.version,
230 HelloEnvelope::MIN_SUPPORTED_VERSION
231 )));
232 }
233 if remote.identity.app_id.len() > MAX_APP_ID_LEN
236 || remote.identity.device_name.len() > MAX_DEVICE_NAME_LEN
237 || remote.identity.device_id.len() > MAX_DEVICE_ID_LEN
238 || remote.identity.tailscale_id.len() > MAX_TAILSCALE_ID_LEN
239 || remote.identity.os.len() > MAX_OS_LEN
240 {
241 return Err(TransportError::HelloMalformed(
242 "identity field exceeds maximum allowed length".to_string(),
243 ));
244 }
245 if remote.identity.app_id != local_app_id {
246 return Err(TransportError::AppMismatch {
247 local: local_app_id.to_string(),
248 remote: remote.identity.app_id.clone(),
249 });
250 }
251 Ok(remote.identity)
252}
253
254#[derive(Debug, Default, serde::Deserialize)]
257struct AuthenticatedIdentity {
258 #[serde(default, rename = "nodeId")]
259 node_id: String,
260}
261
262fn verify_authenticated_identity(
280 claimed: &PeerIdentity,
281 authenticated: &str,
282) -> Result<(), TransportError> {
283 if authenticated.trim().is_empty() {
284 tracing::warn!(
285 claimed_tailscale_id = %claimed.tailscale_id,
286 "ws: no authenticated identity for incoming connection (mock/loopback or WhoIs failure); accepting hello claim unverified"
287 );
288 return Ok(());
289 }
290
291 let parsed = match serde_json::from_str::<AuthenticatedIdentity>(authenticated) {
292 Ok(parsed) => parsed,
293 Err(_) => {
294 tracing::warn!(
295 claimed_tailscale_id = %claimed.tailscale_id,
296 "ws: unparseable authenticated identity; accepting hello claim unverified"
297 );
298 return Ok(());
299 }
300 };
301
302 if parsed.node_id.is_empty() {
303 tracing::warn!(
304 claimed_tailscale_id = %claimed.tailscale_id,
305 "ws: authenticated identity carried no nodeId; accepting hello claim unverified"
306 );
307 return Ok(());
308 }
309
310 if parsed.node_id == claimed.tailscale_id {
311 Ok(())
312 } else {
313 Err(TransportError::IdentityMismatch {
314 claimed: claimed.tailscale_id.clone(),
315 authenticated: parsed.node_id,
316 })
317 }
318}
319
320async fn send_hello(
322 ws: &mut WebSocketStream<TcpStream>,
323 hello: &HelloEnvelope,
324) -> Result<(), TransportError> {
325 let payload =
326 serde_json::to_string(hello).map_err(|e| TransportError::Serialize(e.to_string()))?;
327 ws.send(Message::Text(payload.into()))
328 .await
329 .map_err(|e| TransportError::HandshakeFailed(format!("send hello: {e}")))
330}
331
332async fn client_hello_exchange(
336 ws: &mut WebSocketStream<TcpStream>,
337 local_hello: &HelloEnvelope,
338) -> Result<PeerIdentity, TransportError> {
339 if let Err(e) = send_hello(ws, local_hello).await {
341 close_ws_with_code(ws, CLOSE_HELLO_PROTOCOL, "local send failed").await;
342 return Err(e);
343 }
344
345 let remote = match receive_hello(ws).await {
347 Ok(envelope) => envelope,
348 Err(e) => {
349 match &e {
350 TransportError::HelloTimeout => {
351 close_ws_with_code(ws, CLOSE_HELLO_PROTOCOL, "hello timeout").await;
352 }
353 TransportError::HelloMalformed(_) => {
354 close_ws_with_code(ws, CLOSE_HELLO_PROTOCOL, "malformed hello").await;
355 }
356 _ => {}
357 }
358 return Err(e);
359 }
360 };
361
362 match validate_hello(remote, &local_hello.identity.app_id) {
364 Ok(identity) => Ok(identity),
365 Err(e) => {
366 match &e {
367 TransportError::AppMismatch { .. } => {
368 close_ws_with_code(ws, CLOSE_APP_MISMATCH, "app mismatch").await;
369 }
370 TransportError::HelloMalformed(_) => {
371 close_ws_with_code(ws, CLOSE_HELLO_PROTOCOL, "bad hello").await;
372 }
373 _ => {}
374 }
375 Err(e)
376 }
377 }
378}
379
380async fn server_hello_exchange(
389 ws: &mut WebSocketStream<TcpStream>,
390 local_hello: &HelloEnvelope,
391 authenticated_identity: &str,
392) -> Result<PeerIdentity, TransportError> {
393 let remote = match receive_hello(ws).await {
395 Ok(envelope) => envelope,
396 Err(e) => {
397 match &e {
398 TransportError::HelloTimeout => {
399 close_ws_with_code(ws, CLOSE_HELLO_PROTOCOL, "hello timeout").await;
400 }
401 TransportError::HelloMalformed(_) => {
402 close_ws_with_code(ws, CLOSE_HELLO_PROTOCOL, "malformed hello").await;
403 }
404 _ => {}
405 }
406 return Err(e);
407 }
408 };
409
410 let identity = match validate_hello(remote, &local_hello.identity.app_id) {
412 Ok(identity) => identity,
413 Err(e) => {
414 match &e {
415 TransportError::AppMismatch { .. } => {
416 close_ws_with_code(ws, CLOSE_APP_MISMATCH, "app mismatch").await;
417 }
418 TransportError::HelloMalformed(_) => {
419 close_ws_with_code(ws, CLOSE_HELLO_PROTOCOL, "bad hello").await;
420 }
421 _ => {}
422 }
423 return Err(e);
424 }
425 };
426
427 if let Err(e) = verify_authenticated_identity(&identity, authenticated_identity) {
432 close_ws_with_code(ws, CLOSE_IDENTITY_MISMATCH, "identity mismatch").await;
433 return Err(e);
434 }
435
436 if let Err(e) = send_hello(ws, local_hello).await {
438 close_ws_with_code(ws, CLOSE_HELLO_PROTOCOL, "local send failed").await;
439 return Err(e);
440 }
441
442 Ok(identity)
443}
444
445impl<N: NetworkProvider + 'static> StreamTransport for WebSocketTransport<N> {
446 type Stream = WsFramedStream;
447
448 async fn connect(&self, addr: &PeerAddr) -> Result<Self::Stream, TransportError> {
449 let dial_addr = resolve_dial_addr(addr);
451 tracing::debug!(addr = %dial_addr, port = self.config.port, "ws: dialing peer");
452
453 let tcp_stream = self
454 .network
455 .dial_tcp(&dial_addr, self.config.port)
456 .await
457 .map_err(|e| TransportError::ConnectFailed(format!("dial tcp: {e}")))?;
458
459 let ws_url = format!("ws://{dial_addr}:{}/ws", self.config.port);
461 let ws_config = self.ws_protocol_config();
462 let (mut ws, _response) =
463 tokio_tungstenite::client_async_with_config(ws_url, tcp_stream, Some(ws_config))
464 .await
465 .map_err(|e| TransportError::ConnectFailed(format!("ws upgrade: {e}")))?;
466
467 let local_hello = self.local_hello();
469 let remote_identity = match client_hello_exchange(&mut ws, &local_hello).await {
470 Ok(identity) => identity,
471 Err(e) => {
472 match &e {
473 TransportError::AppMismatch { local, remote } => {
474 tracing::info!(
475 local_app_id = %local,
476 remote_app_id = %remote,
477 "ws: closing connection — app_id mismatch"
478 );
479 }
480 TransportError::HelloTimeout => {
481 tracing::warn!("ws: hello timeout on outgoing connection");
482 }
483 TransportError::HelloMalformed(msg) => {
484 tracing::warn!(error = %msg, "ws: malformed hello on outgoing connection");
485 }
486 _ => {}
487 }
488 return Err(e);
489 }
490 };
491
492 tracing::info!(
493 remote_device_id = %remote_identity.device_id,
494 remote_device_name = %remote_identity.device_name,
495 remote_tailscale_id = %remote_identity.tailscale_id,
496 "ws: connected (hello exchanged)"
497 );
498
499 Ok(WsFramedStream::new(
501 ws,
502 remote_identity.tailscale_id.clone(),
503 Some(remote_identity),
504 dial_addr,
505 self.config.ping_interval,
506 self.config.pong_timeout,
507 ))
508 }
509
510 async fn listen(&self) -> Result<StreamListener<Self::Stream>, TransportError> {
511 let port = self.config.port;
512 tracing::debug!(port, "ws: starting listener");
513
514 let mut tcp_listener = self
516 .network
517 .listen_tcp(port)
518 .await
519 .map_err(|e| TransportError::ListenFailed(format!("listen tcp: {e}")))?;
520
521 let (tx, rx) = tokio::sync::mpsc::channel::<WsFramedStream>(64);
523 let local_hello = self.local_hello();
524 let ping_interval = self.config.ping_interval;
525 let pong_timeout = self.config.pong_timeout;
526 let ws_config = self.ws_protocol_config();
527 let handshake_timeout = self.config.handshake_timeout;
528 let handshake_permits = Arc::new(tokio::sync::Semaphore::new(
529 self.config.max_pending_handshakes,
530 ));
531
532 tokio::spawn(async move {
533 loop {
534 match tcp_listener.incoming.recv().await {
535 Some(incoming) => {
536 let permit = match handshake_permits.clone().try_acquire_owned() {
542 Ok(permit) => permit,
543 Err(_) => {
544 tracing::warn!(
545 remote = %incoming.remote_addr,
546 "ws: max pending handshakes reached; dropping incoming connection"
547 );
548 continue;
549 }
550 };
551
552 let tx = tx.clone();
553 let local_hello = local_hello.clone();
554 let remote_addr = incoming.remote_addr.clone();
555 let authenticated_identity = incoming.remote_identity.clone();
558 let ws_config = ws_config;
559
560 tokio::spawn(async move {
561 let handshake = async {
566 let mut ws = tokio_tungstenite::accept_async_with_config(
567 incoming.stream,
568 Some(ws_config),
569 )
570 .await
571 .map_err(|e| {
572 TransportError::HandshakeFailed(format!("ws upgrade: {e}"))
573 })?;
574 let identity = server_hello_exchange(
575 &mut ws,
576 &local_hello,
577 &authenticated_identity,
578 )
579 .await?;
580 Ok::<(WebSocketStream<TcpStream>, PeerIdentity), TransportError>((
581 ws, identity,
582 ))
583 };
584
585 let (ws, remote_identity) = match tokio::time::timeout(
586 handshake_timeout,
587 handshake,
588 )
589 .await
590 {
591 Ok(Ok(pair)) => pair,
592 Ok(Err(e)) => {
593 match &e {
594 TransportError::AppMismatch { local, remote } => {
595 tracing::info!(
596 remote_addr = %remote_addr,
597 local_app_id = %local,
598 remote_app_id = %remote,
599 "ws: closing incoming connection — app_id mismatch"
600 );
601 }
602 TransportError::HelloTimeout => {
603 tracing::warn!(
604 remote_addr = %remote_addr,
605 "ws: hello timeout on incoming connection"
606 );
607 }
608 TransportError::HelloMalformed(msg) => {
609 tracing::warn!(
610 remote_addr = %remote_addr,
611 error = %msg,
612 "ws: malformed hello on incoming connection"
613 );
614 }
615 TransportError::IdentityMismatch {
616 claimed,
617 authenticated,
618 } => {
619 tracing::warn!(
620 remote_addr = %remote_addr,
621 claimed = %claimed,
622 authenticated = %authenticated,
623 "ws: closing incoming connection — hello identity does not match authenticated Tailscale identity"
624 );
625 }
626 TransportError::HandshakeFailed(msg) => {
627 tracing::warn!(
628 remote = %remote_addr,
629 error = %msg,
630 "ws: upgrade failed"
631 );
632 }
633 _ => {
634 tracing::warn!(
635 remote_addr = %remote_addr,
636 "ws: hello exchange failed: {e}"
637 );
638 }
639 }
640 return;
641 }
642 Err(_) => {
643 tracing::warn!(
644 remote_addr = %remote_addr,
645 timeout = ?handshake_timeout,
646 "ws: handshake timed out before hello; dropping connection"
647 );
648 return;
649 }
650 };
651
652 drop(permit);
656
657 tracing::info!(
658 remote_device_id = %remote_identity.device_id,
659 remote_device_name = %remote_identity.device_name,
660 remote_tailscale_id = %remote_identity.tailscale_id,
661 remote_addr = %remote_addr,
662 "ws: accepted connection (hello exchanged)"
663 );
664
665 let stream = WsFramedStream::new(
666 ws,
667 remote_identity.tailscale_id.clone(),
668 Some(remote_identity),
669 remote_addr,
670 ping_interval,
671 pong_timeout,
672 );
673
674 if tx.send(stream).await.is_err() {
675 tracing::debug!("ws: listener channel closed");
676 }
677 });
678 }
679 None => {
680 tracing::debug!("ws: tcp listener channel closed");
681 break;
682 }
683 }
684 }
685 });
686
687 Ok(StreamListener::new(rx, port))
688 }
689}
690
691pub struct WsFramedStream {
708 write: Arc<Mutex<SplitSink<WebSocketStream<TcpStream>, Message>>>,
710 read: SplitStream<WebSocketStream<TcpStream>>,
712 remote_peer_id: String,
715 remote_identity: Option<PeerIdentity>,
717 remote_addr: String,
719 heartbeat_handle: Option<tokio::task::JoinHandle<()>>,
721 last_pong: Arc<AtomicU64>,
723 closed: Arc<AtomicBool>,
725}
726
727impl std::fmt::Debug for WsFramedStream {
728 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
729 f.debug_struct("WsFramedStream")
730 .field("remote_peer_id", &self.remote_peer_id)
731 .field("remote_addr", &self.remote_addr)
732 .field("closed", &self.closed.load(Ordering::Relaxed))
733 .finish_non_exhaustive()
734 }
735}
736
737#[allow(unsafe_code)]
743unsafe impl Sync for WsFramedStream {}
744
745fn epoch_millis() -> u64 {
747 std::time::SystemTime::now()
748 .duration_since(std::time::UNIX_EPOCH)
749 .unwrap_or_default()
750 .as_millis() as u64
751}
752
753impl WsFramedStream {
754 fn new(
756 ws: WebSocketStream<TcpStream>,
757 remote_peer_id: String,
758 remote_identity: Option<PeerIdentity>,
759 remote_addr: String,
760 ping_interval: Duration,
761 pong_timeout: Duration,
762 ) -> Self {
763 let (write, read) = ws.split();
764 let write = Arc::new(Mutex::new(write));
765 let last_pong = Arc::new(AtomicU64::new(epoch_millis()));
766 let closed = Arc::new(AtomicBool::new(false));
767
768 let hb_write = write.clone();
770 let hb_last_pong = last_pong.clone();
771 let hb_closed = closed.clone();
772 let hb_addr = remote_addr.clone();
773 let heartbeat_handle = tokio::spawn(async move {
774 heartbeat_loop(
775 hb_write,
776 hb_last_pong,
777 hb_closed,
778 ping_interval,
779 pong_timeout,
780 &hb_addr,
781 )
782 .await;
783 });
784
785 Self {
786 write,
787 read,
788 remote_peer_id,
789 remote_identity,
790 remote_addr,
791 heartbeat_handle: Some(heartbeat_handle),
792 last_pong,
793 closed,
794 }
795 }
796
797 pub fn remote_peer_id(&self) -> &str {
802 &self.remote_peer_id
803 }
804
805 pub fn remote_identity(&self) -> Option<&PeerIdentity> {
810 self.remote_identity.as_ref()
811 }
812}
813
814async fn heartbeat_loop(
824 write: Arc<Mutex<SplitSink<WebSocketStream<TcpStream>, Message>>>,
825 last_pong: Arc<AtomicU64>,
826 closed: Arc<AtomicBool>,
827 ping_interval: Duration,
828 pong_timeout: Duration,
829 remote_addr: &str,
830) {
831 let mut interval = tokio::time::interval(ping_interval);
832 interval.tick().await;
834
835 loop {
836 interval.tick().await;
837
838 if closed.load(Ordering::Acquire) {
840 return;
841 }
842
843 let last = last_pong.load(Ordering::Acquire);
845 let now = epoch_millis();
846 let elapsed = Duration::from_millis(now.saturating_sub(last));
847
848 if elapsed > pong_timeout {
849 tracing::warn!(
850 remote = %remote_addr,
851 elapsed = ?elapsed,
852 "heartbeat: pong timeout after {pong_timeout:?}"
853 );
854 closed.store(true, Ordering::Release);
856 let mut w = write.lock().await;
857 let _ = w.close().await;
858 return;
859 }
860
861 {
863 let mut w = write.lock().await;
864 let ping_data = b"truffle-ping".to_vec();
865 if let Err(e) = w.send(Message::Ping(ping_data.into())).await {
866 tracing::debug!(remote = %remote_addr, "heartbeat: ping send failed: {e}");
867 closed.store(true, Ordering::Release);
868 return;
869 }
870 }
871 }
872}
873
874impl FramedStream for WsFramedStream {
875 async fn send(&mut self, data: &[u8]) -> Result<(), TransportError> {
876 if self.closed.load(Ordering::Acquire) {
877 return Err(TransportError::ConnectionClosed(
878 "connection already closed".to_string(),
879 ));
880 }
881 let mut w = self.write.lock().await;
882 w.send(Message::Binary(data.to_vec().into()))
883 .await
884 .map_err(|e| TransportError::WebSocket(format!("send: {e}")))
885 }
886
887 async fn recv(&mut self) -> Result<Option<Vec<u8>>, TransportError> {
888 if self.closed.load(Ordering::Acquire) {
889 return Ok(None);
890 }
891 loop {
892 match self.read.next().await {
893 Some(Ok(Message::Binary(data))) => return Ok(Some(data.to_vec())),
894 Some(Ok(Message::Text(text))) => {
895 return Ok(Some(text.as_bytes().to_vec()));
897 }
898 Some(Ok(Message::Ping(_))) => {
899 continue;
902 }
903 Some(Ok(Message::Pong(_))) => {
904 self.last_pong.store(epoch_millis(), Ordering::Release);
906 continue;
907 }
908 Some(Ok(Message::Close(_))) => {
909 self.closed.store(true, Ordering::Release);
910 return Ok(None);
911 }
912 Some(Ok(Message::Frame(_))) => {
913 continue;
915 }
916 Some(Err(e)) => {
917 self.closed.store(true, Ordering::Release);
918 return Err(TransportError::WebSocket(format!("recv: {e}")));
919 }
920 None => {
921 self.closed.store(true, Ordering::Release);
922 return Ok(None);
923 }
924 }
925 }
926 }
927
928 async fn close(&mut self) -> Result<(), TransportError> {
929 if let Some(handle) = self.heartbeat_handle.take() {
931 handle.abort();
932 }
933
934 self.closed.store(true, Ordering::Release);
935
936 let mut w = self.write.lock().await;
937 w.close()
938 .await
939 .map_err(|e| TransportError::WebSocket(format!("close: {e}")))
940 }
941
942 fn peer_addr(&self) -> String {
943 self.remote_addr.clone()
944 }
945}
946
947impl Drop for WsFramedStream {
948 fn drop(&mut self) {
949 if let Some(handle) = self.heartbeat_handle.take() {
951 handle.abort();
952 }
953 }
954}
955
956#[cfg(test)]
961mod unit_tests {
962 use super::*;
963
964 #[test]
965 fn resolve_dial_addr_prefers_ip() {
966 let addr = PeerAddr {
967 ip: Some("100.64.0.1".parse().unwrap()),
968 hostname: "peer".to_string(),
969 dns_name: Some("peer.tailnet.ts.net".to_string()),
970 };
971 assert_eq!(resolve_dial_addr(&addr), "100.64.0.1");
972 }
973
974 #[test]
975 fn resolve_dial_addr_falls_back_to_dns() {
976 let addr = PeerAddr {
977 ip: None,
978 hostname: "peer".to_string(),
979 dns_name: Some("peer.tailnet.ts.net".to_string()),
980 };
981 assert_eq!(resolve_dial_addr(&addr), "peer.tailnet.ts.net");
982 }
983
984 #[test]
985 fn resolve_dial_addr_falls_back_to_hostname() {
986 let addr = PeerAddr {
987 ip: None,
988 hostname: "peer".to_string(),
989 dns_name: None,
990 };
991 assert_eq!(resolve_dial_addr(&addr), "peer");
992 }
993
994 fn valid_identity() -> crate::session::hello::PeerIdentity {
997 crate::session::hello::PeerIdentity {
998 app_id: "playground".to_string(),
999 device_id: "01J4K9M2Z8AB3RNYQPW6H5TC0X".to_string(),
1000 device_name: "Alice's MacBook".to_string(),
1001 os: "darwin".to_string(),
1002 tailscale_id: "n1234567890.ts-node".to_string(),
1003 }
1004 }
1005
1006 #[test]
1007 fn validate_hello_accepts_normal_identity() {
1008 let envelope = crate::session::hello::HelloEnvelope::new(valid_identity());
1009 let result = validate_hello(envelope, "playground");
1010 assert!(result.is_ok(), "baseline valid hello should validate");
1011 }
1012
1013 #[test]
1014 fn validate_hello_rejects_oversized_app_id() {
1015 let mut identity = valid_identity();
1016 identity.app_id = "a".repeat(MAX_APP_ID_LEN + 1);
1017 let envelope = crate::session::hello::HelloEnvelope::new(identity);
1018 let result = validate_hello(envelope, "playground");
1019 match result {
1020 Err(TransportError::HelloMalformed(msg)) => {
1021 assert!(
1022 msg.contains("maximum allowed length"),
1023 "unexpected error: {msg}"
1024 );
1025 }
1026 other => panic!("expected HelloMalformed, got {other:?}"),
1027 }
1028 }
1029
1030 #[test]
1031 fn validate_hello_rejects_oversized_device_name() {
1032 let mut identity = valid_identity();
1033 identity.device_name = "x".repeat(MAX_DEVICE_NAME_LEN + 1);
1034 let envelope = crate::session::hello::HelloEnvelope::new(identity);
1035 let result = validate_hello(envelope, "playground");
1036 assert!(matches!(result, Err(TransportError::HelloMalformed(_))));
1037 }
1038
1039 #[test]
1040 fn validate_hello_rejects_oversized_device_id() {
1041 let mut identity = valid_identity();
1042 identity.device_id = "x".repeat(MAX_DEVICE_ID_LEN + 1);
1043 let envelope = crate::session::hello::HelloEnvelope::new(identity);
1044 let result = validate_hello(envelope, "playground");
1045 assert!(matches!(result, Err(TransportError::HelloMalformed(_))));
1046 }
1047
1048 #[test]
1049 fn validate_hello_rejects_oversized_tailscale_id() {
1050 let mut identity = valid_identity();
1051 identity.tailscale_id = "x".repeat(MAX_TAILSCALE_ID_LEN + 1);
1052 let envelope = crate::session::hello::HelloEnvelope::new(identity);
1053 let result = validate_hello(envelope, "playground");
1054 assert!(matches!(result, Err(TransportError::HelloMalformed(_))));
1055 }
1056
1057 #[test]
1058 fn validate_hello_rejects_oversized_os() {
1059 let mut identity = valid_identity();
1060 identity.os = "x".repeat(MAX_OS_LEN + 1);
1061 let envelope = crate::session::hello::HelloEnvelope::new(identity);
1062 let result = validate_hello(envelope, "playground");
1063 assert!(matches!(result, Err(TransportError::HelloMalformed(_))));
1064 }
1065
1066 #[test]
1067 fn validate_hello_length_check_runs_before_app_id_mismatch() {
1068 let mut identity = valid_identity();
1073 identity.device_name = "x".repeat(MAX_DEVICE_NAME_LEN + 1);
1074 let envelope = crate::session::hello::HelloEnvelope::new(identity);
1075 let result = validate_hello(envelope, "playground");
1076 assert!(
1077 matches!(result, Err(TransportError::HelloMalformed(_))),
1078 "matched app_id must not let an oversized field through"
1079 );
1080 }
1081
1082 #[test]
1085 fn verify_identity_rejects_mismatched_node_id() {
1086 let identity = valid_identity();
1089 let authenticated = r#"{"dnsName":"mallory.ts.net","nodeId":"real-node-id"}"#;
1090 match verify_authenticated_identity(&identity, authenticated) {
1091 Err(TransportError::IdentityMismatch {
1092 claimed,
1093 authenticated,
1094 }) => {
1095 assert_eq!(claimed, "n1234567890.ts-node");
1096 assert_eq!(authenticated, "real-node-id");
1097 }
1098 other => panic!("expected IdentityMismatch, got {other:?}"),
1099 }
1100 }
1101
1102 #[test]
1103 fn verify_identity_accepts_matching_node_id() {
1104 let identity = valid_identity();
1105 let authenticated = r#"{"dnsName":"alice.ts.net","nodeId":"n1234567890.ts-node"}"#;
1106 assert!(
1107 verify_authenticated_identity(&identity, authenticated).is_ok(),
1108 "matching nodeId must be accepted"
1109 );
1110 }
1111
1112 #[test]
1113 fn verify_identity_accepts_empty_authenticated_string() {
1114 let identity = valid_identity();
1117 assert!(
1118 verify_authenticated_identity(&identity, "").is_ok(),
1119 "empty authenticated identity must fall open"
1120 );
1121 }
1122
1123 #[test]
1124 fn verify_identity_accepts_non_json_legacy_dns_value() {
1125 let identity = valid_identity();
1128 assert!(
1129 verify_authenticated_identity(&identity, "peer.tailnet.ts.net").is_ok(),
1130 "legacy plain-DNS-name authenticated value must fall open"
1131 );
1132 }
1133
1134 #[test]
1135 fn verify_identity_accepts_json_without_node_id() {
1136 let identity = valid_identity();
1138 let authenticated = r#"{"dnsName":"peer.ts.net"}"#;
1139 assert!(
1140 verify_authenticated_identity(&identity, authenticated).is_ok(),
1141 "missing nodeId must fall open"
1142 );
1143 }
1144
1145 #[test]
1146 fn verify_identity_rejects_empty_claim_against_real_node_id() {
1147 let mut identity = valid_identity();
1151 identity.tailscale_id = String::new();
1152 let authenticated = r#"{"nodeId":"real-node-id"}"#;
1153 match verify_authenticated_identity(&identity, authenticated) {
1154 Err(TransportError::IdentityMismatch {
1155 claimed,
1156 authenticated,
1157 }) => {
1158 assert_eq!(claimed, "");
1159 assert_eq!(authenticated, "real-node-id");
1160 }
1161 other => panic!("expected IdentityMismatch, got {other:?}"),
1162 }
1163 }
1164}