跨不同Tokio Runtime实例移动TcpStream是否安全?
跨Tokio Runtime迁移TcpStream是否安全?
我打算开发一款游戏服务器,计划让多个分片(shard)运行在不同线程上,每个线程都拥有一个current_thread Runtime。现在的疑问是:将TcpStream从一个Runtime实例移动到另一个是否安全? 我想通过这种方式实现玩家在不同分片间的迁移。
因为对Tokio Runtime内部机制了解有限,我有这些顾虑:
- 如果用C语言实现,每个线程需要各自的epoll/iocp实例,套接字在线程间移动时需要添加/移除,操作起来并不容易,还可能引发难以排查的隐性bug。
- 虽然如果Tokio用Linux的io_uring(操作可排队且无需注册套接字)的话,这件事可能会很简单,但我不确定实际情况是否如此(毕竟Tokio还要支持Windows的IOCP)。
我写了一个小型测试服务器,编译运行都正常,但还是对内部处理逻辑存疑,没法100%确认这种操作的安全性。查了Tokio文档,只找到mpsc通道可以跨Runtime使用(测试代码里也用了它),但没找到关于套接字的相关说明。
use tokio::{ io::{AsyncReadExt, AsyncWriteExt}, net::{TcpListener, TcpStream}, sync::mpsc, }; use std::{ collections::{HashMap, hash_map::Entry}, sync::{Arc, Mutex}, }; struct TransferMap { inner: Mutex<HashMap<u64, mpsc::UnboundedSender<TcpStream>>>, } impl TransferMap { fn new() -> Self { Self { inner: Mutex::new(HashMap::new()) } } fn add_tx(&self, tx: mpsc::UnboundedSender<TcpStream>) -> u64 { let mut inner = self.inner.lock().unwrap(); loop { let id = rand::random::<u64>(); if let Entry::Vacant(entry) = inner.entry(id) { entry.insert(tx); return id; } } } fn transfer(&self, id: u64, s: TcpStream) { let inner = self.inner.lock().unwrap(); if let Some(tx) = inner.get(&id) { tx.send(s).unwrap(); } } fn transfer_any(&self, s: TcpStream) { let inner = self.inner.lock().unwrap(); if inner.len() > 0 { let index = rand::random::<usize>() % inner.len(); let tx = inner.values().nth(index).unwrap(); tx.send(s).unwrap(); } } } fn spawn_shard(map: &Arc<TransferMap>){ let (connection_tx, connection_rx) = mpsc::unbounded_channel(); let shard_id = map.add_tx(connection_tx); let map = map.clone(); std::thread::Builder::new() .name(format!("shard#{shard_id}")) .spawn(move || shard_main(map, shard_id, connection_rx)) .expect("failed to spawn shard thread"); } fn server_run(map: Arc<TransferMap>) { tokio::runtime::Builder::new_current_thread() .enable_all() .build() .expect("failed to create listener async runtime") .block_on(async { let listener = TcpListener::bind("127.0.0.1:7777").await .expect("failed to bind listener"); println!("server listening on {}", listener.local_addr().unwrap()); loop { match listener.accept().await { Ok((s, _addr)) => map.transfer_any(s), Err(e) => eprintln!("accept error {:?}", e), } } }); } async fn shard_handle_connection( map: Arc<TransferMap>, shard_id: u64, mut s: TcpStream) { let mut buf = [0u8; 1024]; match s.read(&mut buf).await { Ok(n) => { let message = std::str::from_utf8(&buf[..n]).unwrap_or(""); println!("shard#{} message from {}: {}", shard_id, s.peer_addr().unwrap(), message); if let Err(e) = s.write_all(&buf[..n]).await { eprintln!("connection write error: {e}"); } else { map.transfer_any(s); } }, Err(e) => { eprintln!("connection read error: {e}"); }, } } fn shard_main(map: Arc<TransferMap>, shard_id: u64, mut connection_rx: mpsc::UnboundedReceiver<TcpStream>) { tokio::runtime::Builder::new_current_thread() .enable_all() .build() .expect("failed to create shard async runtime") .block_on(async move { loop { match connection_rx.recv().await { Some(s) => { println!("shard#{} handling connection {}", shard_id, s.peer_addr().unwrap()); tokio::spawn(shard_handle_connection(map.clone(), shard_id, s)); }, None => { println!("shard#{shard_id} closed"); break; }, } } // TODO(fusion): Transfer all connections to another shard? }); } pub fn main() { let transfer_map = Arc::new(TransferMap::new()); let nshards = std::thread::available_parallelism().unwrap().get(); println!("spawning {nshards} shards"); for _ in 0..nshards { spawn_shard(&transfer_map); } server_run(transfer_map); }
编辑补充
之后我做了进一步研究,找到了另一种更合理的实现方案:结合LocalSet与current_thread Runtime实现单线程服务器,再在多个线程上对同一个current_thread Runtime调用block_on,这样它们就能共享同一套IO和定时器驱动。
use std::{ collections::{HashMap, hash_map::Entry}, sync::{Arc, Mutex}, thread, }; use tokio::{ io::{AsyncReadExt, AsyncWriteExt}, net::{TcpListener, TcpStream}, runtime::{Builder, Runtime}, sync::mpsc, task, }; struct TransferMap { inner: Mutex<HashMap<u64, mpsc::UnboundedSender<TcpStream>>>, } impl TransferMap { fn new() -> Self { Self { inner: Mutex::new(HashMap::new()) } } fn add_tx(&self, tx: mpsc::UnboundedSender<TcpStream>) -> u64 { let mut inner = self.inner.lock().unwrap(); loop { let id = rand::random::<u64>(); if let Entry::Vacant(entry) = inner.entry(id) { entry.insert(tx); return id; } } } fn transfer(&self, id: u64, s: TcpStream) { let inner = self.inner.lock().unwrap(); if let Some(tx) = inner.get(&id) { tx.send(s).unwrap(); } } fn transfer_any(&self, s: TcpStream) { let inner = self.inner.lock().unwrap(); if inner.len() > 0 { let index = rand::random::<usize>() % inner.len(); let tx = inner.values().nth(index).unwrap(); tx.send(s).unwrap(); } } } async fn shard_handle_connection( map: Arc<TransferMap>, shard_id: u64, mut s: TcpStream){ let addr = s.peer_addr().unwrap(); println!("shard#{} handling connection {:?}", shard_id, addr); let mut buf = [0u8; 1024]; for _ in 0..5 { match s.read(&mut buf).await { Ok(n) => { let message = std::str::from_utf8(&buf[..n]).unwrap_or(""); println!("shard#{} message from {:?}: {}", shard_id, addr, message); if let Err(err) = s.write_all(message.as_bytes()).await { println!("shard#{} connection {:?} write error: {}", shard_id, addr, err); return; } }, Err(err) => { println!("shard#{} connection {:?} read error: {}", shard_id, addr, err); return; }, } } map.transfer_any(s); } async fn shard_main(map: Arc<TransferMap>, shard_id: u64, mut connection_rx: mpsc::UnboundedReceiver<TcpStream>) { loop { match connection_rx.recv().await { Some(s) => { task::spawn_local(shard_handle_connection(Arc::clone(&map), shard_id, s)); }, None => { println!("shard#{shard_id} closed"); break; }, } // TODO(fusion): Transfer all connections to another shard? } } fn spawn_shard(rt: Arc<Runtime>, map: Arc<TransferMap>){ let (connection_tx, connection_rx) = mpsc::unbounded_channel(); let shard_id = map.add_tx(connection_tx); thread::spawn(move || { let local = task::LocalSet::new(); local.block_on(&rt, shard_main(map, shard_id, connection_rx)); }); } pub fn main() { let map = Arc::new(TransferMap::new()); let rt = Arc::new(Builder::new_current_thread() .enable_all() .build() .expect("failed to create async runtime")); let nshards = std::thread::available_parallelism().unwrap().get(); for _ in 0..nshards { spawn_shard(Arc::clone(&rt), Arc::clone(&map)); } rt.block_on(async move { let listener = TcpListener::bind("127.0.0.1:7777").await .expect("failed to bind listener"); loop { match listener.accept().await { Ok((s, _addr)) => map.transfer_any(s), Err(err) => println!("accept error: {:?}", err), } } }); }
内容的提问来源于stack exchange,提问作者fusion32
相关产品推荐
相关产品推荐

