Skip to main content

nautilus_serialization/arrow/display/
trade.rs

1// -------------------------------------------------------------------------------------------------
2//  Copyright (C) 2015-2026 Nautech Systems Pty Ltd. All rights reserved.
3//  https://nautechsystems.io
4//
5//  Licensed under the GNU Lesser General Public License Version 3.0 (the "License");
6//  You may not use this file except in compliance with the License.
7//  You may obtain a copy of the License at https://www.gnu.org/licenses/lgpl-3.0.en.html
8//
9//  Unless required by applicable law or agreed to in writing, software
10//  distributed under the License is distributed on an "AS IS" BASIS,
11//  WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12//  See the License for the specific language governing permissions and
13//  limitations under the License.
14// -------------------------------------------------------------------------------------------------
15
16//! Display-mode Arrow encoder for [`TradeTick`].
17
18use 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/// Returns the display-mode Arrow schema for [`TradeTick`].
34#[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
47/// Encodes trades as a display-friendly Arrow [`RecordBatch`].
48///
49/// Emits `Float64` columns for price and size, `Utf8` columns for the
50/// instrument ID, aggressor side, and trade ID, and `Timestamp(Nanosecond)`
51/// columns for event and init times. Mixed-instrument batches are supported.
52/// Precision is lost on the conversion to `f64`; use
53/// [`crate::arrow::trades_to_arrow_record_batch_bytes`] for catalog storage.
54///
55/// Returns an empty [`RecordBatch`] with the correct schema when `data` is empty.
56///
57/// # Errors
58///
59/// Returns an [`ArrowError`] if the Arrow `RecordBatch` cannot be constructed.
60pub 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}