1use 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
22pub 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
82pub 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
129pub 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}