flare_core/server/transports/
websocket.rs1use crate::common::error::Result;
6use crate::server::config::ServerConfig;
7use crate::server::connection::ConnectionManager;
8use crate::server::transports::Server;
9use crate::server::transports::common::ServerConnectionHelper;
10use crate::server::transports::server_core::ServerCore;
11use crate::transport::connection::Connection;
12use crate::transport::websocket::WebSocketTransport;
13use async_trait::async_trait;
14use std::sync::Arc;
15use tokio::net::{TcpListener, TcpStream};
16use tokio::sync::{Mutex, Semaphore};
17use tokio::time::timeout;
18use tokio_tungstenite::accept_async;
19use tracing::debug;
20
21pub struct WebSocketServer {
25 config: ServerConfig,
26 core: Arc<ServerCore>,
27 is_running: Arc<Mutex<bool>>,
28}
29
30impl WebSocketServer {
31 pub fn new(config: ServerConfig) -> Self {
33 Self::with_connection_manager(config, None)
34 }
35
36 pub fn with_connection_manager(
38 config: ServerConfig,
39 connection_manager: Option<Arc<ConnectionManager>>,
40 ) -> Self {
41 let core = Arc::new(ServerCore::new(&config, connection_manager));
42
43 Self {
44 config,
45 core,
46 is_running: Arc::new(Mutex::new(false)),
47 }
48 }
49
50 pub fn with_shared_core(config: ServerConfig, core: Arc<ServerCore>) -> Self {
52 Self {
53 config,
54 core,
55 is_running: Arc::new(Mutex::new(false)),
56 }
57 }
58}
59
60#[async_trait]
61impl Server for WebSocketServer {
62 async fn start(&mut self) -> Result<()> {
63 let bind_address = self
64 .config
65 .get_protocol_address(&crate::common::config_types::TransportProtocol::WebSocket);
66 let addr = bind_address.parse::<std::net::SocketAddr>().map_err(|e| {
67 crate::common::error::FlareError::protocol_error(format!("Invalid address: {}", e))
68 })?;
69
70 let listener = TcpListener::bind(addr).await.map_err(|e| {
71 crate::common::error::FlareError::connection_failed(format!("Failed to bind: {}", e))
72 })?;
73
74 *self.is_running.lock().await = true;
75
76 self.core.start_heartbeat(&self.config);
78
79 let manager = Arc::clone(&self.core.connection_manager);
81 let config = self.config.clone();
82 let is_running = Arc::clone(&self.is_running);
83 let core = Arc::clone(&self.core);
84 let handshake_limiter = Arc::new(Semaphore::new(config.max_handshake_concurrency.max(1)));
85
86 tokio::spawn(async move {
87 debug!("[WebSocketServer] 开始监听连接");
88 while *is_running.lock().await {
89 match listener.accept().await {
90 Ok((stream, _addr)) => {
91 debug!("[WebSocketServer] 收到新连接");
92 let manager_clone = Arc::clone(&manager);
93 let config_clone = config.clone();
94 let core_clone = Arc::clone(&core);
95 let permit = match Arc::clone(&handshake_limiter).try_acquire_owned() {
96 Ok(permit) => permit,
97 Err(_) => {
98 debug!(
99 "[WebSocketServer] 握手并发已满: {}",
100 config.max_handshake_concurrency
101 );
102 continue;
103 }
104 };
105
106 tokio::spawn(async move {
107 let _permit = permit;
108 handle_websocket_connection(
109 stream,
110 manager_clone,
111 config_clone,
112 core_clone,
113 )
114 .await;
115 });
116 }
117 Err(e) => {
118 debug!("[WebSocketServer] 接受连接失败: {}", e);
119 }
120 }
121 }
122 debug!("[WebSocketServer] 停止监听连接");
123 });
124
125 Ok(())
126 }
127
128 async fn stop(&mut self) -> Result<()> {
129 ServerConnectionHelper::stop_server(&self.core, &self.is_running)
130 .await
131 .map_err(|e| {
132 crate::common::error::FlareError::connection_failed(format!(
133 "停止服务器失败: {}",
134 e
135 ))
136 })
137 }
138
139 fn is_running(&self) -> bool {
140 tokio::task::block_in_place(|| *self.is_running.blocking_lock())
141 }
142}
143
144async fn handle_websocket_connection(
146 stream: TcpStream,
147 manager: Arc<ConnectionManager>,
148 config: ServerConfig,
149 core: Arc<ServerCore>,
150) {
151 let ws_stream = match timeout(config.handshake_timeout, accept_async(stream)).await {
153 Ok(Ok(ws)) => ws,
154 Ok(Err(e)) => {
155 debug!("[WebSocketServer] WebSocket 握手失败: {}", e);
156 return;
157 }
158 Err(_) => {
159 debug!(
160 "[WebSocketServer] WebSocket 握手超时: {:?}",
161 config.handshake_timeout
162 );
163 return;
164 }
165 };
166
167 let transport = WebSocketTransport::from_tcp_stream(ws_stream);
169 let connection: Box<dyn Connection> = Box::new(transport);
170
171 if let Err(e) = ServerConnectionHelper::setup_new_connection(
173 connection,
174 manager.clone(),
175 &config,
176 core.clone(),
177 )
178 .await
179 {
180 debug!("[WebSocketServer] 设置连接失败: {}", e);
181 }
182}