1use std::collections::BTreeMap;
7use std::io::{Read, Seek, SeekFrom, Write};
8use std::path::PathBuf;
9
10use async_trait::async_trait;
11use chrono::{DateTime, Utc};
12
13use crate::error::{Error, Result};
14use crate::session_id::SessionId;
15
16#[async_trait]
17pub trait MessageStore: Send {
18 fn next_sender_seq_num(&self) -> u64;
19 fn next_target_seq_num(&self) -> u64;
20 async fn incr_next_sender_seq_num(&mut self) -> Result<()>;
21 async fn incr_next_target_seq_num(&mut self) -> Result<()>;
22 async fn set_next_sender_seq_num(&mut self, n: u64) -> Result<()>;
23 async fn set_next_target_seq_num(&mut self, n: u64) -> Result<()>;
24 fn creation_time(&self) -> DateTime<Utc>;
25 async fn save_message(&mut self, seq_num: u64, raw: &[u8]) -> Result<()>;
27 async fn save_message_and_incr(&mut self, seq_num: u64, raw: &[u8]) -> Result<()> {
29 self.save_message(seq_num, raw).await?;
30 self.incr_next_sender_seq_num().await
31 }
32 async fn get_messages(&mut self, begin: u64, end: u64) -> Result<Vec<(u64, Vec<u8>)>>;
35 async fn refresh(&mut self) -> Result<()>;
37 async fn reset(&mut self) -> Result<()>;
39}
40
41pub trait MessageStoreFactory: Send + Sync {
42 fn create(&self, session_id: &SessionId) -> Result<Box<dyn MessageStore>>;
43}
44
45#[derive(Debug)]
48pub struct MemoryStore {
49 next_sender: u64,
50 next_target: u64,
51 creation_time: DateTime<Utc>,
52 messages: BTreeMap<u64, Vec<u8>>,
53 capacity: Option<usize>,
57}
58
59impl Default for MemoryStore {
60 fn default() -> Self {
61 Self {
62 next_sender: 1,
63 next_target: 1,
64 creation_time: Utc::now(),
65 messages: BTreeMap::new(),
66 capacity: None,
67 }
68 }
69}
70
71impl MemoryStore {
72 pub fn new() -> Self {
73 Self::default()
74 }
75
76 pub fn with_capacity(capacity: usize) -> Self {
78 Self { capacity: Some(capacity), ..Self::default() }
79 }
80
81 fn trim(&mut self) {
82 if let Some(cap) = self.capacity {
83 while self.messages.len() > cap {
84 let Some((&oldest, _)) = self.messages.iter().next() else { break };
86 self.messages.remove(&oldest);
87 }
88 }
89 }
90}
91
92#[async_trait]
93impl MessageStore for MemoryStore {
94 fn next_sender_seq_num(&self) -> u64 {
95 self.next_sender
96 }
97 fn next_target_seq_num(&self) -> u64 {
98 self.next_target
99 }
100 async fn incr_next_sender_seq_num(&mut self) -> Result<()> {
101 self.next_sender += 1;
102 Ok(())
103 }
104 async fn incr_next_target_seq_num(&mut self) -> Result<()> {
105 self.next_target += 1;
106 Ok(())
107 }
108 async fn set_next_sender_seq_num(&mut self, n: u64) -> Result<()> {
109 self.next_sender = n;
110 Ok(())
111 }
112 async fn set_next_target_seq_num(&mut self, n: u64) -> Result<()> {
113 self.next_target = n;
114 Ok(())
115 }
116 fn creation_time(&self) -> DateTime<Utc> {
117 self.creation_time
118 }
119 async fn save_message(&mut self, seq_num: u64, raw: &[u8]) -> Result<()> {
120 self.messages.insert(seq_num, raw.to_vec());
121 self.trim();
122 Ok(())
123 }
124 async fn get_messages(&mut self, begin: u64, end: u64) -> Result<Vec<(u64, Vec<u8>)>> {
125 Ok(self.messages.range(begin..=end).map(|(k, v)| (*k, v.clone())).collect())
126 }
127 async fn refresh(&mut self) -> Result<()> {
128 Ok(())
129 }
130 async fn reset(&mut self) -> Result<()> {
131 self.next_sender = 1;
132 self.next_target = 1;
133 self.creation_time = Utc::now();
134 self.messages.clear();
135 Ok(())
136 }
137}
138
139#[derive(Debug, Default)]
140pub struct MemoryStoreFactory {
141 capacity: Option<usize>,
142}
143
144impl MemoryStoreFactory {
145 pub fn new() -> Self {
146 Self::default()
147 }
148
149 pub fn with_capacity(capacity: usize) -> Self {
151 Self { capacity: Some(capacity) }
152 }
153}
154
155impl MessageStoreFactory for MemoryStoreFactory {
156 fn create(&self, _session_id: &SessionId) -> Result<Box<dyn MessageStore>> {
157 Ok(Box::new(match self.capacity {
158 Some(cap) => MemoryStore::with_capacity(cap),
159 None => MemoryStore::new(),
160 }))
161 }
162}
163
164pub struct FileStore {
179 prefix: PathBuf,
180 cache: MemoryStore,
181 body: std::fs::File,
182 header: std::fs::File,
183 seqnums: std::fs::File,
184 body_len: u64,
185 offsets: BTreeMap<u64, (u64, u64)>,
186 sync: bool,
187}
188
189impl FileStore {
190 pub fn open(dir: &std::path::Path, session_id: &SessionId) -> Result<Self> {
191 Self::open_with_sync(dir, session_id, false)
192 }
193
194 pub fn open_with_sync(
195 dir: &std::path::Path,
196 session_id: &SessionId,
197 sync: bool,
198 ) -> Result<Self> {
199 std::fs::create_dir_all(dir)?;
200 let prefix = dir.join(session_id.file_prefix());
201 let open = |ext: &str| {
202 std::fs::OpenOptions::new()
203 .read(true)
204 .write(true)
205 .create(true)
206 .truncate(false)
207 .open(prefix.with_extension(ext))
208 };
209 let mut store = Self {
210 cache: MemoryStore::new(),
211 body: open("body")?,
212 header: open("header")?,
213 seqnums: open("seqnums")?,
214 body_len: 0,
215 offsets: BTreeMap::new(),
216 sync,
217 prefix,
218 };
219 store.load()?;
220 Ok(store)
221 }
222
223 async fn fsync(&self, files: &[&std::fs::File]) -> Result<()> {
227 if !self.sync {
228 return Ok(());
229 }
230 let clones: Vec<std::fs::File> =
231 files.iter().map(|f| f.try_clone()).collect::<std::io::Result<_>>()?;
232 tokio::task::spawn_blocking(move || -> std::io::Result<()> {
233 for f in &clones {
234 f.sync_data()?;
235 }
236 Ok(())
237 })
238 .await
239 .map_err(|e| Error::Store(format!("fsync task failed: {e}")))??;
240 Ok(())
241 }
242
243 fn load(&mut self) -> Result<()> {
244 let mut s = String::new();
246 self.seqnums.seek(SeekFrom::Start(0))?;
247 self.seqnums.read_to_string(&mut s)?;
248 if let Some((a, b)) = s.trim().split_once(" : ") {
249 let sender: u64 = a.trim().parse().map_err(|_| corrupt("seqnums"))?;
250 let target: u64 = b.trim().parse().map_err(|_| corrupt("seqnums"))?;
251 self.cache.next_sender = sender;
252 self.cache.next_target = target;
253 }
254 let mut h = String::new();
256 self.header.seek(SeekFrom::Start(0))?;
257 self.header.read_to_string(&mut h)?;
258 for line in h.lines() {
259 let mut parts = line.splitn(3, ',');
260 let (Some(seq), Some(off), Some(len)) = (parts.next(), parts.next(), parts.next())
261 else {
262 return Err(corrupt("header"));
263 };
264 let seq: u64 = seq.parse().map_err(|_| corrupt("header"))?;
265 let off: u64 = off.parse().map_err(|_| corrupt("header"))?;
266 let len: u64 = len.parse().map_err(|_| corrupt("header"))?;
267 self.offsets.insert(seq, (off, len));
268 }
269 self.body_len = self.body.seek(SeekFrom::End(0))?;
270 let session_file = self.prefix.with_extension("session");
272 match std::fs::read_to_string(&session_file) {
273 Ok(ts) => {
274 self.cache.creation_time = ts
275 .trim()
276 .parse::<DateTime<Utc>>()
277 .map_err(|_| corrupt("session"))?;
278 }
279 Err(_) => {
280 self.cache.creation_time = Utc::now();
281 std::fs::write(&session_file, self.cache.creation_time.to_rfc3339())?;
282 }
283 }
284 Ok(())
285 }
286
287 fn write_seqnums(&mut self) -> Result<()> {
288 self.seqnums.seek(SeekFrom::Start(0))?;
289 let line = format!("{:020} : {:020}", self.cache.next_sender, self.cache.next_target);
290 self.seqnums.write_all(line.as_bytes())?;
291 self.seqnums.flush()?;
292 Ok(())
293 }
294}
295
296fn corrupt(which: &str) -> Error {
297 Error::Store(format!("corrupt {which} file"))
298}
299
300#[async_trait]
301impl MessageStore for FileStore {
302 fn next_sender_seq_num(&self) -> u64 {
303 self.cache.next_sender
304 }
305 fn next_target_seq_num(&self) -> u64 {
306 self.cache.next_target
307 }
308 async fn incr_next_sender_seq_num(&mut self) -> Result<()> {
309 self.cache.next_sender += 1;
310 self.write_seqnums()?;
311 self.fsync(&[&self.seqnums]).await
312 }
313 async fn incr_next_target_seq_num(&mut self) -> Result<()> {
314 self.cache.next_target += 1;
315 self.write_seqnums()?;
316 self.fsync(&[&self.seqnums]).await
317 }
318 async fn set_next_sender_seq_num(&mut self, n: u64) -> Result<()> {
319 self.cache.next_sender = n;
320 self.write_seqnums()?;
321 self.fsync(&[&self.seqnums]).await
322 }
323 async fn set_next_target_seq_num(&mut self, n: u64) -> Result<()> {
324 self.cache.next_target = n;
325 self.write_seqnums()?;
326 self.fsync(&[&self.seqnums]).await
327 }
328 fn creation_time(&self) -> DateTime<Utc> {
329 self.cache.creation_time
330 }
331 async fn save_message(&mut self, seq_num: u64, raw: &[u8]) -> Result<()> {
332 let offset = self.body_len;
333 self.body.seek(SeekFrom::End(0))?;
334 self.body.write_all(raw)?;
335 self.body_len += raw.len() as u64;
336
337 self.header.seek(SeekFrom::End(0))?;
338 self.header
339 .write_all(format!("{seq_num},{offset},{}\n", raw.len()).as_bytes())?;
340 self.offsets.insert(seq_num, (offset, raw.len() as u64));
341 self.fsync(&[&self.body, &self.header]).await
344 }
345 async fn get_messages(&mut self, begin: u64, end: u64) -> Result<Vec<(u64, Vec<u8>)>> {
346 let ranges: Vec<(u64, u64, u64)> = self
347 .offsets
348 .range(begin..=end)
349 .map(|(seq, (off, len))| (*seq, *off, *len))
350 .collect();
351 let mut out = Vec::with_capacity(ranges.len());
352 for (seq, off, len) in ranges {
353 let mut buf = vec![0u8; len as usize];
354 self.body.seek(SeekFrom::Start(off))?;
355 self.body.read_exact(&mut buf)?;
356 out.push((seq, buf));
357 }
358 Ok(out)
359 }
360 async fn refresh(&mut self) -> Result<()> {
361 self.offsets.clear();
362 self.cache = MemoryStore::new();
363 self.load()
364 }
365 async fn reset(&mut self) -> Result<()> {
366 self.cache.reset().await?;
367 self.offsets.clear();
368 self.body_len = 0;
369 self.body.set_len(0)?;
370 self.header.set_len(0)?;
371 self.seqnums.set_len(0)?;
372 std::fs::write(self.prefix.with_extension("session"), self.cache.creation_time.to_rfc3339())?;
373 self.write_seqnums()
374 }
375}
376
377pub struct FileStoreFactory {
378 pub path: PathBuf,
379 sync: bool,
380}
381
382impl FileStoreFactory {
383 pub fn new(path: impl Into<PathBuf>) -> Self {
386 Self { path: path.into(), sync: false }
387 }
388
389 pub fn with_sync(mut self, sync: bool) -> Self {
393 self.sync = sync;
394 self
395 }
396}
397
398impl MessageStoreFactory for FileStoreFactory {
399 fn create(&self, session_id: &SessionId) -> Result<Box<dyn MessageStore>> {
400 Ok(Box::new(FileStore::open_with_sync(&self.path, session_id, self.sync)?))
401 }
402}
403
404#[cfg(test)]
405mod tests {
406 use super::*;
407
408 #[tokio::test]
409 async fn memory_store_roundtrip() {
410 let mut s = MemoryStore::new();
411 assert_eq!(s.next_sender_seq_num(), 1);
412 s.save_message_and_incr(1, b"one").await.unwrap();
413 s.save_message_and_incr(2, b"two").await.unwrap();
414 assert_eq!(s.next_sender_seq_num(), 3);
415 let msgs = s.get_messages(1, 10).await.unwrap();
416 assert_eq!(msgs, vec![(1, b"one".to_vec()), (2, b"two".to_vec())]);
417 s.reset().await.unwrap();
418 assert_eq!(s.next_sender_seq_num(), 1);
419 assert!(s.get_messages(1, 10).await.unwrap().is_empty());
420 }
421
422 #[tokio::test]
423 async fn memory_store_capacity_bounds_resend_buffer() {
424 let mut s = MemoryStore::with_capacity(2);
426 for n in 1..=5 {
427 s.save_message_and_incr(n, format!("m{n}").as_bytes()).await.unwrap();
428 }
429 assert_eq!(s.next_sender_seq_num(), 6);
431 let msgs = s.get_messages(1, 10).await.unwrap();
433 assert_eq!(msgs, vec![(4, b"m4".to_vec()), (5, b"m5".to_vec())]);
434 assert!(s.get_messages(1, 3).await.unwrap().is_empty());
436 }
437
438 #[tokio::test]
439 async fn file_store_fsync_roundtrips() {
440 let dir = std::env::temp_dir().join(format!("qft-store-sync-{}", std::process::id()));
443 let _ = std::fs::remove_dir_all(&dir);
444 let sid = SessionId::new("FIX.4.4", "S", "T");
445 {
446 let mut s = FileStore::open_with_sync(&dir, &sid, true).unwrap();
447 s.save_message_and_incr(1, b"synced-one").await.unwrap();
448 s.save_message_and_incr(2, b"synced-two").await.unwrap();
449 }
450 {
451 let mut s = FileStore::open_with_sync(&dir, &sid, true).unwrap();
452 assert_eq!(s.next_sender_seq_num(), 3);
453 let msgs = s.get_messages(1, 5).await.unwrap();
454 assert_eq!(msgs, vec![(1, b"synced-one".to_vec()), (2, b"synced-two".to_vec())]);
455 }
456 let _ = std::fs::remove_dir_all(&dir);
457 }
458
459 #[tokio::test]
460 async fn file_store_persists_across_reopen() {
461 let dir = std::env::temp_dir().join(format!("qft-store-{}", std::process::id()));
462 let _ = std::fs::remove_dir_all(&dir);
463 let sid = SessionId::new("FIX.4.2", "SENDER", "TARGET");
464
465 {
466 let mut s = FileStore::open(&dir, &sid).unwrap();
467 s.save_message_and_incr(1, b"8=FIX.4.2\x01...one").await.unwrap();
468 s.save_message_and_incr(2, b"8=FIX.4.2\x01...two").await.unwrap();
469 s.incr_next_target_seq_num().await.unwrap();
470 }
471 {
472 let mut s = FileStore::open(&dir, &sid).unwrap();
473 assert_eq!(s.next_sender_seq_num(), 3);
474 assert_eq!(s.next_target_seq_num(), 2);
475 let msgs = s.get_messages(2, 5).await.unwrap();
476 assert_eq!(msgs.len(), 1);
477 assert_eq!(msgs[0].0, 2);
478 assert_eq!(msgs[0].1, b"8=FIX.4.2\x01...two".to_vec());
479
480 s.reset().await.unwrap();
481 assert_eq!(s.next_sender_seq_num(), 1);
482 assert!(s.get_messages(1, 10).await.unwrap().is_empty());
483 }
484 let _ = std::fs::remove_dir_all(&dir);
485 }
486}