1#![cfg_attr(docsrs, feature(doc_cfg))]
15#![forbid(unsafe_code)]
16
17use std::path::Path;
18
19use wickra_backtest_core::{BacktestError, Candle, Result};
20
21pub fn load_candles(path: &Path) -> Result<Vec<Candle>> {
25 match path.extension().and_then(|e| e.to_str()) {
26 Some("parquet") => load_parquet(path),
27 Some("jsonl" | "ndjson") => parse_jsonl(&read_text(path)?),
28 Some("json") => parse_json_array(&read_text(path)?),
29 _ => parse_csv(&read_text(path)?),
30 }
31}
32
33fn read_text(path: &Path) -> Result<String> {
34 std::fs::read_to_string(path)
35 .map_err(|e| BacktestError::InvalidData(format!("reading {}: {e}", path.display())))
36}
37
38#[cfg(not(feature = "parquet"))]
43pub fn load_parquet(_path: &Path) -> Result<Vec<Candle>> {
44 Err(BacktestError::InvalidData(
45 "Parquet support is not compiled in; rebuild with the `parquet` feature".into(),
46 ))
47}
48
49#[cfg(feature = "parquet")]
53pub fn load_parquet(path: &Path) -> Result<Vec<Candle>> {
54 use parquet::arrow::arrow_reader::ParquetRecordBatchReaderBuilder;
55
56 let file = std::fs::File::open(path)
57 .map_err(|e| BacktestError::InvalidData(format!("opening {}: {e}", path.display())))?;
58 let builder = ParquetRecordBatchReaderBuilder::try_new(file)
59 .map_err(|e| BacktestError::InvalidData(format!("parquet: {e}")))?;
60
61 let schema = builder.schema().clone();
62 let col = |name: &str| {
63 schema
64 .fields()
65 .iter()
66 .position(|f| f.name().eq_ignore_ascii_case(name))
67 };
68 let require = |name: &str| {
69 col(name).ok_or_else(|| {
70 BacktestError::InvalidData(format!("parquet: missing required column `{name}`"))
71 })
72 };
73 let (i_time, i_open, i_high, i_low, i_close) = (
74 require("time")?,
75 require("open")?,
76 require("high")?,
77 require("low")?,
78 require("close")?,
79 );
80 let i_volume = col("volume");
81
82 let reader = builder
83 .build()
84 .map_err(|e| BacktestError::InvalidData(format!("parquet: {e}")))?;
85
86 let mut out = Vec::new();
87 for batch in reader {
88 let batch = batch.map_err(|e| BacktestError::InvalidData(format!("parquet: {e}")))?;
89 let time = column_i64(&batch, i_time)?;
90 let open = column_f64(&batch, i_open)?;
91 let high = column_f64(&batch, i_high)?;
92 let low = column_f64(&batch, i_low)?;
93 let close = column_f64(&batch, i_close)?;
94 let volume = i_volume.map(|i| column_f64(&batch, i)).transpose()?;
95 for r in 0..batch.num_rows() {
96 out.push(Candle {
97 time: time[r],
98 open: open[r],
99 high: high[r],
100 low: low[r],
101 close: close[r],
102 volume: volume.as_ref().map_or(0.0, |v| v[r]),
103 });
104 }
105 }
106 Ok(out)
107}
108
109#[cfg(feature = "parquet")]
111fn column_f64(batch: &arrow_array::RecordBatch, idx: usize) -> Result<Vec<f64>> {
112 use arrow_array::{
113 cast::AsArray,
114 types::{Float32Type, Float64Type, Int32Type, Int64Type},
115 };
116
117 let array = batch.column(idx);
118 if array.null_count() > 0 {
119 return Err(BacktestError::InvalidData(
120 "parquet: null values are not allowed in OHLCV columns".into(),
121 ));
122 }
123 let dt = array.data_type();
124 if let Some(a) = array.as_primitive_opt::<Float64Type>() {
125 Ok(a.values().to_vec())
126 } else if let Some(a) = array.as_primitive_opt::<Float32Type>() {
127 Ok(a.values().iter().map(|&v| f64::from(v)).collect())
128 } else if let Some(a) = array.as_primitive_opt::<Int64Type>() {
129 Ok(a.values().iter().map(|&v| v as f64).collect())
130 } else if let Some(a) = array.as_primitive_opt::<Int32Type>() {
131 Ok(a.values().iter().map(|&v| f64::from(v)).collect())
132 } else {
133 Err(BacktestError::InvalidData(format!(
134 "parquet: column {idx} has unsupported numeric type {dt:?}"
135 )))
136 }
137}
138
139#[cfg(feature = "parquet")]
141fn column_i64(batch: &arrow_array::RecordBatch, idx: usize) -> Result<Vec<i64>> {
142 Ok(column_f64(batch, idx)?
143 .into_iter()
144 .map(|v| v as i64)
145 .collect())
146}
147
148pub fn parse_csv(content: &str) -> Result<Vec<Candle>> {
151 let mut out = Vec::new();
152 for (i, line) in content.lines().enumerate() {
153 let line = line.trim();
154 if line.is_empty() {
155 continue;
156 }
157 let cols: Vec<&str> = line.split(',').map(str::trim).collect();
158 if i == 0 && cols.first().is_some_and(|c| c.parse::<f64>().is_err()) {
160 continue;
161 }
162 if cols.len() < 5 {
163 return Err(BacktestError::InvalidData(format!(
164 "CSV line {}: expected at least 5 columns (time,o,h,l,c), got {}",
165 i + 1,
166 cols.len()
167 )));
168 }
169 let num = |idx: usize| -> Result<f64> {
170 cols[idx].parse::<f64>().map_err(|_| {
171 BacktestError::InvalidData(format!(
172 "CSV line {}: column {idx} is not a number",
173 i + 1
174 ))
175 })
176 };
177 out.push(Candle {
178 time: num(0)? as i64,
179 open: num(1)?,
180 high: num(2)?,
181 low: num(3)?,
182 close: num(4)?,
183 volume: if cols.len() > 5 { num(5)? } else { 0.0 },
184 });
185 }
186 Ok(out)
187}
188
189pub fn parse_jsonl(content: &str) -> Result<Vec<Candle>> {
191 let mut out = Vec::new();
192 for (i, line) in content.lines().enumerate() {
193 let line = line.trim();
194 if line.is_empty() {
195 continue;
196 }
197 let candle: Candle = serde_json::from_str(line)
198 .map_err(|e| BacktestError::InvalidData(format!("JSONL line {}: {e}", i + 1)))?;
199 out.push(candle);
200 }
201 Ok(out)
202}
203
204pub fn parse_json_array(content: &str) -> Result<Vec<Candle>> {
206 serde_json::from_str(content)
207 .map_err(|e| BacktestError::InvalidData(format!("JSON array: {e}")))
208}
209
210pub fn parse_binance_klines(json: &str) -> Result<Vec<Candle>> {
216 let rows: Vec<Vec<serde_json::Value>> = serde_json::from_str(json)
217 .map_err(|e| BacktestError::InvalidData(format!("binance klines: {e}")))?;
218 let mut out = Vec::with_capacity(rows.len());
219 for (i, row) in rows.iter().enumerate() {
220 if row.len() < 6 {
221 return Err(BacktestError::InvalidData(format!(
222 "binance kline {i}: expected at least 6 fields, got {}",
223 row.len()
224 )));
225 }
226 let time_ms = row[0].as_i64().ok_or_else(|| {
227 BacktestError::InvalidData(format!("binance kline {i}: open time is not an integer"))
228 })?;
229 out.push(Candle {
230 time: time_ms / 1000,
231 open: binance_num(&row[1], i)?,
232 high: binance_num(&row[2], i)?,
233 low: binance_num(&row[3], i)?,
234 close: binance_num(&row[4], i)?,
235 volume: binance_num(&row[5], i)?,
236 });
237 }
238 Ok(out)
239}
240
241fn binance_num(value: &serde_json::Value, kline: usize) -> Result<f64> {
243 match value {
244 serde_json::Value::String(s) => s.parse::<f64>().map_err(|_| {
245 BacktestError::InvalidData(format!("binance kline {kline}: `{s}` is not a number"))
246 }),
247 serde_json::Value::Number(n) => n.as_f64().ok_or_else(|| {
248 BacktestError::InvalidData(format!("binance kline {kline}: non-finite number"))
249 }),
250 _ => Err(BacktestError::InvalidData(format!(
251 "binance kline {kline}: expected a numeric field"
252 ))),
253 }
254}
255
256#[cfg(feature = "binance")]
261pub fn fetch_klines(symbol: &str, interval: &str, limit: u32) -> Result<Vec<Candle>> {
262 let url = format!(
263 "https://api.binance.com/api/v3/klines?symbol={symbol}&interval={interval}&limit={limit}"
264 );
265 let body = ureq::get(&url)
266 .call()
267 .map_err(|e| BacktestError::InvalidData(format!("binance request failed: {e}")))?
268 .into_string()
269 .map_err(|e| BacktestError::InvalidData(format!("binance response: {e}")))?;
270 parse_binance_klines(&body)
271}
272
273fn aggregate(bucket: &[Candle]) -> Candle {
277 let first = &bucket[0];
278 let last = &bucket[bucket.len() - 1];
279 let mut high = first.high;
280 let mut low = first.low;
281 let mut volume = 0.0;
282 for c in bucket {
283 high = high.max(c.high);
284 low = low.min(c.low);
285 volume += c.volume;
286 }
287 Candle {
288 time: first.time,
289 open: first.open,
290 high,
291 low,
292 close: last.close,
293 volume,
294 }
295}
296
297pub fn resample_by_count(candles: &[Candle], count: usize) -> Result<Vec<Candle>> {
301 if count == 0 {
302 return Err(BacktestError::InvalidData(
303 "resample count must be > 0".into(),
304 ));
305 }
306 Ok(candles.chunks(count).map(aggregate).collect())
307}
308
309pub fn resample_by_interval(candles: &[Candle], interval: i64) -> Result<Vec<Candle>> {
314 if interval <= 0 {
315 return Err(BacktestError::InvalidData(
316 "resample interval must be > 0".into(),
317 ));
318 }
319 let mut out: Vec<Candle> = Vec::new();
320 let mut start = 0usize;
321 for i in 0..candles.len() {
322 let bucket = candles[i].time.div_euclid(interval);
323 let next_bucket = candles
324 .get(i + 1)
325 .map(|c| c.time.div_euclid(interval) != bucket);
326 if next_bucket != Some(false) {
327 let mut bar = aggregate(&candles[start..=i]);
329 bar.time = bucket * interval;
330 out.push(bar);
331 start = i + 1;
332 }
333 }
334 Ok(out)
335}
336
337fn edge_candle(index: i64, open_edge: f64, close_edge: f64) -> Candle {
342 Candle {
343 time: index,
344 open: open_edge,
345 high: open_edge.max(close_edge),
346 low: open_edge.min(close_edge),
347 close: close_edge,
348 volume: 0.0,
349 }
350}
351
352pub fn to_renko(candles: &[Candle], box_size: f64) -> Result<Vec<Candle>> {
357 use wickra_core::{BarBuilder, RenkoBars};
358
359 let mut builder =
360 RenkoBars::new(box_size).map_err(|e| BacktestError::InvalidData(e.to_string()))?;
361 let mut out = Vec::new();
362 let mut index: i64 = 0;
363 for candle in candles {
364 for brick in builder.update(candle.to_core()?) {
365 out.push(edge_candle(index, brick.open, brick.close));
366 index += 1;
367 }
368 }
369 Ok(out)
370}
371
372pub fn to_kagi(candles: &[Candle], reversal: f64) -> Result<Vec<Candle>> {
376 use wickra_core::{BarBuilder, KagiBars};
377
378 let mut builder =
379 KagiBars::new(reversal).map_err(|e| BacktestError::InvalidData(e.to_string()))?;
380 let mut out = Vec::new();
381 let mut index: i64 = 0;
382 for candle in candles {
383 for bar in builder.update(candle.to_core()?) {
384 out.push(edge_candle(index, bar.start, bar.end));
385 index += 1;
386 }
387 }
388 Ok(out)
389}
390
391pub fn to_pnf(candles: &[Candle], box_size: f64, reversal: usize) -> Result<Vec<Candle>> {
397 use wickra_core::{BarBuilder, PointAndFigureBars};
398
399 let mut builder = PointAndFigureBars::new(box_size, reversal)
400 .map_err(|e| BacktestError::InvalidData(e.to_string()))?;
401 let mut out = Vec::new();
402 let mut index: i64 = 0;
403 for candle in candles {
404 for col in builder.update(candle.to_core()?) {
405 let (open_edge, close_edge) = if col.direction >= 0 {
406 (col.low, col.high)
407 } else {
408 (col.high, col.low)
409 };
410 out.push(edge_candle(index, open_edge, close_edge));
411 index += 1;
412 }
413 }
414 Ok(out)
415}
416
417#[cfg(test)]
418mod tests {
419 use super::*;
420
421 #[test]
422 fn csv_with_header_and_volume() {
423 let csv = "time,open,high,low,close,volume\n1,10,12,9,11,100\n2,11,13,10,12,200\n";
424 let c = parse_csv(csv).unwrap();
425 assert_eq!(c.len(), 2);
426 assert_eq!(c[0].time, 1);
427 assert!((c[0].close - 11.0).abs() < 1e-9);
428 assert!((c[1].volume - 200.0).abs() < 1e-9);
429 }
430
431 #[test]
432 fn csv_without_header_or_volume() {
433 let csv = "1,10,12,9,11\n2,11,13,10,12\n";
434 let c = parse_csv(csv).unwrap();
435 assert_eq!(c.len(), 2);
436 assert!((c[0].volume).abs() < f64::EPSILON);
437 }
438
439 #[test]
440 fn csv_too_few_columns_errors() {
441 assert!(parse_csv("1,2,3\n").is_err());
442 }
443
444 #[test]
445 fn jsonl_roundtrip() {
446 let jsonl = "{\"time\":1,\"open\":1,\"high\":2,\"low\":0.5,\"close\":1.5}\n{\"time\":2,\"open\":1.5,\"high\":2,\"low\":1,\"close\":1.8,\"volume\":5}\n";
447 let c = parse_jsonl(jsonl).unwrap();
448 assert_eq!(c.len(), 2);
449 assert_eq!(c[1].time, 2);
450 assert!((c[1].volume - 5.0).abs() < 1e-9);
451 }
452
453 #[test]
454 fn json_array() {
455 let json = "[{\"time\":1,\"open\":1,\"high\":2,\"low\":0.5,\"close\":1.5}]";
456 let c = parse_json_array(json).unwrap();
457 assert_eq!(c.len(), 1);
458 }
459
460 fn four_bars() -> Vec<Candle> {
461 parse_csv("0,10,12,9,11,100\n1,11,13,10,12,200\n2,12,14,11,13,300\n3,13,15,12,14,400\n")
463 .unwrap()
464 }
465
466 #[test]
467 fn resample_by_count_aggregates_buckets() {
468 let bars = resample_by_count(&four_bars(), 2).unwrap();
469 assert_eq!(bars.len(), 2);
470 assert_eq!(bars[0].time, 0);
471 assert!((bars[0].open - 10.0).abs() < 1e-9);
472 assert!((bars[0].high - 13.0).abs() < 1e-9); assert!((bars[0].low - 9.0).abs() < 1e-9); assert!((bars[0].close - 12.0).abs() < 1e-9); assert!((bars[0].volume - 300.0).abs() < 1e-9); assert!((bars[1].close - 14.0).abs() < 1e-9);
477 assert!((bars[1].volume - 700.0).abs() < 1e-9);
478 }
479
480 #[test]
481 fn resample_by_count_keeps_trailing_partial_group() {
482 let bars = resample_by_count(&four_bars(), 3).unwrap();
483 assert_eq!(bars.len(), 2); assert!((bars[1].close - 14.0).abs() < 1e-9);
485 assert!((bars[1].volume - 400.0).abs() < 1e-9);
486 }
487
488 #[test]
489 fn resample_by_interval_buckets_on_time() {
490 let bars = resample_by_interval(&four_bars(), 2).unwrap();
491 assert_eq!(bars.len(), 2);
492 assert_eq!(bars[0].time, 0); assert!((bars[0].close - 12.0).abs() < 1e-9);
494 assert_eq!(bars[1].time, 2); assert!((bars[1].close - 14.0).abs() < 1e-9);
496 assert!((bars[1].volume - 700.0).abs() < 1e-9);
497 }
498
499 #[test]
500 fn resample_rejects_zero_step() {
501 assert!(resample_by_count(&four_bars(), 0).is_err());
502 assert!(resample_by_interval(&four_bars(), 0).is_err());
503 }
504
505 #[test]
506 fn binance_klines_parse_with_string_fields() {
507 let json = r#"[
510 [1609459200000,"100.0","102.0","99.5","101.0","1234.5",1609462799999,"0",10,"0","0","0"],
511 [1609462800000,"101.0","103.0","100.0","102.5","2000.0",1609466399999,"0",12,"0","0","0"]
512 ]"#;
513 let c = parse_binance_klines(json).unwrap();
514 assert_eq!(c.len(), 2);
515 assert_eq!(c[0].time, 1_609_459_200); assert!((c[0].open - 100.0).abs() < 1e-9);
517 assert!((c[0].close - 101.0).abs() < 1e-9);
518 assert!((c[1].volume - 2000.0).abs() < 1e-9);
519 }
520
521 #[test]
522 fn binance_klines_reject_short_rows() {
523 assert!(parse_binance_klines(r#"[[1,"1","2","3"]]"#).is_err());
524 assert!(parse_binance_klines("not json").is_err());
525 }
526
527 proptest::proptest! {
528 #[test]
532 fn loaders_never_panic(s in ".*") {
533 let _ = parse_csv(&s);
534 let _ = parse_jsonl(&s);
535 let _ = parse_json_array(&s);
536 }
537 }
538
539 fn rising_closes(prices: &[f64]) -> Vec<Candle> {
540 prices
541 .iter()
542 .zip(0_i64..)
543 .map(|(&p, i)| Candle {
544 time: i,
545 open: p,
546 high: p,
547 low: p,
548 close: p,
549 volume: 0.0,
550 })
551 .collect()
552 }
553
554 #[test]
555 fn renko_builds_up_bricks_on_a_rising_series() {
556 let candles = rising_closes(&[100.0, 101.0, 102.0, 103.0, 104.0, 105.0]);
558 let bricks = to_renko(&candles, 1.0).unwrap();
559 assert_eq!(bricks.len(), 5);
560 assert!((bricks[0].open - 100.0).abs() < 1e-9);
561 assert!((bricks[0].close - 101.0).abs() < 1e-9);
562 assert!(bricks[0].close > bricks[0].open);
564 assert!((bricks[0].high - 101.0).abs() < 1e-9);
565 assert!((bricks[0].low - 100.0).abs() < 1e-9);
566 assert_eq!(bricks[4].time, 4);
567 assert!((bricks[4].close - 105.0).abs() < 1e-9);
568 }
569
570 #[test]
571 fn renko_rejects_non_positive_box() {
572 assert!(to_renko(&rising_closes(&[100.0, 101.0]), 0.0).is_err());
573 }
574
575 #[test]
576 fn kagi_emits_segments_with_edge_prices() {
577 let candles = rising_closes(&[100.0, 103.0, 106.0, 109.0, 104.0, 99.0]);
579 let bars = to_kagi(&candles, 3.0).unwrap();
580 assert!(!bars.is_empty());
581 assert!(bars.iter().all(|b| b.high >= b.low));
583 }
584
585 #[test]
586 fn pnf_columns_open_and_close_on_the_box_edges() {
587 let candles = rising_closes(&[
590 100.0, 101.0, 102.0, 103.0, 104.0, 105.0, 106.0, 105.0, 104.0, 103.0, 100.0,
591 ]);
592 let cols = to_pnf(&candles, 1.0, 3).unwrap();
593 assert!(!cols.is_empty());
594 assert!(cols.iter().all(|c| c.high >= c.low));
597 let rising = cols.iter().find(|c| c.close > c.open).unwrap();
598 assert!((rising.open - rising.low).abs() < 1e-9);
599 assert!((rising.close - rising.high).abs() < 1e-9);
600 }
601
602 #[cfg(feature = "parquet")]
603 #[test]
604 fn parquet_round_trip() {
605 use std::sync::Arc;
606
607 use arrow_array::{Float64Array, Int64Array, RecordBatch};
608 use arrow_schema::{DataType, Field, Schema};
609 use parquet::arrow::ArrowWriter;
610
611 let schema = Arc::new(Schema::new(vec![
612 Field::new("time", DataType::Int64, false),
613 Field::new("open", DataType::Float64, false),
614 Field::new("high", DataType::Float64, false),
615 Field::new("low", DataType::Float64, false),
616 Field::new("close", DataType::Float64, false),
617 Field::new("volume", DataType::Float64, false),
618 ]));
619 let batch = RecordBatch::try_new(
620 schema.clone(),
621 vec![
622 Arc::new(Int64Array::from(vec![1_i64, 2])),
623 Arc::new(Float64Array::from(vec![10.0, 11.0])),
624 Arc::new(Float64Array::from(vec![12.0, 13.0])),
625 Arc::new(Float64Array::from(vec![9.0, 10.0])),
626 Arc::new(Float64Array::from(vec![11.0, 12.0])),
627 Arc::new(Float64Array::from(vec![100.0, 200.0])),
628 ],
629 )
630 .unwrap();
631
632 let path = std::env::temp_dir().join("wkbt_parquet_round_trip.parquet");
633 let file = std::fs::File::create(&path).unwrap();
634 let mut writer = ArrowWriter::try_new(file, schema, None).unwrap();
635 writer.write(&batch).unwrap();
636 writer.close().unwrap();
637
638 let candles = load_candles(&path).unwrap();
639 assert_eq!(candles.len(), 2);
640 assert_eq!(candles[0].time, 1);
641 assert!((candles[0].open - 10.0).abs() < 1e-9);
642 assert!((candles[1].close - 12.0).abs() < 1e-9);
643 assert!((candles[1].volume - 200.0).abs() < 1e-9);
644
645 std::fs::remove_file(&path).ok();
646 }
647
648 #[cfg(not(feature = "parquet"))]
649 #[test]
650 fn parquet_without_feature_errors() {
651 assert!(load_parquet(Path::new("x.parquet")).is_err());
652 }
653}