Skip to main content

phoenix_channel_runtime/
session.rs

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
11/// A sequential, runtime-neutral Phoenix Channels session.
12///
13/// `Session` owns a transport and preserves server events that arrive while a
14/// join, push, leave, or heartbeat reply is being awaited. Applications that
15/// need commands from multiple tasks should own the session in one task and
16/// communicate with it over their executor's channel type.
17pub 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    /// Creates a session using the bounded default Phoenix v2 codec.
26    pub fn new(transport: T) -> Self {
27        Self::with_codec(transport, PhoenixV2Codec::limited(Default::default()))
28    }
29
30    /// Creates a session with an application-provided codec.
31    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    /// Borrows the session's protocol state.
41    pub fn protocol(&self) -> &Protocol {
42        &self.protocol
43    }
44
45    /// Mutably borrows the session's protocol state.
46    pub fn protocol_mut(&mut self) -> &mut Protocol {
47        &mut self.protocol
48    }
49
50    /// Mutably borrows the underlying transport.
51    pub fn transport_mut(&mut self) -> &mut T {
52        &mut self.transport
53    }
54
55    /// Consumes the session and returns its protocol and transport.
56    pub fn into_parts(self) -> (Protocol, T) {
57        (self.protocol, self.transport)
58    }
59
60    /// Sends a join and waits for its correlated reply.
61    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    /// Sends a rejoin with refreshed parameters and waits for its reply.
73    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    /// Sends an application event and waits for its correlated reply.
85    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    /// Leaves a joined topic and waits for its reply.
98    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    /// Sends a Phoenix heartbeat and waits for its acknowledgement.
106    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    /// Returns the next buffered or transport event.
114    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    /// Marks all channels disconnected and returns the requests interrupted by
122    /// transport loss. The caller can then reconnect a transport and call
123    /// `rejoin` with refreshed authentication parameters.
124    pub fn reset_connection(&mut self) -> Vec<ProtocolEvent> {
125        self.protocol.reset_connection()
126    }
127
128    /// Closes the underlying transport.
129    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/// Protocol, codec, or transport failure produced by a [`Session`].
178#[derive(Debug, Error)]
179pub enum SessionError {
180    /// Protocol state rejected an operation or incoming reply.
181    #[error(transparent)]
182    Protocol(#[from] ProtocolError),
183    /// A frame could not be encoded or decoded.
184    #[error(transparent)]
185    Codec(#[from] CodecError),
186    /// The underlying transport operation failed.
187    #[error(transparent)]
188    Transport(#[from] TransportError),
189    /// The transport closed while the session was waiting for an event.
190    #[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}