Skip to main content

sova_ws/
upgrade.rs

1//! WebSocket handshake and session I/O.
2
3use std::future::Future;
4
5use futures_util::{SinkExt, StreamExt};
6use http::header::{
7    CONNECTION, SEC_WEBSOCKET_ACCEPT, SEC_WEBSOCKET_KEY, SEC_WEBSOCKET_VERSION, UPGRADE,
8};
9use http::HeaderMap;
10use hyper_util::rt::TokioIo;
11use sova_core::{Request, Response, UpgradePermit};
12use tokio::sync::mpsc;
13use tokio_tungstenite::tungstenite::protocol::{Role, WebSocketConfig};
14use tokio_tungstenite::tungstenite::{self, Error as WsError, Message};
15use tokio_tungstenite::WebSocketStream;
16
17use crate::hub::{Hub, RoomHandle};
18use crate::WsShared;
19
20type WsIo = TokioIo<hyper::upgrade::Upgraded>;
21
22/// Active WebSocket connection passed to route handlers.
23pub struct WsSession {
24    read: futures_util::stream::SplitStream<WebSocketStream<WsIo>>,
25    out_tx: mpsc::UnboundedSender<Message>,
26    _write_task: tokio::task::JoinHandle<()>,
27    hub: Hub,
28    _permit: UpgradePermit,
29    rooms: Vec<RoomHandle>,
30}
31
32impl WsSession {
33    pub(crate) fn new(
34        stream: WebSocketStream<WsIo>,
35        hub: Hub,
36        permit: UpgradePermit,
37    ) -> Self {
38        let (mut write, read) = stream.split();
39        let (out_tx, mut out_rx) = mpsc::unbounded_channel();
40        let write_task = tokio::spawn(async move {
41            while let Some(msg) = out_rx.recv().await {
42                if write.send(msg).await.is_err() {
43                    break;
44                }
45            }
46        });
47        Self {
48            read,
49            out_tx,
50            _write_task: write_task,
51            hub,
52            _permit: permit,
53            rooms: Vec::new(),
54        }
55    }
56
57    pub fn hub(&self) -> &Hub {
58        &self.hub
59    }
60
61    pub async fn recv(&mut self) -> Option<Result<Message, WsError>> {
62        self.read.next().await
63    }
64
65    pub async fn send(&self, msg: Message) -> Result<(), WsError> {
66        self.out_tx
67            .send(msg)
68            .map_err(|_| WsError::ConnectionClosed)
69    }
70
71    pub fn join(&mut self, room: impl Into<String>) -> RoomHandle {
72        let handle = self.hub.register(&room.into(), self.out_tx.clone()).1;
73        self.rooms.push(handle.clone());
74        handle
75    }
76
77    pub fn leave(&mut self, room: &str) {
78        self.rooms.retain(|h| h.room() != room);
79    }
80}
81
82/// Check `Origin` against an allowlist. Empty allowlist → allow all (dev default).
83pub fn origin_allowed(headers: &HeaderMap, allowed: &[String]) -> bool {
84    if allowed.is_empty() {
85        return true;
86    }
87    let Some(origin) = headers
88        .get("origin")
89        .and_then(|v| v.to_str().ok())
90    else {
91        return false;
92    };
93    allowed.iter().any(|o| o == origin)
94}
95
96fn validate_ws_headers(
97    headers: &HeaderMap,
98) -> Result<(), Box<Response>> {
99    let upgrade_ok = headers
100        .get(UPGRADE)
101        .and_then(|v| v.to_str().ok())
102        .is_some_and(|v| v.eq_ignore_ascii_case("websocket"));
103    if !upgrade_ok {
104        return Err(Box::new(Response::text("Bad Request").status(400)));
105    }
106
107    if headers.get(SEC_WEBSOCKET_KEY).is_none() {
108        return Err(Box::new(Response::text("Bad Request").status(400)));
109    }
110
111    let version_ok = headers
112        .get(SEC_WEBSOCKET_VERSION)
113        .and_then(|v| v.to_str().ok())
114        == Some("13");
115    if !version_ok {
116        return Err(Box::new(Response::text("Bad Request").status(400)));
117    }
118
119    Ok(())
120}
121
122fn ws_key(headers: &HeaderMap) -> Result<&str, Box<Response>> {
123    headers
124        .get(SEC_WEBSOCKET_KEY)
125        .and_then(|v| v.to_str().ok())
126        .ok_or_else(|| Box::new(Response::text("Bad Request").status(400)))
127}
128
129/// Perform a WebSocket upgrade from a route handler.
130pub async fn upgrade_ws<F, Fut>(
131    mut req: Request,
132    handler: F,
133) -> Result<Response, Response>
134where
135    F: FnOnce(WsSession) -> Fut + Send + 'static,
136    Fut: Future<Output = ()> + Send + 'static,
137{
138    let shared = req
139        .try_state::<WsShared>()
140        .ok_or_else(|| Response::text("WebSocket plugin not installed").status(500))?;
141
142    if !origin_allowed(&req.headers, &shared.config.origins) {
143        return Err(Response::text("Forbidden").status(403));
144    }
145
146    validate_ws_headers(&req.headers).map_err(|b| *b)?;
147    let key = ws_key(&req.headers).map_err(|b| *b)?;
148    let accept = tungstenite::handshake::derive_accept_key(key.as_bytes());
149
150    let on_upgrade = match req.on_upgrade() {
151        None => {
152            return Err(Response::text("Upgrade Required")
153                .status(426)
154                .header(UPGRADE.as_str(), "websocket"));
155        }
156        Some(Err(res)) => return Err(res),
157        Some(Ok(up)) => up,
158    };
159
160    let hub = shared.hub.clone();
161    let max_message_size = shared.config.max_message_size;
162    tokio::spawn(async move {
163        let (io, permit) = match on_upgrade.upgrade().await {
164            Ok(v) => v,
165            Err(err) => {
166                tracing::warn!("websocket upgrade failed: {err}");
167                return;
168            }
169        };
170
171        let mut ws_config = WebSocketConfig::default();
172        ws_config.max_message_size = max_message_size;
173        let stream =
174            WebSocketStream::from_raw_socket(TokioIo::new(io), Role::Server, Some(ws_config)).await;
175        let session = WsSession::new(stream, hub, permit);
176        handler(session).await;
177    });
178
179    Ok(Response::empty()
180        .status(101)
181        .header(UPGRADE.as_str(), "websocket")
182        .header(CONNECTION.as_str(), "upgrade")
183        .header(SEC_WEBSOCKET_ACCEPT.as_str(), accept)
184        .header(SEC_WEBSOCKET_VERSION.as_str(), "13"))
185}