修复 register/start_tun 失败后不可重试

- register 拆出 register_impl,Err 时归还 RegistrationContext,
  调用方(CLI/Web 的 5 秒重试循环)不再得到 'can only be called once'
- DeviceIoManager::start_task 改为先创建 tun 设备(唯一可能失败的步骤),
  成功后才消费 TunReceiver/EnhancedOutbound,失败时调用方状态完整可重试
This commit is contained in:
lbl
2026-08-20 21:54:27 +08:00
parent 6489237233
commit 1ff7405681
2 changed files with 46 additions and 36 deletions
+31 -22
View File
@@ -206,13 +206,27 @@ impl NetworkManager {
} }
/// Register with server(s) and start data handling tasks. /// Register with server(s) and start data handling tasks.
/// This method can only be called once.
/// Returns the registration response on success. /// Returns the registration response on success.
/// On connection-level failure the internal state is kept, so the call can be retried.
pub async fn register(&mut self) -> anyhow::Result<RegisterResponse> { pub async fn register(&mut self) -> anyhow::Result<RegisterResponse> {
let Some(mut ctx) = self.registration_context.take() else { let Some(mut ctx) = self.registration_context.take() else {
bail!("register can only be called once"); bail!("register can only be called once");
}; };
match Self::register_impl(&self.app_state, &self.task_group, &mut ctx).await {
Ok(response) => Ok(response),
Err(e) => {
// 注册失败时归还上下文,允许调用方重试
self.registration_context = Some(ctx);
Err(e)
}
}
}
async fn register_impl(
app_state: &AppState,
task_group: &TaskGroup,
ctx: &mut RegistrationContext,
) -> anyhow::Result<RegisterResponse> {
let is_multi_server = ctx.server_managers.len() > 1; let is_multi_server = ctx.server_managers.len() > 1;
let response = if is_multi_server { let response = if is_multi_server {
@@ -251,35 +265,35 @@ impl NetworkManager {
ip: reg_response.ip, ip: reg_response.ip,
prefix_len: reg_response.prefix_len, prefix_len: reg_response.prefix_len,
}; };
self.app_state.network.set(network_addr); app_state.network.set(network_addr);
// 保存服务器版本信息 // 保存服务器版本信息
if !reg_response.server_version.is_empty() { if !reg_response.server_version.is_empty() {
for (index, _) in ctx.server_managers.iter().enumerate() { for (index, _) in ctx.server_managers.iter().enumerate() {
self.app_state app_state
.server_info_collection .server_info_collection
.set_server_version(index as u32, reg_response.server_version.clone()); .set_server_version(index as u32, reg_response.server_version.clone());
} }
} }
// Start data handling tasks for all servers // Start data handling tasks for all servers
for turn_manager in ctx.server_managers { for turn_manager in ctx.server_managers.drain(..) {
let handler_config = Box::new(InboundHandlerConfig { let handler_config = Box::new(InboundHandlerConfig {
network_route: NetworkRoute::new( network_route: NetworkRoute::new(
self.app_state.network.clone(), app_state.network.clone(),
ctx.subnet_external_route.clone(), ctx.subnet_external_route.clone(),
), ),
server_info: self.app_state.server_info_collection.clone(), server_info: app_state.server_info_collection.clone(),
nat_info: self.app_state.nat_info.clone(), nat_info: app_state.nat_info.clone(),
peer_map: self.app_state.peer_map.clone(), peer_map: app_state.peer_map.clone(),
punch_backoff: self.app_state.punch_backoff.clone(), punch_backoff: app_state.punch_backoff.clone(),
puncher: ctx.puncher.clone(), puncher: ctx.puncher.clone(),
packet_crypto: ctx.packet_crypto.clone(), packet_crypto: ctx.packet_crypto.clone(),
packet_compression: ctx.packet_compression.clone(), packet_compression: ctx.packet_compression.clone(),
enhanced_inbound: ctx.enhanced_inbound.clone(), enhanced_inbound: ctx.enhanced_inbound.clone(),
fec_decoder: ctx.fec_decoder.clone(), fec_decoder: ctx.fec_decoder.clone(),
}); });
turn_manager.data_handle_task_connected(&self.task_group, handler_config, network_addr); turn_manager.data_handle_task_connected(task_group, handler_config, network_addr);
} }
Ok(RegisterResponse::Success(network_addr)) Ok(RegisterResponse::Success(network_addr))
@@ -290,29 +304,24 @@ impl NetworkManager {
} }
pub async fn start_tun(&mut self) -> anyhow::Result<()> { pub async fn start_tun(&mut self) -> anyhow::Result<()> {
let Some(receiver) = self.tun_receiver.take() else { if self.tun_receiver.is_none() || self.enhanced_outbound.is_none() {
bail!("start_tun can only be called once"); bail!("start_tun can only be called once");
}; }
let Some(enhanced_outbound) = self.enhanced_outbound.take() else {
bail!("start_tun can only be called once");
};
let mut config = DeviceConfig::default(); let mut config = DeviceConfig::default();
config = config.set_mtu(self.config.mtu.unwrap_or(DEFAULT_MTU)); config = config.set_mtu(self.config.mtu.unwrap_or(DEFAULT_MTU));
if let Some(tun_name) = self.config.tun_name.clone() { if let Some(tun_name) = self.config.tun_name.clone() {
config = config.set_tun_name(tun_name); config = config.set_tun_name(tun_name);
} }
// 失败时 tun_receiver/enhanced_outbound 不会被消耗,可以重试
self.device_io_manager self.device_io_manager
.start_task(config, receiver, enhanced_outbound) .start_task(config, &mut self.tun_receiver, &mut self.enhanced_outbound)
.await .await
} }
#[cfg(unix)] #[cfg(unix)]
pub async fn start_tun_fd(&mut self, tun_fd: Option<i32>) -> anyhow::Result<()> { pub async fn start_tun_fd(&mut self, tun_fd: Option<i32>) -> anyhow::Result<()> {
let Some(receiver) = self.tun_receiver.take() else { if self.tun_receiver.is_none() || self.enhanced_outbound.is_none() {
bail!("start_tun_fd can only be called once"); bail!("start_tun_fd can only be called once");
}; }
let Some(enhanced_outbound) = self.enhanced_outbound.take() else {
bail!("start_tun_fd can only be called once");
};
let mut config = DeviceConfig::default(); let mut config = DeviceConfig::default();
if let Some(tun_fd) = tun_fd { if let Some(tun_fd) = tun_fd {
config = config.set_tun_fd(tun_fd); config = config.set_tun_fd(tun_fd);
@@ -321,7 +330,7 @@ impl NetworkManager {
config = config.set_tun_name(tun_name); config = config.set_tun_name(tun_name);
} }
self.device_io_manager self.device_io_manager
.start_task(config, receiver, enhanced_outbound) .start_task(config, &mut self.tun_receiver, &mut self.enhanced_outbound)
.await .await
} }
#[cfg(not(target_os = "android"))] #[cfg(not(target_os = "android"))]
+15 -14
View File
@@ -76,16 +76,19 @@ impl DeviceIOManager {
pub async fn start_task( pub async fn start_task(
&self, &self,
device_config: DeviceConfig, device_config: DeviceConfig,
receiver: TunReceiver, receiver: &mut Option<TunReceiver>,
enhanced_outbound: EnhancedOutbound, enhanced_outbound: &mut Option<EnhancedOutbound>,
) -> anyhow::Result<()> { ) -> anyhow::Result<()> {
if receiver.is_none() || enhanced_outbound.is_none() {
bail!("device task already started");
}
self.stop_task().await; self.stop_task().await;
let task = create( // 先执行可能失败的 tun 设备创建,成功后才消费 receiver/outbound
&self.task_group, // 保证失败时调用方状态完整、可以重试
device_config, let device = Arc::new(create_tun(device_config)?);
receiver.receiver, let receiver = receiver.take().unwrap();
enhanced_outbound, let enhanced_outbound = enhanced_outbound.take().unwrap();
)?; let task = create(&self.task_group, device, receiver.receiver, enhanced_outbound);
self.device.lock().await.0.replace(task); self.device.lock().await.0.replace(task);
Ok(()) Ok(())
} }
@@ -148,12 +151,10 @@ fn create_tun(config: DeviceConfig) -> anyhow::Result<AsyncDevice> {
} }
fn create( fn create(
task_group: &TaskGroup, task_group: &TaskGroup,
config: DeviceConfig, device: Arc<AsyncDevice>,
receiver: Receiver<TransmissionBytes>, receiver: Receiver<TransmissionBytes>,
enhanced_outbound: EnhancedOutbound, enhanced_outbound: EnhancedOutbound,
) -> anyhow::Result<DeviceTask> { ) -> DeviceTask {
let device = Arc::new(create_tun(config)?);
let device_framed_read = DeviceFramedRead::new(device.clone(), BytesCodec::new()); let device_framed_read = DeviceFramedRead::new(device.clone(), BytesCodec::new());
let device_framed_write = DeviceFramedWrite::new(device.clone(), BytesCodec::new()); let device_framed_write = DeviceFramedWrite::new(device.clone(), BytesCodec::new());
@@ -168,11 +169,11 @@ fn create(
} }
}); });
Ok(DeviceTask { DeviceTask {
device, device,
task_recv, task_recv,
task_send, task_send,
}) }
} }
async fn in_tun_loop( async fn in_tun_loop(