use crate::application::{Application, RejectReason, SessionId};
use crate::connection::{Command, SessionRuntime};
use crate::error::EngineError;
use crate::outbound;
use crate::wire::{
self, MessageFactory, PeerIdentity, PendingMessage, SendingTimeGuard, SendingTimeProblem,
};
use bytes::{Bytes, BytesMut};
use futures_util::{SinkExt, StreamExt};
use ironfix_core::error::EncodeError;
use ironfix_core::message::{MsgType, RawMessage};
use ironfix_session::heartbeat::generate_test_req_id;
use ironfix_session::sequence::{SequenceExhausted, SequenceResult};
use ironfix_session::{
Active, HeartbeatManager, LogoutPending, SequenceManager, Session, SessionConfig,
TestRequestOutcome,
};
use ironfix_store::MessageStore;
use ironfix_transport::FixCodec;
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::net::TcpStream;
use tokio::sync::{mpsc, watch};
use tokio::task::JoinHandle;
use tokio::time::{MissedTickBehavior, interval, timeout};
use tokio_util::codec::Framed;
pub(crate) type FixFramed = Framed<TcpStream, FixCodec>;
pub(crate) const DEFAULT_OUTBOUND_CAPACITY: usize = 1024;
pub(crate) const DEFAULT_APP_QUEUE_CAPACITY: usize = 1024;
pub(crate) const DEFAULT_WRITE_TIMEOUT: Duration = Duration::from_secs(10);
const APP_DRAIN_TIMEOUT: Duration = Duration::from_secs(5);
pub(crate) const TICK_INTERVAL: Duration = Duration::from_millis(100);
const RESEND_PAGE_LIMIT: usize = 256;
#[allow(clippy::too_many_arguments)]
pub(crate) async fn send_handshake_admin<A: Application>(
application: &A,
session_id: &SessionId,
framed: &mut FixFramed,
factory: &mut MessageFactory,
sequences: &SequenceManager,
store: Option<&Arc<dyn MessageStore>>,
write_timeout: Duration,
mut pending: PendingMessage,
) -> Result<(), EngineError> {
application
.to_admin(pending.message_mut(), session_id)
.await;
outbound::check_body(pending.message())?;
send_handshake(
session_id,
framed,
factory,
sequences,
store,
write_timeout,
&pending,
)
.await
}
async fn send_handshake(
session_id: &SessionId,
framed: &mut FixFramed,
factory: &mut MessageFactory,
sequences: &SequenceManager,
store: Option<&Arc<dyn MessageStore>>,
write_timeout: Duration,
pending: &PendingMessage,
) -> Result<(), EngineError> {
let seq = sequences.next_sender_seq().value();
let frame = factory.encode(seq, pending)?;
let allocated = sequences.try_allocate_sender_seq()?.value();
if allocated != seq {
return Err(EngineError::Sequence(format!(
"sender sequence moved from {seq} to {allocated} while a frame was being built"
)));
}
mirror_sender_seq(store, sequences);
persist_outbound(store, session_id, seq, pending.msg_type(), frame).await;
match timeout(write_timeout, framed.send(frame)).await {
Err(_) => Err(EngineError::WriteTimeout(write_timeout)),
Ok(Err(err)) => Err(err.into()),
Ok(Ok(())) => Ok(()),
}
}
fn mirror_sender_seq(store: Option<&Arc<dyn MessageStore>>, sequences: &SequenceManager) {
if let Some(store) = store {
store.set_next_sender_seq(sequences.next_sender_seq().value());
}
}
async fn persist_outbound(
store: Option<&Arc<dyn MessageStore>>,
session_id: &SessionId,
seq: u64,
msg_type: &MsgType,
frame: &[u8],
) {
let Some(store) = store else {
return;
};
if let Err(err) = store.store(seq, msg_type, frame).await {
tracing::warn!(
session = %session_id,
seq,
msg_type = msg_type.as_str(),
error = %err,
"cannot store outbound message: a resend of it will be gap-filled"
);
}
}
pub(crate) fn lock_heartbeat(
runtime: &SessionRuntime,
) -> std::sync::MutexGuard<'_, HeartbeatManager> {
runtime
.heartbeat
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
}
enum Phase {
Active(Session<Active>),
LogoutPending(Session<LogoutPending>),
}
struct SessionClosed {
reason: String,
graceful: bool,
}
fn closed(reason: impl Into<String>, graceful: bool) -> SessionClosed {
SessionClosed {
reason: reason.into(),
graceful,
}
}
fn teardown(phase: Phase) {
match phase {
Phase::Active(session) => {
let _ = session.disconnect();
}
Phase::LogoutPending(session) => {
let _ = session.on_timeout();
}
}
}
fn exhausted(phase: Phase, err: SequenceExhausted) -> SessionClosed {
teardown(phase);
closed(err.to_string(), false)
}
#[derive(Debug)]
struct AppFrame {
frame: Bytes,
seq: u64,
}
#[derive(Debug)]
struct AppRejection {
ref_seq: u64,
ref_msg_type: MsgType,
reason: RejectReason,
}
async fn run_app_dispatcher<A: Application>(
application: Arc<A>,
session_id: SessionId,
mut frames: mpsc::Receiver<AppFrame>,
rejections: mpsc::Sender<AppRejection>,
) {
while let Some(AppFrame { frame, seq }) = frames.recv().await {
let Ok(raw) = wire::decode_frame(&frame) else {
tracing::warn!(
session = %session_id,
seq,
"dropping an application message that no longer decodes"
);
continue;
};
if let Err(reason) = application.from_app(&raw, &session_id).await {
let rejection = AppRejection {
ref_seq: seq,
ref_msg_type: raw.msg_type().clone(),
reason,
};
if rejections.send(rejection).await.is_err() {
break;
}
}
}
}
struct ReactorChannels {
commands: mpsc::Receiver<Command>,
app_rejects: mpsc::Receiver<AppRejection>,
closed: watch::Sender<bool>,
dispatcher: JoinHandle<()>,
}
#[derive(Debug, Clone, Copy)]
pub(crate) struct ResendState {
expected: u64,
high_water: u64,
requested_at: Instant,
attempts: u32,
limit: u32,
timeout: Duration,
}
impl ResendState {
#[must_use]
pub(crate) fn first(expected: u64, high_water: u64, config: &SessionConfig) -> Self {
Self {
expected,
high_water,
requested_at: Instant::now(),
attempts: 1,
limit: config.resend_attempt_limit(),
timeout: config.resend_timeout,
}
}
#[must_use]
fn is_stalled(&self) -> bool {
self.requested_at.elapsed() >= self.timeout
}
#[must_use]
const fn can_retry(&self) -> bool {
self.attempts < self.limit
}
fn record_retry(&mut self) {
self.requested_at = Instant::now();
self.attempts = self.attempts.saturating_add(1);
}
fn record_progress(&mut self, expected: u64) {
self.expected = expected;
self.requested_at = Instant::now();
self.attempts = 1;
}
}
struct Reactor<A: Application> {
factory: MessageFactory,
identity: PeerIdentity,
sending_time: SendingTimeGuard,
runtime: Arc<SessionRuntime>,
config: SessionConfig,
application: Arc<A>,
session_id: SessionId,
store: Option<Arc<dyn MessageStore>>,
resend: Option<ResendState>,
write_timeout: Duration,
app_tx: mpsc::Sender<AppFrame>,
}
async fn run_reactor<A: Application + 'static>(
mut framed: FixFramed,
channels: ReactorChannels,
mut ctx: Reactor<A>,
session: Session<Active>,
) {
let ReactorChannels {
mut commands,
mut app_rejects,
closed: closed_tx,
mut dispatcher,
} = channels;
let mut phase = Some(Phase::Active(session));
let mut commands_open = true;
let mut rejects_open = true;
let mut tick = interval(TICK_INTERVAL);
tick.set_missed_tick_behavior(MissedTickBehavior::Delay);
let outcome = loop {
let Some(current) = phase.take() else {
break closed("internal error: session phase lost", false);
};
let result = tokio::select! {
inbound = framed.next() => match inbound {
Some(Ok(frame)) => ctx.on_frame(&mut framed, current, frame).await,
Some(Err(err)) => {
teardown(current);
Err(closed(format!("codec error: {err}"), false))
}
None => {
teardown(current);
Err(closed("transport closed by peer", false))
}
},
command = commands.recv(), if commands_open => match command {
Some(command) => ctx.on_command(&mut framed, current, command).await,
None => {
commands_open = false;
ctx.on_command(&mut framed, current, Command::Logout).await
}
},
rejection = app_rejects.recv(), if rejects_open => match rejection {
Some(rejection) => {
let ref_msg_type = rejection.ref_msg_type.as_str().to_string();
ctx.send_session_reject(
&mut framed,
current,
rejection.ref_seq,
&ref_msg_type,
&rejection.reason,
)
.await
}
None => {
rejects_open = false;
Ok(current)
}
},
_ = tick.tick() => ctx.on_tick(&mut framed, current).await,
};
match result {
Ok(next) => phase = Some(next),
Err(outcome) => break outcome,
}
ctx.sync_sequences();
};
if ctx.config.reset_on_disconnect || (outcome.graceful && ctx.config.reset_on_logout) {
ctx.runtime.sequences.reset();
if let Some(store) = &ctx.store
&& let Err(err) = store.reset().await
{
tracing::warn!(
session = %ctx.session_id,
error = %err,
"cannot clear the store after a session-closing sequence reset"
);
}
}
ctx.sync_sequences();
drop(ctx.app_tx);
if timeout(APP_DRAIN_TIMEOUT, &mut dispatcher).await.is_err() {
dispatcher.abort();
tracing::warn!(
session = %ctx.session_id,
"application dispatcher did not drain before the session closed; aborting it"
);
}
ctx.application.on_logout(&ctx.session_id).await;
let _ = closed_tx.send(true);
if outcome.graceful {
tracing::info!(session = %ctx.session_id, reason = %outcome.reason, "FIX session closed");
} else {
tracing::warn!(session = %ctx.session_id, reason = %outcome.reason, "FIX session closed");
}
}
enum WriteFailure {
Encode(EncodeError),
Fatal(EngineError),
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Sent {
Yes,
Dropped,
}
impl<A: Application> Reactor<A> {
fn sync_sequences(&self) {
let Some(store) = &self.store else {
return;
};
store.set_next_sender_seq(self.runtime.sequences.next_sender_seq().value());
store.set_next_target_seq(self.runtime.sequences.next_target_seq().value());
}
async fn write_at(
&mut self,
framed: &mut FixFramed,
seq: u64,
pending: &PendingMessage,
persist: bool,
) -> Result<(), WriteFailure> {
let write_timeout = self.write_timeout;
let frame = match self.factory.encode(seq, pending) {
Ok(frame) => frame,
Err(err) => return Err(WriteFailure::Encode(err)),
};
if persist {
persist_outbound(
self.store.as_ref(),
&self.session_id,
seq,
pending.msg_type(),
frame,
)
.await;
}
match timeout(write_timeout, framed.send(frame)).await {
Err(_) => Err(WriteFailure::Fatal(EngineError::WriteTimeout(
write_timeout,
))),
Ok(Err(err)) => Err(WriteFailure::Fatal(err.into())),
Ok(Ok(())) => Ok(()),
}
}
async fn send(
&mut self,
framed: &mut FixFramed,
mut pending: PendingMessage,
) -> Result<Sent, EngineError> {
if !self.prepare(&mut pending).await {
return Ok(Sent::Dropped);
}
let seq = self.runtime.sequences.next_sender_seq().value();
match self.write_at(framed, seq, &pending, true).await {
Ok(()) => {}
Err(WriteFailure::Encode(err)) => {
tracing::warn!(
session = %self.session_id,
error = %err,
"dropping outbound message with no legal wire form; no sequence number spent"
);
return Ok(Sent::Dropped);
}
Err(WriteFailure::Fatal(err)) => return Err(err),
}
let allocated = self.runtime.sequences.try_allocate_sender_seq()?.value();
if allocated != seq {
return Err(EngineError::Sequence(format!(
"sender sequence moved from {seq} to {allocated} while a frame was being built"
)));
}
mirror_sender_seq(self.store.as_ref(), &self.runtime.sequences);
lock_heartbeat(&self.runtime).on_message_sent();
Ok(Sent::Yes)
}
async fn prepare(&mut self, pending: &mut PendingMessage) -> bool {
if pending.msg_type().is_admin() {
self.application
.to_admin(pending.message_mut(), &self.session_id)
.await;
} else {
self.application
.to_app(pending.message_mut(), &self.session_id)
.await;
}
if let Err(err) = outbound::check_body(pending.message()) {
tracing::warn!(
session = %self.session_id,
msg_type = %pending.msg_type(),
error = %err,
"dropping outbound message rejected after the application callback"
);
return false;
}
true
}
async fn send_replay(
&self,
framed: &mut FixFramed,
frame: BytesMut,
) -> Result<(), EngineError> {
let write_timeout = self.write_timeout;
match timeout(write_timeout, framed.send(frame)).await {
Err(_) => return Err(EngineError::WriteTimeout(write_timeout)),
Ok(Err(err)) => return Err(err.into()),
Ok(Ok(())) => {}
}
lock_heartbeat(&self.runtime).on_message_sent();
Ok(())
}
async fn send_gap_fill(
&mut self,
framed: &mut FixFramed,
at_seq: u64,
new_seq: u64,
) -> Result<(), EngineError> {
let mut gap_fill = self.factory.sequence_reset_gap_fill(new_seq);
if !self.prepare(&mut gap_fill).await {
return Ok(());
}
match self.write_at(framed, at_seq, &gap_fill, false).await {
Ok(()) => {
lock_heartbeat(&self.runtime).on_message_sent();
Ok(())
}
Err(WriteFailure::Encode(err)) => {
tracing::warn!(
session = %self.session_id,
error = %err,
"dropping SequenceReset-GapFill with no legal wire form"
);
Ok(())
}
Err(WriteFailure::Fatal(err)) => Err(err),
}
}
async fn on_command(
&mut self,
framed: &mut FixFramed,
phase: Phase,
command: Command,
) -> Result<Phase, SessionClosed> {
match command {
Command::Send(message) => {
if matches!(phase, Phase::LogoutPending(_)) {
tracing::warn!(
session = %self.session_id,
msg_type = %message.msg_type(),
"dropping outbound message: logout pending"
);
return Ok(phase);
}
if let Err(err) = self
.send(framed, PendingMessage::application(message))
.await
{
teardown(phase);
return Err(closed(format!("send failed: {err}"), false));
}
Ok(phase)
}
Command::Logout => match phase {
Phase::LogoutPending(_) => Ok(phase),
Phase::Active(session) => self.begin_logout(framed, session, None).await,
},
}
}
async fn begin_logout(
&mut self,
framed: &mut FixFramed,
session: Session<Active>,
text: Option<&str>,
) -> Result<Phase, SessionClosed> {
let logout = self.factory.logout(text);
match self.send(framed, logout).await {
Err(err) => {
let _ = session.disconnect();
Err(closed(format!("send failed: {err}"), false))
}
Ok(Sent::Dropped) => Ok(Phase::Active(session)),
Ok(Sent::Yes) => {
self.resend = None;
Ok(Phase::LogoutPending(session.initiate_logout()))
}
}
}
async fn on_tick(
&mut self,
framed: &mut FixFramed,
phase: Phase,
) -> Result<Phase, SessionClosed> {
if let Phase::LogoutPending(session) = &phase
&& session.sent_at().elapsed() >= self.config.logout_timeout
{
teardown(phase);
return Err(closed("logout ack timeout", true));
}
if let Some(resend) = self.resend
&& resend.is_stalled()
{
return self.on_resend_stalled(framed, phase, resend).await;
}
let (timed_out, send_test_request, send_heartbeat) = {
let heartbeat = lock_heartbeat(&self.runtime);
(
heartbeat.is_timed_out(),
heartbeat.should_send_test_request(),
heartbeat.should_send_heartbeat(),
)
};
if timed_out {
teardown(phase);
return Err(closed(
"heartbeat timeout: no response to TestRequest",
false,
));
}
if send_test_request {
let test_req_id = generate_test_req_id();
let request = self.factory.test_request(&test_req_id);
match self.send(framed, request).await {
Err(err) => {
teardown(phase);
return Err(closed(format!("send failed: {err}"), false));
}
Ok(Sent::Yes) => lock_heartbeat(&self.runtime).on_test_request_sent(test_req_id),
Ok(Sent::Dropped) => {}
}
} else if send_heartbeat {
let heartbeat = self.factory.heartbeat(None);
if let Err(err) = self.send(framed, heartbeat).await {
teardown(phase);
return Err(closed(format!("send failed: {err}"), false));
}
}
Ok(phase)
}
async fn send_session_reject(
&mut self,
framed: &mut FixFramed,
phase: Phase,
ref_seq: u64,
ref_msg_type: &str,
reason: &RejectReason,
) -> Result<Phase, SessionClosed> {
let reject = self.factory.session_reject(ref_seq, ref_msg_type, reason);
if let Err(err) = self.send(framed, reject).await {
teardown(phase);
return Err(closed(format!("send failed: {err}"), false));
}
Ok(phase)
}
async fn request_resend(
&mut self,
framed: &mut FixFramed,
phase: Phase,
expected: u64,
high_water: u64,
) -> Result<Phase, SessionClosed> {
if let Some(state) = self.resend.as_mut() {
if high_water > state.high_water {
state.high_water = high_water;
}
if state.expected == expected {
return Ok(phase);
}
}
let request = self.factory.resend_request(expected, 0);
match self.send(framed, request).await {
Err(err) => {
teardown(phase);
return Err(closed(format!("send failed: {err}"), false));
}
Ok(Sent::Dropped) => return Ok(phase),
Ok(Sent::Yes) => {}
}
let high_water = self
.resend
.map_or(high_water, |state| state.high_water.max(high_water));
self.resend = Some(ResendState::first(expected, high_water, &self.config));
Ok(phase)
}
fn note_inbound_progress(&mut self) {
let Some(high_water) = self.resend.as_ref().map(|state| state.high_water) else {
return;
};
let next_target = self.runtime.sequences.next_target_seq().value();
if next_target > high_water {
self.resend = None;
} else if let Some(state) = self.resend.as_mut() {
state.record_progress(next_target);
}
}
async fn on_resend_stalled(
&mut self,
framed: &mut FixFramed,
phase: Phase,
resend: ResendState,
) -> Result<Phase, SessionClosed> {
if !resend.can_retry() {
let reason = format!(
"resend of MsgSeqNum {} unanswered after {} request(s)",
resend.expected, resend.attempts
);
tracing::warn!(session = %self.session_id, reason = %reason, "abandoning stalled resend");
return match phase {
Phase::LogoutPending(_) => Ok(phase),
Phase::Active(session) => self.begin_logout(framed, session, Some(&reason)).await,
};
}
let expected = resend.expected;
let request = self.factory.resend_request(expected, 0);
match self.send(framed, request).await {
Err(err) => {
teardown(phase);
return Err(closed(format!("send failed: {err}"), false));
}
Ok(Sent::Dropped) => return Ok(phase),
Ok(Sent::Yes) => {}
}
if let Some(state) = self.resend.as_mut() {
state.record_retry();
}
tracing::debug!(
session = %self.session_id,
expected,
"retrying unanswered ResendRequest"
);
Ok(phase)
}
async fn close_on_too_low(
&mut self,
framed: &mut FixFramed,
phase: Phase,
expected: u64,
received: u64,
) -> SessionClosed {
let reason = format!("MsgSeqNum too low: expected {expected}, received {received}");
let logout = self.factory.logout(Some(&reason));
let _ = self.send(framed, logout).await;
teardown(phase);
closed(reason, false)
}
async fn close_on_identity_mismatch(
&mut self,
framed: &mut FixFramed,
phase: Phase,
ref_seq: u64,
ref_msg_type: &str,
mismatch: wire::IdentityMismatch,
) -> SessionClosed {
let detail = mismatch.to_string();
tracing::warn!(session = %self.session_id, detail = %detail, "inbound identity mismatch");
let reason = RejectReason::new(9, detail.clone()).with_ref_tag(mismatch.tag);
let reject = self.factory.session_reject(ref_seq, ref_msg_type, &reason);
let _ = self.send(framed, reject).await;
let logout = self.factory.logout(Some(&detail));
let _ = self.send(framed, logout).await;
teardown(phase);
closed(detail, false)
}
async fn close_on_sending_time(
&mut self,
framed: &mut FixFramed,
phase: Phase,
ref_seq: u64,
ref_msg_type: &str,
problem: SendingTimeProblem,
) -> SessionClosed {
let detail = problem.to_string();
tracing::warn!(session = %self.session_id, detail = %detail, "inbound SendingTime problem");
let reason = problem.reject_reason();
let reject = self.factory.session_reject(ref_seq, ref_msg_type, &reason);
let _ = self.send(framed, reject).await;
let logout = self.factory.logout(Some(&detail));
let _ = self.send(framed, logout).await;
teardown(phase);
closed(detail, false)
}
fn consume_if_in_sequence(&self, gap_fill: bool, seq: u64) -> Result<bool, SequenceExhausted> {
if !gap_fill || !self.runtime.sequences.validate_incoming(seq).is_ok() {
return Ok(false);
}
self.runtime
.sequences
.try_increment_target_seq()
.map(|_| true)
}
fn consume_rejected_fill(&mut self, gap_fill: bool, seq: u64) -> Result<(), SequenceExhausted> {
if self.consume_if_in_sequence(gap_fill, seq)? {
self.note_inbound_progress();
}
Ok(())
}
async fn on_sequence_reset(
&mut self,
framed: &mut FixFramed,
phase: Phase,
raw: &RawMessage<'_>,
seq: u64,
) -> Result<Phase, SessionClosed> {
let gap_fill = match raw.get_field_str(123) {
Some("Y") => true,
Some("N") | None => false,
Some(other) => {
let reason = RejectReason::new(
6,
format!("GapFillFlag (123) must be Y or N, got '{other}'"),
)
.with_ref_tag(123);
return self
.send_session_reject(
framed,
phase,
seq,
MsgType::SequenceReset.as_str(),
&reason,
)
.await;
}
};
if gap_fill {
match self.runtime.sequences.validate_incoming(seq) {
SequenceResult::Gap { expected, received } => {
return self.request_resend(framed, phase, expected, received).await;
}
SequenceResult::TooLow { expected, received } => {
if raw.get_field_str(43) == Some("Y") {
return Ok(phase);
}
return Err(self
.close_on_too_low(framed, phase, expected, received)
.await);
}
SequenceResult::Ok => {}
}
}
if let Err(reason) = self.application.from_admin(raw, &self.session_id).await {
if let Err(err) = self.consume_rejected_fill(gap_fill, seq) {
return Err(exhausted(phase, err));
}
return self
.send_session_reject(framed, phase, seq, MsgType::SequenceReset.as_str(), &reason)
.await;
}
let Some(new_seq) = raw.get_field_str(36).and_then(|s| s.parse::<u64>().ok()) else {
let reason = RejectReason::new(1, "SequenceReset without a valid NewSeqNo (36)")
.with_ref_tag(36);
if let Err(err) = self.consume_rejected_fill(gap_fill, seq) {
return Err(exhausted(phase, err));
}
return self
.send_session_reject(framed, phase, seq, MsgType::SequenceReset.as_str(), &reason)
.await;
};
if gap_fill {
if new_seq <= seq {
if let Err(err) = self.runtime.sequences.try_increment_target_seq() {
return Err(exhausted(phase, err));
}
self.note_inbound_progress();
let reason = RejectReason::new(
5,
format!("GapFill NewSeqNo {new_seq} does not advance past MsgSeqNum {seq}"),
)
.with_ref_tag(36);
return self
.send_session_reject(
framed,
phase,
seq,
MsgType::SequenceReset.as_str(),
&reason,
)
.await;
}
} else {
let expected = self.runtime.sequences.next_target_seq().value();
if new_seq < expected {
let reason = RejectReason::new(
5,
format!("SequenceReset NewSeqNo {new_seq} is below the expected {expected}"),
)
.with_ref_tag(36);
return self
.send_session_reject(
framed,
phase,
seq,
MsgType::SequenceReset.as_str(),
&reason,
)
.await;
}
}
let expected_before = self.runtime.sequences.next_target_seq().value();
self.runtime.sequences.set_target_seq(new_seq);
if new_seq > expected_before {
self.note_inbound_progress();
}
Ok(phase)
}
async fn on_resend_request(
&mut self,
framed: &mut FixFramed,
phase: Phase,
raw: &RawMessage<'_>,
seq: u64,
) -> Result<Phase, SessionClosed> {
let ref_msg_type = MsgType::ResendRequest.as_str();
let Some(begin_seq) = raw.get_field_str(7).and_then(|s| s.parse::<u64>().ok()) else {
let reason = RejectReason::new(1, "ResendRequest without a valid BeginSeqNo (7)")
.with_ref_tag(7);
return self
.send_session_reject(framed, phase, seq, ref_msg_type, &reason)
.await;
};
let Some(end_seq) = raw.get_field_str(16).and_then(|s| s.parse::<u64>().ok()) else {
let reason = RejectReason::new(1, "ResendRequest without a valid EndSeqNo (16)")
.with_ref_tag(16);
return self
.send_session_reject(framed, phase, seq, ref_msg_type, &reason)
.await;
};
let next_sender = self.runtime.sequences.next_sender_seq().value();
if begin_seq == 0 || begin_seq >= next_sender {
let reason = RejectReason::new(
5,
format!("ResendRequest BeginSeqNo {begin_seq} is outside the sent range 1..{next_sender}"),
)
.with_ref_tag(7);
return self
.send_session_reject(framed, phase, seq, ref_msg_type, &reason)
.await;
}
if end_seq != 0 && end_seq < begin_seq {
let reason = RejectReason::new(
5,
format!("ResendRequest EndSeqNo {end_seq} is below BeginSeqNo {begin_seq}"),
)
.with_ref_tag(16);
return self
.send_session_reject(framed, phase, seq, ref_msg_type, &reason)
.await;
}
let new_seq = match end_seq {
0 => next_sender,
_ => end_seq
.checked_add(1)
.map_or(next_sender, |bound| bound.min(next_sender)),
};
let cursor = match self.replay_range(framed, begin_seq, new_seq).await {
Ok(cursor) => cursor,
Err(err) => {
teardown(phase);
return Err(closed(format!("resend failed: {err}"), false));
}
};
if cursor < new_seq
&& let Err(err) = self.send_gap_fill(framed, cursor, new_seq).await
{
teardown(phase);
return Err(closed(format!("send failed: {err}"), false));
}
Ok(phase)
}
async fn replay_range(
&mut self,
framed: &mut FixFramed,
begin_seq: u64,
new_seq: u64,
) -> Result<u64, EngineError> {
let Some(store) = self.store.clone() else {
return Ok(begin_seq);
};
let Some(end_seq) = new_seq.checked_sub(1) else {
return Ok(begin_seq);
};
if end_seq < begin_seq {
return Ok(begin_seq);
}
let mut reply_cursor = begin_seq;
let mut read_from = begin_seq;
loop {
let page = match store.get_page(read_from, end_seq, RESEND_PAGE_LIMIT).await {
Ok(page) => page,
Err(err) => {
tracing::warn!(
session = %self.session_id,
read_from,
end_seq,
error = %err,
"cannot read the resend range from the store: gap-filling the rest"
);
return Ok(reply_cursor);
}
};
let page_len = page.len();
let Some(furthest) = page.last().map(|message| message.seq_num()) else {
break;
};
for message in page {
let seq = message.seq_num();
if seq < reply_cursor || seq > end_seq {
continue;
}
if message.msg_type().is_admin() {
continue;
}
let frame = match wire::resend_frame(message.payload()) {
Ok(frame) => frame,
Err(err) => {
tracing::warn!(
session = %self.session_id,
seq,
error = %err,
"cannot rebuild a stored message as a resend: gap-filling it instead"
);
continue;
}
};
if seq > reply_cursor {
self.send_gap_fill(framed, reply_cursor, seq).await?;
}
self.send_replay(framed, frame).await?;
let Some(next) = seq.checked_add(1) else {
return Ok(new_seq);
};
reply_cursor = next;
}
if page_len < RESEND_PAGE_LIMIT {
break;
}
let Some(next_read) = furthest.checked_add(1) else {
break;
};
read_from = next_read;
if read_from > end_seq {
break;
}
tokio::task::yield_now().await;
}
Ok(reply_cursor)
}
async fn on_frame(
&mut self,
framed: &mut FixFramed,
phase: Phase,
frame: BytesMut,
) -> Result<Phase, SessionClosed> {
let frame = frame.freeze();
let raw = match wire::decode_frame(&frame) {
Ok(raw) => raw,
Err(err) => {
tracing::warn!(session = %self.session_id, error = %err, "dropping undecodable frame");
return Ok(phase);
}
};
let msg_type = raw.msg_type().clone();
let Some(seq) = wire::header_seq_num(&raw) else {
tracing::warn!(
session = %self.session_id,
msg_type = %msg_type,
"dropping message without a valid header MsgSeqNum (34)"
);
return Ok(phase);
};
if let Err(mismatch) = self.identity.validate(&raw) {
return Err(self
.close_on_identity_mismatch(framed, phase, seq, msg_type.as_str(), mismatch)
.await);
}
if let Err(problem) = self.sending_time.validate(&raw) {
return Err(self
.close_on_sending_time(framed, phase, seq, msg_type.as_str(), problem)
.await);
}
let test_req_id = raw.get_field_str(112);
let outcome = lock_heartbeat(&self.runtime)
.on_message_received(msg_type == MsgType::Heartbeat, test_req_id);
match outcome {
TestRequestOutcome::Confirmed => tracing::debug!(
session = %self.session_id,
"TestRequest answered by Heartbeat with matching TestReqID"
),
TestRequestOutcome::SupersededByTraffic => tracing::debug!(
session = %self.session_id,
msg_type = %msg_type,
"pending TestRequest cleared by inbound traffic"
),
TestRequestOutcome::NonePending => {}
}
if msg_type == MsgType::SequenceReset {
return self.on_sequence_reset(framed, phase, &raw, seq).await;
}
match self.runtime.sequences.validate_incoming(seq) {
SequenceResult::Ok => {
if msg_type.is_admin() {
if let Err(err) = self.runtime.sequences.try_increment_target_seq() {
return Err(exhausted(phase, err));
}
self.note_inbound_progress();
self.dispatch_admin(framed, phase, &raw, seq).await
} else {
let phase = self.enqueue_app(phase, seq, &frame)?;
if let Err(err) = self.runtime.sequences.try_increment_target_seq() {
return Err(exhausted(phase, err));
}
self.note_inbound_progress();
Ok(phase)
}
}
SequenceResult::TooLow { expected, received } => {
if raw.get_field_str(43) == Some("Y") {
return Ok(phase);
}
Err(self
.close_on_too_low(framed, phase, expected, received)
.await)
}
SequenceResult::Gap { expected, received } => {
let phase = self
.request_resend(framed, phase, expected, received)
.await?;
if msg_type.is_app() {
return Ok(phase);
}
self.dispatch_admin(framed, phase, &raw, seq).await
}
}
}
fn enqueue_app(&self, phase: Phase, seq: u64, frame: &Bytes) -> Result<Phase, SessionClosed> {
let queued = AppFrame {
frame: frame.clone(),
seq,
};
match self.app_tx.try_send(queued) {
Ok(()) => Ok(phase),
Err(mpsc::error::TrySendError::Full(_)) => {
teardown(phase);
Err(closed(
"application queue full: the application is not consuming inbound \
messages fast enough",
false,
))
}
Err(mpsc::error::TrySendError::Closed(_)) => {
teardown(phase);
Err(closed("application dispatcher stopped", false))
}
}
}
async fn dispatch_admin(
&mut self,
framed: &mut FixFramed,
phase: Phase,
raw: &RawMessage<'_>,
seq: u64,
) -> Result<Phase, SessionClosed> {
let msg_type = raw.msg_type();
if let Err(reason) = self.application.from_admin(raw, &self.session_id).await {
return self
.send_session_reject(framed, phase, seq, msg_type.as_str(), &reason)
.await;
}
match msg_type {
MsgType::Heartbeat => Ok(phase),
MsgType::TestRequest => {
let test_req_id = raw.get_field_str(112).filter(|id| !id.is_empty());
let heartbeat = self.factory.heartbeat(test_req_id);
if let Err(err) = self.send(framed, heartbeat).await {
teardown(phase);
return Err(closed(format!("send failed: {err}"), false));
}
Ok(phase)
}
MsgType::ResendRequest => self.on_resend_request(framed, phase, raw, seq).await,
MsgType::Logon => {
tracing::warn!(
session = %self.session_id,
"ignoring unexpected Logon on established session"
);
Ok(phase)
}
MsgType::Reject => {
tracing::warn!(
session = %self.session_id,
text = raw.get_field_str(58).unwrap_or(""),
"session-level Reject received"
);
Ok(phase)
}
MsgType::Logout => match phase {
Phase::LogoutPending(session) => {
let _ = session.on_logout_ack();
Err(closed("logout complete", true))
}
Phase::Active(session) => {
let logout = self.factory.logout(None);
let _ = self.send(framed, logout).await;
let _ = session.disconnect();
let text = raw
.get_field_str(58)
.unwrap_or("logout initiated by counterparty");
Err(closed(format!("logout by counterparty: {text}"), true))
}
},
_ => Ok(phase),
}
}
}
pub(crate) struct SessionParams<A: Application> {
pub(crate) framed: FixFramed,
pub(crate) session: Session<Active>,
pub(crate) runtime: Arc<SessionRuntime>,
pub(crate) factory: MessageFactory,
pub(crate) identity: PeerIdentity,
pub(crate) sending_time: SendingTimeGuard,
pub(crate) config: SessionConfig,
pub(crate) application: Arc<A>,
pub(crate) session_id: SessionId,
pub(crate) store: Option<Arc<dyn MessageStore>>,
pub(crate) resend: Option<ResendState>,
pub(crate) write_timeout: Duration,
pub(crate) outbound_capacity: usize,
pub(crate) app_queue_capacity: usize,
}
pub(crate) fn spawn_session<A, G>(
params: SessionParams<A>,
guard: G,
) -> (mpsc::Sender<Command>, watch::Receiver<bool>)
where
A: Application + 'static,
G: Send + 'static,
{
let SessionParams {
framed,
session,
runtime,
factory,
identity,
sending_time,
config,
application,
session_id,
store,
resend,
write_timeout,
outbound_capacity,
app_queue_capacity,
} = params;
let (command_tx, command_rx) = mpsc::channel(outbound_capacity);
let (closed_tx, closed_rx) = watch::channel(false);
let (app_tx, app_rx) = mpsc::channel(app_queue_capacity);
let (reject_tx, reject_rx) = mpsc::channel(app_queue_capacity);
let dispatcher = tokio::spawn(run_app_dispatcher(
Arc::clone(&application),
session_id.clone(),
app_rx,
reject_tx,
));
let reactor = Reactor {
factory,
identity,
sending_time,
runtime,
config,
application,
session_id,
store,
resend,
write_timeout,
app_tx,
};
reactor.sync_sequences();
let channels = ReactorChannels {
commands: command_rx,
app_rejects: reject_rx,
closed: closed_tx,
dispatcher,
};
tokio::spawn(async move {
let _guard = guard;
run_reactor(framed, channels, reactor, session).await;
});
(command_tx, closed_rx)
}