use std::{marker::PhantomData, range::Range};
use crate::{
aggregation::Aggregator,
binned_method::{BinnedDataError, BinnedPoint, FlexBinnedData},
instrument::InstrumentSpec,
timestamp::{Timestamp, Timestamped, TradeTimestamped},
};
pub struct FlexBinnedTimePointData<
'instrument_data,
'aggregated_data,
IS: InstrumentSpec,
T: TradeTimestamped
+ Aggregator<'instrument_data, 'aggregated_data, T, AggregatorData, IS>
+ Ord
+ 'aggregated_data,
AggregatorData,
const SECTOR_STRIDE_NS: u64,
> {
flex_binned_data: FlexBinnedData<T>,
_instrument_data: PhantomData<&'instrument_data ()>,
_aggregated_data: PhantomData<&'aggregated_data ()>,
_is: PhantomData<IS>,
_aggregator_data: PhantomData<AggregatorData>,
}
impl<
'instrument_data,
'aggregated_data,
IS: InstrumentSpec,
T: TradeTimestamped
+ Aggregator<'instrument_data, 'aggregated_data, T, AggregatorData, IS>
+ Ord
+ 'aggregated_data,
AggregatorData,
const SECTOR_STRIDE_NS: u64,
> BinnedPoint<T>
for FlexBinnedTimePointData<
'instrument_data,
'aggregated_data,
IS,
T,
AggregatorData,
SECTOR_STRIDE_NS,
>
{
fn insert(
&mut self,
data: impl Iterator<Item = T>,
) {
for data in data {
let sector_idx =
data.trade_timestamp()
.as_utc_nanos() / SECTOR_STRIDE_NS as i128;
let sector_idx = sector_idx.cast_unsigned() as u64;
self.flex_binned_data
.insert_in_sector(sector_idx, [data].into_iter())
.unwrap();
}
}
fn get<'a>(
&'a self,
timestamp: &Timestamp,
) -> Result<impl Iterator<Item = &'a T>, BinnedDataError>
where
T: 'a,
{
let sector_idx = timestamp.as_utc_nanos() / SECTOR_STRIDE_NS as i128;
let sector_idx = sector_idx.cast_unsigned() as u64;
let sector_data = self
.flex_binned_data
.get_sector(sector_idx)?;
Ok(sector_data
.iter()
.filter(move |t| {
t.trade_timestamp()
.timestamp() == timestamp
}))
}
fn get_first(
&self,
timestamp: &Timestamp,
) -> Result<Option<&T>, BinnedDataError> {
let sector_idx = timestamp.as_utc_nanos() / SECTOR_STRIDE_NS as i128;
let sector_idx = sector_idx.cast_unsigned() as u64;
let sector_data = self
.flex_binned_data
.get_sector(sector_idx)?;
Ok(sector_data.iter().find(|t| {
t.trade_timestamp()
.timestamp() == timestamp
}))
}
fn iter<'a>(
&'a self,
range: Range<&Timestamp>,
) -> Result<impl Iterator<Item = &'a T>, BinnedDataError>
where
T: 'a,
{
let start_sector_idx = range.start.as_utc_nanos() / SECTOR_STRIDE_NS as i128;
let last_sector_idx = range.end.as_utc_nanos() / SECTOR_STRIDE_NS as i128;
debug_assert!(start_sector_idx <= last_sector_idx);
let start_sector_idx = start_sector_idx.cast_unsigned() as u64;
let last_sector_idx = last_sector_idx.cast_unsigned() as u64;
self.flex_binned_data
.iter(
start_sector_idx,
last_sector_idx,
|t: &T| {
t.trade_timestamp()
.timestamp() >= range.start
},
|t: &T| {
t.trade_timestamp()
.timestamp() < range.end
},
)
.map_err(|_| BinnedDataError::InternalError)
}
}
impl<
'instrument_data,
'aggregated_data,
IS: InstrumentSpec,
T: TradeTimestamped
+ Aggregator<'instrument_data, 'aggregated_data, T, AggregatorData, IS>
+ Ord
+ 'aggregated_data,
AggregatorData,
const SECTOR_STRIDE_NS: u64,
> Default
for FlexBinnedTimePointData<
'instrument_data,
'aggregated_data,
IS,
T,
AggregatorData,
SECTOR_STRIDE_NS,
>
{
fn default() -> Self {
Self {
flex_binned_data: FlexBinnedData::default(),
_instrument_data: PhantomData::default(),
_aggregated_data: PhantomData::default(),
_is: PhantomData::default(),
_aggregator_data: PhantomData::default(),
}
}
}