1use async_io::Async;
2use axum::Router;
3use bevy_defer::{AccessResult, AsyncCommandsExtension, AsyncExecutor, AsyncWorld};
4use bevy_ecs::prelude::*;
5use bevy_log::{debug, error, info, warn};
6use std::{
7 collections::HashMap,
8 net::{IpAddr, TcpListener},
9 time::Duration,
10};
11
12use super::{ServerStatus, TaskType};
13use crate::{WebPort, WebServer, WebServerError, WebServerResult};
14
15#[derive(Default, Resource)]
17pub struct WebServerManager(HashMap<WebPort, WebServer>);
18
19impl WebServerManager {
20 const SHUTDOWN_CHECK_INTERVAL_MS: u64 = 100;
21
22 pub fn cleanup_finished_tasks(mut manager: ResMut<Self>) {
23 for (port, server) in manager.iter_mut() {
24 let finished_count = server.task_store().finished_task_count();
25 if finished_count > 0 {
26 debug!(
27 "Cleaning up {} finished tasks on port {}",
28 finished_count, port
29 );
30 }
31
32 server.task_store_mut().cleanup_finished_tasks();
33 }
34 }
35
36 pub fn check_retry_servers(mut manager: ResMut<Self>) {
37 let servers_to_retry: Vec<_> = manager
38 .iter_mut()
39 .filter(|(_, server)| {
40 server.status() == crate::server::ServerStatus::Retrying && server.should_retry()
41 })
42 .map(|(port, _)| *port)
43 .collect();
44
45 if !servers_to_retry.is_empty() {
46 debug!("Found {} servers ready to retry", servers_to_retry.len());
47 manager.set_changed();
49 }
50 }
51
52 pub fn changed(mut manager: ResMut<Self>, async_executor: NonSend<AsyncExecutor>) {
53 if !manager.is_changed() {
54 return;
55 }
56
57 let servers_to_start: Vec<_> = manager
58 .iter_mut()
59 .filter(|(_, server)| server.status().can_start())
60 .filter_map(|(port, server)| {
61 if !server.task_store().contains_key(&TaskType::Server) {
62 if server.status() == crate::server::ServerStatus::Retrying {
64 if server.should_retry() {
65 debug!("Retry time reached for server on port {}", port);
66 server.set_status(crate::server::ServerStatus::Starting);
67 Some(*port)
68 } else {
69 None }
71 } else {
72 server.set_status(crate::server::ServerStatus::Starting);
73 Some(*port)
74 }
75 } else {
76 None
77 }
78 })
79 .collect();
80
81 for port in servers_to_start {
82 debug!(" - Starting server on port {}", port);
83 if let Err(err) = manager.start_server(&port, &async_executor) {
84 error!("Failed to start server on port {}: {}", port, err);
85
86 if let Some(server) = manager.0.get_mut(&port) {
87 server.schedule_retry();
88 }
89 }
90 }
91 }
92
93 pub fn add_server(&mut self, server: WebServer) -> WebServerResult<()> {
94 let port = server.port();
95 let ip = server.ip();
96 if self.0.contains_key(&port) {
97 return Err(WebServerError::server_already_running(port));
98 }
99
100 match Self::test_bind(ip, port) {
102 Ok(_) => {
103 self.0.insert(port, server);
105 }
106 Err(bind_error) => {
107 warn!(
109 "Initial bind test failed for {}:{}, server will retry: {}",
110 ip, port, bind_error
111 );
112 let mut server = server;
113 server.set_error(bind_error.to_string());
114 server.schedule_retry();
115 self.0.insert(port, server);
116 }
117 }
118
119 Ok(())
120 }
121
122 pub fn remove_server(&mut self, port: &WebPort) {
123 if let Some(mut server) = self.0.remove(port) {
124 server.stop();
125 }
126 }
127
128 pub fn stop_server(&mut self, port: &WebPort) {
129 if let Some(server) = self.0.get_mut(port) {
130 server.stop();
131 }
132 }
133
134 pub fn server_error(&self, port: &WebPort) -> Option<&str> {
136 self.0.get(port).and_then(|server| server.last_error())
137 }
138
139 pub fn server_failed(&self, port: &WebPort) -> bool {
141 self.0
142 .get(port)
143 .map(|server| server.status() == ServerStatus::Failed)
144 .unwrap_or(false)
145 }
146
147 pub fn server_status_report(&self) -> Vec<(WebPort, ServerStatus, Option<String>)> {
149 self.0
150 .iter()
151 .map(|(port, server)| {
152 (
153 *port,
154 server.status(),
155 server.last_error().map(|s| s.to_string()),
156 )
157 })
158 .collect()
159 }
160
161 pub fn stop_all(&mut self) {
162 for (_, server) in self.0.iter_mut() {
163 server.stop();
164 }
165 self.0.clear();
166 }
167
168 pub fn graceful_shutdown(&mut self, port: &WebPort) {
170 if let Some(server) = self.0.get_mut(port) {
171 if server.status() != ServerStatus::ShuttingDown {
173 server.graceful_shutdown();
174 }
175 }
176 }
177
178 pub fn graceful_shutdown_with_timeout(
181 &mut self,
182 port: &WebPort,
183 timeout: Duration,
184 commands: &mut Commands,
185 ) {
186 if !self.0.contains_key(port) {
187 warn!("Cannot shutdown server on port {}: server not found", port);
188 return;
189 }
190
191 info!(
192 "Initiating graceful shutdown for server on port {} with timeout {:?}",
193 port, timeout
194 );
195
196 if let Some(server) = self.0.get_mut(port) {
198 server.set_status(ServerStatus::ShuttingDown);
199 }
200
201 self.graceful_shutdown(&port);
203
204 let port = *port;
205
206 commands.spawn_task(async move || Self::shutdown_server(port, timeout).await);
208 }
209
210 pub async fn graceful_shutdown_server(&mut self, port: &WebPort, timeout: Duration) -> bool {
211 if let Some(server) = self.0.get_mut(port) {
212 server.graceful_shutdown_with_timeout(timeout).await
213 } else {
214 false
215 }
216 }
217
218 pub async fn graceful_shutdown_all(&mut self, timeout: Duration) -> HashMap<WebPort, bool> {
219 let mut results = HashMap::new();
220 let ports: Vec<WebPort> = self.0.keys().copied().collect();
221
222 for port in &ports {
223 if let Some(server) = self.0.get_mut(port) {
224 server.graceful_shutdown();
225 }
226 }
227
228 for port in ports {
229 if let Some(server) = self.0.get_mut(&port) {
230 let completed_gracefully = server.graceful_shutdown_with_timeout(timeout).await;
231 results.insert(port, completed_gracefully);
232 }
233 }
234
235 self.0.clear();
236 results
237 }
238
239 pub fn has_server(&self, port: &WebPort) -> bool {
240 self.0.contains_key(port)
241 }
242
243 pub fn ports(&self) -> Vec<WebPort> {
244 self.0.keys().copied().collect()
245 }
246
247 pub fn len(&self) -> usize {
248 self.0.len()
249 }
250
251 pub fn shutdown_requested(&self, port: &WebPort) -> bool {
252 self.0
253 .get(port)
254 .map(|server| server.shutdown_requested())
255 .unwrap_or(false)
256 }
257
258 pub fn active_connections(&self, port: &WebPort) -> usize {
259 self.0
260 .get(port)
261 .map(|server| server.count_active_connections())
262 .unwrap_or(0)
263 }
264
265 pub fn shutdown_status(&self) -> HashMap<WebPort, (bool, usize)> {
266 self.0
267 .iter()
268 .map(|(port, server)| {
269 let shutdown_requested = server.shutdown_requested();
270 let active_connections = server.count_active_connections();
271 (*port, (shutdown_requested, active_connections))
272 })
273 .collect()
274 }
275
276 pub fn router(&self, port: &WebPort) -> Option<&Router> {
277 self.0.get(port).map(|server| server.router())
278 }
279
280 pub fn router_mut(&mut self, port: &WebPort) -> Option<&mut Router> {
281 self.0.get_mut(port).map(|server| server.router_mut())
282 }
283
284 pub fn set_router(&mut self, port: &WebPort, router: Router) {
285 if let Some(server) = self.0.get_mut(port) {
286 *server.router_mut() = router;
287 } else {
288 error!("No server found on port {}", port);
289 }
290 }
291
292 pub fn iter(&self) -> impl Iterator<Item = (&WebPort, &WebServer)> {
293 self.0.iter()
294 }
295
296 pub fn iter_mut(&mut self) -> impl Iterator<Item = (&WebPort, &mut WebServer)> {
297 self.0.iter_mut()
298 }
299
300 pub(crate) fn get_server(&self, port: &WebPort) -> Option<&WebServer> {
301 self.0.get(port)
302 }
303 pub(crate) fn get_server_mut(&mut self, port: &WebPort) -> Option<&mut WebServer> {
304 self.0.get_mut(port)
305 }
306
307 pub fn start_server(
308 &mut self,
309 port: &WebPort,
310 executor: &AsyncExecutor,
311 ) -> WebServerResult<()> {
312 let port = *port;
313
314 if let Some(server) = self.0.get(&port) {
315 if server.task_store().contains_key(&TaskType::Server) {
316 debug!("Server on port {} already has a running task", port);
317 return Err(WebServerError::server_already_running(port));
318 }
319 }
320
321 let server = self
322 .0
323 .get_mut(&port)
324 .ok_or_else(|| WebServerError::server_not_found(port))?;
325
326 server.clear_error();
328 server.set_status(ServerStatus::Starting);
329
330 let server_task = executor.spawn_task({
332 async move {
333 if let Err(err) = WebServer::run(port).await {
334 error!("bevy_webserver on port {} failed with: {}", port, err);
335 let _ = AsyncWorld
337 .resource::<WebServerManager>()
338 .get_mut(|manager| {
339 if let Some(server) = manager.get_server_mut(&port) {
340 server.set_error(err.to_string());
341 if err.to_string().contains("already in use")
343 || err.to_string().contains("bind")
344 {
345 server.schedule_retry();
346 } else {
347 server.set_status(crate::server::ServerStatus::Failed);
349 }
350 }
351 Ok::<(), bevy_defer::AccessError>(())
352 });
353 } else {
354 let _ = AsyncWorld
356 .resource::<WebServerManager>()
357 .get_mut(|manager| {
358 if let Some(server) = manager.get_server_mut(&port) {
359 server.set_status(crate::server::ServerStatus::Running);
360 }
361 Ok::<(), bevy_defer::AccessError>(())
362 });
363 }
364 Ok(())
365 }
366 });
367
368 let server = self
369 .0
370 .get_mut(&port)
371 .ok_or_else(|| WebServerError::server_not_found(port))?;
372
373 server
374 .task_store_mut()
375 .insert(TaskType::Server, server_task);
376
377 Ok(())
378 }
379
380 async fn shutdown_server(port: WebPort, timeout: Duration) -> AccessResult {
382 let start_time = std::time::Instant::now();
383 loop {
384 let shutdown_result = AsyncWorld.run(|world| {
385 let manager = world.resource::<WebServerManager>();
386
387 if manager.has_server(&port) {
389 let active_connections = manager.active_connections(&port);
390 let elapsed = start_time.elapsed();
391
392 if elapsed >= timeout {
393 if active_connections > 0 {
395 warn!("⚠️ Graceful shutdown timeout reached after {:?}, {} connections still active", elapsed, active_connections);
396 } else {
397 info!("✅ Server on port {} shutdown gracefully within timeout", port);
398 }
399 Some(()) } else if active_connections == 0 {
401 info!("✅ Server on port {} shutdown gracefully in {:?}", port, elapsed);
402 Some(()) } else {
404 None
406 }
407 } else {
408 info!("✅ Server on port {} shutdown completed", port);
410 Some(()) }
412 });
413
414 if shutdown_result.is_some() {
415 break;
416 }
417
418 AsyncWorld
420 .sleep(Duration::from_millis(Self::SHUTDOWN_CHECK_INTERVAL_MS))
421 .await;
422 }
423
424 AsyncWorld.run(|world| {
426 let mut manager = world.resource_mut::<WebServerManager>();
427 manager.remove_server(&port);
428 });
429 Ok(())
430 }
431
432 pub async fn wait_for_server_start(
435 &self,
436 port: &WebPort,
437 timeout: Duration,
438 ) -> WebServerResult<()> {
439 let start_time = std::time::Instant::now();
440
441 loop {
442 if let Some(server) = self.0.get(port) {
443 match server.status() {
444 ServerStatus::Running => return Ok(()),
445 ServerStatus::Failed => {
446 if let Some(error) = server.last_error() {
447 return Err(WebServerError::io_error(
448 format!("server startup on port {}", port),
449 std::io::Error::new(std::io::ErrorKind::Other, error.to_string()),
450 ));
451 } else {
452 return Err(WebServerError::io_error(
453 format!("server startup on port {}", port),
454 std::io::Error::new(
455 std::io::ErrorKind::Other,
456 "Server failed to start",
457 ),
458 ));
459 }
460 }
461 ServerStatus::Starting => {
462 if start_time.elapsed() > timeout {
464 return Err(WebServerError::timeout(
465 format!("starting server on port {}", port),
466 timeout.as_millis() as u64,
467 ));
468 }
469 }
470 _ => {
471 return Err(WebServerError::config_error(
472 "server_status",
473 format!(
474 "Server on port {} has unexpected status: {:?}",
475 port,
476 server.status()
477 ),
478 ));
479 }
480 }
481 } else {
482 return Err(WebServerError::server_not_found(*port));
483 }
484
485 AsyncWorld.sleep(Duration::from_millis(10)).await;
487 }
488 }
489
490 pub fn test_bind(ip: IpAddr, port: WebPort) -> WebServerResult<()> {
492 debug!("Testing bind on {}:{}", ip, port);
493
494 match TcpListener::bind(("0.0.0.0", port)) {
496 Ok(listener) => {
497 drop(listener);
499 }
500 Err(_) => {
501 let error_msg = format!("Port {} is already in use", port);
503 error!("{}:{}: {}", ip, port, error_msg);
504 return Err(WebServerError::bind_failed(
505 ip,
506 port,
507 std::io::Error::new(std::io::ErrorKind::AddrInUse, error_msg),
508 ));
509 }
510 }
511
512 let listener = Async::<TcpListener>::bind((ip, port)).map_err(|e| {
515 error!("Test bind failed on {}:{}: {}", ip, port, e);
516 WebServerError::bind_failed(ip, port, e)
517 })?;
518
519 drop(listener);
520 Ok(())
521 }
522}