1use crate::error::{Error, Result};
5use crate::field_map::{FieldMap, TagValue, write_tag_value};
6use crate::tags;
7use crate::value::{FixEncode, UtcTimestamp};
8
9pub type Tag = i32;
13pub const SOH: u8 = 0x01;
14
15#[derive(Debug, Clone, Default)]
16pub struct Message {
17 pub header: FieldMap,
18 pub body: FieldMap,
19 pub trailer: FieldMap,
20 structure_error: Option<Tag>,
24}
25
26impl Message {
27 pub fn new() -> Self {
28 Self::default()
29 }
30
31 pub fn with_type(msg_type: &str) -> Self {
34 let mut m = Self::default();
35 m.header.set(tags::MSG_TYPE, msg_type);
36 m
37 }
38
39 pub fn msg_type(&self) -> Result<String> {
40 Ok(self.header.get_string(tags::MSG_TYPE)?)
41 }
42
43 pub fn is_admin(&self) -> bool {
44 self.header
45 .get_raw(tags::MSG_TYPE)
46 .map(|t| matches!(t, b"0" | b"1" | b"2" | b"3" | b"4" | b"5" | b"A" | b"n"))
47 .unwrap_or(false)
48 }
49
50 pub fn seq_num(&self) -> Result<u64> {
51 Ok(self.header.get::<u64>(tags::MSG_SEQ_NUM)?)
52 }
53
54 pub fn poss_dup(&self) -> bool {
55 self.header.get_raw(tags::POSS_DUP_FLAG) == Some(b"Y")
56 }
57
58 pub fn structure_error(&self) -> Option<Tag> {
60 self.structure_error
61 }
62
63 pub fn parse(raw: &[u8], validate_length_checksum: bool) -> Result<Self> {
71 let mut msg = Self::default();
72 let mut pos = 0usize;
73 let mut section = 0u8;
75 let mut field_index = 0usize;
76 let mut body_start = None;
77 let mut checksum_field_start = None;
78 let mut pending_data: Option<(Tag, usize)> = None;
79
80 while pos < raw.len() {
81 let field_start = pos;
82 let eq = raw[pos..]
84 .iter()
85 .position(|&b| b == b'=')
86 .ok_or_else(|| Error::Parse("field without '='".into()))?
87 + pos;
88 let tag: Tag = std::str::from_utf8(&raw[pos..eq])
89 .ok()
90 .and_then(|s| s.parse().ok())
91 .ok_or_else(|| {
92 Error::Parse(format!(
93 "invalid tag {:?}",
94 String::from_utf8_lossy(&raw[pos..eq.min(pos + 16)])
95 ))
96 })?;
97
98 let val_start = eq + 1;
100 let val_end = match pending_data.take() {
101 Some((data_tag, len)) if data_tag == tag => {
102 let end = val_start + len;
103 if end > raw.len() || raw.get(end) != Some(&SOH) {
104 return Err(Error::Parse(format!(
105 "data field {tag} shorter than its declared length {len}"
106 )));
107 }
108 end
109 }
110 _ => raw[val_start..]
111 .iter()
112 .position(|&b| b == SOH)
113 .map(|p| val_start + p)
114 .ok_or_else(|| Error::Parse(format!("field {tag} not SOH-terminated")))?,
115 };
116 let value = raw[val_start..val_end].to_vec();
117 pos = val_end + 1;
118
119 if let Some(data_tag) = tags::data_tag_for_length_tag(tag) {
121 if let Ok(len) = std::str::from_utf8(&value)
122 .map_err(|_| ())
123 .and_then(|s| s.parse::<usize>().map_err(|_| ()))
124 {
125 pending_data = Some((data_tag, len));
126 }
127 }
128
129 match field_index {
131 0 if tag != tags::BEGIN_STRING => {
132 return Err(Error::Parse("first field is not BeginString(8)".into()));
133 }
134 1 if tag != tags::BODY_LENGTH => {
135 return Err(Error::Parse("second field is not BodyLength(9)".into()));
136 }
137 2 if tag != tags::MSG_TYPE => {
138 return Err(Error::Parse("third field is not MsgType(35)".into()));
139 }
140 _ => {}
141 }
142 field_index += 1;
143
144 let tv = TagValue { tag, value };
148 if tag == tags::CHECK_SUM {
149 checksum_field_start = Some(field_start);
150 section = 2;
151 msg.trailer.push_tag_value(tv);
152 } else if tags::is_header_tag(tag) {
153 if section != 0 {
154 msg.structure_error.get_or_insert(tag);
155 }
156 msg.header.push_tag_value(tv);
157 } else if tags::is_trailer_tag(tag) {
158 section = 2;
159 msg.trailer.push_tag_value(tv);
160 } else {
161 if section == 0 {
162 section = 1;
163 body_start = Some(field_start);
164 } else if section == 2 {
165 msg.structure_error.get_or_insert(tag);
167 }
168 if body_start.is_none() {
169 body_start = Some(field_start);
170 }
171 msg.body.push_tag_value(tv);
172 }
173 }
174
175 if validate_length_checksum {
176 let checksum_start = checksum_field_start
177 .ok_or_else(|| Error::Parse("message has no CheckSum(10)".into()))?;
178 let declared_len: usize = msg
179 .header
180 .get(tags::BODY_LENGTH)
181 .map_err(|_| Error::Parse("missing/invalid BodyLength(9)".into()))?;
182 let body_from = field_end_of(raw, 1)?;
185 let actual_len = checksum_start - body_from;
186 if declared_len != actual_len {
187 return Err(Error::Parse(format!(
188 "BodyLength mismatch: declared {declared_len}, actual {actual_len}"
189 )));
190 }
191 let declared_sum: u32 = msg
192 .trailer
193 .get(tags::CHECK_SUM)
194 .map_err(|_| Error::Parse("missing/invalid CheckSum(10)".into()))?;
195 let actual_sum = checksum(&raw[..checksum_start]);
196 if declared_sum != actual_sum {
197 return Err(Error::Parse(format!(
198 "CheckSum mismatch: declared {declared_sum:03}, actual {actual_sum:03}"
199 )));
200 }
201 }
202
203 Ok(msg)
204 }
205
206 pub fn to_bytes(&self) -> Vec<u8> {
215 let mut inner = Vec::with_capacity(self.header.wire_len() + self.body.wire_len() + 64);
216
217 if let Some(v) = self.header.get_raw(tags::MSG_TYPE) {
219 write_tag_value(&mut inner, tags::MSG_TYPE, v);
220 }
221 let mut header: Vec<_> = self
222 .header
223 .fields()
224 .iter()
225 .filter(|f| !matches!(f.tag, tags::BEGIN_STRING | tags::BODY_LENGTH | tags::MSG_TYPE))
226 .collect();
227 header.sort_by_key(|f| f.tag);
228 for f in header {
229 write_tag_value(&mut inner, f.tag, &f.value);
230 }
231 self.body.write_to(&mut inner);
232 let mut trailer: Vec<_> =
233 self.trailer.fields().iter().filter(|f| f.tag != tags::CHECK_SUM).collect();
234 trailer.sort_by_key(|f| std::cmp::Reverse(f.tag));
235 for f in trailer {
236 write_tag_value(&mut inner, f.tag, &f.value);
237 }
238
239 let begin_string = self.header.get_raw(tags::BEGIN_STRING).unwrap_or(b"FIX.4.4");
240 let mut out = Vec::with_capacity(inner.len() + 32);
241 write_tag_value(&mut out, tags::BEGIN_STRING, begin_string);
242 write_tag_value(&mut out, tags::BODY_LENGTH, inner.len().to_string().as_bytes());
243 out.extend_from_slice(&inner);
244 let sum = checksum(&out);
245 write_tag_value(&mut out, tags::CHECK_SUM, format!("{sum:03}").as_bytes());
246 out
247 }
248
249 pub fn set_header(&mut self, tag: Tag, value: impl FixEncode) {
251 self.header.set(tag, value);
252 }
253
254 pub fn set(&mut self, tag: Tag, value: impl FixEncode) {
256 self.body.set(tag, value);
257 }
258
259 pub fn stamp_sending_time(&mut self, ts: UtcTimestamp) {
261 self.header.set(tags::SENDING_TIME, ts);
262 }
263}
264
265fn field_end_of(raw: &[u8], n: usize) -> Result<usize> {
267 let mut seen = 0usize;
268 for (i, &b) in raw.iter().enumerate() {
269 if b == SOH {
270 if seen == n {
271 return Ok(i + 1);
272 }
273 seen += 1;
274 }
275 }
276 Err(Error::Parse(format!("message has fewer than {} fields", n + 1)))
277}
278
279pub fn checksum(bytes: &[u8]) -> u32 {
281 bytes.iter().map(|&b| b as u32).sum::<u32>() % 256
282}
283
284#[cfg(test)]
285pub(crate) fn build_raw(fields: &[(Tag, &str)]) -> Vec<u8> {
286 let mut m = Message::new();
288 for &(tag, val) in fields {
289 if tags::is_header_tag(tag) {
290 m.header.push(tag, val);
291 } else if tags::is_trailer_tag(tag) {
292 m.trailer.push(tag, val);
293 } else {
294 m.body.push(tag, val);
295 }
296 }
297 m.to_bytes()
298}
299
300#[cfg(test)]
301mod tests {
302 use super::*;
303
304 fn sample() -> Vec<u8> {
305 build_raw(&[
306 (8, "FIX.4.2"),
307 (35, "D"),
308 (34, "2"),
309 (49, "TW"),
310 (52, "20240515-19:49:56.659"),
311 (56, "ISLD"),
312 (11, "100"),
313 (21, "1"),
314 (40, "1"),
315 (54, "1"),
316 (55, "TSLA"),
317 ])
318 }
319
320 #[test]
321 fn parse_roundtrip() {
322 let raw = sample();
323 let msg = Message::parse(&raw, true).unwrap();
324 assert_eq!(msg.msg_type().unwrap(), "D");
325 assert_eq!(msg.seq_num().unwrap(), 2);
326 assert_eq!(msg.header.get_string(49).unwrap(), "TW");
327 assert_eq!(msg.body.get_string(55).unwrap(), "TSLA");
328 assert!(msg.structure_error().is_none());
329 assert!(!msg.is_admin());
330 assert_eq!(msg.to_bytes(), raw);
331 }
332
333 #[test]
334 fn bad_checksum_rejected() {
335 let mut raw = sample();
336 let n = raw.len();
337 raw[n - 3] = b'9'; assert!(Message::parse(&raw, true).is_err());
339 assert!(Message::parse(&raw, false).is_ok());
341 }
342
343 #[test]
344 fn bad_body_length_rejected() {
345 let raw = sample();
346 let s = String::from_utf8(raw).unwrap();
347 let tampered = s.replacen("9=", "9=9", 1); assert!(Message::parse(tampered.as_bytes(), true).is_err());
349 }
350
351 #[test]
352 fn leading_field_order_enforced() {
353 let raw = b"8=FIX.4.2\x0135=D\x019=5\x0110=000\x01";
355 assert!(Message::parse(raw, false).is_err());
356 }
357
358 #[test]
359 fn data_field_with_soh_survives() {
360 let mut m = Message::new();
361 m.header.push(8, "FIX.4.2");
362 m.header.push(35, "B");
363 m.body.push(95, 5usize);
364 m.body.set_raw(96, b"a\x01b\x01c".to_vec());
365 m.body.push(58, "after");
366 let raw = m.to_bytes();
367
368 let parsed = Message::parse(&raw, true).unwrap();
369 assert_eq!(parsed.body.get_raw(96).unwrap(), b"a\x01b\x01c");
370 assert_eq!(parsed.body.get_string(58).unwrap(), "after");
371 }
372
373 #[test]
374 fn structure_error_recorded() {
375 let raw = build_raw(&[(8, "FIX.4.2"), (35, "D"), (55, "TSLA")]);
377 let s = String::from_utf8(raw).unwrap();
378 let spliced = s.replace("55=TSLA\x01", "55=TSLA\x0149=LATE\x01");
380 let msg = Message::parse(spliced.as_bytes(), false).unwrap();
381 assert_eq!(msg.structure_error(), Some(49));
382 }
383}