1use std::collections::HashMap;
17
18use arrow::{datatypes::Schema, error::ArrowError, record_batch::RecordBatch};
19use nautilus_model::reports::{
20 ExecutionMassStatus, FillReport, OrderStatusReport, PositionStatusReport,
21};
22
23use super::{
24 ArrowSchemaProvider, DecodeTypedFromRecordBatch, EncodeToRecordBatch, EncodingError,
25 KEY_INSTRUMENT_ID,
26 json::{JsonFieldSpec, decode_batch, encode_batch, metadata_for_type, schema_for_type},
27};
28
29const ORDER_STATUS_REPORT_FIELDS: &[JsonFieldSpec] = &[
30 JsonFieldSpec::utf8("account_id", false),
31 JsonFieldSpec::utf8("instrument_id", false),
32 JsonFieldSpec::utf8("client_order_id", true),
33 JsonFieldSpec::utf8("venue_order_id", false),
34 JsonFieldSpec::utf8("order_side", false),
35 JsonFieldSpec::utf8("order_type", false),
36 JsonFieldSpec::utf8("time_in_force", false),
37 JsonFieldSpec::utf8("order_status", false),
38 JsonFieldSpec::utf8("quantity", false),
39 JsonFieldSpec::utf8("filled_qty", false),
40 JsonFieldSpec::utf8("report_id", false),
41 JsonFieldSpec::timestamp("ts_accepted", false),
42 JsonFieldSpec::timestamp("ts_last", false),
43 JsonFieldSpec::timestamp("ts_init", false),
44 JsonFieldSpec::utf8("order_list_id", true),
45 JsonFieldSpec::utf8("venue_position_id", true),
46 JsonFieldSpec::utf8_json("linked_order_ids", true),
47 JsonFieldSpec::utf8("parent_order_id", true),
48 JsonFieldSpec::utf8("contingency_type", false),
49 JsonFieldSpec::timestamp("expire_time", true),
50 JsonFieldSpec::utf8("price", true),
51 JsonFieldSpec::utf8("activation_price", true),
52 JsonFieldSpec::utf8("trigger_price", true),
53 JsonFieldSpec::utf8("trigger_type", true),
54 JsonFieldSpec::utf8("limit_offset", true),
55 JsonFieldSpec::utf8("trailing_offset", true),
56 JsonFieldSpec::utf8("trailing_offset_type", false),
57 JsonFieldSpec::utf8("avg_px", true),
58 JsonFieldSpec::utf8("display_qty", true),
59 JsonFieldSpec::boolean("post_only", false),
60 JsonFieldSpec::boolean("reduce_only", false),
61 JsonFieldSpec::utf8("cancel_reason", true),
62 JsonFieldSpec::timestamp("ts_triggered", true),
63];
64
65const FILL_REPORT_FIELDS: &[JsonFieldSpec] = &[
66 JsonFieldSpec::utf8("account_id", false),
67 JsonFieldSpec::utf8("instrument_id", false),
68 JsonFieldSpec::utf8("venue_order_id", false),
69 JsonFieldSpec::utf8("trade_id", false),
70 JsonFieldSpec::utf8("order_side", false),
71 JsonFieldSpec::utf8("last_qty", false),
72 JsonFieldSpec::utf8("last_px", false),
73 JsonFieldSpec::utf8("commission", false),
74 JsonFieldSpec::utf8("liquidity_side", false),
75 JsonFieldSpec::utf8("report_id", false),
76 JsonFieldSpec::timestamp("ts_event", false),
77 JsonFieldSpec::timestamp("ts_init", false),
78 JsonFieldSpec::utf8("client_order_id", true),
79 JsonFieldSpec::utf8("venue_position_id", true),
80];
81
82const POSITION_STATUS_REPORT_FIELDS: &[JsonFieldSpec] = &[
83 JsonFieldSpec::utf8("account_id", false),
84 JsonFieldSpec::utf8("instrument_id", false),
85 JsonFieldSpec::utf8("position_side", false),
86 JsonFieldSpec::utf8("quantity", false),
87 JsonFieldSpec::utf8("signed_decimal_qty", false),
88 JsonFieldSpec::utf8("report_id", false),
89 JsonFieldSpec::timestamp("ts_last", false),
90 JsonFieldSpec::timestamp("ts_init", false),
91 JsonFieldSpec::utf8("venue_position_id", true),
92 JsonFieldSpec::utf8("avg_px_open", true),
93];
94
95const EXECUTION_MASS_STATUS_FIELDS: &[JsonFieldSpec] = &[
96 JsonFieldSpec::utf8("client_id", false),
97 JsonFieldSpec::utf8("account_id", false),
98 JsonFieldSpec::utf8("venue", false),
99 JsonFieldSpec::utf8("report_id", false),
100 JsonFieldSpec::timestamp("ts_init", false),
101 JsonFieldSpec::utf8_json("order_reports", false),
102 JsonFieldSpec::utf8_json("fill_reports", false),
103 JsonFieldSpec::utf8_json("position_reports", false),
104];
105
106fn instrument_metadata(type_name: &'static str, instrument_id: &str) -> HashMap<String, String> {
107 let mut metadata = metadata_for_type(type_name);
108 metadata.insert(KEY_INSTRUMENT_ID.to_string(), instrument_id.to_string());
109 metadata
110}
111
112macro_rules! impl_report_arrow {
113 ($type:ty, $type_name:expr, $fields:expr) => {
114 impl ArrowSchemaProvider for $type {
115 fn get_schema(metadata: Option<HashMap<String, String>>) -> Schema {
116 schema_for_type($type_name, metadata, $fields)
117 }
118 }
119
120 impl EncodeToRecordBatch for $type {
121 fn encode_batch<T>(
122 metadata: &HashMap<String, String>,
123 data: &[T],
124 ) -> Result<RecordBatch, ArrowError>
125 where
126 T: std::borrow::Borrow<Self>,
127 {
128 encode_batch(
129 $type_name,
130 metadata,
131 data.iter().map(std::borrow::Borrow::borrow),
132 $fields,
133 )
134 }
135
136 fn metadata(&self) -> HashMap<String, String> {
137 instrument_metadata($type_name, &self.instrument_id.to_string())
138 }
139 }
140
141 impl DecodeTypedFromRecordBatch for $type {
142 fn decode_typed_batch(
143 metadata: &HashMap<String, String>,
144 record_batch: RecordBatch,
145 ) -> Result<Vec<Self>, EncodingError> {
146 decode_batch(metadata, &record_batch, $fields, Some($type_name))
147 }
148 }
149 };
150}
151
152impl_report_arrow!(
153 OrderStatusReport,
154 "OrderStatusReport",
155 ORDER_STATUS_REPORT_FIELDS
156);
157impl_report_arrow!(FillReport, "FillReport", FILL_REPORT_FIELDS);
158impl_report_arrow!(
159 PositionStatusReport,
160 "PositionStatusReport",
161 POSITION_STATUS_REPORT_FIELDS
162);
163
164impl ArrowSchemaProvider for ExecutionMassStatus {
165 fn get_schema(metadata: Option<HashMap<String, String>>) -> Schema {
166 schema_for_type(
167 "ExecutionMassStatus",
168 metadata,
169 EXECUTION_MASS_STATUS_FIELDS,
170 )
171 }
172}
173
174impl EncodeToRecordBatch for ExecutionMassStatus {
175 fn encode_batch<T>(
176 metadata: &HashMap<String, String>,
177 data: &[T],
178 ) -> Result<RecordBatch, ArrowError>
179 where
180 T: std::borrow::Borrow<Self>,
181 {
182 encode_batch(
183 "ExecutionMassStatus",
184 metadata,
185 data.iter().map(std::borrow::Borrow::borrow),
186 EXECUTION_MASS_STATUS_FIELDS,
187 )
188 }
189
190 fn metadata(&self) -> HashMap<String, String> {
191 metadata_for_type("ExecutionMassStatus")
192 }
193}
194
195impl DecodeTypedFromRecordBatch for ExecutionMassStatus {
196 fn decode_typed_batch(
197 metadata: &HashMap<String, String>,
198 record_batch: RecordBatch,
199 ) -> Result<Vec<Self>, EncodingError> {
200 decode_batch(
201 metadata,
202 &record_batch,
203 EXECUTION_MASS_STATUS_FIELDS,
204 Some("ExecutionMassStatus"),
205 )
206 }
207}
208
209#[cfg(test)]
210mod tests {
211 use std::str::FromStr;
212
213 use nautilus_core::{UUID4, UnixNanos};
214 use nautilus_model::{
215 enums::{OrderSide, OrderStatus, OrderType, PositionSide, TimeInForce},
216 identifiers::{AccountId, ClientOrderId, InstrumentId, PositionId, VenueOrderId},
217 reports::{OrderStatusReport, PositionStatusReport},
218 types::{Price, Quantity},
219 };
220 use rstest::rstest;
221 use rust_decimal::Decimal;
222
223 use super::*;
224
225 #[rstest]
226 fn test_order_status_report_round_trip() {
227 let report = OrderStatusReport::new(
228 AccountId::from("SIM-001"),
229 InstrumentId::from("AUDUSD.SIM"),
230 Some(ClientOrderId::from("O-19700101-000000-001-001-1")),
231 VenueOrderId::from("1"),
232 OrderSide::Buy.into(),
233 OrderType::Limit,
234 TimeInForce::Gtc,
235 OrderStatus::Accepted,
236 Quantity::from("100"),
237 Quantity::from("25"),
238 UnixNanos::from(1_000_000_000),
239 UnixNanos::from(2_000_000_000),
240 UnixNanos::from(3_000_000_000),
241 None,
242 )
243 .with_linked_order_ids([ClientOrderId::from("O-19700101-000000-001-001-2")]);
244 let report = OrderStatusReport {
245 activation_price: Some(Price::from("1.05000")),
246 limit_offset: Some(Decimal::from_str("0.123456789123456789").unwrap()),
247 trailing_offset: Some(Decimal::from_str("0.987654321987654321").unwrap()),
248 avg_px: Some(Decimal::from_str("1.23456789123456789").unwrap()),
249 ..report
250 };
251
252 let metadata = report.metadata();
253 let batch =
254 OrderStatusReport::encode_batch(&metadata, std::slice::from_ref(&report)).unwrap();
255 let decoded =
256 OrderStatusReport::decode_typed_batch(batch.schema().metadata(), batch).unwrap();
257
258 assert_eq!(decoded, vec![report]);
259 }
260
261 #[rstest]
262 fn test_position_status_report_round_trip_preserves_decimal_precision() {
263 let report = PositionStatusReport {
264 account_id: AccountId::from("SIM-001"),
265 instrument_id: InstrumentId::from("AUDUSD.SIM"),
266 position_side: PositionSide::Long,
267 quantity: Quantity::from("100.25"),
268 signed_decimal_qty: Decimal::from_str("100.250000000123456789").unwrap(),
269 report_id: UUID4::default(),
270 ts_last: UnixNanos::from(1_000_000_000),
271 ts_init: UnixNanos::from(2_000_000_000),
272 venue_position_id: Some(PositionId::from("P-001")),
273 avg_px_open: Some(Decimal::from_str("1.23456789123456789").unwrap()),
274 };
275 let metadata = report.metadata();
276 let batch =
277 PositionStatusReport::encode_batch(&metadata, std::slice::from_ref(&report)).unwrap();
278 let decoded =
279 PositionStatusReport::decode_typed_batch(batch.schema().metadata(), batch).unwrap();
280
281 assert_eq!(decoded, vec![report]);
282 }
283}