Skip to main content

yrs_axum_ws/
lib.rs

1use axum::extract::ws::{Message, WebSocket};
2use futures_util::Sink;
3use futures_util::stream::{SplitSink, SplitStream};
4use yrs_tokio::signaling::Message as SignalingMessage;
5use yrs_tokio::{
6    YrsExchange, YrsSink, YrsStream, impl_yrs_signal_stream, to_signaling_message, yrs_common_sink,
7};
8
9#[derive(YrsStream)]
10pub struct YrsStream(SplitStream<WebSocket>);
11#[derive(YrsExchange)]
12pub struct YrsSignalStream(SplitStream<WebSocket>);
13
14impl_yrs_signal_stream!(YrsSignalStream, item => to_signaling_message!(item));
15
16#[derive(YrsSink)]
17pub struct YrsSink(SplitSink<WebSocket, Message>);
18#[yrs_common_sink]
19impl Sink<SignalingMessage> for YrsSink {}
20
21#[cfg(test)]
22mod test {
23    use crate::{YrsSink, YrsStream};
24    use axum::Router;
25    use axum::extract::ws::WebSocket;
26    use axum::extract::{State, WebSocketUpgrade};
27    use axum::response::Response;
28    use axum::routing::any;
29    use futures_util::{SinkExt, StreamExt};
30    use std::net::SocketAddr;
31    use std::str::FromStr;
32    use std::sync::Arc;
33    use tokio::sync::Mutex;
34    use tokio::task::JoinHandle;
35    use yrs::updates::encoder::Encode;
36    use yrs::{GetString, Text, Transact};
37    use yrs_tokio::broadcast::BroadcastGroup;
38    use yrs_tokio::yrs_common_test;
39
40    #[yrs_common_test]
41    async fn start_server(
42        addr: &str,
43        bcast: Arc<BroadcastGroup>,
44    ) -> Result<JoinHandle<()>, Box<dyn std::error::Error>> {
45        let _ = tracing_subscriber::fmt::try_init();
46        let addr = SocketAddr::from_str(addr)?;
47
48        let app = Router::new()
49            .route("/my-room", any(ws_handler))
50            .with_state(bcast);
51
52        Ok(tokio::spawn(async move {
53            let listener = tokio::net::TcpListener::bind(addr).await.unwrap();
54            axum::serve(listener, app).await.unwrap();
55        }))
56    }
57
58    async fn ws_handler(
59        ws: WebSocketUpgrade,
60        State(bcast): State<Arc<BroadcastGroup>>,
61    ) -> Response {
62        ws.on_upgrade(move |socket| peer(socket, bcast))
63    }
64
65    async fn peer(ws: WebSocket, bcast: Arc<BroadcastGroup>) {
66        let (sink, stream) = ws.split();
67        let sink = Arc::new(Mutex::new(YrsSink::from(sink)));
68        let stream = YrsStream::from(stream);
69
70        let sub = bcast.subscribe(sink, stream);
71        match sub.completed().await {
72            Ok(_) => println!("broadcasting for channel finished successfully"),
73            Err(e) => eprintln!("broadcasting for channel finished abruptly: {}", e),
74        }
75    }
76}