From 2e66d111944eef4cbea0a6421eaaa1adf43383a0 Mon Sep 17 00:00:00 2001 From: VirxEC Date: Sat, 15 Aug 2026 22:18:55 -0400 Subject: [PATCH] Refactor event loop I/O for proper mio readiness handling on Windows --- rlbot/src/agents/bot.rs | 144 ++++++++++++++++++++++---- rlbot/src/lib.rs | 32 ++++-- rlbot/tests/event_loop.rs | 208 ++++++++++++++++++++++++++++++++++++++ 3 files changed, 355 insertions(+), 29 deletions(-) create mode 100644 rlbot/tests/event_loop.rs diff --git a/rlbot/src/agents/bot.rs b/rlbot/src/agents/bot.rs index 0ecc649..86ddc20 100644 --- a/rlbot/src/agents/bot.rs +++ b/rlbot/src/agents/bot.rs @@ -1,8 +1,15 @@ -use std::{io::ErrorKind, sync::Arc, thread}; +use std::{ + io::{self, Read, Write}, + sync::Arc, + thread, +}; use mio::Interest; -use crate::{RLBotConnection, RLBotError, StartingInfo, flat::*, pkanal, util::PacketQueue}; +use crate::{ + RLBotConnection, RLBotError, StartingInfo, flat::*, parse_core_message, pkanal, + util::PacketQueue, +}; use super::AgentError; @@ -132,33 +139,67 @@ pub fn run_bot_agents( connection.send_packet(InitComplete {})?; - // Main loop, broadcast packet to all of the bots, then wait for all of the outgoing vecs + // Main loop. Do all socket I/O through `mio_stream`, the handle registered + // with mio. On Windows, mio re-arms the readiness event only after I/O on + // the registered handle returns `WouldBlock`. let mut events = mio::Events::with_capacity(128); + let mut read_buf: Vec = Vec::with_capacity(1024); + let mut out_buf: Vec = Vec::new(); + let mut writable_registered = false; + 'main: loop { poll.poll(&mut events, None) .expect("couldn't poll with mio"); for event in &events { match event.token() { - INCOMING => 'incoming: loop { - let packet = match connection.recv_packet() { - Ok(x) => x, - Err(RLBotError::Connection(e)) if e.kind() == ErrorKind::WouldBlock => { - break 'incoming; - } - Err(e) => Err(e)?, - }; - let packet = Arc::new(packet); - - for (incoming_sender, _) in &threads { - if incoming_sender.send(packet.clone()).is_err() { - return Err(AgentError::AgentPanic); + INCOMING => { + if event.is_writable() && !out_buf.is_empty() { + match flush_pending(&mut mio_stream, &mut out_buf) { + Ok(true) => { + poll.registry() + .reregister(&mut mio_stream, INCOMING, Interest::READABLE) + .expect("couldn't reregister tcp stream"); + writable_registered = false; + } + Ok(false) => {} + Err(e) => return Err(RLBotError::Connection(e).into()), } } + if event.is_readable() { + 'incoming: loop { + match drain_socket(&mut mio_stream, &mut read_buf) { + Ok(false) => {} + Ok(true) => break 'incoming, + Err(e) => return Err(RLBotError::Connection(e).into()), + } + } - if matches!(&*packet, CoreMessage::DisconnectSignal(_)) { - break 'main; + // Broadcast each complete packet. + 'packets: loop { + if read_buf.len() < 2 { + break 'packets; + } + let data_len = u16::from_be_bytes([read_buf[0], read_buf[1]]); + let frame_len = data_len as usize + 2; + if read_buf.len() < frame_len { + break 'packets; + } + let frame: Vec = read_buf.drain(..frame_len).collect(); + + let packet = Arc::new(parse_core_message(&frame[2..])?); + + for (incoming_sender, _) in &threads { + if incoming_sender.send(packet.clone()).is_err() { + return Err(AgentError::AgentPanic); + } + } + + if matches!(&*packet, CoreMessage::DisconnectSignal(_)) { + break 'main; + } + } } - }, + } OUTGOING => 'outgoing: loop { let Ok(maybe_msgs) = outgoing_recver.try_recv() else { break 'main; @@ -168,7 +209,22 @@ pub fn run_bot_agents( break 'outgoing; }; - connection.send_packets_enum(p.into_iter())?; + out_buf.extend(connection.build_interface_messages(p.into_iter())); + + // Send the queued bytes when the socket becomes writable. + // Only arm writability when there are bytes to send: a + // writable edge on an empty buffer would be consumed + // without flushing, and then never re-armed. + if !out_buf.is_empty() && !writable_registered { + poll.registry() + .reregister( + &mut mio_stream, + INCOMING, + Interest::READABLE | Interest::WRITABLE, + ) + .expect("couldn't reregister tcp stream"); + writable_registered = true; + } }, _ => unreachable!(), } @@ -182,6 +238,54 @@ pub fn run_bot_agents( Ok(()) } +/// Read bytes from core through the mio-registered handle. +/// +/// Returns `Ok(true)` when the socket would block. +fn drain_socket(mio_stream: &mut mio::net::TcpStream, read_buf: &mut Vec) -> io::Result { + let mut scratch = [0u8; 8192]; + match mio_stream.read(&mut scratch) { + Ok(n) if n > 0 => { + read_buf.extend_from_slice(&scratch[..n]); + Ok(false) + } + // A zero-byte read means core closed the connection. Return an + // error so the loop does not spin forever. + Ok(_) => Err(io::Error::new( + io::ErrorKind::UnexpectedEof, + "connection to core closed", + )), + Err(e) if e.kind() == io::ErrorKind::WouldBlock => Ok(true), + Err(e) => Err(e), + } +} + +/// Write queued bytes to core through the mio-registered handle. +/// +/// Returns `Ok(true)` when all bytes are sent. +fn flush_pending(mio_stream: &mut mio::net::TcpStream, out_buf: &mut Vec) -> io::Result { + let mut sent = 0; + while sent < out_buf.len() { + match mio_stream.write(&out_buf[sent..]) { + Ok(0) => { + return Err(io::Error::new( + io::ErrorKind::WriteZero, + "socket stopped accepting data", + )); + } + Ok(n) => sent += n, + Err(e) if e.kind() == io::ErrorKind::WouldBlock => break, + Err(e) => return Err(e), + } + } + if sent == out_buf.len() { + out_buf.clear(); + Ok(true) + } else { + out_buf.drain(..sent); + Ok(false) + } +} + fn run_bot_agent( incoming_recver: kanal::Receiver>, team: u32, diff --git a/rlbot/src/lib.rs b/rlbot/src/lib.rs index acadf3f..f786918 100644 --- a/rlbot/src/lib.rs +++ b/rlbot/src/lib.rs @@ -117,17 +117,25 @@ pub struct RLBotConnection { } impl RLBotConnection { - pub(crate) fn send_packets_enum( + /// Build the bytes for outgoing packets without touching the socket. + pub(crate) fn build_interface_messages( &mut self, packets: impl Iterator, - ) -> Result<(), RLBotError> { - let to_write = packets + ) -> Vec { + packets // convert Packet to Vec that RLBotServer can understand .flat_map(|x| { build_packet_payload(GenericMessage::from(x), &mut self.builder) .expect("failed to build packet") }) - .collect::>(); + .collect::>() + } + + pub(crate) fn send_packets_enum( + &mut self, + packets: impl Iterator, + ) -> Result<(), RLBotError> { + let to_write = self.build_interface_messages(packets); self.stream.write_all(&to_write)?; self.stream.flush()?; @@ -159,11 +167,7 @@ impl RLBotConnection { self.stream.read_exact(buf)?; - let packet_ref: CorePacketRef = - CorePacketRef::read_as_root(buf).map_err(PacketParseError::InvalidFlatbuffer)?; - let packet: CorePacket = packet_ref.try_into().unwrap(); - - Ok(packet.message) + parse_core_message(buf) } /// Sets the TCP connection to core to be non-blocking. @@ -217,6 +221,16 @@ impl RLBotConnection { } } +/// Parse a flatbuffer payload into a [`CoreMessage`]. +/// The payload excludes the 2-byte length prefix. +pub(crate) fn parse_core_message(buf: &[u8]) -> Result { + let packet_ref: CorePacketRef = + CorePacketRef::read_as_root(buf).map_err(PacketParseError::InvalidFlatbuffer)?; + let packet: CorePacket = packet_ref.try_into().unwrap(); + + Ok(packet.message) +} + #[derive(Error, Debug)] pub enum PacketBuildError { #[error("Payload too large {0}, couldn't fit in u16")] diff --git a/rlbot/tests/event_loop.rs b/rlbot/tests/event_loop.rs new file mode 100644 index 0000000..a7e4809 --- /dev/null +++ b/rlbot/tests/event_loop.rs @@ -0,0 +1,208 @@ +use std::{ + io::{Read, Write}, + net::{SocketAddr, TcpListener}, + sync::{Arc, mpsc}, + thread, + time::Duration, +}; + +use rlbot::{ + RLBotConnection, + agents::{BotAgent, run_bot_agents}, + flat::{ + ControllableInfo, ControllableTeamInfo, ControllerState, CoreMessage, CorePacket, + FieldInfo, GamePacket, InterfaceMessage, InterfacePacket, InterfacePacketRef, + MatchConfiguration, PlayerInput, + }, + util::PacketQueue, +}; +use rlbot_flat::planus::{self, ReadAsRoot}; + +/// A bot that sends one PlayerInput for every GamePacket it receives. +struct TestBot; + +impl BotAgent for TestBot { + fn new( + _team: u32, + _controllable_info: ControllableInfo, + _match_configuration: Arc, + _field_info: Arc, + _packet_queue: &mut PacketQueue, + ) -> Self { + TestBot + } + + fn tick(&mut self, _game_packet: &GamePacket, packet_queue: &mut PacketQueue) { + packet_queue.push(PlayerInput { + player_index: 0, + controller_state: ControllerState { + throttle: 1.0, + steer: 0.0, + pitch: 0.0, + yaw: 0.0, + roll: 0.0, + jump: false, + boost: false, + handbrake: false, + use_item: false, + }, + }); + } +} + +/// Serialize a CoreMessage into its flatbuffer payload. +fn core_payload(msg: CoreMessage) -> Vec { + let mut builder = planus::Builder::with_capacity(1024); + let packet: CorePacket = msg.into(); + builder.finish(packet, None).to_vec() +} + +/// Prefix a payload with the 2-byte big-endian length. +fn frame(payload: &[u8]) -> Vec { + let mut out = Vec::with_capacity(payload.len() + 2); + out.extend_from_slice(&u16::try_from(payload.len()).unwrap().to_be_bytes()); + out.extend_from_slice(payload); + out +} + +/// Read one complete frame from the stream. +fn read_frame(stream: &mut std::net::TcpStream) -> Result, String> { + let mut len_buf = [0u8; 2]; + stream.read_exact(&mut len_buf).map_err(|e| e.to_string())?; + let len = u16::from_be_bytes(len_buf) as usize; + let mut payload = vec![0u8; len]; + stream.read_exact(&mut payload).map_err(|e| e.to_string())?; + Ok(payload) +} + +fn parse_interface(payload: &[u8]) -> Result { + let packet_ref = InterfacePacketRef::read_as_root(payload).map_err(|e| e.to_string())?; + let packet: InterfacePacket = packet_ref + .try_into() + .map_err(|e: planus::Error| e.to_string())?; + Ok(packet.message) +} + +/// A mock of RLBot core. Speaks the socket protocol: each frame is a 2-byte +/// big-endian length followed by a flatbuffer payload. +fn run_mock_core(listener: TcpListener, num_game_packets: usize) -> Result<(), String> { + let (mut stream, _) = listener.accept().map_err(|e| e.to_string())?; + stream + .set_read_timeout(Some(Duration::from_secs(10))) + .map_err(|e| e.to_string())?; + + // The client sends its ConnectionSettings first. + match parse_interface(&read_frame(&mut stream)?)? { + InterfaceMessage::ConnectionSettings(_) => {} + _ => return Err("expected ConnectionSettings".into()), + } + + // Then it waits for the starting info. + write_frame( + &mut stream, + core_payload(CoreMessage::ControllableTeamInfo(Box::new( + ControllableTeamInfo { + team: 0, + controllables: vec![ControllableInfo { + index: 0, + identifier: 0, + }], + }, + ))), + )?; + write_frame( + &mut stream, + core_payload(CoreMessage::MatchConfiguration(Box::default())), + )?; + write_frame( + &mut stream, + core_payload(CoreMessage::FieldInfo(Box::default())), + )?; + + // The client signals readiness. + match parse_interface(&read_frame(&mut stream)?)? { + InterfaceMessage::InitComplete(_) => {} + _ => return Err("expected InitComplete".into()), + } + + // Send game packets. Split the first frame to exercise partial reads, + // and send the last two in a single write to exercise multiple frames + // per read. + for _ in 0..num_game_packets.saturating_sub(2) { + write_frame( + &mut stream, + core_payload(CoreMessage::GamePacket(Box::default())), + )?; + } + let mut tail = frame(&core_payload(CoreMessage::GamePacket(Box::default()))); + tail.extend(frame(&core_payload( + CoreMessage::GamePacket(Box::default()), + ))); + stream.write_all(&tail).map_err(|e| e.to_string())?; + stream.flush().map_err(|e| e.to_string())?; + + // The client answers each game packet with one PlayerInput. + for _ in 0..num_game_packets { + match parse_interface(&read_frame(&mut stream)?)? { + InterfaceMessage::PlayerInput(_) => {} + _ => return Err("expected PlayerInput".into()), + } + } + + // Disconnect and expect the client to close the connection. + write_frame( + &mut stream, + core_payload(CoreMessage::DisconnectSignal(Box::default())), + )?; + loop { + match read_frame(&mut stream) { + Ok(_) => {} + Err(_) => return Ok(()), + } + } +} + +fn write_frame(stream: &mut std::net::TcpStream, payload: Vec) -> Result<(), String> { + stream + .write_all(&frame(&payload)) + .map_err(|e| e.to_string())?; + stream.flush().map_err(|e| e.to_string()) +} + +#[test] +fn event_loop_round_trip() { + let listener = TcpListener::bind("127.0.0.1:0").unwrap(); + let addr: SocketAddr = listener.local_addr().unwrap(); + + let (tx, rx) = mpsc::channel::<(&'static str, Result<(), String>)>(); + + let server_tx = tx.clone(); + let server_handle = thread::spawn(move || { + let _ = server_tx.send(("server", run_mock_core(listener, 3))); + }); + + let client_tx = tx.clone(); + let client_handle = thread::spawn(move || { + let result = (|| -> Result<(), String> { + let connection = RLBotConnection::new(&addr.to_string()).map_err(|e| e.to_string())?; + run_bot_agents::("test-bot".to_string(), false, false, connection) + .map_err(|e| e.to_string()) + })(); + let _ = client_tx.send(("client", result)); + }); + + let mut results = Vec::new(); + while results.len() < 2 { + let msg = rx + .recv_timeout(Duration::from_secs(30)) + .expect("event loop hung"); + results.push(msg); + } + + for (who, result) in results { + assert!(result.is_ok(), "{who} failed: {result:?}"); + } + + server_handle.join().unwrap(); + client_handle.join().unwrap(); +}