1use std::net::SocketAddr;
4use std::sync::Arc;
5use std::time::Duration;
6
7use tokio::io::{AsyncReadExt, AsyncWriteExt};
8use tokio::net::TcpStream;
9use tokio::task::JoinHandle;
10use tokio::time::timeout;
11
12use truefix_core::{Field, Message, decode, frame_length};
13use truefix_session::{Application, Role, SessionConfig, SessionId};
14use truefix_transport::AcceptorBuilder;
15
16#[derive(Debug, Clone)]
19pub struct ExpectMsg {
20 pub msg_type: String,
22 pub fields: Vec<(u32, String)>,
24 pub fields_absent: Vec<u32>,
27 pub exact: bool,
35}
36
37pub const ALWAYS_ALLOWED_TAGS: &[u32] = &[
42 8, 9, 35, 34, 49, 56, 52, 43, 122, 10, ];
53
54impl ExpectMsg {
55 pub fn of(msg_type: &str) -> Self {
57 Self {
58 msg_type: msg_type.to_owned(),
59 fields: Vec::new(),
60 fields_absent: Vec::new(),
61 exact: false,
62 }
63 }
64
65 #[must_use]
67 pub fn field(mut self, tag: u32, value: &str) -> Self {
68 self.fields.push((tag, value.to_owned()));
69 self
70 }
71
72 #[must_use]
74 pub fn without_field(mut self, tag: u32) -> Self {
75 self.fields_absent.push(tag);
76 self
77 }
78
79 #[must_use]
82 pub fn exact(mut self) -> Self {
83 self.exact = true;
84 self
85 }
86}
87
88#[derive(Debug, Clone)]
90pub enum Step {
91 Send(Message),
93 SendRaw(Vec<u8>),
95 Expect(ExpectMsg),
97 ExpectDisconnect,
99}
100
101#[derive(Debug, Clone, Default)]
103pub struct SessionTweaks {
104 pub enable_next_expected: bool,
106 pub enable_last_processed: bool,
108 pub check_latency: bool,
110 pub resend_chunk_size: u32,
112 pub executor_app: bool,
114 pub reject_garbled: bool,
116 pub validate_fields_out_of_order: bool,
118 pub disconnect_on_error: bool,
121 pub fixed_identity: Option<(String, String)>,
127}
128
129#[derive(Debug, Clone)]
131pub struct Scenario {
132 pub name: String,
134 pub versions: Vec<String>,
136 pub steps: Vec<Step>,
138 pub tweaks: SessionTweaks,
140}
141
142#[derive(Debug, Clone)]
144pub struct ScenarioResult {
145 pub name: String,
147 pub version: String,
149 pub outcome: Result<(), String>,
151}
152
153#[must_use]
157pub fn per_scenario_report(results: &[ScenarioResult]) -> String {
158 results
159 .iter()
160 .map(|r| match &r.outcome {
161 Ok(()) => format!("PASS {} [{}]", r.name, r.version),
162 Err(reason) => format!("FAIL {} [{}]: {reason}", r.name, r.version),
163 })
164 .collect::<Vec<_>>()
165 .join("\n")
166}
167
168pub const FLAT_DICTIONARY_VERSIONS: &[&str] =
187 &["FIX.4.0", "FIX.4.1", "FIX.4.2", "FIX.4.3", "FIX.4.4"];
188
189pub fn dictionary_for_version(version: &str) -> Option<truefix_dict::DataDictionary> {
190 match version {
191 "FIX.4.0" => truefix_dict::load_fix40().ok(),
192 "FIX.4.1" => truefix_dict::load_fix41().ok(),
193 "FIX.4.2" => truefix_dict::load_fix42().ok(),
194 "FIX.4.3" => truefix_dict::load_fix43().ok(),
195 "FIX.4.4" => truefix_dict::load_fix44().ok(),
196 _ => None,
197 }
198}
199
200pub async fn start_acceptor(
203 version: &str,
204 tweaks: &SessionTweaks,
205) -> std::io::Result<(SocketAddr, JoinHandle<()>)> {
206 struct AtApp {
209 monitor: Option<truefix_transport::Monitor>,
210 }
211 #[async_trait::async_trait]
212 impl Application for AtApp {
213 async fn on_logon(&self, _s: &SessionId) {}
214 async fn from_app(
215 &self,
216 message: &Message,
217 id: &SessionId,
218 ) -> Result<(), truefix_core::BusinessReject> {
219 if let Some(monitor) = &self.monitor
220 && message.msg_type() == Some("D")
221 {
222 let clordid = message.body.get(11).and_then(|f| f.as_str().ok());
224 if clordid == Some("LOGOUT") {
225 monitor.force_logout(id).await;
226 } else {
227 monitor.send_app(id, execution_report(message)).await;
228 }
229 }
230 Ok(())
231 }
232 async fn to_app(
239 &self,
240 message: &mut Message,
241 _id: &SessionId,
242 ) -> Result<(), truefix_core::DoNotSend> {
243 let is_veto_sentinel =
244 message.body.get(11).and_then(|f| f.as_str().ok()) == Some("VETO-RESEND");
245 let is_resend = message.header.get(43).and_then(|f| f.as_str().ok()) == Some("Y");
246 if is_veto_sentinel && is_resend {
247 Err(truefix_core::DoNotSend)
248 } else {
249 Ok(())
250 }
251 }
252 }
253
254 let mut template = SessionConfig::new(
255 wire_begin_string(version),
256 "SERVER",
257 "CLIENT",
258 Role::Acceptor,
259 );
260 template.heartbeat_interval = 30;
261 template.check_latency = tweaks.check_latency;
263 template.enable_next_expected_msg_seq_num = tweaks.enable_next_expected;
264 template.enable_last_msg_seq_num_processed = tweaks.enable_last_processed;
265 template.resend_request_chunk_size = tweaks.resend_chunk_size;
266 template.reject_garbled_message = tweaks.reject_garbled;
267 template.disconnect_on_error = tweaks.disconnect_on_error;
268
269 let validation_opts = truefix_dict::ValidationOptions {
271 validate_fields_out_of_order: tweaks.validate_fields_out_of_order,
272 ..truefix_dict::ValidationOptions::default()
273 };
274 let validator = dictionary_for_version(version).map(|dict| (dict, validation_opts));
275 let monitor = tweaks.executor_app.then(truefix_transport::Monitor::new);
276 let services = truefix_transport::Services {
277 validator,
278 monitor: monitor.clone(),
279 ..truefix_transport::Services::default()
280 };
281
282 let acceptor = AcceptorBuilder::bind(
283 "127.0.0.1:0".parse().unwrap_or_else(|_| unreachable_addr()),
284 Arc::new(AtApp { monitor }),
285 )
286 .await?
287 .with_dynamic_template(template)
288 .with_services(services);
289 let addr = acceptor.local_addr()?;
290 let handle = acceptor.serve();
291 Ok((addr, handle))
292}
293
294pub async fn start_fixed_identity_acceptor(
303 version: &str,
304 sender: &str,
305 target: &str,
306 tweaks: &SessionTweaks,
307) -> std::io::Result<(SocketAddr, JoinHandle<()>)> {
308 struct FixedIdentityApp;
309 #[async_trait::async_trait]
310 impl Application for FixedIdentityApp {
311 async fn on_logon(&self, _s: &SessionId) {}
312 }
313
314 let mut config = SessionConfig::new(wire_begin_string(version), sender, target, Role::Acceptor);
315 config.heartbeat_interval = 30;
316 config.check_latency = tweaks.check_latency;
317 config.enable_next_expected_msg_seq_num = tweaks.enable_next_expected;
318 config.enable_last_msg_seq_num_processed = tweaks.enable_last_processed;
319 config.resend_request_chunk_size = tweaks.resend_chunk_size;
320 config.reject_garbled_message = tweaks.reject_garbled;
321 config.disconnect_on_error = tweaks.disconnect_on_error;
322
323 let validation_opts = truefix_dict::ValidationOptions {
324 validate_fields_out_of_order: tweaks.validate_fields_out_of_order,
325 ..truefix_dict::ValidationOptions::default()
326 };
327 let validator = dictionary_for_version(version).map(|dict| (dict, validation_opts));
328 let services = truefix_transport::Services {
329 validator,
330 ..truefix_transport::Services::default()
331 };
332
333 let acceptor = AcceptorBuilder::bind(
334 "127.0.0.1:0".parse().unwrap_or_else(|_| unreachable_addr()),
335 Arc::new(FixedIdentityApp),
336 )
337 .await?
338 .with_session(config)
339 .with_services(services);
340 let addr = acceptor.local_addr()?;
341 let handle = acceptor.serve();
342 Ok((addr, handle))
343}
344
345fn execution_report(order: &Message) -> Message {
348 let mut m = Message::new();
349 m.header.set(Field::string(35, "8"));
350 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] {
355 if let Some(f) = order.body.get(tag) {
356 m.body.set(Field::new(tag, f.value_bytes().to_vec()));
357 }
358 }
359 m
360}
361
362fn unreachable_addr() -> SocketAddr {
363 SocketAddr::from(([127, 0, 0, 1], 0))
364}
365
366pub async fn run_scenario(scenario: &Scenario, addr: SocketAddr) -> Result<(), String> {
368 let mut stream = TcpStream::connect(addr)
369 .await
370 .map_err(|e| format!("connect: {e}"))?;
371 let mut buf: Vec<u8> = Vec::new();
372
373 for (i, step) in scenario.steps.iter().enumerate() {
374 match step {
375 Step::Send(msg) => {
376 stream
377 .write_all(&msg.encode())
378 .await
379 .map_err(|e| format!("step {i}: send: {e}"))?;
380 }
381 Step::SendRaw(bytes) => {
382 stream
383 .write_all(bytes)
384 .await
385 .map_err(|e| format!("step {i}: send raw: {e}"))?;
386 }
387 Step::Expect(expect) => {
388 let msg = match read_message(&mut stream, &mut buf, Duration::from_secs(3)).await {
390 ReadMessageOutcome::Message(msg) => msg,
391 ReadMessageOutcome::TimedOut => {
392 return Err(format!(
393 "step {i}: expected {} but timed out",
394 expect.msg_type
395 ));
396 }
397 ReadMessageOutcome::DecodeFailed(error) => {
398 return Err(format!(
399 "step {i}: expected {} but got an undecodable message: {error}",
400 expect.msg_type
401 ));
402 }
403 ReadMessageOutcome::CleanEof => {
404 return Err(format!(
405 "step {i}: expected {} but the peer disconnected",
406 expect.msg_type
407 ));
408 }
409 ReadMessageOutcome::ReadFailed(error) => {
410 return Err(format!("step {i}: read failed: {error}"));
411 }
412 };
413 check_match(&msg, expect).map_err(|e| format!("step {i}: {e}"))?;
414 }
415 Step::ExpectDisconnect => {
416 match read_message(&mut stream, &mut buf, Duration::from_secs(3)).await {
417 ReadMessageOutcome::CleanEof => {}
418 ReadMessageOutcome::Message(_) => {
419 return Err(format!("step {i}: expected disconnect but got a message"));
420 }
421 ReadMessageOutcome::TimedOut => {
422 return Err(format!("step {i}: expected disconnect but timed out"));
423 }
424 ReadMessageOutcome::DecodeFailed(error) => {
425 return Err(format!(
426 "step {i}: expected disconnect but got an undecodable message: {error}"
427 ));
428 }
429 ReadMessageOutcome::ReadFailed(error) => {
430 return Err(format!(
431 "step {i}: expected disconnect but read failed: {error}"
432 ));
433 }
434 }
435 }
436 }
437 }
438 match read_message(&mut stream, &mut buf, Duration::from_millis(25)).await {
443 ReadMessageOutcome::TimedOut | ReadMessageOutcome::CleanEof => Ok(()),
444 ReadMessageOutcome::Message(msg) => Err(format!(
445 "scenario complete but an extra, unrequested message arrived: {msg:?}"
446 )),
447 ReadMessageOutcome::DecodeFailed(error) => Err(format!(
448 "scenario complete but extra, undecodable bytes arrived: {error}"
449 )),
450 ReadMessageOutcome::ReadFailed(error) => Err(format!(
451 "scenario complete but the trailing read failed: {error}"
452 )),
453 }
454}
455
456pub async fn run_report(scenarios: &[Scenario]) -> Vec<ScenarioResult> {
459 let mut results = Vec::new();
460 for s in scenarios {
463 for version in &s.versions {
464 let started = match &s.tweaks.fixed_identity {
465 Some((sender, target)) => {
466 start_fixed_identity_acceptor(version, sender, target, &s.tweaks).await
467 }
468 None => start_acceptor(version, &s.tweaks).await,
469 };
470 let outcome = match started {
471 Ok((addr, handle)) => {
472 let outcome = run_scenario(s, addr).await;
473 handle.abort();
474 outcome
475 }
476 Err(e) => Err(format!("could not start acceptor: {e}")),
477 };
478 results.push(ScenarioResult {
479 name: s.name.clone(),
480 version: version.clone(),
481 outcome,
482 });
483 }
484 }
485 results
486}
487
488fn check_match(msg: &Message, expect: &ExpectMsg) -> Result<(), String> {
489 if msg.msg_type() != Some(expect.msg_type.as_str()) {
490 return Err(format!(
491 "expected MsgType {:?}, got {:?}",
492 expect.msg_type,
493 msg.msg_type()
494 ));
495 }
496 for (tag, want) in &expect.fields {
497 let got = field_value(msg, *tag);
498 if got.as_deref() != Some(want.as_str()) {
499 return Err(format!("tag {tag}: expected {want:?}, got {got:?}"));
500 }
501 }
502 for tag in &expect.fields_absent {
503 if let Some(got) = field_value(msg, *tag) {
504 return Err(format!("tag {tag}: expected absent, got {got:?}"));
505 }
506 }
507 if expect.exact {
508 let allowed = |tag: u32| {
509 ALWAYS_ALLOWED_TAGS.contains(&tag) || expect.fields.iter().any(|(t, _)| *t == tag)
510 };
511 for field in msg
512 .header
513 .fields()
514 .chain(msg.body.fields())
515 .chain(msg.trailer.fields())
516 {
517 if !allowed(field.tag()) {
518 return Err(format!(
519 "unexpected extra tag {}: {:?}",
520 field.tag(),
521 field.as_str().ok()
522 ));
523 }
524 }
525 }
526 Ok(())
527}
528
529fn field_value(msg: &Message, tag: u32) -> Option<String> {
530 let field = msg
531 .header
532 .get(tag)
533 .or_else(|| msg.body.get(tag))
534 .or_else(|| msg.trailer.get(tag))?;
535 field.as_str().ok().map(str::to_owned)
536}
537
538#[derive(Debug)]
540pub enum ReadMessageOutcome {
541 Message(Message),
543 TimedOut,
545 DecodeFailed(truefix_core::DecodeError),
547 CleanEof,
549 ReadFailed(std::io::Error),
551}
552
553pub async fn read_message(
555 stream: &mut TcpStream,
556 buf: &mut Vec<u8>,
557 wait: Duration,
558) -> ReadMessageOutcome {
559 loop {
560 if let Ok(Some(total)) = frame_length(buf) {
561 let raw: Vec<u8> = buf.drain(..total).collect();
562 return match decode(&raw) {
563 Ok(message) => ReadMessageOutcome::Message(message),
564 Err(error) => ReadMessageOutcome::DecodeFailed(error),
565 };
566 }
567 let mut chunk = [0u8; 4096];
568 match timeout(wait, stream.read(&mut chunk)).await {
569 Ok(Ok(0)) => return ReadMessageOutcome::CleanEof,
570 Ok(Err(error)) => return ReadMessageOutcome::ReadFailed(error),
571 Err(_) => return ReadMessageOutcome::TimedOut,
572 Ok(Ok(n)) => {
573 if let Some(slice) = chunk.get(..n) {
574 buf.extend_from_slice(slice);
575 }
576 }
577 }
578 }
579}
580
581pub fn wire_begin_string(version: &str) -> &str {
589 if version == "FIX.Latest" {
590 "FIX.5.0SP2"
591 } else {
592 version
593 }
594}
595
596pub fn client_message(version: &str, msg_type: &str, seq: i64) -> Message {
598 let mut m = Message::new();
599 m.header.set(Field::string(8, wire_begin_string(version)));
600 m.header.set(Field::string(35, msg_type));
601 m.header.set(Field::int(34, seq));
602 m.header.set(Field::string(49, "CLIENT"));
603 m.header.set(Field::string(56, "SERVER"));
604 m.header.set(Field::string(52, "20240101-00:00:00"));
605 m
606}