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}