use crate::error::FinError;
use crate::ohlcv::OhlcvBar;
use crate::types::{NanoTimestamp, Price, Quantity, Symbol};
use arrow_array::cast::AsArray;
use arrow_array::Array;
use arrow_array::types::{Decimal128Type, TimestampNanosecondType, UInt64Type};
use arrow_array::{
ArrayRef, Decimal128Array, RecordBatch, StringArray, TimestampNanosecondArray, UInt64Array,
};
use arrow_schema::{DataType, Field, Schema, TimeUnit};
use rust_decimal::Decimal;
use std::sync::Arc;
const PRECISION: u8 = 38;
fn err(e: impl std::fmt::Display) -> FinError {
FinError::InvalidInput(format!("arrow: {e}"))
}
fn decimal_column(values: &[Decimal]) -> Result<(ArrayRef, u8), FinError> {
let scale = values.iter().map(Decimal::scale).max().unwrap_or(0);
let limit = 10i128.pow(u32::from(PRECISION));
let mut out = Vec::with_capacity(values.len());
for d in values {
let factor = 10i128.pow(scale - d.scale());
let v = d
.mantissa()
.checked_mul(factor)
.filter(|v| v.abs() < limit)
.ok_or_else(|| err(format!("{d} does not fit Decimal128({PRECISION}, {scale})")))?;
out.push(v);
}
#[allow(clippy::cast_possible_truncation, clippy::cast_possible_wrap)]
let scale_i8 = scale as i8; let arr = Decimal128Array::from(out).with_precision_and_scale(PRECISION, scale_i8).map_err(err)?;
Ok((Arc::new(arr), scale as u8))
}
fn ts_type() -> DataType {
DataType::Timestamp(TimeUnit::Nanosecond, Some("UTC".into()))
}
pub fn bars_to_record_batch(bars: &[OhlcvBar]) -> Result<RecordBatch, FinError> {
let col = |f: fn(&OhlcvBar) -> Decimal| bars.iter().map(f).collect::<Vec<_>>();
let (open, so) = decimal_column(&col(|b| b.open.value()))?;
let (high, sh) = decimal_column(&col(|b| b.high.value()))?;
let (low, sl) = decimal_column(&col(|b| b.low.value()))?;
let (close, sc) = decimal_column(&col(|b| b.close.value()))?;
let (volume, sv) = decimal_column(&col(|b| b.volume.value()))?;
#[allow(clippy::cast_possible_wrap)]
let dec = |s: u8| DataType::Decimal128(PRECISION, s as i8);
let schema = Schema::new(vec![
Field::new("symbol", DataType::Utf8, false),
Field::new("ts_open", ts_type(), false),
Field::new("ts_close", ts_type(), false),
Field::new("open", dec(so), false),
Field::new("high", dec(sh), false),
Field::new("low", dec(sl), false),
Field::new("close", dec(sc), false),
Field::new("volume", dec(sv), false),
Field::new("tick_count", DataType::UInt64, false),
]);
let columns: Vec<ArrayRef> = vec![
Arc::new(StringArray::from_iter_values(bars.iter().map(|b| b.symbol.as_str()))),
Arc::new(TimestampNanosecondArray::from_iter_values(bars.iter().map(|b| b.ts_open.nanos())).with_timezone("UTC")),
Arc::new(TimestampNanosecondArray::from_iter_values(bars.iter().map(|b| b.ts_close.nanos())).with_timezone("UTC")),
open,
high,
low,
close,
volume,
Arc::new(UInt64Array::from_iter_values(bars.iter().map(|b| b.tick_count))),
];
RecordBatch::try_new(Arc::new(schema), columns).map_err(err)
}
fn column<'a>(batch: &'a RecordBatch, name: &str) -> Result<&'a ArrayRef, FinError> {
batch.column_by_name(name).ok_or_else(|| err(format!("missing column `{name}`")))
}
fn decimals(batch: &RecordBatch, name: &str) -> Result<Vec<Decimal>, FinError> {
let col = column(batch, name)?;
let DataType::Decimal128(_, scale) = col.data_type() else {
return Err(err(format!("column `{name}` is {}, expected Decimal128", col.data_type())));
};
let scale = u32::try_from(*scale).map_err(|_| err(format!("negative scale in `{name}`")))?;
let arr = col.as_primitive_opt::<Decimal128Type>().ok_or_else(|| err(format!("column `{name}` is not Decimal128")))?;
arr.iter()
.map(|v| {
let v = v.ok_or_else(|| err(format!("null in `{name}`")))?;
Decimal::try_from_i128_with_scale(v, scale).map_err(err)
})
.collect()
}
pub fn record_batch_to_bars(batch: &RecordBatch) -> Result<Vec<OhlcvBar>, FinError> {
let sym = column(batch, "symbol")?.as_string_opt::<i32>().ok_or_else(|| err("`symbol` is not Utf8"))?;
let ts = |name: &str| -> Result<Vec<i64>, FinError> {
let c = column(batch, name)?
.as_primitive_opt::<TimestampNanosecondType>()
.ok_or_else(|| err(format!("`{name}` is not Timestamp(Nanosecond)")))?;
c.iter().map(|v| v.ok_or_else(|| err(format!("null in `{name}`")))).collect()
};
let (ts_open, ts_close) = (ts("ts_open")?, ts("ts_close")?);
let (open, high, low, close, volume) = (
decimals(batch, "open")?,
decimals(batch, "high")?,
decimals(batch, "low")?,
decimals(batch, "close")?,
decimals(batch, "volume")?,
);
let ticks = column(batch, "tick_count")?
.as_primitive_opt::<UInt64Type>()
.ok_or_else(|| err("`tick_count` is not UInt64"))?;
(0..batch.num_rows())
.map(|i| {
let symbol = sym.is_valid(i).then(|| sym.value(i)).ok_or_else(|| err("null in `symbol`"))?;
Ok(OhlcvBar {
symbol: Symbol::new(symbol)?,
open: Price::new(open[i])?,
high: Price::new(high[i])?,
low: Price::new(low[i])?,
close: Price::new(close[i])?,
volume: Quantity::new(volume[i])?,
ts_open: NanoTimestamp::new(ts_open[i]),
ts_close: NanoTimestamp::new(ts_close[i]),
tick_count: ticks.is_valid(i).then(|| ticks.value(i)).ok_or_else(|| err("null in `tick_count`"))?,
})
})
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
use rust_decimal_macros::dec;
fn bar(sym: &str, o: Decimal, c: Decimal, v: Decimal, t: i64) -> OhlcvBar {
OhlcvBar {
symbol: Symbol::new(sym).unwrap(),
open: Price::new(o).unwrap(),
high: Price::new(o.max(c) + dec!(1)).unwrap(),
low: Price::new(o.min(c) / dec!(2)).unwrap(),
close: Price::new(c).unwrap(),
volume: Quantity::new(v).unwrap(),
ts_open: NanoTimestamp::new(t),
ts_close: NanoTimestamp::new(t + 59_999_999_999),
tick_count: 7,
}
}
#[test]
fn round_trip_keeps_every_digit_across_mixed_scales() {
let bars = vec![
bar("BTC-USD", dec!(64250.5), dec!(64251), dec!(0.00012345), 0),
bar("ETH-USD", dec!(3186.17), dec!(3190.9412345678), dec!(1200), 60_000_000_000),
bar("SHIB-USD", dec!(0.000012345678901234), dec!(0.000012345), dec!(98765432109876543210), -60_000_000_000),
];
let batch = bars_to_record_batch(&bars).unwrap();
assert_eq!(batch.num_rows(), 3);
let back = record_batch_to_bars(&batch).unwrap();
for (a, b) in bars.iter().zip(&back) {
assert_eq!(a.symbol, b.symbol);
assert_eq!(a.ts_open, b.ts_open);
assert_eq!(a.tick_count, b.tick_count);
for (x, y) in [(a.open, b.open), (a.high, b.high), (a.low, b.low), (a.close, b.close)] {
assert_eq!(x.value(), y.value());
}
assert_eq!(a.volume.value(), b.volume.value());
}
}
#[test]
fn schema_uses_exact_decimal_and_utc_timestamps() {
let batch = bars_to_record_batch(&[bar("X", dec!(1.25), dec!(1.5), dec!(3), 0)]).unwrap();
let s = batch.schema();
assert_eq!(s.field_with_name("open").unwrap().data_type(), &DataType::Decimal128(38, 2));
assert_eq!(s.field_with_name("close").unwrap().data_type(), &DataType::Decimal128(38, 1));
assert_eq!(s.field_with_name("ts_open").unwrap().data_type(), &ts_type());
}
#[test]
fn empty_input_gives_an_empty_batch() {
let batch = bars_to_record_batch(&[]).unwrap();
assert_eq!(batch.num_rows(), 0);
assert!(record_batch_to_bars(&batch).unwrap().is_empty());
}
#[test]
fn invalid_rows_are_rejected_on_import() {
let good = bars_to_record_batch(&[bar("X", dec!(1), dec!(2), dec!(1), 0)]).unwrap();
let mut cols: Vec<ArrayRef> = good.columns().to_vec();
let idx = good.schema().index_of("close").unwrap();
cols[idx] = Arc::new(Decimal128Array::from(vec![0i128]).with_precision_and_scale(38, 0).unwrap());
let mut fields: Vec<Field> = good.schema().fields().iter().map(|f| f.as_ref().clone()).collect();
fields[idx] = Field::new("close", DataType::Decimal128(38, 0), false);
let bad = RecordBatch::try_new(Arc::new(Schema::new(fields)), cols).unwrap();
assert!(matches!(record_batch_to_bars(&bad), Err(FinError::InvalidPrice(_))));
}
}