Skip to main content

flare_core/server/transports/
websocket.rs

1//! WebSocket 服务端实现
2//!
3//! 专注于 WebSocket 协议层面的连接处理,连接管理和心跳检测由 ServerCore 统一管理
4
5use 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
21/// WebSocket 服务端
22///
23/// 专注于 WebSocket 协议层面的连接处理
24pub struct WebSocketServer {
25    config: ServerConfig,
26    core: Arc<ServerCore>,
27    is_running: Arc<Mutex<bool>>,
28}
29
30impl WebSocketServer {
31    /// 创建新的 WebSocket 服务端
32    pub fn new(config: ServerConfig) -> Self {
33        Self::with_connection_manager(config, None)
34    }
35
36    /// 使用指定的连接管理器创建 WebSocket 服务端
37    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    /// 使用指定的 ServerCore 创建 WebSocket 服务端(用于共享 ServerCore)
51    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        // 启动心跳检测
77        self.core.start_heartbeat(&self.config);
78
79        // 准备共享资源
80        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
144/// 处理 WebSocket 连接(内部函数)
145async fn handle_websocket_connection(
146    stream: TcpStream,
147    manager: Arc<ConnectionManager>,
148    config: ServerConfig,
149    core: Arc<ServerCore>,
150) {
151    // 建立 WebSocket 连接
152    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    // 创建传输层连接
168    let transport = WebSocketTransport::from_tcp_stream(ws_stream);
169    let connection: Box<dyn Connection> = Box::new(transport);
170
171    // 使用公共模块设置连接
172    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}