use std::net::SocketAddr;
use std::sync::Arc;
use std::time::Duration;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpStream;
use tokio::task::JoinHandle;
use tokio::time::timeout;
use truefix_core::{Field, Message, decode, frame_length};
use truefix_session::{Application, Role, SessionConfig, SessionId};
use truefix_transport::AcceptorBuilder;
#[derive(Debug, Clone)]
pub struct ExpectMsg {
pub msg_type: String,
pub fields: Vec<(u32, String)>,
pub fields_absent: Vec<u32>,
pub exact: bool,
}
pub const ALWAYS_ALLOWED_TAGS: &[u32] = &[
8, 9, 35, 34, 49, 56, 52, 43, 122, 10, ];
impl ExpectMsg {
pub fn of(msg_type: &str) -> Self {
Self {
msg_type: msg_type.to_owned(),
fields: Vec::new(),
fields_absent: Vec::new(),
exact: false,
}
}
#[must_use]
pub fn field(mut self, tag: u32, value: &str) -> Self {
self.fields.push((tag, value.to_owned()));
self
}
#[must_use]
pub fn without_field(mut self, tag: u32) -> Self {
self.fields_absent.push(tag);
self
}
#[must_use]
pub fn exact(mut self) -> Self {
self.exact = true;
self
}
}
#[derive(Debug, Clone)]
pub enum Step {
Send(Message),
SendRaw(Vec<u8>),
Expect(ExpectMsg),
ExpectDisconnect,
}
#[derive(Debug, Clone, Default)]
pub struct SessionTweaks {
pub enable_next_expected: bool,
pub enable_last_processed: bool,
pub check_latency: bool,
pub resend_chunk_size: u32,
pub executor_app: bool,
pub reject_garbled: bool,
pub validate_fields_out_of_order: bool,
pub disconnect_on_error: bool,
pub fixed_identity: Option<(String, String)>,
}
#[derive(Debug, Clone)]
pub struct Scenario {
pub name: String,
pub versions: Vec<String>,
pub steps: Vec<Step>,
pub tweaks: SessionTweaks,
}
#[derive(Debug, Clone)]
pub struct ScenarioResult {
pub name: String,
pub version: String,
pub outcome: Result<(), String>,
}
#[must_use]
pub fn per_scenario_report(results: &[ScenarioResult]) -> String {
results
.iter()
.map(|r| match &r.outcome {
Ok(()) => format!("PASS {} [{}]", r.name, r.version),
Err(reason) => format!("FAIL {} [{}]: {reason}", r.name, r.version),
})
.collect::<Vec<_>>()
.join("\n")
}
pub const FLAT_DICTIONARY_VERSIONS: &[&str] =
&["FIX.4.0", "FIX.4.1", "FIX.4.2", "FIX.4.3", "FIX.4.4"];
pub fn dictionary_for_version(version: &str) -> Option<truefix_dict::DataDictionary> {
match version {
"FIX.4.0" => truefix_dict::load_fix40().ok(),
"FIX.4.1" => truefix_dict::load_fix41().ok(),
"FIX.4.2" => truefix_dict::load_fix42().ok(),
"FIX.4.3" => truefix_dict::load_fix43().ok(),
"FIX.4.4" => truefix_dict::load_fix44().ok(),
_ => None,
}
}
pub async fn start_acceptor(
version: &str,
tweaks: &SessionTweaks,
) -> std::io::Result<(SocketAddr, JoinHandle<()>)> {
struct AtApp {
monitor: Option<truefix_transport::Monitor>,
}
#[async_trait::async_trait]
impl Application for AtApp {
async fn on_logon(&self, _s: &SessionId) {}
async fn from_app(
&self,
message: &Message,
id: &SessionId,
) -> Result<(), truefix_core::BusinessReject> {
if let Some(monitor) = &self.monitor
&& message.msg_type() == Some("D")
{
let clordid = message.body.get(11).and_then(|f| f.as_str().ok());
if clordid == Some("LOGOUT") {
monitor.force_logout(id).await;
} else {
monitor.send_app(id, execution_report(message)).await;
}
}
Ok(())
}
async fn to_app(
&self,
message: &mut Message,
_id: &SessionId,
) -> Result<(), truefix_core::DoNotSend> {
let is_veto_sentinel =
message.body.get(11).and_then(|f| f.as_str().ok()) == Some("VETO-RESEND");
let is_resend = message.header.get(43).and_then(|f| f.as_str().ok()) == Some("Y");
if is_veto_sentinel && is_resend {
Err(truefix_core::DoNotSend)
} else {
Ok(())
}
}
}
let mut template = SessionConfig::new(
wire_begin_string(version),
"SERVER",
"CLIENT",
Role::Acceptor,
);
template.heartbeat_interval = 30;
template.check_latency = tweaks.check_latency;
template.enable_next_expected_msg_seq_num = tweaks.enable_next_expected;
template.enable_last_msg_seq_num_processed = tweaks.enable_last_processed;
template.resend_request_chunk_size = tweaks.resend_chunk_size;
template.reject_garbled_message = tweaks.reject_garbled;
template.disconnect_on_error = tweaks.disconnect_on_error;
let validation_opts = truefix_dict::ValidationOptions {
validate_fields_out_of_order: tweaks.validate_fields_out_of_order,
..truefix_dict::ValidationOptions::default()
};
let validator = dictionary_for_version(version).map(|dict| (dict, validation_opts));
let monitor = tweaks.executor_app.then(truefix_transport::Monitor::new);
let services = truefix_transport::Services {
validator,
monitor: monitor.clone(),
..truefix_transport::Services::default()
};
let acceptor = AcceptorBuilder::bind(
"127.0.0.1:0".parse().unwrap_or_else(|_| unreachable_addr()),
Arc::new(AtApp { monitor }),
)
.await?
.with_dynamic_template(template)
.with_services(services);
let addr = acceptor.local_addr()?;
let handle = acceptor.serve();
Ok((addr, handle))
}
pub async fn start_fixed_identity_acceptor(
version: &str,
sender: &str,
target: &str,
tweaks: &SessionTweaks,
) -> std::io::Result<(SocketAddr, JoinHandle<()>)> {
struct FixedIdentityApp;
#[async_trait::async_trait]
impl Application for FixedIdentityApp {
async fn on_logon(&self, _s: &SessionId) {}
}
let mut config = SessionConfig::new(wire_begin_string(version), sender, target, Role::Acceptor);
config.heartbeat_interval = 30;
config.check_latency = tweaks.check_latency;
config.enable_next_expected_msg_seq_num = tweaks.enable_next_expected;
config.enable_last_msg_seq_num_processed = tweaks.enable_last_processed;
config.resend_request_chunk_size = tweaks.resend_chunk_size;
config.reject_garbled_message = tweaks.reject_garbled;
config.disconnect_on_error = tweaks.disconnect_on_error;
let validation_opts = truefix_dict::ValidationOptions {
validate_fields_out_of_order: tweaks.validate_fields_out_of_order,
..truefix_dict::ValidationOptions::default()
};
let validator = dictionary_for_version(version).map(|dict| (dict, validation_opts));
let services = truefix_transport::Services {
validator,
..truefix_transport::Services::default()
};
let acceptor = AcceptorBuilder::bind(
"127.0.0.1:0".parse().unwrap_or_else(|_| unreachable_addr()),
Arc::new(FixedIdentityApp),
)
.await?
.with_session(config)
.with_services(services);
let addr = acceptor.local_addr()?;
let handle = acceptor.serve();
Ok((addr, handle))
}
fn execution_report(order: &Message) -> Message {
let mut m = Message::new();
m.header.set(Field::string(35, "8"));
m.body.set(Field::string(37, "ORDER-1")); m.body.set(Field::string(17, "EXEC-1")); m.body.set(Field::string(150, "0")); m.body.set(Field::string(39, "0")); for tag in [11u32, 55, 54, 38] {
if let Some(f) = order.body.get(tag) {
m.body.set(Field::new(tag, f.value_bytes().to_vec()));
}
}
m
}
fn unreachable_addr() -> SocketAddr {
SocketAddr::from(([127, 0, 0, 1], 0))
}
pub async fn run_scenario(scenario: &Scenario, addr: SocketAddr) -> Result<(), String> {
let mut stream = TcpStream::connect(addr)
.await
.map_err(|e| format!("connect: {e}"))?;
let mut buf: Vec<u8> = Vec::new();
for (i, step) in scenario.steps.iter().enumerate() {
match step {
Step::Send(msg) => {
stream
.write_all(&msg.encode())
.await
.map_err(|e| format!("step {i}: send: {e}"))?;
}
Step::SendRaw(bytes) => {
stream
.write_all(bytes)
.await
.map_err(|e| format!("step {i}: send raw: {e}"))?;
}
Step::Expect(expect) => {
let msg = match read_message(&mut stream, &mut buf, Duration::from_secs(3)).await {
ReadMessageOutcome::Message(msg) => msg,
ReadMessageOutcome::TimedOut => {
return Err(format!(
"step {i}: expected {} but timed out",
expect.msg_type
));
}
ReadMessageOutcome::DecodeFailed(error) => {
return Err(format!(
"step {i}: expected {} but got an undecodable message: {error}",
expect.msg_type
));
}
ReadMessageOutcome::CleanEof => {
return Err(format!(
"step {i}: expected {} but the peer disconnected",
expect.msg_type
));
}
ReadMessageOutcome::ReadFailed(error) => {
return Err(format!("step {i}: read failed: {error}"));
}
};
check_match(&msg, expect).map_err(|e| format!("step {i}: {e}"))?;
}
Step::ExpectDisconnect => {
match read_message(&mut stream, &mut buf, Duration::from_secs(3)).await {
ReadMessageOutcome::CleanEof => {}
ReadMessageOutcome::Message(_) => {
return Err(format!("step {i}: expected disconnect but got a message"));
}
ReadMessageOutcome::TimedOut => {
return Err(format!("step {i}: expected disconnect but timed out"));
}
ReadMessageOutcome::DecodeFailed(error) => {
return Err(format!(
"step {i}: expected disconnect but got an undecodable message: {error}"
));
}
ReadMessageOutcome::ReadFailed(error) => {
return Err(format!(
"step {i}: expected disconnect but read failed: {error}"
));
}
}
}
}
}
match read_message(&mut stream, &mut buf, Duration::from_millis(25)).await {
ReadMessageOutcome::TimedOut | ReadMessageOutcome::CleanEof => Ok(()),
ReadMessageOutcome::Message(msg) => Err(format!(
"scenario complete but an extra, unrequested message arrived: {msg:?}"
)),
ReadMessageOutcome::DecodeFailed(error) => Err(format!(
"scenario complete but extra, undecodable bytes arrived: {error}"
)),
ReadMessageOutcome::ReadFailed(error) => Err(format!(
"scenario complete but the trailing read failed: {error}"
)),
}
}
pub async fn run_report(scenarios: &[Scenario]) -> Vec<ScenarioResult> {
let mut results = Vec::new();
for s in scenarios {
for version in &s.versions {
let started = match &s.tweaks.fixed_identity {
Some((sender, target)) => {
start_fixed_identity_acceptor(version, sender, target, &s.tweaks).await
}
None => start_acceptor(version, &s.tweaks).await,
};
let outcome = match started {
Ok((addr, handle)) => {
let outcome = run_scenario(s, addr).await;
handle.abort();
outcome
}
Err(e) => Err(format!("could not start acceptor: {e}")),
};
results.push(ScenarioResult {
name: s.name.clone(),
version: version.clone(),
outcome,
});
}
}
results
}
fn check_match(msg: &Message, expect: &ExpectMsg) -> Result<(), String> {
if msg.msg_type() != Some(expect.msg_type.as_str()) {
return Err(format!(
"expected MsgType {:?}, got {:?}",
expect.msg_type,
msg.msg_type()
));
}
for (tag, want) in &expect.fields {
let got = field_value(msg, *tag);
if got.as_deref() != Some(want.as_str()) {
return Err(format!("tag {tag}: expected {want:?}, got {got:?}"));
}
}
for tag in &expect.fields_absent {
if let Some(got) = field_value(msg, *tag) {
return Err(format!("tag {tag}: expected absent, got {got:?}"));
}
}
if expect.exact {
let allowed = |tag: u32| {
ALWAYS_ALLOWED_TAGS.contains(&tag) || expect.fields.iter().any(|(t, _)| *t == tag)
};
for field in msg
.header
.fields()
.chain(msg.body.fields())
.chain(msg.trailer.fields())
{
if !allowed(field.tag()) {
return Err(format!(
"unexpected extra tag {}: {:?}",
field.tag(),
field.as_str().ok()
));
}
}
}
Ok(())
}
fn field_value(msg: &Message, tag: u32) -> Option<String> {
let field = msg
.header
.get(tag)
.or_else(|| msg.body.get(tag))
.or_else(|| msg.trailer.get(tag))?;
field.as_str().ok().map(str::to_owned)
}
#[derive(Debug)]
pub enum ReadMessageOutcome {
Message(Message),
TimedOut,
DecodeFailed(truefix_core::DecodeError),
CleanEof,
ReadFailed(std::io::Error),
}
pub async fn read_message(
stream: &mut TcpStream,
buf: &mut Vec<u8>,
wait: Duration,
) -> ReadMessageOutcome {
loop {
if let Ok(Some(total)) = frame_length(buf) {
let raw: Vec<u8> = buf.drain(..total).collect();
return match decode(&raw) {
Ok(message) => ReadMessageOutcome::Message(message),
Err(error) => ReadMessageOutcome::DecodeFailed(error),
};
}
let mut chunk = [0u8; 4096];
match timeout(wait, stream.read(&mut chunk)).await {
Ok(Ok(0)) => return ReadMessageOutcome::CleanEof,
Ok(Err(error)) => return ReadMessageOutcome::ReadFailed(error),
Err(_) => return ReadMessageOutcome::TimedOut,
Ok(Ok(n)) => {
if let Some(slice) = chunk.get(..n) {
buf.extend_from_slice(slice);
}
}
}
}
}
pub fn wire_begin_string(version: &str) -> &str {
if version == "FIX.Latest" {
"FIX.5.0SP2"
} else {
version
}
}
pub fn client_message(version: &str, msg_type: &str, seq: i64) -> Message {
let mut m = Message::new();
m.header.set(Field::string(8, wire_begin_string(version)));
m.header.set(Field::string(35, msg_type));
m.header.set(Field::int(34, seq));
m.header.set(Field::string(49, "CLIENT"));
m.header.set(Field::string(56, "SERVER"));
m.header.set(Field::string(52, "20240101-00:00:00"));
m
}