use czsc_core::objects::{bar::RawBar, freq::Freq};
use crate::bar_generator::{BarGenerator, nan_ohlcv_field};
use crate::errors::UtilsError;
use crate::freq_data::infer_market_from_bars;
pub fn resample_bars(
bars: &[RawBar],
target_freq: Freq,
drop_unfinished: bool,
) -> Result<Vec<RawBar>, UtilsError> {
if bars.is_empty() {
return Ok(Vec::new());
}
validate_batch_invariants(bars)?;
let base_freq = bars[0].freq;
let market = infer_market_from_bars(bars, base_freq);
let max_count = bars.len().saturating_add(1);
let bg = BarGenerator::new(base_freq, vec![target_freq], max_count, market)?;
for bar in bars {
bg.update_bar(bar)?;
}
let mut out: Vec<RawBar> = bg
.freq_bars
.get(&target_freq)
.map(|lock| lock.read().iter().cloned().collect())
.unwrap_or_default();
if drop_unfinished {
let last_base_dt = bars.last().map(|b| b.dt);
let last_target_dt = out.last().map(|b| b.dt);
if let (Some(lb), Some(lt)) = (last_base_dt, last_target_dt)
&& lb < lt
{
out.pop();
}
}
Ok(out)
}
fn validate_batch_invariants(bars: &[RawBar]) -> Result<(), UtilsError> {
debug_assert!(!bars.is_empty(), "caller must short-circuit empty input");
let first_symbol = &bars[0].symbol;
let first_freq = bars[0].freq;
let mut prev_dt = bars[0].dt;
check_no_nan(&bars[0], 0)?;
for (idx, bar) in bars.iter().enumerate().skip(1) {
if bar.symbol != *first_symbol {
return Err(UtilsError::Unexpected(anyhow::anyhow!(
"resample_bars: bars 列表混合了多个 symbol(bars[0]={}, bars[{}]={}),\
batch 重采样不做 symbol 分组,请先按 symbol 拆分输入",
first_symbol,
idx,
bar.symbol
)));
}
if bar.freq != first_freq {
return Err(UtilsError::Unexpected(anyhow::anyhow!(
"resample_bars: bars 列表混合了多个 freq(bars[0]={}, bars[{}]={}),\
batch 重采样要求输入同频率",
first_freq,
idx,
bar.freq
)));
}
if bar.dt <= prev_dt {
let reason = if bar.dt == prev_dt {
"存在重复 dt"
} else {
"存在乱序 dt"
};
return Err(UtilsError::Unexpected(anyhow::anyhow!(
"resample_bars: bars[{}].dt={} 与 bars[{}].dt={} {}(要求严格单调递增),\
请在调用前去重并按 dt 升序排列",
idx,
bar.dt.format("%Y-%m-%d %H:%M:%S"),
idx - 1,
prev_dt.format("%Y-%m-%d %H:%M:%S"),
reason,
)));
}
check_no_nan(bar, idx)?;
prev_dt = bar.dt;
}
Ok(())
}
fn check_no_nan(bar: &RawBar, idx: usize) -> Result<(), UtilsError> {
if let Some(field) = nan_ohlcv_field(bar) {
return Err(UtilsError::Unexpected(anyhow::anyhow!(
"resample_bars: bars[{}].{} = NaN(dt={}),\
batch 模式拒绝 NaN OHLCV 输入以避免桶聚合静默污染——请在上游 \
dropna 或填充后再调用",
idx,
field,
bar.dt.format("%Y-%m-%d %H:%M:%S"),
)));
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use chrono::{NaiveDateTime, TimeZone, Utc};
use czsc_core::objects::bar::RawBarBuilder;
fn build_ashare_1min_bars(n: usize) -> Vec<RawBar> {
let start = Utc.from_utc_datetime(
&NaiveDateTime::parse_from_str("2024-12-12 09:31:00", "%Y-%m-%d %H:%M:%S").unwrap(),
);
(0..n)
.map(|i| {
let price = 100.0 + i as f64;
RawBarBuilder::default()
.symbol("000001.XSHG")
.id(i as i32)
.dt(start + chrono::Duration::minutes(i as i64))
.freq(Freq::F1)
.open(price)
.close(price + 0.1)
.high(price + 0.5)
.low(price - 0.5)
.vol(1_000.0 * (i as f64 + 1.0))
.amount(10_000.0 * (i as f64 + 1.0))
.build()
.unwrap()
})
.collect()
}
fn dt_str(bar: &RawBar) -> String {
bar.dt.format("%Y-%m-%d %H:%M:%S").to_string()
}
#[test]
fn empty_input_returns_empty() {
let out = resample_bars(&[], Freq::F5, true).unwrap();
assert!(out.is_empty());
}
#[test]
fn one_minute_to_five_minute_complete_bucket() {
let bars = build_ashare_1min_bars(5);
let out = resample_bars(&bars, Freq::F5, true).unwrap();
assert_eq!(out.len(), 1, "5 根 1min 恰好凑成 1 根 5min");
assert_eq!(dt_str(&out[0]), "2024-12-12 09:35:00");
assert_eq!(out[0].freq, Freq::F5);
assert_eq!(out[0].open, bars[0].open);
assert_eq!(out[0].close, bars[4].close);
assert_eq!(out[0].high, bars[4].high);
assert_eq!(out[0].low, bars[0].low);
let vol_sum: f64 = bars.iter().map(|b| b.vol).sum();
let amt_sum: f64 = bars.iter().map(|b| b.amount).sum();
assert_eq!(out[0].vol, vol_sum);
assert_eq!(out[0].amount, amt_sum);
}
#[test]
fn drop_unfinished_drops_partial_tail() {
let bars = build_ashare_1min_bars(7);
let kept = resample_bars(&bars, Freq::F5, true).unwrap();
assert_eq!(kept.len(), 1);
assert_eq!(dt_str(&kept[0]), "2024-12-12 09:35:00");
let all = resample_bars(&bars, Freq::F5, false).unwrap();
assert_eq!(all.len(), 2);
assert_eq!(dt_str(&all[0]), "2024-12-12 09:35:00");
assert_eq!(dt_str(&all[1]), "2024-12-12 09:40:00");
assert_eq!(all[1].open, bars[5].open);
assert_eq!(all[1].close, bars[6].close);
assert_eq!(all[1].high, bars[5].high.max(bars[6].high));
assert_eq!(all[1].low, bars[5].low.min(bars[6].low));
assert_eq!(all[1].vol, bars[5].vol + bars[6].vol);
assert_eq!(all[1].amount, bars[5].amount + bars[6].amount);
}
#[test]
fn base_equals_target_is_identity_in_count() {
let bars = build_ashare_1min_bars(3);
let out = resample_bars(&bars, Freq::F1, true).unwrap();
assert_eq!(out.len(), 3);
for (a, b) in bars.iter().zip(out.iter()) {
assert_eq!(a.dt, b.dt);
assert_eq!(a.open, b.open);
assert_eq!(a.close, b.close);
assert_eq!(a.high, b.high);
assert_eq!(a.low, b.low);
assert_eq!(a.vol, b.vol);
}
}
#[test]
fn one_minute_to_daily_keeps_partial_known_limitation() {
let bars = build_ashare_1min_bars(2);
let out = resample_bars(&bars, Freq::D, true).unwrap();
assert_eq!(out.len(), 1);
assert_eq!(out[0].freq, Freq::D);
assert_eq!(dt_str(&out[0]), "2024-12-12 00:00:00");
assert_eq!(out[0].open, bars[0].open);
assert_eq!(out[0].close, bars[1].close);
}
#[test]
fn mixed_symbol_returns_err() {
let mut bars = build_ashare_1min_bars(3);
bars[1].symbol = "OTHER.XSHG".into();
let err = resample_bars(&bars, Freq::F5, true).unwrap_err();
assert!(
err.to_string().contains("symbol"),
"错误信息应当点名 symbol 不一致:{err}"
);
}
#[test]
fn mixed_freq_returns_err() {
let mut bars = build_ashare_1min_bars(3);
bars[1].freq = Freq::F5;
let err = resample_bars(&bars, Freq::F5, true).unwrap_err();
assert!(
err.to_string().contains("freq"),
"错误信息应当点名 freq 不一致:{err}"
);
}
#[test]
fn duplicate_dt_returns_err() {
let mut bars = build_ashare_1min_bars(3);
bars[2].dt = bars[1].dt;
let err = resample_bars(&bars, Freq::F5, true).unwrap_err();
assert!(
err.to_string().contains("重复"),
"错误信息应当点名重复 dt:{err}"
);
}
#[test]
fn out_of_order_dt_returns_err() {
let mut bars = build_ashare_1min_bars(3);
bars.swap(1, 2);
let err = resample_bars(&bars, Freq::F5, true).unwrap_err();
assert!(
err.to_string().contains("乱序"),
"错误信息应当点名乱序:{err}"
);
}
#[test]
fn nan_ohlcv_returns_err() {
for field in ["open", "close", "high", "low", "vol", "amount"] {
let mut bars = build_ashare_1min_bars(3);
match field {
"open" => bars[1].open = f64::NAN,
"close" => bars[1].close = f64::NAN,
"high" => bars[1].high = f64::NAN,
"low" => bars[1].low = f64::NAN,
"vol" => bars[1].vol = f64::NAN,
"amount" => bars[1].amount = f64::NAN,
_ => unreachable!(),
}
let err = resample_bars(&bars, Freq::F5, true).unwrap_err();
let msg = err.to_string();
assert!(
msg.contains(field) && msg.contains("NaN"),
"字段 {field} 的 NaN 错误信息应当点名字段与 NaN:{err}"
);
}
}
#[test]
fn nan_at_first_bar_returns_err() {
let mut bars = build_ashare_1min_bars(3);
bars[0].vol = f64::NAN;
let err = resample_bars(&bars, Freq::F5, true).unwrap_err();
assert!(
err.to_string().contains("bars[0]"),
"错误信息应当点名 bars[0]:{err}"
);
}
}