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, TICK_INTERVAL, lock_heartbeat, send_handshake_admin, spawn_session,
};
use crate::wire::{self, MessageFactory, PeerIdentity, SendingTimeGuard, UnsupportedVersion};
use futures_util::StreamExt;
use ironfix_core::error::DecodeError;
use ironfix_core::message::MsgType;
use ironfix_core::version::FixVersion;
use ironfix_session::config::SessionConfigError;
use ironfix_session::heartbeat::negotiate_interval;
use ironfix_session::sequence::{SequenceCounter, SequenceResult};
use ironfix_session::{Disconnected, HeartbeatManager, SequenceManager, Session, SessionConfig};
use ironfix_store::MessageStore;
use ironfix_transport::FixCodec;
use std::num::NonZeroU64;
use std::sync::{Arc, Mutex};
use std::time::Duration;
use tokio::net::{TcpStream, ToSocketAddrs};
use tokio::time::timeout;
use tokio_util::codec::Framed;
pub struct Initiator<A: Application = NoOpApplication> {
config: SessionConfig,
application: Arc<A>,
session_id: SessionId,
version: Result<FixVersion, UnsupportedVersion>,
config_check: Result<(), SessionConfigError>,
connect_timeout: Duration,
write_timeout: Duration,
initial_sequences: Option<(u64, u64)>,
outbound_capacity: usize,
app_queue_capacity: usize,
store: Option<Arc<dyn MessageStore>>,
}
impl<A: Application> std::fmt::Debug for Initiator<A> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Initiator")
.field("config", &self.config)
.field("session_id", &self.session_id)
.field("version", &self.version)
.field("connect_timeout", &self.connect_timeout)
.field("write_timeout", &self.write_timeout)
.field("initial_sequences", &self.initial_sequences)
.field("outbound_capacity", &self.outbound_capacity)
.field("app_queue_capacity", &self.app_queue_capacity)
.field("store", &self.store.is_some())
.finish_non_exhaustive()
}
}
impl<A: Application + 'static> Initiator<A> {
#[must_use]
pub fn new(config: SessionConfig, application: Arc<A>) -> Self {
let mut session_id = SessionId::new(
config.begin_string.clone(),
config.sender_comp_id.as_str(),
config.target_comp_id.as_str(),
);
if let Some(sub) = &config.sender_sub_id {
session_id = session_id.with_sender_sub_id(sub.clone());
}
if let Some(sub) = &config.target_sub_id {
session_id = session_id.with_target_sub_id(sub.clone());
}
let version = wire::wire_version(&config.begin_string);
let config_check = config.validate();
Self {
config,
application,
session_id,
version,
config_check,
connect_timeout: Duration::from_secs(30),
write_timeout: DEFAULT_WRITE_TIMEOUT,
initial_sequences: None,
outbound_capacity: DEFAULT_OUTBOUND_CAPACITY,
app_queue_capacity: DEFAULT_APP_QUEUE_CAPACITY,
store: None,
}
}
#[must_use]
pub fn with_store(mut self, store: Arc<dyn MessageStore>) -> Self {
self.store = Some(store);
self
}
#[must_use]
pub fn with_connect_timeout(mut self, timeout: Duration) -> Self {
self.connect_timeout = timeout;
self
}
#[must_use]
pub fn with_write_timeout(mut self, timeout: Duration) -> Self {
self.write_timeout = timeout.max(TICK_INTERVAL);
self
}
#[must_use]
pub fn with_app_queue_capacity(mut self, capacity: usize) -> Self {
self.app_queue_capacity = capacity.max(1);
self
}
#[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
}
async fn seed_sequences(&self) -> Result<SequenceManager, EngineError> {
if self.config.reset_on_logon {
if let Some(store) = &self.store
&& let Err(err) = store.reset().await
{
tracing::error!(
session = %self.session_id,
error = %err,
"cannot reset the store for a ResetSeqNumFlag logon: refusing the session \
rather than starting at 1 with the previous stream still filed"
);
return Err(EngineError::Store(err));
}
return Ok(SequenceManager::new());
}
if let Some((sender, target)) = self.initial_sequences {
return Ok(SequenceManager::with_initial(
NonZeroU64::new(sender).unwrap_or(NonZeroU64::MIN),
NonZeroU64::new(target).unwrap_or(NonZeroU64::MIN),
));
}
let Some(store) = &self.store else {
return Ok(SequenceManager::new());
};
if let Err(err) = store.refresh().await {
tracing::error!(
session = %self.session_id,
error = %err,
"cannot refresh the store: refusing the session rather than starting against \
unknown counters and risking a reused MsgSeqNum"
);
return Err(EngineError::Store(err));
}
Ok(SequenceManager::with_initial(
NonZeroU64::new(store.next_sender_seq().max(1)).unwrap_or(NonZeroU64::MIN),
NonZeroU64::new(store.next_target_seq().max(1)).unwrap_or(NonZeroU64::MIN),
))
}
#[must_use]
pub fn config(&self) -> &SessionConfig {
&self.config
}
#[must_use]
pub fn session_id(&self) -> &SessionId {
&self.session_id
}
pub async fn connect(&self, addr: impl ToSocketAddrs) -> Result<Connection, EngineError> {
if let Err(err) = &self.config_check {
return Err(EngineError::Config(err.clone()));
}
let version = match &self.version {
Ok(version) => *version,
Err(err) => {
return Err(EngineError::UnsupportedVersion {
version: err.version.clone(),
detail: err.detail.clone(),
});
}
};
if let Some((sender, target)) = self.initial_sequences
&& !self.config.reset_on_logon
{
NonZeroU64::new(sender).ok_or(EngineError::InvalidInitialSequence {
counter: SequenceCounter::Sender,
})?;
NonZeroU64::new(target).ok_or(EngineError::InvalidInitialSequence {
counter: SequenceCounter::Target,
})?;
}
let session_id = self.session_id.clone();
self.application.on_create(&session_id).await;
let session = Session::<Disconnected>::new(session_id.to_string()).connect();
let stream = match timeout(self.connect_timeout, TcpStream::connect(addr)).await {
Err(_) => {
let _ = session.disconnect();
return Err(EngineError::ConnectTimeout(self.connect_timeout));
}
Ok(Err(err)) => {
let _ = session.disconnect();
return Err(err.into());
}
Ok(Ok(stream)) => stream,
};
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.seed_sequences().await {
Ok(sequences) => sequences,
Err(err) => {
let _ = session.disconnect();
return Err(err);
}
};
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 write_timeout = self.write_timeout;
let store = self.store.as_ref();
let logon = factory.logon(
self.config.heartbeat_interval_secs(),
self.config.reset_on_logon,
);
if let Err(err) = send_handshake_admin(
self.application.as_ref(),
&session_id,
&mut framed,
&mut factory,
&runtime.sequences,
store,
write_timeout,
logon,
)
.await
{
let _ = session.disconnect();
return Err(err);
}
lock_heartbeat(&runtime).on_message_sent();
let session = session.send_logon();
let ack = match timeout(self.config.logon_timeout, framed.next()).await {
Err(_) => {
let _ = session.on_logon_reject();
return Err(EngineError::LogonTimeout(self.config.logon_timeout));
}
Ok(None) => {
let _ = session.on_logon_reject();
return Err(EngineError::Closed);
}
Ok(Some(Err(err))) => {
let _ = session.on_logon_reject();
return Err(err.into());
}
Ok(Some(Ok(frame))) => frame,
};
let mut pending_resend: Option<(u64, u64)> = None;
{
let raw = wire::decode_frame(&ack)?;
match raw.msg_type() {
MsgType::Logon => {}
MsgType::Logout | MsgType::Reject => {
let reason = raw
.get_field_str(58)
.unwrap_or("logon rejected by counterparty")
.to_string();
let _ = session.on_logon_reject();
return Err(EngineError::LogonRejected { reason });
}
other => {
let msg_type = other.as_str().to_string();
let _ = session.on_logon_reject();
return Err(EngineError::UnexpectedMessage { msg_type });
}
}
match raw.begin_string() {
Ok(begin_string) if begin_string == version.begin_string() => {}
Ok(begin_string) => {
let received = begin_string.to_string();
let _ = session.on_logon_reject();
return Err(EngineError::BeginStringMismatch {
expected: version.begin_string().to_string(),
received,
});
}
Err(err) => {
let _ = session.on_logon_reject();
return Err(err.into());
}
}
let Some(ack_seq) = wire::header_seq_num(&raw) else {
let _ = session.on_logon_reject();
return Err(EngineError::Sequence(
"Logon ack 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(ack_seq, MsgType::Logon.as_str(), &reason);
let _ = send_handshake_admin(
self.application.as_ref(),
&session_id,
&mut framed,
&mut factory,
&runtime.sequences,
store,
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,
store,
write_timeout,
logout,
)
.await;
let _ = session.on_logon_reject();
return Err(EngineError::IdentityMismatch { detail });
}
if let Err(problem) = sending_time.validate(&raw) {
let detail = problem.to_string();
let reason = problem.reject_reason();
let reject = factory.session_reject(ack_seq, MsgType::Logon.as_str(), &reason);
let _ = send_handshake_admin(
self.application.as_ref(),
&session_id,
&mut framed,
&mut factory,
&runtime.sequences,
store,
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,
store,
write_timeout,
logout,
)
.await;
let _ = session.on_logon_reject();
return Err(EngineError::SendingTime { detail });
}
if let Err(reason) = self.application.from_admin(&raw, &session_id).await {
let logout = factory.logout(Some(&reason.text));
if let Err(err) = send_handshake_admin(
self.application.as_ref(),
&session_id,
&mut framed,
&mut factory,
&runtime.sequences,
store,
write_timeout,
logout,
)
.await
{
tracing::warn!(
session = %session_id,
error = %err,
"cannot send Logout after from_admin rejection"
);
}
let _ = session.on_logon_reject();
return Err(EngineError::LogonRejected {
reason: reason.text,
});
}
let secs: u64 = match raw.get_field_as::<u64>(108) {
Ok(secs) => secs,
Err(err) => {
let (code, detail) = if matches!(err, DecodeError::MissingRequiredField { .. })
{
(
1,
"Logon acknowledgement omitted the required HeartBtInt (108)"
.to_string(),
)
} else {
(6, err.to_string())
};
let reason = RejectReason::new(code, detail.clone()).with_ref_tag(108);
let reject = factory.session_reject(ack_seq, MsgType::Logon.as_str(), &reason);
let _ = send_handshake_admin(
self.application.as_ref(),
&session_id,
&mut framed,
&mut factory,
&runtime.sequences,
store,
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,
store,
write_timeout,
logout,
)
.await;
let _ = session.on_logon_reject();
return Err(EngineError::HeartbeatInterval { detail });
}
};
let negotiated = match negotiate_interval(self.config.heartbeat_interval, secs) {
Ok(interval) => interval,
Err(err) => {
let detail = err.to_string();
let reason = RejectReason::new(5, detail.clone()).with_ref_tag(108);
let reject = factory.session_reject(ack_seq, MsgType::Logon.as_str(), &reason);
let _ = send_handshake_admin(
self.application.as_ref(),
&session_id,
&mut framed,
&mut factory,
&runtime.sequences,
store,
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,
store,
write_timeout,
logout,
)
.await;
let _ = session.on_logon_reject();
return Err(EngineError::HeartbeatInterval { detail });
}
};
{
let mut heartbeat = lock_heartbeat(&runtime);
heartbeat.on_message_received(false, None);
if negotiated != heartbeat.interval() {
tracing::info!(
session = %session_id,
heartbeat_secs = negotiated.as_secs(),
"using heartbeat interval confirmed by counterparty"
);
*heartbeat = HeartbeatManager::new(negotiated);
}
}
if raw.get_field_str(141) == Some("Y") {
if ack_seq != 1 {
let _ = session.on_logon_reject();
return Err(EngineError::Sequence(format!(
"Logon ack set ResetSeqNumFlag=Y but carried MsgSeqNum {ack_seq}, not 1"
)));
}
tracing::info!(
session = %session_id,
"counterparty set ResetSeqNumFlag on the Logon ack: resetting sequence numbers"
);
runtime.sequences.set_target_seq(1);
runtime.sequences.set_sender_seq(2);
}
match runtime.sequences.validate_incoming(ack_seq) {
SequenceResult::Ok => {
if let Err(err) = runtime.sequences.try_increment_target_seq() {
let _ = session.on_logon_reject();
return Err(err.into());
}
}
SequenceResult::TooLow { expected, received } => {
let _ = session.on_logon_reject();
return Err(EngineError::Sequence(format!(
"logon ack MsgSeqNum too low: expected {expected}, received {received}"
)));
}
SequenceResult::Gap { expected, received } => {
pending_resend = Some((expected, received));
}
}
}
let session = session.on_logon_ack();
self.application.on_logon(&session_id).await;
tracing::info!(session = %session_id, "FIX session established");
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,
store,
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: self.store.clone(),
resend,
write_timeout,
outbound_capacity: self.outbound_capacity,
app_queue_capacity: self.app_queue_capacity,
};
let (command_tx, closed_rx) = spawn_session(params, ());
Ok(Connection {
session_id,
commands: command_tx,
closed: closed_rx,
runtime,
})
}
}