1use std::{
22 collections::{HashMap, VecDeque},
23 net::SocketAddr,
24 sync::Arc,
25 time::Instant,
26};
27
28use ana_gotatun::{
29 noise::{Tunn, TunnResult, handshake::parse_handshake_anon, rate_limiter::RateLimiter},
30 packet::{Packet, WgKind},
31 x25519,
32};
33
34pub struct SnapTunServer<T: SnapTunAuthorization> {
85 static_private: x25519::StaticSecret,
86 static_public: x25519::PublicKey,
87 active_tunnels: HashMap<SocketAddr, ActiveTunnel>,
88 rate_limiter: Arc<RateLimiter>,
89 authz: Arc<T>,
90}
91
92struct ActiveTunnel {
93 peer_static: x25519::PublicKey,
94 tunn: Tunn,
95}
96
97pub enum HandleIncomingPacketResult<S> {
103 Result {
106 result: TunnResult,
109 },
110 Forwarded {
113 packet: Packet,
115 processed_at: Instant,
117 session_data: Arc<S>,
119 },
120}
121
122impl<S> HandleIncomingPacketResult<S> {
123 pub fn into_result(self) -> TunnResult {
125 match self {
126 HandleIncomingPacketResult::Result { result } => result,
127 HandleIncomingPacketResult::Forwarded { packet, .. } => {
128 TunnResult::WriteToTunnel(packet)
129 }
130 }
131 }
132}
133
134pub struct HandleOutgoingPacketResult<S> {
137 pub network_packet: Option<WgKind>,
140 pub processed_at: Instant,
143 pub session_data: Arc<S>,
145}
146
147impl<S> HandleOutgoingPacketResult<S> {
148 pub fn into_packet(self) -> Option<WgKind> {
150 self.network_packet
151 }
152}
153
154impl<T: SnapTunAuthorization> SnapTunServer<T> {
155 fn incoming_packet_result(
158 result: TunnResult,
159 session_data: Arc<T::SessionData>,
160 now: Instant,
161 ) -> HandleIncomingPacketResult<T::SessionData> {
162 match result {
163 TunnResult::WriteToTunnel(packet) => {
164 HandleIncomingPacketResult::Forwarded {
165 packet,
166 processed_at: now,
167 session_data,
168 }
169 }
170 result => HandleIncomingPacketResult::Result { result },
171 }
172 }
173
174 fn outgoing_packet_result(
175 network_packet: Option<WgKind>,
176 now: Instant,
177 session_data: Arc<T::SessionData>,
178 ) -> HandleOutgoingPacketResult<T::SessionData> {
179 HandleOutgoingPacketResult {
180 network_packet,
181 processed_at: now,
182 session_data,
183 }
184 }
185
186 pub fn new(
188 static_private: x25519::StaticSecret,
189 rate_limiter: Arc<RateLimiter>,
190 authz: Arc<T>,
191 ) -> Self {
192 let static_public = x25519::PublicKey::from(&static_private);
193 Self {
194 static_private,
195 static_public,
196 active_tunnels: Default::default(),
197 rate_limiter,
198 authz,
199 }
200 }
201
202 pub fn handle_incoming_packet(
216 &mut self,
217 packet: Packet,
218 from: SocketAddr,
219 send_to_network: &mut VecDeque<WgKind>,
220 ) -> TunnResult {
221 self.handle_incoming_packet_with_session(packet, from, send_to_network)
222 .into_result()
223 }
224
225 #[tracing::instrument(skip_all, fields(remote = %from))]
232 pub fn handle_incoming_packet_with_session(
233 &mut self,
234 packet: Packet,
235 from: SocketAddr,
236 send_to_network: &mut VecDeque<WgKind>,
237 ) -> HandleIncomingPacketResult<T::SessionData> {
238 let parsed_packet = match self.rate_limiter.verify_packet(from.ip(), packet) {
239 Ok(p) => p,
240 Err(TunnResult::WriteToNetwork(c)) => {
241 tracing::debug!(remote = ?from, "rate limiter issued cookie reply");
242 send_to_network.push_back(c);
243 return HandleIncomingPacketResult::Result {
244 result: TunnResult::Done,
245 };
246 }
247 Err(e) => {
248 tracing::debug!(remote = ?from, err = ?e, "rate limiter rejected packet");
249 return HandleIncomingPacketResult::Result { result: e };
250 }
251 };
252 let packet_now = Instant::now();
256
257 use std::collections::hash_map::Entry;
258
259 use ana_gotatun::noise::errors::WireGuardError;
260 match (self.active_tunnels.entry(from), parsed_packet) {
261 (Entry::Occupied(mut occupied_entry), p) => {
262 let active_tunnel = occupied_entry.get_mut();
263 let Some(session_data) = self
270 .authz
271 .is_authorized(packet_now, active_tunnel.peer_static.as_bytes())
272 else {
273 tracing::debug!(remote = ?from, peer_static = ?active_tunnel.peer_static, "rejected packet from unauthorized peer");
274 return HandleIncomingPacketResult::Result {
275 result: TunnResult::Err(WireGuardError::UnexpectedPacket),
276 };
277 };
278 let result = Self::handle_incoming_and_drain_queue(
279 send_to_network,
280 p,
281 &mut active_tunnel.tunn,
282 );
283 Self::incoming_packet_result(result, session_data, packet_now)
284 }
285 (e, WgKind::HandshakeInit(wg_init)) => {
286 let peer = match parse_handshake_anon(
287 &self.static_private,
288 &self.static_public,
289 &wg_init,
290 ) {
291 Ok(v) => v,
292 Err(e) => {
293 tracing::debug!(remote = ?from, err = ?e, "failed to parse handshake init");
294 return HandleIncomingPacketResult::Result {
295 result: TunnResult::from(e),
296 };
297 }
298 };
299
300 let Some(session_data) = self
306 .authz
307 .is_authorized(packet_now, &peer.peer_static_public)
308 else {
309 tracing::debug!(remote = ?from, "rejected handshake from unauthorized peer");
310 return HandleIncomingPacketResult::Result {
311 result: TunnResult::Err(WireGuardError::UnexpectedPacket),
312 };
313 };
314 tracing::debug!(remote = ?from, "accepted new handshake, inserting tunnel");
315 let peer_static = x25519::PublicKey::from(peer.peer_static_public);
316 let mut tunn = Tunn::new(
317 self.static_private.clone(),
318 peer_static,
319 None,
320 None,
321 0,
322 self.rate_limiter.clone(),
323 from,
324 );
325 let res = Self::handle_incoming_and_drain_queue(
326 send_to_network,
327 WgKind::HandshakeInit(wg_init),
328 &mut tunn,
329 );
330 let handled = Self::incoming_packet_result(res, session_data.clone(), packet_now);
333 e.insert_entry(ActiveTunnel { peer_static, tunn });
334 handled
335 }
336 (_, _p) => {
337 tracing::debug!(remote = ?from, "received unexpected packet kind for new entry");
338 HandleIncomingPacketResult::Result {
339 result: TunnResult::Err(WireGuardError::InvalidPacket),
340 }
341 }
342 }
343 }
344
345 pub fn handle_outgoing_packet(&mut self, packet: Packet, to: SocketAddr) -> Option<WgKind> {
348 self.handle_outgoing_packet_with_session(packet, to)
349 .and_then(HandleOutgoingPacketResult::into_packet)
350 }
351
352 #[tracing::instrument(skip_all, fields(remote = %to))]
359 pub fn handle_outgoing_packet_with_session(
360 &mut self,
361 packet: Packet,
362 to: SocketAddr,
363 ) -> Option<HandleOutgoingPacketResult<T::SessionData>> {
364 let Some(active_tunnel) = self.active_tunnels.get_mut(&to) else {
365 tracing::error!(to=?to, "No tunnel for outgoing packet found.");
366 return None;
367 };
368 let packet_now = Instant::now();
369 let Some(session_data) = self
370 .authz
371 .is_authorized(packet_now, active_tunnel.peer_static.as_bytes())
372 else {
373 tracing::debug!(remote = ?to, peer_static = ?active_tunnel.peer_static, "dropping outgoing packet for unauthorized peer");
374 return None;
375 };
376 Some(Self::outgoing_packet_result(
377 active_tunnel
378 .tunn
379 .handle_outgoing_packet(packet.into_bytes()),
380 packet_now,
381 session_data,
382 ))
383 }
384
385 pub fn update_timers(&mut self) -> Vec<(SocketAddr, WgKind)> {
391 let mut res = vec![];
392 self.active_tunnels.retain(|k, active_tunnel| {
393 match active_tunnel.tunn.update_timers() {
394 Ok(Some(wg)) => res.push((*k, wg)),
395 Ok(None) => {},
396 Err(e) => tracing::error!(err=?e, remote_sockaddr=?k, "error when updating timers on tunnel"),
397 }
398
399 !active_tunnel.tunn.is_expired()
400 });
401 res
402 }
403
404 fn handle_incoming_and_drain_queue(
405 q: &mut VecDeque<WgKind>,
406 p: WgKind,
407 tunn: &mut Tunn,
408 ) -> TunnResult {
409 let r = match tunn.handle_incoming_packet(p) {
410 TunnResult::WriteToNetwork(p) => {
411 q.push_back(p);
412 TunnResult::Done
413 }
414 TunnResult::WriteToTunnel(p) if p.is_empty() => TunnResult::Done,
416 r => r,
417 };
418 for p in tunn.get_queued_packets() {
419 q.push_back(p);
420 }
421 r
422 }
423}
424
425pub trait SnapTunAuthorization: Send + Sync {
427 type SessionData: Clone + Send + Sync + 'static;
429
430 fn is_authorized(&self, now: Instant, identity: &[u8; 32]) -> Option<Arc<Self::SessionData>>;
432}
433
434#[cfg(test)]
435mod tests {
436 use std::{
437 collections::{HashMap, VecDeque},
438 net::SocketAddr,
439 sync::{Arc, Mutex},
440 };
441
442 use ana_gotatun::{
443 noise::{Tunn, TunnResult, rate_limiter::RateLimiter},
444 packet::{IpNextProtocol, Packet, WgKind},
445 x25519,
446 };
447 use zerocopy::IntoBytes;
448
449 use crate::{
450 scion_packet::{Scion, ScionHeader},
451 server::{
452 HandleIncomingPacketResult, HandleOutgoingPacketResult, SnapTunAuthorization,
453 SnapTunServer,
454 },
455 };
456
457 type ResultT = Result<(), Box<dyn std::error::Error>>;
458
459 struct TrivialAuthz;
460
461 impl SnapTunAuthorization for TrivialAuthz {
462 type SessionData = ();
463
464 fn is_authorized(
465 &self,
466 _now: std::time::Instant,
467 _ident: &[u8; 32],
468 ) -> Option<Arc<Self::SessionData>> {
469 Some(Arc::new(()))
470 }
471 }
472
473 #[derive(Debug, Clone, PartialEq, Eq)]
474 struct MutableSessionData {
475 jti: &'static str,
476 pssid: &'static str,
477 tags: Vec<(&'static str, &'static str)>,
478 }
479
480 #[derive(Default)]
481 struct MutableAuthz {
482 sessions: Mutex<HashMap<[u8; 32], Arc<MutableSessionData>>>,
483 }
484
485 impl MutableAuthz {
486 fn set_session_data(&self, identity: [u8; 32], session_data: MutableSessionData) {
487 self.sessions
488 .lock()
489 .unwrap()
490 .insert(identity, Arc::new(session_data));
491 }
492 }
493
494 impl SnapTunAuthorization for MutableAuthz {
495 type SessionData = MutableSessionData;
496
497 fn is_authorized(
498 &self,
499 _now: std::time::Instant,
500 ident: &[u8; 32],
501 ) -> Option<Arc<Self::SessionData>> {
502 self.sessions.lock().unwrap().get(ident).cloned()
503 }
504 }
505
506 fn test_packet<const N: usize>(payload: [u8; N]) -> Packet {
507 let packet = Scion {
508 header: ScionHeader::new(
509 0,
510 0xAA,
511 0xABCDE,
512 payload.len() as _,
513 IpNextProtocol::Udp,
514 7,
515 0x0123_4567_89AB_CDEF,
516 0xFEDC_BA98_7654_3210,
517 ),
518 payload,
519 };
520 Packet::copy_from(packet.as_bytes())
521 }
522
523 fn establish_tunnel<T: SnapTunAuthorization>(
524 snaptun_server: &mut SnapTunServer<T>,
525 tunn_client: &mut Tunn,
526 packet: &Packet,
527 sockaddr_client: SocketAddr,
528 send_to_network: &mut VecDeque<WgKind>,
529 ) {
530 let Some(WgKind::HandshakeInit(hs_init)) =
531 tunn_client.handle_outgoing_packet(Packet::copy_from(packet))
532 else {
533 panic!("expected handshake init")
534 };
535
536 snaptun_server.handle_incoming_packet(
537 Packet::copy_from(hs_init.as_bytes()),
538 sockaddr_client,
539 send_to_network,
540 );
541 dispatch_one(tunn_client, send_to_network);
542 }
543
544 #[test]
545 fn connect_with_multiple_clients() -> ResultT {
546 let sockaddr_client0: SocketAddr = "192.168.1.1:1234".parse().unwrap();
547 let static_client0 = x25519::StaticSecret::from([0u8; 32]);
548 let sockaddr_client1: SocketAddr = "192.168.1.2:4321".parse().unwrap();
549 let static_client1 = x25519::StaticSecret::from([1u8; 32]);
550 let sockaddr_server: SocketAddr = "10.0.0.1:5001".parse().unwrap();
551 let static_server = x25519::StaticSecret::from([2u8; 32]);
552 let static_server_public = x25519::PublicKey::from(&static_server);
553
554 let rate_limiter = Arc::new(RateLimiter::new(&static_server_public, 100));
555 let mut snaptun_server =
556 SnapTunServer::new(static_server, rate_limiter.clone(), Arc::new(TrivialAuthz));
557
558 let mut send_to_network = VecDeque::<WgKind>::new();
559
560 let test_packet0 = test_packet([b'T', b'E', b'S', b'T', b'0']);
561 let test_packet1 = test_packet([b'T', b'E', b'S', b'T', b'1']);
562
563 let mut tunn_client0 = Tunn::new(
564 static_client0,
565 static_server_public,
566 None,
567 None,
568 0,
569 rate_limiter.clone(),
570 sockaddr_server,
571 );
572
573 let mut tunn_client1 = Tunn::new(
574 static_client1,
575 static_server_public,
576 None,
577 None,
578 0,
579 rate_limiter,
580 sockaddr_server,
581 );
582
583 establish_tunnel(
585 &mut snaptun_server,
586 &mut tunn_client0,
587 &test_packet0,
588 sockaddr_client0,
589 &mut send_to_network,
590 );
591 assert_eq!(
592 tunn_client0.get_initiator_remote_sockaddr(),
593 Some(sockaddr_client0)
594 );
595
596 establish_tunnel(
598 &mut snaptun_server,
599 &mut tunn_client1,
600 &test_packet1,
601 sockaddr_client1,
602 &mut send_to_network,
603 );
604 assert_eq!(
605 tunn_client1.get_initiator_remote_sockaddr(),
606 Some(sockaddr_client1)
607 );
608
609 let Some(WgKind::Data(p)) = tunn_client0.get_queued_packets().next() else {
611 panic!("expected packet to be queued");
612 };
613
614 let TunnResult::WriteToTunnel(p) = snaptun_server.handle_incoming_packet(
615 Packet::copy_from(p.as_bytes()),
616 sockaddr_client0,
617 &mut send_to_network,
618 ) else {
619 panic!("Expected packet to be processed")
620 };
621 assert_eq!(p.as_bytes(), test_packet0.as_bytes());
622
623 let Some(WgKind::Data(p1)) = tunn_client1.get_queued_packets().next() else {
627 panic!("expected packet to be queued");
628 };
629
630 let TunnResult::WriteToTunnel(p1) = snaptun_server.handle_incoming_packet(
631 Packet::copy_from(p1.as_bytes()),
632 sockaddr_client1,
633 &mut send_to_network,
634 ) else {
635 panic!("expected packet to be received on server side");
636 };
637 assert_eq!(p1.as_bytes(), test_packet1.as_bytes());
638
639 let res = snaptun_server.handle_outgoing_packet(p, sockaddr_client1);
641 let Some(p @ WgKind::Data(_)) = res else {
642 panic!("expected packet to be sent back to client")
643 };
644
645 let TunnResult::WriteToTunnel(p) = tunn_client1.handle_incoming_packet(p) else {
646 panic!("expected packet to be sent back to client")
647 };
648
649 assert_eq!(p.as_bytes(), test_packet0.as_bytes());
650
651 Ok(())
652 }
653
654 #[test]
655 fn outgoing_packet_with_session_returns_active_session() {
656 let sockaddr_client: SocketAddr = "192.168.1.1:1234".parse().unwrap();
657 let static_client = x25519::StaticSecret::from([0u8; 32]);
658 let sockaddr_server: SocketAddr = "10.0.0.1:5001".parse().unwrap();
659 let static_server = x25519::StaticSecret::from([2u8; 32]);
660 let static_server_public = x25519::PublicKey::from(&static_server);
661
662 let rate_limiter = Arc::new(RateLimiter::new(&static_server_public, 100));
663 let mut snaptun_server =
664 SnapTunServer::new(static_server, rate_limiter.clone(), Arc::new(TrivialAuthz));
665 let mut send_to_network = VecDeque::<WgKind>::new();
666
667 let test_packet = test_packet([b'T', b'E', b'S', b'T']);
668
669 let mut tunn_client = Tunn::new(
670 static_client,
671 static_server_public,
672 None,
673 None,
674 0,
675 rate_limiter,
676 sockaddr_server,
677 );
678
679 establish_tunnel(
680 &mut snaptun_server,
681 &mut tunn_client,
682 &test_packet,
683 sockaddr_client,
684 &mut send_to_network,
685 );
686
687 let Some(WgKind::Data(client_data)) = tunn_client.get_queued_packets().next() else {
688 panic!("expected packet to be queued");
689 };
690 let TunnResult::WriteToTunnel(server_plaintext) = snaptun_server.handle_incoming_packet(
691 Packet::copy_from(client_data.as_bytes()),
692 sockaddr_client,
693 &mut send_to_network,
694 ) else {
695 panic!("expected packet to be processed")
696 };
697
698 let handled = snaptun_server
699 .handle_outgoing_packet_with_session(server_plaintext, sockaddr_client)
700 .expect("expected packet to be encapsulated");
701 let HandleOutgoingPacketResult {
702 network_packet: Some(WgKind::Data(encapsulated)),
703 processed_at: _,
704 session_data,
705 } = handled
706 else {
707 panic!("expected encapsulated data packet")
708 };
709 assert_eq!(session_data.as_ref(), &());
710
711 let TunnResult::WriteToTunnel(plaintext) =
712 tunn_client.handle_incoming_packet(WgKind::Data(encapsulated))
713 else {
714 panic!("expected packet to be delivered back to client")
715 };
716 assert_eq!(plaintext.as_bytes(), test_packet.as_bytes());
717 }
718
719 #[test]
720 fn established_tunnel_refreshes_session_data_for_later_packets() {
721 let sockaddr_client: SocketAddr = "192.168.1.1:1234".parse().unwrap();
722 let static_client = x25519::StaticSecret::from([0u8; 32]);
723 let client_identity = x25519::PublicKey::from(&static_client);
724 let sockaddr_server: SocketAddr = "10.0.0.1:5001".parse().unwrap();
725 let static_server = x25519::StaticSecret::from([2u8; 32]);
726 let static_server_public = x25519::PublicKey::from(&static_server);
727 let rate_limiter = Arc::new(RateLimiter::new(&static_server_public, 100));
728 let authz = Arc::new(MutableAuthz::default());
729 let original_session = MutableSessionData {
730 jti: "original-jti",
731 pssid: "original-pssid",
732 tags: vec![("subject_id", "subject-1"), ("scope", "basic")],
733 };
734 authz.set_session_data(*client_identity.as_bytes(), original_session);
735
736 let mut snaptun_server =
737 SnapTunServer::new(static_server, rate_limiter.clone(), authz.clone());
738 let mut send_to_network = VecDeque::<WgKind>::new();
739 let test_packet = test_packet([b'T', b'E', b'S', b'T']);
740
741 let mut tunn_client = Tunn::new(
742 static_client,
743 static_server_public,
744 None,
745 None,
746 0,
747 rate_limiter,
748 sockaddr_server,
749 );
750
751 establish_tunnel(
752 &mut snaptun_server,
753 &mut tunn_client,
754 &test_packet,
755 sockaddr_client,
756 &mut send_to_network,
757 );
758
759 let refreshed_session = MutableSessionData {
760 jti: "refreshed-jti",
761 pssid: "refreshed-pssid",
762 tags: vec![("subject_id", "subject-2"), ("scope", "premium")],
763 };
764 authz.set_session_data(*client_identity.as_bytes(), refreshed_session.clone());
765
766 let Some(WgKind::Data(client_data)) = tunn_client.get_queued_packets().next() else {
767 panic!("expected packet to be queued");
768 };
769 let HandleIncomingPacketResult::Forwarded {
770 packet: server_plaintext,
771 session_data,
772 ..
773 } = snaptun_server.handle_incoming_packet_with_session(
774 Packet::copy_from(client_data.as_bytes()),
775 sockaddr_client,
776 &mut send_to_network,
777 )
778 else {
779 panic!("expected forwarded packet with refreshed session data")
780 };
781 assert_eq!(session_data.as_ref(), &refreshed_session);
782
783 let Some(HandleOutgoingPacketResult {
784 network_packet: Some(WgKind::Data(encapsulated)),
785 processed_at: _,
786 session_data,
787 }) = snaptun_server.handle_outgoing_packet_with_session(server_plaintext, sockaddr_client)
788 else {
789 panic!("expected encapsulated data packet with refreshed session data")
790 };
791 assert_eq!(session_data.as_ref(), &refreshed_session);
792
793 let TunnResult::WriteToTunnel(plaintext) =
794 tunn_client.handle_incoming_packet(WgKind::Data(encapsulated))
795 else {
796 panic!("expected packet to be delivered back to client")
797 };
798 assert_eq!(plaintext.as_bytes(), test_packet.as_bytes());
799 }
800
801 #[test]
802 fn outgoing_packet_with_session_returns_none_without_tunnel() {
803 let sockaddr_client: SocketAddr = "192.168.1.1:1234".parse().unwrap();
804 let static_server = x25519::StaticSecret::from([2u8; 32]);
805 let static_server_public = x25519::PublicKey::from(&static_server);
806 let rate_limiter = Arc::new(RateLimiter::new(&static_server_public, 100));
807 let mut snaptun_server =
808 SnapTunServer::new(static_server, rate_limiter, Arc::new(TrivialAuthz));
809
810 let payload = [b'T', b'E', b'S', b'T'];
811 let test_packet = Scion {
812 header: ScionHeader::new(
813 0,
814 0xAA,
815 0xABCDE,
816 payload.len() as _,
817 IpNextProtocol::Udp,
818 7,
819 0x0123_4567_89AB_CDEF,
820 0xFEDC_BA98_7654_3210,
821 ),
822 payload,
823 };
824
825 assert!(
826 snaptun_server
827 .handle_outgoing_packet_with_session(
828 Packet::copy_from(test_packet.as_bytes()),
829 sockaddr_client
830 )
831 .is_none()
832 );
833 }
834
835 fn dispatch_one(tunn: &mut Tunn, packets: &mut VecDeque<WgKind>) -> TunnResult {
836 if let Some(packet) = packets.pop_front() {
837 return tunn.handle_incoming_packet(packet);
838 }
839 TunnResult::Done
840 }
841}