nautilus_serialization/arrow/display/
trade.rs1use std::sync::Arc;
19
20use arrow::{
21 array::{Float64Builder, StringBuilder, TimestampNanosecondBuilder},
22 datatypes::Schema,
23 error::ArrowError,
24 record_batch::RecordBatch,
25};
26use nautilus_model::data::TradeTick;
27
28use super::{
29 float64_field, price_to_f64, quantity_to_f64, timestamp_field, unix_nanos_to_i64, utf8_field,
30};
31use crate::arrow::timestamp_data_type;
32
33#[must_use]
35pub fn trades_schema() -> Schema {
36 Schema::new(vec![
37 utf8_field("instrument_id", false),
38 float64_field("price", false),
39 float64_field("size", false),
40 utf8_field("aggressor_side", false),
41 utf8_field("trade_id", false),
42 timestamp_field("ts_event", false),
43 timestamp_field("ts_init", false),
44 ])
45}
46
47pub fn encode_trades(data: &[TradeTick]) -> Result<RecordBatch, ArrowError> {
61 let mut instrument_id_builder = StringBuilder::new();
62 let mut price_builder = Float64Builder::with_capacity(data.len());
63 let mut size_builder = Float64Builder::with_capacity(data.len());
64 let mut aggressor_side_builder = StringBuilder::new();
65 let mut trade_id_builder = StringBuilder::new();
66 let mut ts_event_builder =
67 TimestampNanosecondBuilder::with_capacity(data.len()).with_data_type(timestamp_data_type());
68 let mut ts_init_builder =
69 TimestampNanosecondBuilder::with_capacity(data.len()).with_data_type(timestamp_data_type());
70
71 for trade in data {
72 instrument_id_builder.append_value(trade.instrument_id.to_string());
73 price_builder.append_value(price_to_f64(&trade.price));
74 size_builder.append_value(quantity_to_f64(&trade.size));
75 aggressor_side_builder.append_value(format!("{}", trade.aggressor_side));
76 trade_id_builder.append_value(trade.trade_id.to_string());
77 ts_event_builder.append_value(unix_nanos_to_i64(trade.ts_event.as_u64()));
78 ts_init_builder.append_value(unix_nanos_to_i64(trade.ts_init.as_u64()));
79 }
80
81 RecordBatch::try_new(
82 Arc::new(trades_schema()),
83 vec![
84 Arc::new(instrument_id_builder.finish()),
85 Arc::new(price_builder.finish()),
86 Arc::new(size_builder.finish()),
87 Arc::new(aggressor_side_builder.finish()),
88 Arc::new(trade_id_builder.finish()),
89 Arc::new(ts_event_builder.finish()),
90 Arc::new(ts_init_builder.finish()),
91 ],
92 )
93}
94
95#[cfg(test)]
96mod tests {
97 use arrow::{
98 array::{Array, Float64Array, StringArray, TimestampNanosecondArray},
99 datatypes::{DataType, TimeUnit},
100 };
101 use nautilus_model::{
102 enums::AggressorSide,
103 identifiers::{InstrumentId, TradeId},
104 types::{Price, Quantity},
105 };
106 use rstest::rstest;
107
108 use super::*;
109
110 fn make_trade(
111 instrument_id: &str,
112 price: &str,
113 aggressor_side: AggressorSide,
114 trade_id: &str,
115 ts: u64,
116 ) -> TradeTick {
117 TradeTick {
118 instrument_id: InstrumentId::from(instrument_id),
119 price: Price::from(price),
120 size: Quantity::from(1_000),
121 aggressor_side,
122 trade_id: TradeId::new(trade_id),
123 ts_event: ts.into(),
124 ts_init: (ts + 1).into(),
125 }
126 }
127
128 #[rstest]
129 fn test_encode_trades_schema() {
130 let batch = encode_trades(&[]).unwrap();
131 let fields = batch.schema().fields().clone();
132 assert_eq!(fields.len(), 7);
133 assert_eq!(fields[0].name(), "instrument_id");
134 assert_eq!(fields[0].data_type(), &DataType::Utf8);
135 assert_eq!(fields[1].name(), "price");
136 assert_eq!(fields[1].data_type(), &DataType::Float64);
137 assert_eq!(fields[2].name(), "size");
138 assert_eq!(fields[2].data_type(), &DataType::Float64);
139 assert_eq!(fields[3].name(), "aggressor_side");
140 assert_eq!(fields[3].data_type(), &DataType::Utf8);
141 assert_eq!(fields[4].name(), "trade_id");
142 assert_eq!(fields[4].data_type(), &DataType::Utf8);
143 assert_eq!(fields[5].name(), "ts_event");
144 assert_eq!(
145 fields[5].data_type(),
146 &DataType::Timestamp(TimeUnit::Nanosecond, Some("UTC".into()))
147 );
148 assert_eq!(fields[6].name(), "ts_init");
149 assert_eq!(
150 fields[6].data_type(),
151 &DataType::Timestamp(TimeUnit::Nanosecond, Some("UTC".into()))
152 );
153 }
154
155 #[rstest]
156 fn test_encode_trades_values() {
157 let trades = vec![
158 make_trade("AAPL.XNAS", "100.10", AggressorSide::Buy, "T-1", 1_000),
159 make_trade("AAPL.XNAS", "100.20", AggressorSide::Sell, "T-2", 2_000),
160 ];
161 let batch = encode_trades(&trades).unwrap();
162
163 assert_eq!(batch.num_rows(), 2);
164
165 let instrument_id_col = batch
166 .column(0)
167 .as_any()
168 .downcast_ref::<StringArray>()
169 .unwrap();
170 let price_col = batch
171 .column(1)
172 .as_any()
173 .downcast_ref::<Float64Array>()
174 .unwrap();
175 let size_col = batch
176 .column(2)
177 .as_any()
178 .downcast_ref::<Float64Array>()
179 .unwrap();
180 let aggressor_col = batch
181 .column(3)
182 .as_any()
183 .downcast_ref::<StringArray>()
184 .unwrap();
185 let trade_id_col = batch
186 .column(4)
187 .as_any()
188 .downcast_ref::<StringArray>()
189 .unwrap();
190 let ts_event_col = batch
191 .column(5)
192 .as_any()
193 .downcast_ref::<TimestampNanosecondArray>()
194 .unwrap();
195 let ts_init_col = batch
196 .column(6)
197 .as_any()
198 .downcast_ref::<TimestampNanosecondArray>()
199 .unwrap();
200
201 assert_eq!(instrument_id_col.value(0), "AAPL.XNAS");
202 assert!((price_col.value(0) - 100.10).abs() < 1e-9);
203 assert!((price_col.value(1) - 100.20).abs() < 1e-9);
204 assert!((size_col.value(0) - 1_000.0).abs() < 1e-9);
205 assert_eq!(aggressor_col.value(0), format!("{}", AggressorSide::Buy));
206 assert_eq!(aggressor_col.value(1), format!("{}", AggressorSide::Sell));
207 assert_eq!(trade_id_col.value(0), "T-1");
208 assert_eq!(trade_id_col.value(1), "T-2");
209 assert_eq!(ts_event_col.value(0), 1_000);
210 assert_eq!(ts_init_col.value(1), 2_001);
211 }
212
213 #[rstest]
214 fn test_encode_trades_empty() {
215 let batch = encode_trades(&[]).unwrap();
216 assert_eq!(batch.num_rows(), 0);
217 }
218
219 #[rstest]
220 fn test_encode_trades_mixed_instruments() {
221 let trades = vec![
222 make_trade("AAPL.XNAS", "100.10", AggressorSide::Buy, "A-1", 1),
223 make_trade("MSFT.XNAS", "250.00", AggressorSide::Sell, "M-1", 2),
224 ];
225 let batch = encode_trades(&trades).unwrap();
226 let instrument_id_col = batch
227 .column(0)
228 .as_any()
229 .downcast_ref::<StringArray>()
230 .unwrap();
231 assert_eq!(instrument_id_col.value(0), "AAPL.XNAS");
232 assert_eq!(instrument_id_col.value(1), "MSFT.XNAS");
233 }
234}