Skip to main content

easyfix_session/
acceptor.rs

1use std::{
2    cell::{Cell, RefCell},
3    collections::HashMap,
4    future::Future,
5    io,
6    net::SocketAddr,
7    pin::Pin,
8    rc::Rc,
9    task::{Context, Poll},
10};
11
12use easyfix_messages::fields::{FixString, SeqNum, SessionStatus};
13use futures::{self, Stream};
14use pin_project::pin_project;
15use tokio::{
16    io::{AsyncRead, AsyncWrite},
17    net::TcpListener,
18    task::JoinHandle,
19};
20use tracing::{Instrument, error, info, info_span, instrument, warn};
21
22use crate::{
23    DisconnectReason, Settings,
24    application::{AsEvent, Emitter, EventStream, events_channel},
25    io::{PendingLogout, acceptor_connection, supervise_connection},
26    messages_storage::MessagesStorage,
27    session::Session,
28    session_id::SessionId,
29    session_state::State as SessionState,
30    settings::SessionSettings,
31};
32
33#[derive(Debug, thiserror::Error)]
34pub enum AcceptorError {
35    #[error("Unknown session")]
36    UnknownSession,
37    #[error("Session active")]
38    SessionActive,
39}
40
41#[allow(async_fn_in_trait)]
42pub trait Connection {
43    async fn accept(
44        &mut self,
45    ) -> Result<
46        (
47            impl AsyncRead + Unpin + 'static,
48            impl AsyncWrite + Unpin + 'static,
49            SocketAddr,
50        ),
51        io::Error,
52    >;
53}
54
55pub struct TcpConnection {
56    listener: TcpListener,
57}
58
59impl TcpConnection {
60    pub async fn new(socket_addr: impl Into<SocketAddr>) -> Result<TcpConnection, io::Error> {
61        let socket_addr = socket_addr.into();
62        let listener = TcpListener::bind(&socket_addr).await?;
63        Ok(TcpConnection { listener })
64    }
65}
66
67impl Connection for TcpConnection {
68    async fn accept(
69        &mut self,
70    ) -> Result<
71        (
72            impl AsyncRead + Unpin + 'static,
73            impl AsyncWrite + Unpin + 'static,
74            SocketAddr,
75        ),
76        io::Error,
77    > {
78        let (tcp_stream, peer_addr) = self.listener.accept().await?;
79        tcp_stream.set_nodelay(true)?;
80        let (reader, writer) = tcp_stream.into_split();
81        Ok((reader, writer, peer_addr))
82    }
83}
84
85type SessionMapInternal<S> = HashMap<SessionId, (SessionSettings, Rc<RefCell<SessionState<S>>>)>;
86
87pub struct SessionsMap<S> {
88    map: SessionMapInternal<S>,
89    message_storage_builder: Box<dyn Fn(&SessionId) -> S>,
90}
91
92impl<S: MessagesStorage> SessionsMap<S> {
93    fn new(message_storage_builder: Box<dyn Fn(&SessionId) -> S>) -> SessionsMap<S> {
94        SessionsMap {
95            map: HashMap::new(),
96            message_storage_builder,
97        }
98    }
99
100    pub fn register_session(&mut self, session_id: SessionId, session_settings: SessionSettings) {
101        let storage = (self.message_storage_builder)(&session_id);
102        self.map.insert(
103            session_id.clone(),
104            (
105                session_settings,
106                Rc::new(RefCell::new(SessionState::new(storage))),
107            ),
108        );
109    }
110
111    pub(crate) fn get_session(
112        &self,
113        session_id: &SessionId,
114    ) -> Option<(SessionSettings, Rc<RefCell<SessionState<S>>>)> {
115        self.map.get(session_id).cloned()
116    }
117
118    fn contains(&self, session_id: &SessionId) -> bool {
119        self.map.contains_key(session_id)
120    }
121}
122
123pub struct SessionTask<S> {
124    settings: Settings,
125    sessions: Rc<RefCell<SessionsMap<S>>>,
126    active_sessions: Rc<RefCell<ActiveSessionsMap<S>>>,
127    emitter: Emitter,
128    enabled: Rc<Cell<bool>>,
129}
130
131impl<S> Clone for SessionTask<S> {
132    fn clone(&self) -> Self {
133        Self {
134            settings: self.settings.clone(),
135            sessions: self.sessions.clone(),
136            active_sessions: self.active_sessions.clone(),
137            emitter: self.emitter.clone(),
138            enabled: self.enabled.clone(),
139        }
140    }
141}
142
143impl<S: MessagesStorage + 'static> SessionTask<S> {
144    fn new(
145        settings: Settings,
146        sessions: Rc<RefCell<SessionsMap<S>>>,
147        active_sessions: Rc<RefCell<ActiveSessionsMap<S>>>,
148        emitter: Emitter,
149        enabled: Rc<Cell<bool>>,
150    ) -> SessionTask<S> {
151        SessionTask {
152            settings,
153            sessions,
154            active_sessions,
155            emitter,
156            enabled,
157        }
158    }
159
160    pub async fn run(
161        self,
162        peer_addr: SocketAddr,
163        reader: impl AsyncRead + Unpin + 'static,
164        writer: impl AsyncWrite + Unpin + 'static,
165    ) {
166        let span = info_span!("connection", %peer_addr);
167
168        span.in_scope(|| {
169            info!("New connection");
170        });
171
172        if self.enabled.get() {
173            let pending_logout = PendingLogout::default();
174            supervise_connection(
175                acceptor_connection(
176                    reader,
177                    writer,
178                    self.settings,
179                    self.sessions,
180                    self.active_sessions,
181                    self.emitter.clone(),
182                    self.enabled,
183                    pending_logout.clone(),
184                ),
185                pending_logout,
186                &self.emitter,
187            )
188            .instrument(span.clone())
189            .await;
190        } else {
191            span.in_scope(|| warn!("Acceptor is disabled"))
192        }
193
194        span.in_scope(|| {
195            info!("Connection closed");
196        });
197    }
198}
199
200pub(crate) type ActiveSessionsMap<S> = HashMap<SessionId, Rc<Session<S>>>;
201
202#[pin_project]
203pub struct Acceptor<S> {
204    sessions: Rc<RefCell<SessionsMap<S>>>,
205    active_sessions: Rc<RefCell<ActiveSessionsMap<S>>>,
206    session_task: SessionTask<S>,
207    #[pin]
208    event_stream: EventStream,
209    enabled: Rc<Cell<bool>>,
210}
211
212impl<S: MessagesStorage + 'static> Acceptor<S> {
213    pub fn new(
214        settings: Settings,
215        message_storage_builder: Box<dyn Fn(&SessionId) -> S>,
216    ) -> Acceptor<S> {
217        let (emitter, event_stream) = events_channel();
218        let sessions = Rc::new(RefCell::new(SessionsMap::new(message_storage_builder)));
219        let active_sessions = Rc::new(RefCell::new(HashMap::new()));
220        let enabled = Rc::new(Cell::new(true));
221        let session_task = SessionTask::new(
222            settings,
223            sessions.clone(),
224            active_sessions.clone(),
225            emitter,
226            enabled.clone(),
227        );
228
229        Acceptor {
230            sessions,
231            active_sessions,
232            session_task,
233            event_stream,
234            enabled,
235        }
236    }
237
238    pub fn enable(&self) {
239        info!("acceptor enabled");
240        self.enabled.set(true);
241    }
242
243    pub fn disable(&self) {
244        info!("acceptor disabled");
245        self.enabled.set(false);
246        for (_, session) in self.active_sessions.borrow_mut().drain() {
247            session.disconnect(
248                &mut session.state().borrow_mut(),
249                DisconnectReason::ApplicationForcedDisconnect,
250            );
251        }
252    }
253
254    pub fn disable_with_logout(
255        &self,
256        session_status: Option<SessionStatus>,
257        reason: Option<FixString>,
258    ) {
259        info!("acceptor disabled with logout");
260        self.enabled.set(false);
261        for (_, session) in self.active_sessions.borrow_mut().drain() {
262            let mut state = session.state().borrow_mut();
263            session.send_logout(&mut state, session_status, reason.clone());
264            session.disconnect(&mut state, DisconnectReason::ApplicationForcedDisconnect);
265        }
266    }
267
268    pub fn register_session(&mut self, session_id: SessionId, session_settings: SessionSettings) {
269        self.sessions
270            .borrow_mut()
271            .register_session(session_id, session_settings);
272    }
273
274    pub fn sessions_map(&self) -> Rc<RefCell<SessionsMap<S>>> {
275        self.sessions.clone()
276    }
277
278    pub fn start(&self, connection: impl Connection + 'static) -> JoinHandle<()> {
279        tokio::task::spawn_local(Self::server_task(connection, self.session_task.clone()))
280    }
281
282    pub fn is_session_active(&self, session_id: &SessionId) -> Result<bool, AcceptorError> {
283        if self.active_sessions.borrow().contains_key(session_id) {
284            Ok(true)
285        } else if self.sessions.borrow().contains(session_id) {
286            Ok(false)
287        } else {
288            Err(AcceptorError::UnknownSession)
289        }
290    }
291
292    pub fn logout(
293        &self,
294        session_id: &SessionId,
295        session_status: Option<SessionStatus>,
296        reason: Option<FixString>,
297    ) -> Result<(), AcceptorError> {
298        if let Some(session) = self.active_sessions.borrow().get(session_id) {
299            session.send_logout(&mut session.state().borrow_mut(), session_status, reason);
300            Ok(())
301        } else if self.sessions.borrow().contains(session_id) {
302            // Already logged out
303            Ok(())
304        } else {
305            Err(AcceptorError::UnknownSession)
306        }
307    }
308
309    pub fn disconnect(&self, session_id: &SessionId) -> Result<(), AcceptorError> {
310        if let Some(session) = self.active_sessions.borrow_mut().remove(session_id) {
311            session.disconnect(
312                &mut session.state().borrow_mut(),
313                DisconnectReason::ApplicationForcedDisconnect,
314            );
315            Ok(())
316        } else if self.sessions.borrow().contains(session_id) {
317            // Already disconnected
318            Ok(())
319        } else {
320            Err(AcceptorError::UnknownSession)
321        }
322    }
323
324    pub fn disconnect_with_logout(
325        &self,
326        session_id: &SessionId,
327        session_status: Option<SessionStatus>,
328        reason: Option<FixString>,
329    ) -> Result<(), AcceptorError> {
330        if let Some(session) = self.active_sessions.borrow().get(session_id) {
331            session.send_logout(&mut session.state().borrow_mut(), session_status, reason);
332            session.disconnect(
333                &mut session.state().borrow_mut(),
334                DisconnectReason::ApplicationForcedDisconnect,
335            );
336            Ok(())
337        } else if self.sessions.borrow().contains(session_id) {
338            // Already logged out
339            Ok(())
340        } else {
341            Err(AcceptorError::UnknownSession)
342        }
343    }
344
345    /// Force reset of the session
346    ///
347    /// Functionally equivalent to `reset_on_logon/logout/disconnect` settings,
348    /// but triggered manually.
349    ///
350    /// Returns [`AcceptorError::SessionActive`] if the session is still active.
351    /// In that case, call [Self::disconnect] or [Self::logout] first and wait
352    /// for the session to fully terminate before retrying.
353    #[instrument(skip_all, fields(session_id=%session_id) ret)]
354    pub fn reset(&self, session_id: &SessionId) -> Result<(), AcceptorError> {
355        if self.active_sessions.borrow().contains_key(session_id) {
356            Err(AcceptorError::SessionActive)
357        } else if let Some((_, session_state)) = self.sessions.borrow().get_session(session_id) {
358            session_state.borrow_mut().reset();
359            Ok(())
360        } else {
361            Err(AcceptorError::UnknownSession)
362        }
363    }
364
365    // TODO: temporary solution, remove when diconnect will be synchronized
366    #[instrument(skip_all, fields(session_id=%session_id) ret)]
367    pub fn force_reset(&self, session_id: &SessionId) -> Result<(), AcceptorError> {
368        if let Some(session) = self.active_sessions.borrow().get(session_id) {
369            session.state().borrow_mut().reset();
370            Ok(())
371        } else if let Some((_, session_state)) = self.sessions.borrow().get_session(session_id) {
372            session_state.borrow_mut().reset();
373            Ok(())
374        } else {
375            Err(AcceptorError::UnknownSession)
376        }
377    }
378
379    /// Sender seq_num getter
380    #[instrument(skip_all, fields(session_id=%session_id) ret)]
381    pub fn next_sender_msg_seq_num(&self, session_id: &SessionId) -> Result<SeqNum, AcceptorError> {
382        if let Some(session) = self.active_sessions.borrow().get(session_id) {
383            Ok(session.state().borrow().next_sender_msg_seq_num())
384        } else if let Some((_, session_state)) = self.sessions.borrow().get_session(session_id) {
385            Ok(session_state.borrow().next_sender_msg_seq_num())
386        } else {
387            Err(AcceptorError::UnknownSession)
388        }
389    }
390
391    /// Override sender's next seq_num
392    #[instrument(skip_all, fields(session_id=%session_id, seq_num) ret)]
393    pub fn set_next_sender_msg_seq_num(
394        &self,
395        session_id: &SessionId,
396        seq_num: SeqNum,
397    ) -> Result<(), AcceptorError> {
398        if let Some(session) = self.active_sessions.borrow().get(session_id) {
399            session
400                .state()
401                .borrow_mut()
402                .set_next_sender_msg_seq_num(seq_num);
403            Ok(())
404        } else if let Some((_, session_state)) = self.sessions.borrow().get_session(session_id) {
405            session_state
406                .borrow_mut()
407                .set_next_sender_msg_seq_num(seq_num);
408            Ok(())
409        } else {
410            Err(AcceptorError::UnknownSession)
411        }
412    }
413
414    async fn server_task(mut connection: impl Connection, session_task: SessionTask<S>) {
415        info!("Acceptor started");
416        loop {
417            match connection.accept().await {
418                Ok((reader, writer, peer_addr)) => {
419                    tokio::task::spawn_local(session_task.clone().run(peer_addr, reader, writer));
420                }
421                Err(err) => error!("server task failed to accept incoming connection: {err}"),
422            }
423        }
424    }
425
426    pub fn session_task(&self) -> SessionTask<S> {
427        self.session_task.clone()
428    }
429
430    pub fn run_session_task(
431        &self,
432        peer_addr: SocketAddr,
433        reader: impl AsyncRead + Unpin + 'static,
434        writer: impl AsyncWrite + Unpin + 'static,
435    ) -> impl Future<Output = ()> {
436        self.session_task.clone().run(peer_addr, reader, writer)
437    }
438}
439
440impl<S: MessagesStorage> Stream for Acceptor<S> {
441    type Item = impl AsEvent;
442
443    fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
444        Pin::new(&mut self.event_stream).poll_next(cx)
445    }
446
447    fn size_hint(&self) -> (usize, Option<usize>) {
448        self.event_stream.size_hint()
449    }
450}