diff --git a/rnex-core/src/executables/regular_backend.rs b/rnex-core/src/executables/regular_backend.rs index 7f62864..f618560 100644 --- a/rnex-core/src/executables/regular_backend.rs +++ b/rnex-core/src/executables/regular_backend.rs @@ -1,8 +1,14 @@ use std::sync::{Arc, atomic::AtomicU32}; +use tokio::sync::{Mutex, mpsc::channel}; + use crate::{ executables::common::new_simple_backend, - nex::{matchmake::MatchmakeManager, remote_console::RemoteConsole, user::User}, + nex::{ + matchmake::MatchmakeManager, + remote_console::RemoteConsole, + user::{ConnectionTicket, User}, + }, rmc::protocols::RmcPureRemoteObject, }; @@ -21,13 +27,31 @@ pub async fn start_regular_backend() { new_simple_backend(move |c, r| { let mmm = mmm.clone(); - Arc::new_cyclic(move |this| User { - this: this.clone(), - ip: c.prudpsock_addr, - pid: c.pid, - remote: RemoteConsole::new(r), - matchmake_manager: mmm, - station_url: Default::default(), + Arc::new_cyclic(move |this| { + let (join_tickets_stage1_sender, join_tickets_stage1_recv) = + channel::(100); + let join_tickets_stage1_recv = Mutex::new(join_tickets_stage1_recv); + + let (join_tickets_stage2_sender, join_tickets_stage2_recv) = + channel::(100); + let join_tickets_stage2_recv = Mutex::new(join_tickets_stage2_recv); + let cid = mmm.next_cid(); + + User { + cid, + this: this.clone(), + ip: c.prudpsock_addr, + pid: c.pid, + remote: RemoteConsole::new(r), + matchmake_manager: mmm, + station_url: Default::default(), + join_tickets_stage1_recv, + join_tickets_stage1_sender, + join_tickets_stage2_recv, + join_tickets_stage2_sender, + self_join_ticket_requesters: Default::default(), + remote_join_ticket_requesters: Default::default(), + } }) }) .await; diff --git a/rnex-core/src/nex/matchmake.rs b/rnex-core/src/nex/matchmake.rs index 1b12f04..ef699e1 100644 --- a/rnex-core/src/nex/matchmake.rs +++ b/rnex-core/src/nex/matchmake.rs @@ -1,4 +1,4 @@ -use log::info; +use log::{info, warn}; use rand::random; use rnex_core::PID; use rnex_core::nex::user::User; @@ -19,7 +19,9 @@ use std::sync::atomic::Ordering::Relaxed; use std::sync::{Arc, Weak}; use std::time::Duration; use tokio::sync::{Mutex, RwLock}; -use tokio::time::sleep; +use tokio::time::{sleep, timeout}; + +use crate::rmc::protocols::nat_traversal::RemoteNatTraversalConsole; pub struct MatchmakeManager { //pub gid_counter: AtomicU32, @@ -313,6 +315,12 @@ impl ExtendedMatchmakeSession { >= self.session.gathering.minimum_participants as _ } + #[inline] + pub fn get_host(&self) -> Option> { + self.get_active_players() + .find(|v| v.pid == self.session.gathering.host_pid) + } + #[inline] pub fn is_reachable(&self) -> bool { self.get_active_players() @@ -340,6 +348,89 @@ impl ExtendedMatchmakeSession { self.is_reachable() && is_open } + pub async fn is_joinable_by(&self, user: Arc) -> bool { + let Some(host) = self.get_host() else { + return false; + }; + + let Some(user_station_url) = user.station_url.read().await.first().cloned() else { + return false; + }; + let Some(host_station_url) = host.station_url.read().await.first().cloned() else { + return false; + }; + + let mut tickets_requesters = host.remote_join_ticket_requesters.lock().await; + tickets_requesters.insert(user.cid, Arc::downgrade(&user)); + drop(tickets_requesters); + + host.remote + .request_probe_initiation(user_station_url.to_string()) + .await; + + let Some(_) = timeout(Duration::from_secs(5), async { + loop { + let mut stage1_recv = user.join_tickets_stage1_recv.lock().await; + + let Some(ticket) = stage1_recv.recv().await else { + return None; + }; + + if ticket.cid != host.cid { + user.join_tickets_stage1_sender.send(ticket).await.ok(); + drop(stage1_recv); + warn!("got incorrect ticket sleeping for 500 millis whilest leaving ticket reciever open for use"); + + sleep(Duration::from_millis(500)).await; + continue; + } + + return Some(ticket); + } + }) + .await + .ok() + .flatten() else { + return false; + }; + + let mut ticket_requesters = user.self_join_ticket_requesters.lock().await; + ticket_requesters.insert(host.cid); + drop(ticket_requesters); + + user.remote + .request_probe_initiation(host_station_url.to_string()) + .await; + + let Some(stage2_ticket) = timeout(Duration::from_secs(5), async { + loop { + let mut stage2_recv = user.join_tickets_stage2_recv.lock().await; + + let Some(ticket) = stage2_recv.recv().await else { + return None; + }; + + if ticket.cid != host.cid { + user.join_tickets_stage2_sender.send(ticket).await.ok(); + drop(stage2_recv); + warn!("got incorrect ticket sleeping for 500 millis whilest leaving ticket reciever open for use"); + + sleep(Duration::from_millis(500)).await; + continue; + } + + return Some(ticket); + } + }) + .await + .ok() + .flatten() else { + return false; + }; + + stage2_ticket.result + } + pub fn matches_criteria( &self, search_criteria: &MatchmakeSessionSearchCriteria, diff --git a/rnex-core/src/nex/user.rs b/rnex-core/src/nex/user.rs index d00c5df..5d8a3d4 100644 --- a/rnex-core/src/nex/user.rs +++ b/rnex-core/src/nex/user.rs @@ -1,3 +1,5 @@ +use futures::future::join_all; +use log::warn; use rnex_core::PID; use rnex_core::define_rmc_proto; use rnex_core::kerberos::KerberosDateTime; @@ -33,8 +35,12 @@ use rnex_core::rmc::structures::matchmake::{ AutoMatchmakeParam, CreateMatchmakeSessionParam, JoinMatchmakeSessionParam, MatchmakeSession, }; use serde::{Deserialize, Serialize}; +use std::collections::HashMap; +use std::collections::HashSet; use std::env; use std::str::FromStr; +use tokio::sync::mpsc::Receiver; +use tokio::sync::mpsc::Sender; use cfg_if::cfg_if; use log::{error, info}; @@ -54,6 +60,7 @@ use rnex_core::rmc::structures::ranking::UploadCompetitionData; use std::sync::{Arc, Weak}; use tokio::sync::{Mutex, RwLock}; +use crate::kerberos::Ticket; use crate::rmc::protocols::message_delivery::RemoteMessageDeliveryNoResponse; use crate::rmc::protocols::messaging::UserMessage; use crate::rmc::structures::matchmake::Gathering; @@ -90,14 +97,29 @@ cfg_if! { } } +/// Connection tickets are allowances to join a specific lobby, they are given out as soon as nat checks pass, +/// there are 2 stages of tickets because both sides have to do nat checking before we let the player join +/// the lobby +pub struct ConnectionTicket { + pub cid: u32, + pub result: bool, +} + #[rmc_struct(UserProtocol)] pub struct User { pub pid: PID, + pub cid: u32, pub ip: PRUDPSockAddr, pub this: Weak, pub remote: RemoteConsole, pub station_url: RwLock>, pub matchmake_manager: Arc, + pub remote_join_ticket_requesters: Mutex>>, + pub self_join_ticket_requesters: Mutex>, + pub join_tickets_stage1_sender: Sender, + pub join_tickets_stage1_recv: Mutex>, + pub join_tickets_stage2_sender: Sender, + pub join_tickets_stage2_recv: Mutex>, } impl Secure for User { @@ -105,8 +127,7 @@ impl Secure for User { &self, station_urls: Vec, ) -> Result<(QResult, u32, StationUrl), ErrorCode> { - let cid = self.matchmake_manager.next_cid(); - + let cid = self.cid; println!("{:?}", station_urls); let mut users = self.matchmake_manager.users.write().await; @@ -364,6 +385,21 @@ impl MatchmakeExtension for User { if bool_matched_criteria { println!("matched session: {:?}", session); + let is_joinable_by_all = join_all( + joining_players + .iter() + .filter_map(|f| f.upgrade()) + .map(|v| session.is_joinable_by(v)), + ) + .await + .iter() + .copied() + .fold(true, |a, b| a || b); + if is_joinable_by_all { + warn!( + "tripped unreachable host detection for one of the users who were trying to join" + ); + } session .add_players(&joining_players, param.join_message) .await; @@ -709,10 +745,33 @@ impl NatTraversal for User { async fn report_nat_traversal_result( &self, - _cid: u32, - _result: bool, + cid: u32, + result: bool, _rtt: u32, ) -> Result<(), ErrorCode> { + if let Some(user) = self + .remote_join_ticket_requesters + .lock() + .await + .remove(&cid) + .map(|u| u.upgrade()) + .flatten() + { + user.join_tickets_stage1_sender + .send(ConnectionTicket { + cid: self.cid, + result, + }) + .await + .ok(); + } + if let Some(user) = self.self_join_ticket_requesters.lock().await.take(&cid) { + self.join_tickets_stage2_sender + .send(ConnectionTicket { cid, result }) + .await + .ok(); + } + Ok(()) }