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 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 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 Ok(())
340 } else {
341 Err(AcceptorError::UnknownSession)
342 }
343 }
344
345 #[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 #[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 #[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 #[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}