mod beacon;
pub use beacon::*;
use crate::sim::{INPUT_BYTES, InputMsg};
use std::io;
use std::net::{SocketAddr, ToSocketAddrs, UdpSocket};
pub const PROTOCOL_VERSION: u8 = 8;
pub const MAX_PEERS: usize = 9;
const MAX_DATAGRAM: usize = 1024;
pub const MAX_BEACH_BYTES: usize = 832;
pub struct UdpTransport {
socket: UdpSocket,
peers: Vec<SocketAddr>,
accept_new: bool,
max_peers: usize,
}
pub fn local_ip() -> Option<std::net::IpAddr> {
let probe = UdpSocket::bind(("0.0.0.0", 0)).ok()?;
probe.connect(("203.0.113.1", 9)).ok()?;
let ip = probe.local_addr().ok()?.ip();
(!ip.is_unspecified() && !ip.is_loopback()).then_some(ip)
}
impl UdpTransport {
pub fn host(port: u16) -> io::Result<UdpTransport> {
let socket = UdpSocket::bind(("0.0.0.0", port))?;
socket.set_nonblocking(true)?;
Ok(UdpTransport {
socket,
peers: Vec::new(),
accept_new: true,
max_peers: MAX_PEERS,
})
}
pub fn join(host: impl ToSocketAddrs) -> io::Result<UdpTransport> {
let socket = UdpSocket::bind(("0.0.0.0", 0))?;
socket.set_nonblocking(true)?;
let peer = host
.to_socket_addrs()?
.next()
.ok_or_else(|| io::Error::new(io::ErrorKind::InvalidInput, "unresolvable host"))?;
Ok(UdpTransport {
socket,
peers: vec![peer],
accept_new: false,
max_peers: 1,
})
}
pub fn connected(&self) -> bool {
!self.peers.is_empty()
}
pub fn peer_count(&self) -> usize {
self.peers.len()
}
pub fn local_addr(&self) -> io::Result<SocketAddr> {
self.socket.local_addr()
}
pub fn peer_addr(&self) -> Option<SocketAddr> {
match self.peers.as_slice() {
[only] => Some(*only),
_ => None,
}
}
pub fn send(&self, msg: NetMsg) {
let bytes = msg.encode();
for peer in &self.peers {
let _ = self.socket.send_to(&bytes, peer);
}
}
#[cfg(test)]
fn send_raw(&self, bytes: &[u8]) {
for peer in &self.peers {
let _ = self.socket.send_to(bytes, peer);
}
}
pub fn forget(&mut self, index: usize) {
if index < self.peers.len() {
self.peers.remove(index);
}
}
pub fn send_to(&self, index: usize, msg: NetMsg) {
debug_assert!(
index < self.peers.len(),
"sending to peer {index} of {}",
self.peers.len()
);
if let Some(peer) = self.peers.get(index) {
let _ = self.socket.send_to(&msg.encode(), peer);
}
}
pub fn recv_all(&mut self) -> Vec<(NetMsg, usize)> {
let mut out = Vec::new();
let mut buf = [0u8; MAX_DATAGRAM];
loop {
match self.socket.recv_from(&mut buf) {
Ok((len, from)) => {
let Some(msg) = NetMsg::decode(&buf[..len]) else {
if NetMsg::peek_version(&buf[..len])
.is_some_and(|version| version != PROTOCOL_VERSION)
{
let reply = NetMsg::Incompatible {
version: PROTOCOL_VERSION,
};
let _ = self.socket.send_to(&reply.encode(), from);
}
continue;
};
let index = match self.peers.iter().position(|p| *p == from) {
Some(index) => index,
None if self.accept_new && self.peers.len() < self.max_peers => {
self.peers.push(from);
self.peers.len() - 1
}
None => continue,
};
out.push((msg, index));
}
Err(e) if e.kind() == io::ErrorKind::WouldBlock => break,
Err(_) => break,
}
}
out
}
}
mod msg;
mod wire;
pub use msg::*;
pub use wire::*;
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_mismatched_greeting_is_answered_not_swallowed() {
let mut host = UdpTransport::host(0).expect("bind");
let port = host.local_addr().expect("addr").port();
let mut joiner = UdpTransport::join(("127.0.0.1", port)).expect("join");
let mut hello = NetMsg::hello("Anna").encode();
hello[1] = PROTOCOL_VERSION.wrapping_add(1); let mut answer = None;
for _ in 0..40 {
joiner.send_raw(&hello);
std::thread::sleep(std::time::Duration::from_millis(5));
assert!(host.recv_all().is_empty(), "nothing the host can act on");
if let Some((msg, _)) = joiner.recv_all().into_iter().next() {
answer = Some(msg);
break;
}
}
assert_eq!(
answer,
Some(NetMsg::Incompatible {
version: PROTOCOL_VERSION
})
);
assert_eq!(host.peer_count(), 0, "and it claimed no seat at the table");
}
#[test]
fn host_accepts_peers_up_to_its_cap() {
let mut host = UdpTransport::host(0).expect("bind");
let port = host.local_addr().expect("addr").port();
let joiners: Vec<UdpTransport> = (0..MAX_PEERS + 2)
.map(|_| UdpTransport::join(("127.0.0.1", port)).expect("join"))
.collect();
for joiner in &joiners {
joiner.send(NetMsg::hello(""));
}
for _ in 0..40 {
std::thread::sleep(std::time::Duration::from_millis(5));
let _ = host.recv_all();
if host.peer_count() == MAX_PEERS {
break;
}
}
assert_eq!(
host.peer_count(),
MAX_PEERS,
"the ones past the cap are refused"
);
}
}