1use std::collections::VecDeque;
2
3use serde_json::Value;
4use thiserror::Error;
5
6use crate::{
7 Codec, CodecError, Frame, Payload, PhoenixV2Codec, Protocol, ProtocolError, ProtocolEvent,
8 Transport, TransportError, TransportEvent,
9};
10
11pub struct Session<T> {
18 protocol: Protocol,
19 transport: T,
20 codec: Box<dyn Codec>,
21 buffered_events: VecDeque<ProtocolEvent>,
22}
23
24impl<T: Transport> Session<T> {
25 pub fn new(transport: T) -> Self {
27 Self::with_codec(transport, PhoenixV2Codec::limited(Default::default()))
28 }
29
30 pub fn with_codec(transport: T, codec: impl Codec + 'static) -> Self {
32 Self {
33 protocol: Protocol::new(),
34 transport,
35 codec: Box::new(codec),
36 buffered_events: VecDeque::new(),
37 }
38 }
39
40 pub fn protocol(&self) -> &Protocol {
42 &self.protocol
43 }
44
45 pub fn protocol_mut(&mut self) -> &mut Protocol {
47 &mut self.protocol
48 }
49
50 pub fn transport_mut(&mut self) -> &mut T {
52 &mut self.transport
53 }
54
55 pub fn into_parts(self) -> (Protocol, T) {
57 (self.protocol, self.transport)
58 }
59
60 pub async fn join(
62 &mut self,
63 topic: impl Into<String>,
64 params: Value,
65 ) -> Result<ProtocolEvent, SessionError> {
66 let outbound = self.protocol.join(topic, params)?;
67 let reference = outbound.reference.clone();
68 self.send(outbound.frame).await?;
69 self.wait_for_reference(&reference).await
70 }
71
72 pub async fn rejoin(
74 &mut self,
75 topic: impl Into<String>,
76 refreshed_params: Value,
77 ) -> Result<ProtocolEvent, SessionError> {
78 let outbound = self.protocol.rejoin(topic, refreshed_params)?;
79 let reference = outbound.reference.clone();
80 self.send(outbound.frame).await?;
81 self.wait_for_reference(&reference).await
82 }
83
84 pub async fn push(
86 &mut self,
87 topic: &str,
88 event: impl Into<String>,
89 payload: impl Into<Payload>,
90 ) -> Result<ProtocolEvent, SessionError> {
91 let outbound = self.protocol.push(topic, event, payload)?;
92 let reference = outbound.reference.clone();
93 self.send(outbound.frame).await?;
94 self.wait_for_reference(&reference).await
95 }
96
97 pub async fn leave(&mut self, topic: &str) -> Result<ProtocolEvent, SessionError> {
99 let outbound = self.protocol.leave(topic)?;
100 let reference = outbound.reference.clone();
101 self.send(outbound.frame).await?;
102 self.wait_for_reference(&reference).await
103 }
104
105 pub async fn heartbeat(&mut self) -> Result<ProtocolEvent, SessionError> {
107 let outbound = self.protocol.heartbeat();
108 let reference = outbound.reference.clone();
109 self.send(outbound.frame).await?;
110 self.wait_for_reference(&reference).await
111 }
112
113 pub async fn next_event(&mut self) -> Result<ProtocolEvent, SessionError> {
115 if let Some(event) = self.buffered_events.pop_front() {
116 return Ok(event);
117 }
118 self.receive_event().await
119 }
120
121 pub fn reset_connection(&mut self) -> Vec<ProtocolEvent> {
125 self.protocol.reset_connection()
126 }
127
128 pub async fn close(&mut self) -> Result<(), SessionError> {
130 self.transport.close().await?;
131 Ok(())
132 }
133
134 async fn send(&mut self, frame: Frame) -> Result<(), SessionError> {
135 self.transport.send(self.codec.encode(&frame)?).await?;
136 Ok(())
137 }
138
139 async fn wait_for_reference(
140 &mut self,
141 expected_reference: &str,
142 ) -> Result<ProtocolEvent, SessionError> {
143 loop {
144 let event = self.receive_event().await?;
145 if event_reference(&event) == Some(expected_reference) {
146 return Ok(event);
147 }
148 self.buffered_events.push_back(event);
149 }
150 }
151
152 async fn receive_event(&mut self) -> Result<ProtocolEvent, SessionError> {
153 let message = match self.transport.receive().await? {
154 TransportEvent::Message(message) => message,
155 TransportEvent::Closed(close) => return Err(SessionError::ConnectionClosed(close)),
156 };
157 Ok(self.protocol.receive(self.codec.decode(message)?)?)
158 }
159}
160
161fn event_reference(event: &ProtocolEvent) -> Option<&str> {
162 match event {
163 ProtocolEvent::Joined { reference, .. }
164 | ProtocolEvent::JoinError { reference, .. }
165 | ProtocolEvent::Left { reference, .. }
166 | ProtocolEvent::Reply { reference, .. }
167 | ProtocolEvent::HeartbeatAck { reference, .. }
168 | ProtocolEvent::RequestInterrupted { reference, .. } => Some(reference),
169 ProtocolEvent::Message(_)
170 | ProtocolEvent::ChannelClosed { .. }
171 | ProtocolEvent::ChannelError { .. }
172 | ProtocolEvent::StaleMessage(_)
173 | ProtocolEvent::UnmatchedReply(_) => None,
174 }
175}
176
177#[derive(Debug, Error)]
179pub enum SessionError {
180 #[error(transparent)]
182 Protocol(#[from] ProtocolError),
183 #[error(transparent)]
185 Codec(#[from] CodecError),
186 #[error(transparent)]
188 Transport(#[from] TransportError),
189 #[error("WebSocket connection closed: {0:?}")]
191 ConnectionClosed(crate::TransportClose),
192}
193
194#[cfg(test)]
195mod tests {
196 use std::{cell::RefCell, collections::VecDeque, rc::Rc};
197
198 use futures::future::LocalBoxFuture;
199 use serde_json::json;
200
201 use super::*;
202 use crate::WireMessage;
203
204 #[derive(Default)]
205 struct MockState {
206 incoming: VecDeque<WireMessage>,
207 sent: Vec<WireMessage>,
208 closed: bool,
209 }
210
211 struct MockTransport {
212 state: Rc<RefCell<MockState>>,
213 }
214
215 impl MockTransport {
216 fn with_incoming(
217 messages: impl IntoIterator<Item = WireMessage>,
218 ) -> (Self, Rc<RefCell<MockState>>) {
219 let state = Rc::new(RefCell::new(MockState {
220 incoming: messages.into_iter().collect(),
221 ..MockState::default()
222 }));
223 (
224 Self {
225 state: state.clone(),
226 },
227 state,
228 )
229 }
230 }
231
232 impl Transport for MockTransport {
233 fn send<'a>(
234 &'a mut self,
235 message: WireMessage,
236 ) -> LocalBoxFuture<'a, Result<(), TransportError>> {
237 self.state.borrow_mut().sent.push(message);
238 Box::pin(async { Ok(()) })
239 }
240
241 fn receive<'a>(&'a mut self) -> LocalBoxFuture<'a, Result<TransportEvent, TransportError>> {
242 let message = self.state.borrow_mut().incoming.pop_front();
243 Box::pin(async move {
244 Ok(message.map_or_else(
245 || TransportEvent::Closed(crate::TransportClose::connection_ended()),
246 TransportEvent::Message,
247 ))
248 })
249 }
250
251 fn close<'a>(&'a mut self) -> LocalBoxFuture<'a, Result<(), TransportError>> {
252 self.state.borrow_mut().closed = true;
253 Box::pin(async { Ok(()) })
254 }
255 }
256
257 fn text(frame: &str) -> WireMessage {
258 WireMessage::Text(frame.into())
259 }
260
261 #[test]
262 fn joins_and_sends_a_v2_frame() {
263 futures::executor::block_on(async {
264 let (transport, state) = MockTransport::with_incoming([text(
265 r#"["1","1","room:lobby","phx_reply",{"status":"ok","response":{"ready":true}}]"#,
266 )]);
267 let mut session = Session::new(transport);
268
269 let event = session
270 .join("room:lobby", json!({"token": "abc"}))
271 .await
272 .unwrap();
273
274 assert!(matches!(
275 event,
276 ProtocolEvent::Joined { response, .. } if response == json!({"ready": true})
277 ));
278 let sent = state.borrow().sent[0].clone();
279 let WireMessage::Text(sent) = sent else {
280 panic!("expected text frame")
281 };
282 let frame = Frame::decode_text(&sent).unwrap();
283 assert_eq!(frame.event, "phx_join");
284 assert_eq!(frame.payload, json!({"token": "abc"}));
285 });
286 }
287
288 #[test]
289 fn buffers_broadcasts_received_while_waiting_for_a_reply() {
290 futures::executor::block_on(async {
291 let (transport, _) = MockTransport::with_incoming([
292 text(r#"["1",null,"room:lobby","new_message",{"body":"early"}]"#),
293 text(r#"["1","1","room:lobby","phx_reply",{"status":"ok","response":{}}]"#),
294 ]);
295 let mut session = Session::new(transport);
296
297 session.join("room:lobby", json!({})).await.unwrap();
298 let event = session.next_event().await.unwrap();
299
300 assert!(matches!(
301 event,
302 ProtocolEvent::Message(Frame { event, .. }) if event == "new_message"
303 ));
304 });
305 }
306}