use crate::application::{Application, NoOpApplication, RejectReason, SessionId};
use crate::connection::{Connection, SessionRuntime};
use crate::error::EngineError;
use crate::reactor::{
DEFAULT_APP_QUEUE_CAPACITY, DEFAULT_OUTBOUND_CAPACITY, DEFAULT_WRITE_TIMEOUT, ResendState,
SessionParams, lock_heartbeat, send_handshake_admin, spawn_session,
};
use crate::wire::{self, MessageFactory, PeerIdentity, SendingTimeGuard, UnsupportedVersion};
use futures_util::StreamExt;
use ironfix_core::message::MsgType;
use ironfix_core::version::FixVersion;
use ironfix_session::sequence::SequenceResult;
use ironfix_session::{Disconnected, HeartbeatManager, SequenceManager, Session, SessionConfig};
use ironfix_transport::FixCodec;
use std::collections::HashSet;
use std::num::NonZeroU64;
use std::sync::{Arc, Mutex, PoisonError};
use std::time::Duration;
use tokio::net::{TcpListener, TcpStream};
use tokio::time::timeout;
use tokio_util::codec::Framed;
struct AdmissionGuard {
admitted: Arc<Mutex<HashSet<SessionId>>>,
session_id: SessionId,
}
impl Drop for AdmissionGuard {
fn drop(&mut self) {
let mut admitted = self.admitted.lock().unwrap_or_else(PoisonError::into_inner);
admitted.remove(&self.session_id);
}
}
#[derive(Debug)]
pub struct Acceptor<A: Application = NoOpApplication> {
config: SessionConfig,
application: Arc<A>,
session_id: SessionId,
version: Result<FixVersion, UnsupportedVersion>,
initial_sequences: Option<(u64, u64)>,
outbound_capacity: usize,
admitted: Arc<Mutex<HashSet<SessionId>>>,
}
impl<A: Application + 'static> Acceptor<A> {
#[must_use]
pub fn new(config: SessionConfig, application: Arc<A>) -> Self {
let session_id = wire::session_id_from_config(&config);
let version = wire::wire_version(&config.begin_string);
Self {
config,
application,
session_id,
version,
initial_sequences: None,
outbound_capacity: DEFAULT_OUTBOUND_CAPACITY,
admitted: Arc::new(Mutex::new(HashSet::new())),
}
}
fn try_admit(&self) -> Option<AdmissionGuard> {
let mut admitted = self.admitted.lock().unwrap_or_else(PoisonError::into_inner);
if !admitted.insert(self.session_id.clone()) {
return None;
}
Some(AdmissionGuard {
admitted: Arc::clone(&self.admitted),
session_id: self.session_id.clone(),
})
}
#[must_use]
pub fn with_initial_sequences(mut self, sender_seq: u64, target_seq: u64) -> Self {
self.initial_sequences = Some((sender_seq, target_seq));
self
}
#[must_use]
pub fn with_outbound_capacity(mut self, capacity: usize) -> Self {
self.outbound_capacity = capacity.max(1);
self
}
#[must_use]
pub fn config(&self) -> &SessionConfig {
&self.config
}
#[must_use]
pub fn session_id(&self) -> &SessionId {
&self.session_id
}
pub async fn accept(&self, listener: &TcpListener) -> Result<Connection, EngineError> {
let (stream, _addr) = listener.accept().await?;
self.serve(stream).await
}
pub async fn serve(&self, stream: TcpStream) -> Result<Connection, EngineError> {
let version = match &self.version {
Ok(version) => *version,
Err(err) => {
return Err(EngineError::UnsupportedVersion {
version: err.version.clone(),
detail: err.detail.clone(),
});
}
};
let session_id = self.session_id.clone();
self.application.on_create(&session_id).await;
let session = Session::<Disconnected>::new(session_id.to_string()).accept();
let _ = stream.set_nodelay(true);
let codec = FixCodec::new()
.with_max_message_size(self.config.max_message_size)
.with_checksum_validation(self.config.validate_checksum);
let mut framed = Framed::new(stream, codec);
let sequences = match self.initial_sequences {
Some((sender, target)) if !self.config.reset_on_logon => {
SequenceManager::with_initial(
NonZeroU64::new(sender).unwrap_or(NonZeroU64::MIN),
NonZeroU64::new(target).unwrap_or(NonZeroU64::MIN),
)
}
_ => SequenceManager::new(),
};
let runtime = Arc::new(SessionRuntime {
sequences,
heartbeat: Mutex::new(HeartbeatManager::new(self.config.heartbeat_interval)),
});
let mut factory = MessageFactory::new(&self.config, version);
let identity = PeerIdentity::new(&self.config);
let sending_time = SendingTimeGuard::new(&self.config);
let logon_frame = match timeout(self.config.logon_timeout, framed.next()).await {
Err(_) => {
let _ = session.disconnect();
return Err(EngineError::LogonTimeout(self.config.logon_timeout));
}
Ok(None) => {
let _ = session.disconnect();
return Err(EngineError::Closed);
}
Ok(Some(Err(err))) => {
let _ = session.disconnect();
return Err(err.into());
}
Ok(Some(Ok(frame))) => frame,
};
let session = session.on_logon_received();
let mut pending_resend: Option<(u64, u64)> = None;
let admission;
{
let raw = match wire::decode_frame(&logon_frame) {
Ok(raw) => raw,
Err(err) => {
let _ = session.reject_logon();
return Err(err.into());
}
};
match raw.msg_type() {
MsgType::Logon => {}
other => {
let msg_type = other.as_str().to_string();
let _ = session.reject_logon();
return Err(EngineError::UnexpectedMessage { msg_type });
}
}
let inbound_begin_string = raw.begin_string().unwrap_or_default();
if inbound_begin_string != version.begin_string() {
let detail = format!(
"Logon BeginString (8) '{}' does not match the configured '{}'",
inbound_begin_string.chars().take(16).collect::<String>(),
version.begin_string()
);
let logout = factory.logout(Some(&detail));
let _ = send_handshake_admin(
self.application.as_ref(),
&session_id,
&mut framed,
&mut factory,
&runtime.sequences,
None,
DEFAULT_WRITE_TIMEOUT,
logout,
)
.await;
let _ = session.reject_logon();
return Err(EngineError::LogonRejected { reason: detail });
}
let encrypt_method = raw.get_field_str(98);
if encrypt_method != Some("0") {
let shown = encrypt_method
.map_or_else(|| "<absent>".to_string(), |v| v.chars().take(16).collect());
let detail =
format!("unsupported EncryptMethod (98) '{shown}'; only 0 (None) is supported");
let logout = factory.logout(Some(&detail));
let _ = send_handshake_admin(
self.application.as_ref(),
&session_id,
&mut framed,
&mut factory,
&runtime.sequences,
None,
DEFAULT_WRITE_TIMEOUT,
logout,
)
.await;
let _ = session.reject_logon();
return Err(EngineError::LogonRejected { reason: detail });
}
let logon_seq = match wire::header_seq_num(&raw) {
Some(seq) => seq,
None => {
let _ = session.reject_logon();
return Err(EngineError::Sequence(
"Logon has no valid MsgSeqNum (34) in the standard header".to_string(),
));
}
};
if let Err(mismatch) = identity.validate(&raw) {
let detail = mismatch.to_string();
let reason = RejectReason::new(9, detail.clone()).with_ref_tag(mismatch.tag);
let reject = factory.session_reject(logon_seq, MsgType::Logon.as_str(), &reason);
let _ = send_handshake_admin(
self.application.as_ref(),
&session_id,
&mut framed,
&mut factory,
&runtime.sequences,
None,
DEFAULT_WRITE_TIMEOUT,
reject,
)
.await;
let logout = factory.logout(Some(&detail));
let _ = send_handshake_admin(
self.application.as_ref(),
&session_id,
&mut framed,
&mut factory,
&runtime.sequences,
None,
DEFAULT_WRITE_TIMEOUT,
logout,
)
.await;
let _ = session.reject_logon();
return Err(EngineError::IdentityMismatch { detail });
}
admission = match self.try_admit() {
Some(guard) => guard,
None => {
let detail = format!("session already active for {session_id}");
tracing::warn!(session = %session_id, "refusing duplicate concurrent Logon");
let logout = factory.logout(Some(&detail));
let _ = send_handshake_admin(
self.application.as_ref(),
&session_id,
&mut framed,
&mut factory,
&runtime.sequences,
None,
DEFAULT_WRITE_TIMEOUT,
logout,
)
.await;
let _ = session.reject_logon();
return Err(EngineError::LogonRejected { reason: detail });
}
};
if let Err(problem) = sending_time.validate(&raw) {
let detail = problem.to_string();
let reason = problem.reject_reason();
let reject = factory.session_reject(logon_seq, MsgType::Logon.as_str(), &reason);
let _ = send_handshake_admin(
self.application.as_ref(),
&session_id,
&mut framed,
&mut factory,
&runtime.sequences,
None,
DEFAULT_WRITE_TIMEOUT,
reject,
)
.await;
let logout = factory.logout(Some(&detail));
let _ = send_handshake_admin(
self.application.as_ref(),
&session_id,
&mut framed,
&mut factory,
&runtime.sequences,
None,
DEFAULT_WRITE_TIMEOUT,
logout,
)
.await;
let _ = session.reject_logon();
return Err(EngineError::SendingTime { detail });
}
if let Err(reason) = self.application.from_admin(&raw, &session_id).await {
let logout = factory.logout(Some(&reason.text));
let _ = send_handshake_admin(
self.application.as_ref(),
&session_id,
&mut framed,
&mut factory,
&runtime.sequences,
None,
DEFAULT_WRITE_TIMEOUT,
logout,
)
.await;
let _ = session.reject_logon();
return Err(EngineError::LogonRejected {
reason: reason.text,
});
}
let requested_heartbeat_secs = match wire::parse_heartbeat_interval(&raw) {
Ok(secs) => secs,
Err(problem) => {
let detail = problem.to_string();
tracing::warn!(session = %session_id, detail = %detail, "rejecting Logon HeartBtInt");
let logout = factory.logout(Some(&detail));
let _ = send_handshake_admin(
self.application.as_ref(),
&session_id,
&mut framed,
&mut factory,
&runtime.sequences,
None,
DEFAULT_WRITE_TIMEOUT,
logout,
)
.await;
let _ = session.reject_logon();
return Err(EngineError::LogonRejected { reason: detail });
}
};
{
let mut heartbeat = lock_heartbeat(&runtime);
heartbeat.on_message_received(false, None);
if Duration::from_secs(requested_heartbeat_secs) != heartbeat.interval() {
tracing::info!(
session = %session_id,
heartbeat_secs = requested_heartbeat_secs,
"honoring HeartBtInt requested by counterparty"
);
*heartbeat =
HeartbeatManager::new(Duration::from_secs(requested_heartbeat_secs));
}
}
let reply_heartbeat_secs = requested_heartbeat_secs;
let inbound_reset = raw.get_field_str(141) == Some("Y");
let reset = inbound_reset || self.config.reset_on_logon;
if inbound_reset && logon_seq != 1 {
let _ = session.reject_logon();
return Err(EngineError::Sequence(format!(
"Logon set ResetSeqNumFlag=Y but carried MsgSeqNum {logon_seq}, not 1"
)));
}
if reset {
tracing::info!(
session = %session_id,
inbound_reset,
"resetting sequence numbers on Logon"
);
runtime.sequences.reset();
}
match runtime.sequences.validate_incoming(logon_seq) {
SequenceResult::Ok => {
if let Err(err) = runtime.sequences.try_increment_target_seq() {
let _ = session.reject_logon();
return Err(err.into());
}
}
SequenceResult::TooLow { expected, received } => {
let detail = format!(
"logon MsgSeqNum too low: expected {expected}, received {received}"
);
let logout = factory.logout(Some(&detail));
let _ = send_handshake_admin(
self.application.as_ref(),
&session_id,
&mut framed,
&mut factory,
&runtime.sequences,
None,
DEFAULT_WRITE_TIMEOUT,
logout,
)
.await;
let _ = session.reject_logon();
return Err(EngineError::Sequence(detail));
}
SequenceResult::Gap { expected, received } => {
pending_resend = Some((expected, received));
}
}
let logon_ack = factory.logon(reply_heartbeat_secs, reset);
if let Err(err) = send_handshake_admin(
self.application.as_ref(),
&session_id,
&mut framed,
&mut factory,
&runtime.sequences,
None,
DEFAULT_WRITE_TIMEOUT,
logon_ack,
)
.await
{
let _ = session.reject_logon();
return Err(err);
}
lock_heartbeat(&runtime).on_message_sent();
}
let session = session.accept_logon();
self.application.on_logon(&session_id).await;
tracing::info!(session = %session_id, "FIX session established (acceptor)");
if let Some((expected, _high_water)) = pending_resend {
let request = factory.resend_request(expected, 0);
send_handshake_admin(
self.application.as_ref(),
&session_id,
&mut framed,
&mut factory,
&runtime.sequences,
None,
DEFAULT_WRITE_TIMEOUT,
request,
)
.await?;
lock_heartbeat(&runtime).on_message_sent();
}
let resend = pending_resend
.map(|(expected, high_water)| ResendState::first(expected, high_water, &self.config));
let params = SessionParams {
framed,
session,
runtime: Arc::clone(&runtime),
factory,
identity,
sending_time,
config: self.config.clone(),
application: Arc::clone(&self.application),
session_id: session_id.clone(),
store: None,
resend,
write_timeout: DEFAULT_WRITE_TIMEOUT,
outbound_capacity: self.outbound_capacity,
app_queue_capacity: DEFAULT_APP_QUEUE_CAPACITY,
};
let (command_tx, closed_rx) = spawn_session(params, admission);
Ok(Connection {
session_id,
commands: command_tx,
closed: closed_rx,
runtime,
})
}
}