chore: prepare v1.1.0 release

Update version numbers across Rust, Python, and documentation files to 1.1.0. Enhance the .gitignore to include macOS dSYM files and plans directory. Introduce new dependencies in the Rust core library and update the README to reflect recent performance benchmarks and backtesting engine capabilities. Add new artifacts to the benchmarks manifest and improve documentation for the backtesting engine API.
This commit is contained in:
Pratik Bhadane
2026-03-30 12:45:52 +05:30
parent 2d776b6f90
commit 436954138f
174 changed files with 29297 additions and 10773 deletions
+17
View File
@@ -0,0 +1,17 @@
# Local development build configuration for ferro-ta.
#
# Enables target-cpu=native so the compiler can emit instructions for the
# host machine (AVX2, NEON, etc.). This primarily benefits release builds
# where LTO and codegen-units=1 are active (see Cargo.toml [profile.release]).
#
# Cargo config.toml does not support per-profile rustflags, so this applies
# to both debug and release profiles. The impact on debug builds is negligible.
#
# WASM targets are excluded so wasm-pack / wasm32-unknown-unknown builds
# are unaffected.
#
# CI may override RUSTFLAGS or use a separate .cargo/config.toml to produce
# portable binaries for distribution.
[target.'cfg(not(target_arch = "wasm32"))']
rustflags = ["-C", "target-cpu=native"]
+6
View File
@@ -7,6 +7,12 @@ wasm/target/
*.pyd
*.dll
# macOS dSYM debug symbols (generated by maturin develop)
*.dSYM/
#
/plans/
# Maturin / wheel build outputs
dist/
*.egg-info/
Generated
+4 -2
View File
@@ -207,7 +207,7 @@ checksum = "48c757948c5ede0e46177b7add2e67155f70e33c07fea8284df6576da70b3719"
[[package]]
name = "ferro_ta"
version = "1.0.6"
version = "1.2.0"
dependencies = [
"criterion",
"ferro_ta_core",
@@ -222,9 +222,11 @@ dependencies = [
[[package]]
name = "ferro_ta_core"
version = "1.0.6"
version = "1.2.0"
dependencies = [
"criterion",
"serde",
"serde_json",
"wide",
]
+2 -2
View File
@@ -5,7 +5,7 @@ resolver = "2"
[package]
name = "ferro_ta"
version = "1.0.6"
version = "1.2.0"
edition = "2021"
description = "Rust-powered Python technical analysis library with a TA-Lib-compatible API"
license = "MIT"
@@ -30,7 +30,7 @@ ndarray = "0.16"
rayon = "1.10"
log = "0.4"
pyo3-log = "0.12"
ferro_ta_core = { path = "crates/ferro_ta_core", version = "1.0.6" }
ferro_ta_core = { path = "crates/ferro_ta_core", version = "1.2.0", features = ["serde"] }
[dev-dependencies]
criterion = { version = "0.8", features = ["html_reports"] }
+4 -3
View File
@@ -31,9 +31,10 @@ The latest checked-in TA-Lib comparison artifact uses contiguous `float64`
arrays at 10k and 100k bars on an `Apple M3 Max`, `CPython 3.13.5`, and `Rust
1.91.1`.
- `ferro-ta` is ahead outside the tie band on 6 of 12 indicators at both 10k and 100k bars.
- Strong public wins in the latest 100k-bar artifact include `SMA` (`2.28x`), `BBANDS` (`2.34x`), `MFI` (`3.04x`), and `WMA` (`2.39x`).
- TA-Lib still wins or ties on parts of the suite, including `STOCH`, `ADX`, and some current `EMA` / `RSI` / `ATR` runs.
- `ferro-ta` achieves competitive parity with TA-Lib, winning on 7 of 12 tested indicators at 100k bars (5 of 12 at 10k bars).
- Strong performance wins at 100k bars include `MFI` (`3.25×`), `WMA` (`2.20×`), `BBANDS` (`1.97×`), and `SMA` (`1.93×`) vs TA-Lib.
- TA-Lib maintains performance advantages on `STOCH` and `ADX`; `EMA`, `ATR`, and `OBV` are statistical ties.
- Compared to pure-Python libraries like Tulipy, `ferro-ta` provides 150-350x speedups through Rust-optimized implementations.
See the benchmark methodology and artifacts:
@@ -0,0 +1,153 @@
{
"metadata": {
"suite": "backtest",
"runtime": {
"generated_at_utc": "2026-03-27T16:31:53.866252+00:00",
"python_version": "3.13.5",
"python_implementation": "CPython",
"python_executable": "/Users/pratikbhadane/Work/Projects/ferro-ta/.venv/bin/python3",
"platform": "macOS-26.3.1-arm64-arm-64bit-Mach-O",
"system": "Darwin",
"release": "25.3.0",
"machine": "arm64",
"processor": "arm",
"cpu_model": "Apple M3 Max",
"cpu_count_logical": 14,
"total_memory_bytes": 38654705664
},
"git": {
"commit": "2d776b6f908fd1a4f30a696972b7df5e5fe2ca00",
"dirty": true,
"branch": "main"
},
"build": {
"rustc": "rustc 1.93.1 (01f6ddf75 2026-02-11)\nbinary: rustc\ncommit-hash: 01f6ddf7588f42ae2d7eb0a2f21d44e8e96674cf\ncommit-date: 2026-02-11\nhost: aarch64-apple-darwin\nrelease: 1.93.1\nLLVM version: 21.1.8",
"cargo": "cargo 1.93.1 (083ac5135 2025-12-15)",
"cargo_release_profile": {
"lto": true,
"codegen-units": 1
},
"rustflags": null,
"cargo_build_rustflags": null,
"maturin_flags": null
},
"packages": {
"numpy": "2.2.6",
"ferro-ta": "1.0.6"
}
},
"results": {
"backtest_core_single": [
{
"n_bars": 10000,
"ferro_ta_ms": 0.024,
"ferro_ta_mbars_s": 415.9388,
"vectorbt_ms": 1.2843,
"speedup_vs_vectorbt": 53.4187
},
{
"n_bars": 100000,
"ferro_ta_ms": 0.1964,
"ferro_ta_mbars_s": 509.1209,
"vectorbt_ms": 3.047,
"speedup_vs_vectorbt": 15.5129
}
],
"backtest_ohlcv_core": [
{
"n_bars": 10000,
"ferro_ta_ms": 0.0573,
"ferro_ta_mbars_s": 174.4166
},
{
"n_bars": 100000,
"ferro_ta_ms": 0.7068,
"ferro_ta_mbars_s": 141.4927
}
],
"performance_metrics": [
{
"n_bars": 10000,
"ferro_ta_ms": 0.2182,
"numpy_partial_ms": 0.0496,
"speedup_vs_numpy": 0.2272,
"note": "numpy_partial only computes sharpe+max_dd (2/23 metrics)"
},
{
"n_bars": 100000,
"ferro_ta_ms": 3.0303,
"numpy_partial_ms": 0.351,
"speedup_vs_numpy": 0.1158,
"note": "numpy_partial only computes sharpe+max_dd (2/23 metrics)"
}
],
"multi_asset": [
{
"n_bars": 10000,
"n_assets": 50,
"parallel_ms": 2.4245,
"serial_ms": 4.1751,
"loop_ms": 2.0349,
"parallel_speedup_vs_loop": 0.8393,
"parallel_speedup_vs_serial": 1.722
},
{
"n_bars": 100000,
"n_assets": 50,
"parallel_ms": 24.0349,
"serial_ms": 47.9311,
"loop_ms": 24.7476,
"parallel_speedup_vs_loop": 1.0297,
"parallel_speedup_vs_serial": 1.9942
}
],
"monte_carlo": [
{
"n_bars": 10000,
"n_sims": 500,
"ferro_ta_ms": 3.862,
"numpy_loop_ms": 51.1589,
"speedup_vs_numpy": 13.2469
},
{
"n_bars": 100000,
"n_sims": 500,
"ferro_ta_ms": 26.0019,
"numpy_loop_ms": 310.582,
"speedup_vs_numpy": 11.9446
}
],
"engine_full_pipeline": [
{
"n_bars": 10000,
"ferro_ta_ms": 0.4402,
"description": "Full pipeline: signals + OHLCV fill + 23 metrics + trades + drawdown"
},
{
"n_bars": 100000,
"ferro_ta_ms": 4.445,
"description": "Full pipeline: signals + OHLCV fill + 23 metrics + trades + drawdown"
}
],
"walk_forward_indices": [
{
"n_bars": 10000,
"train_bars": 2000,
"test_bars": 500,
"ferro_ta_us": 0.333
},
{
"n_bars": 100000,
"train_bars": 20000,
"test_bars": 5000,
"ferro_ta_us": 0.292
}
],
"kelly_fraction": [
{
"n_calls": 1000,
"ferro_ta_us": 86.458
}
]
}
}
@@ -57,6 +57,11 @@
"path": "benchmarks/artifacts/latest/wasm.json",
"size_bytes": 935,
"sha256": "f31fd871990c44e24a2259d618ae40a52866d20b95aa6047af3d38b9371c2ab7"
},
"bench_backtest": {
"path": "benchmarks/artifacts/latest/bench_backtest_results.json",
"size_bytes": 4022,
"sha256": "acf27cd5d5077aff51194e31936aba2b9304a8a62d993b2ec496d6f347545316"
}
}
}
+425
View File
@@ -0,0 +1,425 @@
"""
ferro_ta backtesting engine speed benchmark.
Measures throughput for single-asset, multi-asset, and analytics functions
across multiple bar sizes. Optional competitor comparison (vectorbt, backtrader)
is guarded behind try/except.
Usage:
python benchmarks/bench_backtest.py
python benchmarks/bench_backtest.py --sizes 10000 100000
python benchmarks/bench_backtest.py --skip-competitors --json benchmarks/artifacts/bench_backtest_results.json
"""
from __future__ import annotations
import argparse
import json
import time
from pathlib import Path
from typing import Any
import numpy as np
from ferro_ta._ferro_ta import (
backtest_core,
backtest_multi_asset_core,
backtest_ohlcv_core,
compute_performance_metrics,
kelly_fraction,
monte_carlo_bootstrap,
walk_forward_indices,
)
from ferro_ta.analysis.backtest import BacktestEngine
try:
from benchmarks.metadata import benchmark_metadata
except ModuleNotFoundError: # pragma: no cover
from metadata import benchmark_metadata # type: ignore[no-redef]
# Optional competitors -------------------------------------------------------
try:
import vectorbt as vbt # type: ignore[import]
VECTORBT_AVAILABLE = True
except ImportError:
VECTORBT_AVAILABLE = False
vbt = None # type: ignore[assignment]
try:
import backtrader as bt # type: ignore[import]
BACKTRADER_AVAILABLE = True
except ImportError:
BACKTRADER_AVAILABLE = False
bt = None # type: ignore[assignment]
# ---------------------------------------------------------------------------
N_WARMUP = 1
N_RUNS = 5
DEFAULT_SIZES = [10_000, 100_000, 1_000_000]
N_ASSETS = 50
N_SIMS = 500
# ---------------------------------------------------------------------------
# Timer helper
# ---------------------------------------------------------------------------
def _time_fn(
fn, *args, n_warmup: int = N_WARMUP, n_runs: int = N_RUNS, **kwargs
) -> float:
for _ in range(n_warmup):
fn(*args, **kwargs)
times: list[float] = []
for _ in range(n_runs):
t0 = time.perf_counter()
fn(*args, **kwargs)
times.append(time.perf_counter() - t0)
return float(np.median(times))
# ---------------------------------------------------------------------------
# Data generators
# ---------------------------------------------------------------------------
def _make_ohlcv(n: int, seed: int = 0) -> tuple[np.ndarray, ...]:
rng = np.random.default_rng(seed)
close = np.cumprod(1 + rng.standard_normal(n) * 0.01) * 100.0
high = close + rng.uniform(0.1, 1.5, n)
low = close - rng.uniform(0.1, 1.5, n)
open_ = close + rng.standard_normal(n) * 0.3
return open_, high, low, close
def _make_signals(n: int, seed: int = 1) -> np.ndarray:
rng = np.random.default_rng(seed)
raw = np.sign(rng.standard_normal(n))
raw[raw == 0] = 1.0
return raw.astype(np.float64)
# ---------------------------------------------------------------------------
# Benchmark functions
# ---------------------------------------------------------------------------
def bench_backtest_core_single(n: int) -> dict[str, Any]:
_, _, _, close = _make_ohlcv(n)
signals = _make_signals(n)
t_ferro = _time_fn(backtest_core, close, signals)
row: dict[str, Any] = {
"n_bars": n,
"ferro_ta_ms": round(t_ferro * 1000, 4),
"ferro_ta_mbars_s": round(n / t_ferro / 1e6, 4),
}
if VECTORBT_AVAILABLE:
import pandas as pd # noqa: PLC0415
close_s = pd.Series(close)
sig_s = pd.Series(signals.astype(bool))
def _vbt():
pf = vbt.Portfolio.from_signals(close_s, sig_s, ~sig_s, freq="1D")
return pf.total_return()
t_vbt = _time_fn(_vbt)
row["vectorbt_ms"] = round(t_vbt * 1000, 4)
row["speedup_vs_vectorbt"] = round(t_vbt / t_ferro, 4)
return row
def bench_backtest_ohlcv_core(n: int) -> dict[str, Any]:
open_, high, low, close = _make_ohlcv(n)
signals = _make_signals(n)
t_ferro = _time_fn(
backtest_ohlcv_core,
open_,
high,
low,
close,
signals,
fill_mode="market_open",
stop_loss_pct=0.02,
take_profit_pct=0.04,
)
return {
"n_bars": n,
"ferro_ta_ms": round(t_ferro * 1000, 4),
"ferro_ta_mbars_s": round(n / t_ferro / 1e6, 4),
}
def bench_performance_metrics(n: int) -> dict[str, Any]:
rng = np.random.default_rng(42)
returns = rng.standard_normal(n) * 0.01
equity = np.cumprod(1 + returns)
t_ferro = _time_fn(compute_performance_metrics, returns, equity)
def _numpy_sharpe():
mean_r = np.mean(returns)
std_r = np.std(returns, ddof=1)
_ = mean_r / std_r * np.sqrt(252)
rolling_max = np.maximum.accumulate(equity)
drawdown = (equity - rolling_max) / rolling_max
_ = float(drawdown.min())
t_numpy = _time_fn(_numpy_sharpe)
return {
"n_bars": n,
"ferro_ta_ms": round(t_ferro * 1000, 4),
"numpy_partial_ms": round(t_numpy * 1000, 4),
"speedup_vs_numpy": round(t_numpy / t_ferro, 4),
"note": "numpy_partial only computes sharpe+max_dd (2/23 metrics)",
}
def bench_multi_asset(n: int, n_assets: int = N_ASSETS) -> dict[str, Any]:
rng = np.random.default_rng(7)
close_2d = np.ascontiguousarray(
np.cumprod(1 + rng.standard_normal((n, n_assets)) * 0.01, axis=0) * 100.0
)
weights_2d = np.full((n, n_assets), 1.0 / n_assets)
t_parallel = _time_fn(
backtest_multi_asset_core, close_2d, weights_2d, parallel=True
)
t_serial = _time_fn(backtest_multi_asset_core, close_2d, weights_2d, parallel=False)
def _numpy_loop():
results = []
for j in range(n_assets):
col = np.ascontiguousarray(close_2d[:, j])
sig = np.ones(n)
_, _, sr, _ = backtest_core(col, sig)
results.append(sr)
return np.stack(results, axis=1)
t_loop = _time_fn(_numpy_loop)
return {
"n_bars": n,
"n_assets": n_assets,
"parallel_ms": round(t_parallel * 1000, 4),
"serial_ms": round(t_serial * 1000, 4),
"loop_ms": round(t_loop * 1000, 4),
"parallel_speedup_vs_loop": round(t_loop / t_parallel, 4),
"parallel_speedup_vs_serial": round(t_serial / t_parallel, 4),
}
def bench_monte_carlo(n: int, n_sims: int = N_SIMS) -> dict[str, Any]:
rng = np.random.default_rng(3)
returns = rng.standard_normal(n) * 0.01
t_ferro = _time_fn(monte_carlo_bootstrap, returns, n_sims=n_sims, seed=42)
def _numpy_mc():
out = np.empty((n_sims, n))
for i in range(n_sims):
idx = np.random.choice(len(returns), size=len(returns), replace=True)
out[i] = np.cumprod(1 + returns[idx])
return out
t_numpy = _time_fn(_numpy_mc)
return {
"n_bars": n,
"n_sims": n_sims,
"ferro_ta_ms": round(t_ferro * 1000, 4),
"numpy_loop_ms": round(t_numpy * 1000, 4),
"speedup_vs_numpy": round(t_numpy / t_ferro, 4),
}
def bench_engine_pipeline(n: int) -> dict[str, Any]:
_, high, low, open_ = _make_ohlcv(n)
_, _, _, close = _make_ohlcv(n, seed=10)
engine = (
BacktestEngine()
.with_commission(0.001)
.with_slippage(5.0)
.with_ohlcv(high=high, low=low, open_=open_)
.with_stop_loss(0.02)
.with_take_profit(0.04)
)
t_ferro = _time_fn(engine.run, close, "sma_crossover")
return {
"n_bars": n,
"ferro_ta_ms": round(t_ferro * 1000, 4),
"description": "Full pipeline: signals + OHLCV fill + 23 metrics + trades + drawdown",
}
def bench_walk_forward_indices(n: int) -> dict[str, Any]:
train = max(n // 5, 100)
test = max(n // 20, 20)
t = _time_fn(walk_forward_indices, n, train, test)
return {
"n_bars": n,
"train_bars": train,
"test_bars": test,
"ferro_ta_us": round(t * 1_000_000, 4),
}
def bench_kelly_fraction() -> dict[str, Any]:
win_rates = np.linspace(0.3, 0.7, 1000)
avg_wins = np.linspace(0.01, 0.05, 1000)
avg_losses = np.linspace(0.005, 0.03, 1000)
def _loop():
for w, a, b in zip(win_rates, avg_wins, avg_losses):
kelly_fraction(w, a, b)
t = _time_fn(_loop)
return {"n_calls": 1000, "ferro_ta_us": round(t * 1_000_000, 4)}
# ---------------------------------------------------------------------------
# Runner
# ---------------------------------------------------------------------------
def run_all(
sizes: list[int],
skip_competitors: bool,
n_assets: int,
n_sims: int,
) -> dict[str, Any]:
results: dict[str, list[dict[str, Any]]] = {
"backtest_core_single": [],
"backtest_ohlcv_core": [],
"performance_metrics": [],
"multi_asset": [],
"monte_carlo": [],
"engine_full_pipeline": [],
"walk_forward_indices": [],
}
for n in sizes:
print(f"\n--- {n:,} bars ---")
r = bench_backtest_core_single(n)
results["backtest_core_single"].append(r)
print(
f" backtest_core_single: {r['ferro_ta_ms']:.2f} ms ({r['ferro_ta_mbars_s']:.2f} M bars/s)"
)
r = bench_backtest_ohlcv_core(n)
results["backtest_ohlcv_core"].append(r)
print(
f" backtest_ohlcv_core: {r['ferro_ta_ms']:.2f} ms ({r['ferro_ta_mbars_s']:.2f} M bars/s)"
)
r = bench_performance_metrics(n)
results["performance_metrics"].append(r)
print(
f" performance_metrics: {r['ferro_ta_ms']:.2f} ms (numpy partial: {r['numpy_partial_ms']:.2f} ms, {r['speedup_vs_numpy']:.2f}x)"
)
r = bench_multi_asset(n, n_assets)
results["multi_asset"].append(r)
print(
f" multi_asset ({n_assets}): parallel={r['parallel_ms']:.1f} ms serial={r['serial_ms']:.1f} ms loop={r['loop_ms']:.1f} ms ({r['parallel_speedup_vs_loop']:.2f}x vs loop)"
)
r = bench_monte_carlo(n, n_sims)
results["monte_carlo"].append(r)
print(
f" monte_carlo ({n_sims} sims): {r['ferro_ta_ms']:.2f} ms (numpy: {r['numpy_loop_ms']:.2f} ms, {r['speedup_vs_numpy']:.2f}x)"
)
r = bench_engine_pipeline(n)
results["engine_full_pipeline"].append(r)
print(f" engine_full_pipeline: {r['ferro_ta_ms']:.2f} ms")
r = bench_walk_forward_indices(n)
results["walk_forward_indices"].append(r)
print(f" walk_forward_indices: {r['ferro_ta_us']:.1f} µs")
kelly_row = bench_kelly_fraction()
results["kelly_fraction"] = [kelly_row]
print(f"\n kelly_fraction (1k calls): {kelly_row['ferro_ta_us']:.1f} µs")
return {
"metadata": benchmark_metadata("backtest"),
"results": results,
}
# ---------------------------------------------------------------------------
# CLI
# ---------------------------------------------------------------------------
def main() -> int:
parser = argparse.ArgumentParser(
description="Benchmark ferro-ta backtesting engine."
)
parser.add_argument(
"--sizes",
type=int,
nargs="+",
default=DEFAULT_SIZES,
metavar="N",
help="Bar counts to benchmark (default: 10000 100000 1000000)",
)
parser.add_argument(
"--skip-competitors",
action="store_true",
help="Skip optional competitor benchmarks",
)
parser.add_argument(
"--assets",
type=int,
default=N_ASSETS,
help="Number of assets for multi-asset benchmark",
)
parser.add_argument(
"--sims",
type=int,
default=N_SIMS,
help="Number of simulations for Monte Carlo benchmark",
)
parser.add_argument(
"--json", dest="json_path", help="Write JSON results to this path"
)
args = parser.parse_args()
print(
f"ferro-ta backtest benchmark | sizes={args.sizes} | assets={args.assets} | sims={args.sims}"
)
print("=" * 72)
payload = run_all(
sizes=args.sizes,
skip_competitors=args.skip_competitors,
n_assets=args.assets,
n_sims=args.sims,
)
if args.json_path:
json_path = Path(args.json_path)
json_path.parent.mkdir(parents=True, exist_ok=True)
json_path.write_text(json.dumps(payload, indent=2), encoding="utf-8")
print(f"\nWrote JSON results to {json_path}")
return 0
if __name__ == "__main__":
raise SystemExit(main())
+1 -1
View File
@@ -1,5 +1,5 @@
{% set name = "ferro-ta" %}
{% set version = "1.0.6" %}
{% set version = "1.2.0" %}
package:
name: {{ name|lower }}
+4 -1
View File
@@ -1,6 +1,6 @@
[package]
name = "ferro_ta_core"
version = "1.0.6"
version = "1.2.0"
edition = "2021"
description = "Pure Rust core indicator library — no PyO3, no numpy dependency"
license = "MIT"
@@ -17,6 +17,8 @@ crate-type = ["lib"]
[dependencies]
wide = { version = "1.1.1", optional = true }
serde = { version = "1.0", features = ["derive"], optional = true }
serde_json = { version = "1.0", optional = true }
[dev-dependencies]
criterion = { version = "0.8", features = ["html_reports"] }
@@ -28,3 +30,4 @@ harness = false
[features]
wide = ["dep:wide"]
simd = ["wide"]
serde = ["dep:serde", "dep:serde_json"]
+1 -1
View File
@@ -13,7 +13,7 @@ PyO3, NumPy, or Python runtime dependency, which makes it a good fit for:
```toml
[dependencies]
ferro_ta_core = "1.0.6"
ferro_ta_core = "1.2.0"
```
## Design
+346
View File
@@ -0,0 +1,346 @@
//! Tick / Trade Aggregation Pipeline — pure Rust, no PyO3.
//!
//! Aggregates raw tick/trade data into OHLCV bars:
//! - **tick bars** — fixed number of ticks per bar
//! - **volume bars** — fixed volume threshold per bar
//! - **time bars** — label-based grouping (labels from Python timestamps)
/// OHLCV 5-tuple return type alias.
type Ohlcv5 = (Vec<f64>, Vec<f64>, Vec<f64>, Vec<f64>, Vec<f64>);
/// OHLCV 5-tuple plus labels return type alias.
type Ohlcv5AndLabels = (Vec<f64>, Vec<f64>, Vec<f64>, Vec<f64>, Vec<f64>, Vec<i64>);
// ---------------------------------------------------------------------------
// aggregate_tick_bars
// ---------------------------------------------------------------------------
/// Aggregate tick/trade data into tick bars (every N ticks become one bar).
///
/// Returns `(open, high, low, close, volume)` where volume = sum of sizes.
///
/// # Panics
/// Panics if `ticks_per_bar == 0`, arrays are empty, or lengths differ.
pub fn aggregate_tick_bars(
price: &[f64],
size: &[f64],
ticks_per_bar: usize,
) -> Ohlcv5 {
assert!(ticks_per_bar >= 1, "ticks_per_bar must be >= 1");
let n = price.len();
assert!(n > 0 && size.len() == n, "price and size must be non-empty and equal length");
let n_bars = n.div_ceil(ticks_per_bar);
let mut out_open = Vec::with_capacity(n_bars);
let mut out_high = Vec::with_capacity(n_bars);
let mut out_low = Vec::with_capacity(n_bars);
let mut out_close = Vec::with_capacity(n_bars);
let mut out_vol = Vec::with_capacity(n_bars);
let mut i = 0;
while i < n {
let end = (i + ticks_per_bar).min(n);
let bar_p = &price[i..end];
let bar_s = &size[i..end];
let bar_open = bar_p[0];
let bar_high = bar_p.iter().cloned().fold(f64::NEG_INFINITY, f64::max);
let bar_low = bar_p.iter().cloned().fold(f64::INFINITY, f64::min);
let bar_close = *bar_p.last().expect("slice cannot be empty");
let bar_vol: f64 = bar_s.iter().sum();
out_open.push(bar_open);
out_high.push(bar_high);
out_low.push(bar_low);
out_close.push(bar_close);
out_vol.push(bar_vol);
i = end;
}
(out_open, out_high, out_low, out_close, out_vol)
}
// ---------------------------------------------------------------------------
// aggregate_volume_bars_ticks
// ---------------------------------------------------------------------------
/// Aggregate tick data into volume bars (fixed volume threshold).
///
/// Accumulates ticks until cumulative size >= `volume_threshold`, then emits
/// a bar. Any remaining partial bar is also emitted.
///
/// Returns `(open, high, low, close, volume)`.
///
/// # Panics
/// Panics if `volume_threshold <= 0`, arrays are empty, or lengths differ.
pub fn aggregate_volume_bars_ticks(
price: &[f64],
size: &[f64],
volume_threshold: f64,
) -> Ohlcv5 {
assert!(volume_threshold > 0.0, "volume_threshold must be > 0");
let n = price.len();
assert!(n > 0 && size.len() == n, "price and size must be non-empty and equal length");
let mut out_open: Vec<f64> = Vec::new();
let mut out_high: Vec<f64> = Vec::new();
let mut out_low: Vec<f64> = Vec::new();
let mut out_close: Vec<f64> = Vec::new();
let mut out_vol: Vec<f64> = Vec::new();
let mut bar_open = price[0];
let mut bar_high = price[0];
let mut bar_low = price[0];
let mut bar_close = price[0];
let mut bar_vol = size[0];
for i in 1..n {
bar_high = bar_high.max(price[i]);
bar_low = bar_low.min(price[i]);
bar_close = price[i];
bar_vol += size[i];
if bar_vol >= volume_threshold {
out_open.push(bar_open);
out_high.push(bar_high);
out_low.push(bar_low);
out_close.push(bar_close);
out_vol.push(bar_vol);
if i + 1 < n {
bar_open = price[i + 1];
bar_high = price[i + 1];
bar_low = price[i + 1];
bar_close = price[i + 1];
bar_vol = size[i + 1];
} else {
bar_vol = 0.0;
}
}
}
// Push remaining partial bar
if bar_vol > 0.0 {
out_open.push(bar_open);
out_high.push(bar_high);
out_low.push(bar_low);
out_close.push(bar_close);
out_vol.push(bar_vol);
}
(out_open, out_high, out_low, out_close, out_vol)
}
// ---------------------------------------------------------------------------
// aggregate_time_bars
// ---------------------------------------------------------------------------
/// Aggregate tick data into time bars using pre-computed integer bucket labels.
///
/// Each tick is assigned a `label` (e.g. unix_ts // period_secs). Ticks with
/// the same label are accumulated into one bar. Labels must be non-decreasing.
///
/// Returns `(open, high, low, close, volume, unique_labels)`.
///
/// # Panics
/// Panics if arrays are empty or have unequal lengths.
pub fn aggregate_time_bars(
price: &[f64],
size: &[f64],
labels: &[i64],
) -> Ohlcv5AndLabels {
let n = price.len();
assert!(
n > 0 && size.len() == n && labels.len() == n,
"price, size, and labels must be non-empty and equal length"
);
let mut out_open: Vec<f64> = Vec::new();
let mut out_high: Vec<f64> = Vec::new();
let mut out_low: Vec<f64> = Vec::new();
let mut out_close: Vec<f64> = Vec::new();
let mut out_vol: Vec<f64> = Vec::new();
let mut out_labels: Vec<i64> = Vec::new();
let mut cur_label = labels[0];
let mut bar_open = price[0];
let mut bar_high = price[0];
let mut bar_low = price[0];
let mut bar_close = price[0];
let mut bar_vol = size[0];
for i in 1..n {
if labels[i] != cur_label {
out_open.push(bar_open);
out_high.push(bar_high);
out_low.push(bar_low);
out_close.push(bar_close);
out_vol.push(bar_vol);
out_labels.push(cur_label);
cur_label = labels[i];
bar_open = price[i];
bar_high = price[i];
bar_low = price[i];
bar_close = price[i];
bar_vol = size[i];
} else {
bar_high = bar_high.max(price[i]);
bar_low = bar_low.min(price[i]);
bar_close = price[i];
bar_vol += size[i];
}
}
out_open.push(bar_open);
out_high.push(bar_high);
out_low.push(bar_low);
out_close.push(bar_close);
out_vol.push(bar_vol);
out_labels.push(cur_label);
(out_open, out_high, out_low, out_close, out_vol, out_labels)
}
// ---------------------------------------------------------------------------
// Tests
// ---------------------------------------------------------------------------
#[cfg(test)]
mod tests {
use super::*;
// -- aggregate_tick_bars -------------------------------------------------
#[test]
fn test_tick_bars_exact_division() {
let price = [10.0, 11.0, 12.0, 13.0, 14.0, 15.0];
let size = [1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
let (o, h, l, c, v) = aggregate_tick_bars(&price, &size, 3);
assert_eq!(o.len(), 2);
// Bar 0: ticks 0..3
assert!((o[0] - 10.0).abs() < 1e-10);
assert!((h[0] - 12.0).abs() < 1e-10);
assert!((l[0] - 10.0).abs() < 1e-10);
assert!((c[0] - 12.0).abs() < 1e-10);
assert!((v[0] - 6.0).abs() < 1e-10);
// Bar 1: ticks 3..6
assert!((o[1] - 13.0).abs() < 1e-10);
assert!((h[1] - 15.0).abs() < 1e-10);
assert!((l[1] - 13.0).abs() < 1e-10);
assert!((c[1] - 15.0).abs() < 1e-10);
assert!((v[1] - 15.0).abs() < 1e-10);
}
#[test]
fn test_tick_bars_partial_last_bar() {
let price = [10.0, 11.0, 12.0, 13.0, 14.0];
let size = [1.0, 2.0, 3.0, 4.0, 5.0];
let (o, _h, _l, c, v) = aggregate_tick_bars(&price, &size, 3);
assert_eq!(o.len(), 2);
// Partial bar: ticks 3..5
assert!((o[1] - 13.0).abs() < 1e-10);
assert!((c[1] - 14.0).abs() < 1e-10);
assert!((v[1] - 9.0).abs() < 1e-10);
}
#[test]
fn test_tick_bars_single_tick() {
let (o, h, l, c, v) = aggregate_tick_bars(&[42.0], &[100.0], 5);
assert_eq!(o.len(), 1);
assert!((o[0] - 42.0).abs() < 1e-10);
assert!((h[0] - 42.0).abs() < 1e-10);
assert!((l[0] - 42.0).abs() < 1e-10);
assert!((c[0] - 42.0).abs() < 1e-10);
assert!((v[0] - 100.0).abs() < 1e-10);
}
#[test]
#[should_panic(expected = "ticks_per_bar must be >= 1")]
fn test_tick_bars_zero_ticks() {
aggregate_tick_bars(&[1.0], &[1.0], 0);
}
// -- aggregate_volume_bars_ticks -----------------------------------------
#[test]
fn test_volume_bars_ticks_basic() {
let price = [10.0, 11.0, 12.0, 13.0, 14.0];
let size = [30.0, 40.0, 50.0, 20.0, 60.0];
// threshold=70: bar0 = ticks 0+1 (vol=70), bar1 = tick2 (vol=50) + tick3 (vol=70),
// then tick4 as partial
let (o, h, l, c, v) = aggregate_volume_bars_ticks(&price, &size, 70.0);
// First bar: 30+40=70 >= 70
assert!((o[0] - 10.0).abs() < 1e-10);
assert!((c[0] - 11.0).abs() < 1e-10);
assert!((v[0] - 70.0).abs() < 1e-10);
assert!((h[0] - 11.0).abs() < 1e-10);
assert!((l[0] - 10.0).abs() < 1e-10);
assert!(v.len() >= 2);
}
#[test]
fn test_volume_bars_ticks_single() {
let (o, _h, _l, _c, v) = aggregate_volume_bars_ticks(&[5.0], &[10.0], 100.0);
assert_eq!(o.len(), 1);
assert!((v[0] - 10.0).abs() < 1e-10);
}
#[test]
#[should_panic(expected = "volume_threshold must be > 0")]
fn test_volume_bars_ticks_zero_threshold() {
aggregate_volume_bars_ticks(&[1.0], &[1.0], 0.0);
}
// -- aggregate_time_bars -------------------------------------------------
#[test]
fn test_time_bars_basic() {
let price = [10.0, 11.0, 12.0, 13.0, 14.0];
let size = [1.0, 2.0, 3.0, 4.0, 5.0];
let labels: [i64; 5] = [0, 0, 1, 1, 1];
let (o, h, l, c, v, out_lbl) = aggregate_time_bars(&price, &size, &labels);
assert_eq!(o.len(), 2);
assert_eq!(out_lbl, vec![0, 1]);
// Group 0: ticks 0,1
assert!((o[0] - 10.0).abs() < 1e-10);
assert!((h[0] - 11.0).abs() < 1e-10);
assert!((l[0] - 10.0).abs() < 1e-10);
assert!((c[0] - 11.0).abs() < 1e-10);
assert!((v[0] - 3.0).abs() < 1e-10);
// Group 1: ticks 2,3,4
assert!((o[1] - 12.0).abs() < 1e-10);
assert!((h[1] - 14.0).abs() < 1e-10);
assert!((l[1] - 12.0).abs() < 1e-10);
assert!((c[1] - 14.0).abs() < 1e-10);
assert!((v[1] - 12.0).abs() < 1e-10);
}
#[test]
fn test_time_bars_all_same_label() {
let price = [5.0, 6.0, 4.0];
let size = [10.0, 20.0, 30.0];
let labels: [i64; 3] = [42, 42, 42];
let (o, h, l, c, v, out_lbl) = aggregate_time_bars(&price, &size, &labels);
assert_eq!(o.len(), 1);
assert_eq!(out_lbl, vec![42]);
assert!((o[0] - 5.0).abs() < 1e-10);
assert!((h[0] - 6.0).abs() < 1e-10);
assert!((l[0] - 4.0).abs() < 1e-10);
assert!((c[0] - 4.0).abs() < 1e-10);
assert!((v[0] - 60.0).abs() < 1e-10);
}
#[test]
fn test_time_bars_each_tick_own_label() {
let price = [10.0, 20.0, 30.0];
let size = [1.0, 2.0, 3.0];
let labels: [i64; 3] = [0, 1, 2];
let (o, _h, _l, _c, v, out_lbl) = aggregate_time_bars(&price, &size, &labels);
assert_eq!(o.len(), 3);
assert_eq!(out_lbl, vec![0, 1, 2]);
assert!((v[0] - 1.0).abs() < 1e-10);
assert!((v[1] - 2.0).abs() < 1e-10);
assert!((v[2] - 3.0).abs() < 1e-10);
}
#[test]
#[should_panic(expected = "price, size, and labels must be non-empty and equal length")]
fn test_time_bars_empty() {
aggregate_time_bars(&[], &[], &[]);
}
}
+130
View File
@@ -0,0 +1,130 @@
//! Alerts — condition evaluation helpers.
//!
//! - `check_threshold` — fires when a series crosses above/below a level
//! - `check_cross` — fires when *fast* crosses above or below *slow*
//! - `collect_alert_bars` — returns indices of bars where a mask is non-zero
/// Fire an alert when `series` crosses a threshold level.
///
/// `direction`: `1` = cross above, `-1` = cross below.
///
/// Returns a `Vec<i8>` with `1` at crossing bars, `0` elsewhere.
/// Element 0 is always 0.
pub fn check_threshold(series: &[f64], level: f64, direction: i32) -> Vec<i8> {
let n = series.len();
let mut out = vec![0i8; n];
if n < 2 {
return out;
}
for i in 1..n {
let prev = series[i - 1];
let curr = series[i];
if prev.is_nan() || curr.is_nan() {
continue;
}
if direction == 1 {
if prev <= level && curr > level {
out[i] = 1;
}
} else if direction == -1 {
if prev >= level && curr < level {
out[i] = 1;
}
}
}
out
}
/// Detect cross-over / cross-under events between two series.
///
/// Returns `Vec<i8>`: `1` = bullish cross (fast above slow), `-1` = bearish, `0` = none.
/// Element 0 is always 0.
pub fn check_cross(fast: &[f64], slow: &[f64]) -> Vec<i8> {
let n = fast.len();
let mut out = vec![0i8; n];
if n < 2 {
return out;
}
for i in 1..n {
let fp = fast[i - 1];
let fc = fast[i];
let sp = slow[i - 1];
let sc = slow[i];
if fp.is_nan() || fc.is_nan() || sp.is_nan() || sc.is_nan() {
continue;
}
if fp <= sp && fc > sc {
out[i] = 1;
} else if fp >= sp && fc < sc {
out[i] = -1;
}
}
out
}
/// Collect bar indices where `mask` is non-zero.
pub fn collect_alert_bars(mask: &[i8]) -> Vec<i64> {
mask.iter()
.enumerate()
.filter(|(_, &v)| v != 0)
.map(|(i, _)| i as i64)
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_check_threshold_cross_above() {
let series = vec![10.0, 20.0, 30.0, 40.0, 50.0];
let result = check_threshold(&series, 25.0, 1);
assert_eq!(result, vec![0, 0, 1, 0, 0]);
}
#[test]
fn test_check_threshold_cross_below() {
let series = vec![50.0, 40.0, 30.0, 20.0, 10.0];
let result = check_threshold(&series, 25.0, -1);
assert_eq!(result, vec![0, 0, 0, 1, 0]);
}
#[test]
fn test_check_cross_bullish() {
let fast = vec![1.0, 2.0, 5.0];
let slow = vec![3.0, 3.0, 3.0];
let result = check_cross(&fast, &slow);
assert_eq!(result, vec![0, 0, 1]);
}
#[test]
fn test_check_cross_bearish() {
let fast = vec![5.0, 4.0, 1.0];
let slow = vec![3.0, 3.0, 3.0];
let result = check_cross(&fast, &slow);
assert_eq!(result, vec![0, 0, -1]);
}
#[test]
fn test_collect_alert_bars() {
let mask = vec![0i8, 1, 0, -1, 0, 1];
let result = collect_alert_bars(&mask);
assert_eq!(result, vec![1, 3, 5]);
}
#[test]
fn test_empty() {
assert_eq!(check_threshold(&[], 0.0, 1), Vec::<i8>::new());
assert_eq!(check_cross(&[], &[]), Vec::<i8>::new());
assert_eq!(collect_alert_bars(&[]), Vec::<i64>::new());
}
#[test]
fn test_nan_handling() {
let series = vec![10.0, f64::NAN, 30.0, 40.0];
let result = check_threshold(&series, 25.0, 1);
// NaN bars are skipped
assert_eq!(result[1], 0);
assert_eq!(result[2], 0); // prev is NaN
}
}
+329
View File
@@ -0,0 +1,329 @@
//! Performance attribution and trade analysis — pure Rust, no PyO3.
//!
//! Functions
//! ---------
//! - `trade_stats` — win rate, avg win/loss, profit factor, avg hold
//! - `monthly_contribution` — group bar returns by month index and sum
//! - `signal_attribution` — group bar returns by signal label and sum
//! - `extract_trades` — extract trade pnl and hold durations from positions
use std::collections::HashMap;
// ---------------------------------------------------------------------------
// trade_stats
// ---------------------------------------------------------------------------
/// Compute trade-level statistics from trade PnL and hold durations.
///
/// Returns `(win_rate, avg_win, avg_loss, profit_factor, avg_hold_bars)`.
///
/// - **win_rate** : fraction of trades with PnL > 0
/// - **avg_win** : mean PnL of winning trades (0 if none)
/// - **avg_loss** : mean PnL of losing trades (negative; 0 if none)
/// - **profit_factor** : gross profit / |gross loss| (inf if no losses)
/// - **avg_hold_bars** : mean hold duration across all trades
///
/// # Panics
/// Panics if `pnl` is empty or `pnl.len() != hold_bars.len()`.
pub fn trade_stats(pnl: &[f64], hold_bars: &[f64]) -> (f64, f64, f64, f64, f64) {
let n = pnl.len();
assert!(n > 0, "pnl must be non-empty");
assert_eq!(n, hold_bars.len(), "pnl and hold_bars must have equal length");
let mut wins: Vec<f64> = Vec::new();
let mut losses: Vec<f64> = Vec::new();
for &v in pnl.iter() {
if v > 0.0 {
wins.push(v);
} else if v < 0.0 {
losses.push(v);
}
}
let win_rate = wins.len() as f64 / n as f64;
let avg_win = if wins.is_empty() {
0.0
} else {
wins.iter().sum::<f64>() / wins.len() as f64
};
let avg_loss = if losses.is_empty() {
0.0
} else {
losses.iter().sum::<f64>() / losses.len() as f64
};
let gross_profit: f64 = wins.iter().sum();
let gross_loss: f64 = losses.iter().map(|v| v.abs()).sum();
let profit_factor = if gross_loss == 0.0 {
f64::INFINITY
} else {
gross_profit / gross_loss
};
let avg_hold = hold_bars.iter().sum::<f64>() / n as f64;
(win_rate, avg_win, avg_loss, profit_factor, avg_hold)
}
// ---------------------------------------------------------------------------
// monthly_contribution
// ---------------------------------------------------------------------------
/// Group per-bar returns by month index and sum each month's contribution.
///
/// Returns `(months, contributions)` where `months` is sorted unique month
/// indices and `contributions` is the corresponding total return per month.
/// NaN returns are skipped.
///
/// # Panics
/// Panics if `bar_returns.len() != month_index.len()`.
pub fn monthly_contribution(bar_returns: &[f64], month_index: &[i64]) -> (Vec<i64>, Vec<f64>) {
let n = bar_returns.len();
assert_eq!(
n,
month_index.len(),
"bar_returns and month_index must have equal length"
);
let mut map: HashMap<i64, f64> = HashMap::new();
for i in 0..n {
if !bar_returns[i].is_nan() {
*map.entry(month_index[i]).or_insert(0.0) += bar_returns[i];
}
}
let mut months: Vec<i64> = map.keys().copied().collect();
months.sort_unstable();
let contributions: Vec<f64> = months.iter().map(|m| map[m]).collect();
(months, contributions)
}
// ---------------------------------------------------------------------------
// signal_attribution
// ---------------------------------------------------------------------------
/// Attribute per-bar returns to each signal label.
///
/// Returns `(labels, contributions)` where `labels` is sorted unique signal
/// labels and `contributions` is the corresponding total return per label.
/// NaN returns are skipped.
///
/// # Panics
/// Panics if `bar_returns.len() != signal_labels.len()`.
pub fn signal_attribution(bar_returns: &[f64], signal_labels: &[i64]) -> (Vec<i64>, Vec<f64>) {
let n = bar_returns.len();
assert_eq!(
n,
signal_labels.len(),
"bar_returns and signal_labels must have equal length"
);
let mut map: HashMap<i64, f64> = HashMap::new();
for i in 0..n {
if !bar_returns[i].is_nan() {
*map.entry(signal_labels[i]).or_insert(0.0) += bar_returns[i];
}
}
let mut labels: Vec<i64> = map.keys().copied().collect();
labels.sort_unstable();
let contributions: Vec<f64> = labels.iter().map(|l| map[l]).collect();
(labels, contributions)
}
// ---------------------------------------------------------------------------
// extract_trades
// ---------------------------------------------------------------------------
/// Extract trade-level PnL and hold durations from positions and strategy returns.
///
/// A trade is a maximal contiguous run of non-zero position values with the
/// same sign/magnitude. Returns `(pnl, hold_durations)`.
///
/// # Panics
/// Panics if `positions.len() != strategy_returns.len()`.
pub fn extract_trades(positions: &[f64], strategy_returns: &[f64]) -> (Vec<f64>, Vec<f64>) {
let n = positions.len();
assert_eq!(
n,
strategy_returns.len(),
"positions and strategy_returns must have equal length"
);
let mut pnl = Vec::<f64>::new();
let mut hold = Vec::<f64>::new();
let mut i = 0usize;
while i < n {
if positions[i] == 0.0 {
i += 1;
continue;
}
let mut j = i + 1;
while j < n && positions[j] == positions[i] {
j += 1;
}
let mut trade_pnl = 0.0_f64;
for v in strategy_returns.iter().take(j).skip(i) {
trade_pnl += *v;
}
pnl.push(trade_pnl);
hold.push((j - i) as f64);
i = j;
}
(pnl, hold)
}
// ---------------------------------------------------------------------------
// Tests
// ---------------------------------------------------------------------------
#[cfg(test)]
mod tests {
use super::*;
// -- trade_stats ---------------------------------------------------------
#[test]
fn test_trade_stats_basic() {
let pnl = [100.0, -50.0, 200.0, -30.0, 150.0];
let hold = [5.0, 3.0, 7.0, 2.0, 6.0];
let (wr, aw, al, pf, ah) = trade_stats(&pnl, &hold);
// 3 wins out of 5
assert!((wr - 0.6).abs() < 1e-10);
// avg win = (100+200+150)/3
assert!((aw - 150.0).abs() < 1e-10);
// avg loss = (-50 + -30)/2 = -40
assert!((al - (-40.0)).abs() < 1e-10);
// profit_factor = 450 / 80
assert!((pf - 5.625).abs() < 1e-10);
// avg hold = (5+3+7+2+6)/5 = 4.6
assert!((ah - 4.6).abs() < 1e-10);
}
#[test]
fn test_trade_stats_all_wins() {
let pnl = [10.0, 20.0];
let hold = [1.0, 2.0];
let (wr, _aw, al, pf, _ah) = trade_stats(&pnl, &hold);
assert!((wr - 1.0).abs() < 1e-10);
assert!((al - 0.0).abs() < 1e-10);
assert!(pf.is_infinite());
}
#[test]
fn test_trade_stats_all_losses() {
let pnl = [-10.0, -20.0];
let hold = [1.0, 2.0];
let (wr, aw, _al, pf, _ah) = trade_stats(&pnl, &hold);
assert!((wr - 0.0).abs() < 1e-10);
assert!((aw - 0.0).abs() < 1e-10);
assert!((pf - 0.0).abs() < 1e-10);
}
#[test]
#[should_panic(expected = "pnl must be non-empty")]
fn test_trade_stats_empty() {
trade_stats(&[], &[]);
}
// -- monthly_contribution ------------------------------------------------
#[test]
fn test_monthly_contribution_basic() {
let returns = [0.01, 0.02, -0.01, 0.03, -0.02];
let months = [0, 0, 1, 1, 2];
let (m, c) = monthly_contribution(&returns, &months);
assert_eq!(m, vec![0, 1, 2]);
assert!((c[0] - 0.03).abs() < 1e-10);
assert!((c[1] - 0.02).abs() < 1e-10);
assert!((c[2] - (-0.02)).abs() < 1e-10);
}
#[test]
fn test_monthly_contribution_nan_skipped() {
let returns = [0.01, f64::NAN, 0.03];
let months = [0, 0, 1];
let (m, c) = monthly_contribution(&returns, &months);
assert_eq!(m, vec![0, 1]);
assert!((c[0] - 0.01).abs() < 1e-10);
assert!((c[1] - 0.03).abs() < 1e-10);
}
#[test]
fn test_monthly_contribution_empty() {
let (m, c) = monthly_contribution(&[], &[]);
assert!(m.is_empty());
assert!(c.is_empty());
}
// -- signal_attribution --------------------------------------------------
#[test]
fn test_signal_attribution_basic() {
let returns = [0.05, -0.02, 0.03, 0.01];
let labels = [1, -1, 2, 1];
let (l, c) = signal_attribution(&returns, &labels);
assert_eq!(l, vec![-1, 1, 2]);
assert!((c[0] - (-0.02)).abs() < 1e-10);
assert!((c[1] - 0.06).abs() < 1e-10); // 0.05 + 0.01
assert!((c[2] - 0.03).abs() < 1e-10);
}
#[test]
fn test_signal_attribution_nan_skipped() {
let returns = [0.05, f64::NAN];
let labels = [1, 2];
let (l, c) = signal_attribution(&returns, &labels);
assert_eq!(l, vec![1]);
assert!((c[0] - 0.05).abs() < 1e-10);
}
// -- extract_trades ------------------------------------------------------
#[test]
fn test_extract_trades_basic() {
// positions: flat, long, long, flat, short, short
let positions = [0.0, 1.0, 1.0, 0.0, -1.0, -1.0];
let strat_ret = [0.0, 0.01, 0.02, 0.0, -0.01, 0.03];
let (pnl, hold) = extract_trades(&positions, &strat_ret);
assert_eq!(pnl.len(), 2);
assert_eq!(hold.len(), 2);
// First trade: bars 1..3 => 0.01 + 0.02 = 0.03
assert!((pnl[0] - 0.03).abs() < 1e-10);
assert!((hold[0] - 2.0).abs() < 1e-10);
// Second trade: bars 4..6 => -0.01 + 0.03 = 0.02
assert!((pnl[1] - 0.02).abs() < 1e-10);
assert!((hold[1] - 2.0).abs() < 1e-10);
}
#[test]
fn test_extract_trades_all_flat() {
let positions = [0.0, 0.0, 0.0];
let strat_ret = [0.01, 0.02, 0.03];
let (pnl, hold) = extract_trades(&positions, &strat_ret);
assert!(pnl.is_empty());
assert!(hold.is_empty());
}
#[test]
fn test_extract_trades_empty() {
let (pnl, hold) = extract_trades(&[], &[]);
assert!(pnl.is_empty());
assert!(hold.is_empty());
}
#[test]
fn test_extract_trades_single_bar_trade() {
let positions = [0.0, 1.0, 0.0];
let strat_ret = [0.0, 0.05, 0.0];
let (pnl, hold) = extract_trades(&positions, &strat_ret);
assert_eq!(pnl.len(), 1);
assert!((pnl[0] - 0.05).abs() < 1e-10);
assert!((hold[0] - 1.0).abs() < 1e-10);
}
}
File diff suppressed because it is too large Load Diff
+642
View File
@@ -0,0 +1,642 @@
//! Pure-Rust batch operations — apply indicators across multiple series
//! (columns) sequentially. The PyO3 wrapper can add Rayon parallelism on top.
//!
//! Input convention: `data[j]` is column *j* (one time-series). All columns
//! must have the same length.
use crate::{momentum, overlap, statistic, volatility};
// ---------------------------------------------------------------------------
// helpers
// ---------------------------------------------------------------------------
/// Validate that every column in `data` has the same length. Returns `Ok(n)`
/// where `n` is the common length, or `Err` with a message.
fn validate_columns(data: &[Vec<f64>]) -> Result<usize, String> {
if data.is_empty() {
return Ok(0);
}
let n = data[0].len();
for (idx, col) in data.iter().enumerate() {
if col.len() != n {
return Err(format!(
"column 0 has length {n}, but column {idx} has length {}",
col.len()
));
}
}
Ok(n)
}
fn validate_hlc_columns(
high: &[Vec<f64>],
low: &[Vec<f64>],
close: &[Vec<f64>],
) -> Result<(usize, usize), String> {
let n_series = high.len();
if low.len() != n_series || close.len() != n_series {
return Err(format!(
"high has {} columns, low has {}, close has {} — must be equal",
n_series,
low.len(),
close.len()
));
}
if n_series == 0 {
return Ok((0, 0));
}
let n = high[0].len();
for (idx, (h, (l, c))) in high
.iter()
.zip(low.iter().zip(close.iter()))
.enumerate()
{
if h.len() != n || l.len() != n || c.len() != n {
return Err(format!(
"column {idx}: high len={}, low len={}, close len={} — must all be {n}",
h.len(),
l.len(),
c.len()
));
}
}
Ok((n, n_series))
}
// ---------------------------------------------------------------------------
// rolling linear regression (self-contained so core has no PyO3 dep)
// ---------------------------------------------------------------------------
fn linreg(window: &[f64]) -> (f64, f64) {
let n = window.len() as f64;
let sum_x: f64 = (0..window.len()).map(|i| i as f64).sum();
let sum_y: f64 = window.iter().sum();
let sum_xy: f64 = window.iter().enumerate().map(|(i, &y)| i as f64 * y).sum();
let sum_x2: f64 = (0..window.len()).map(|i| (i as f64).powi(2)).sum();
let denom = n * sum_x2 - sum_x * sum_x;
let slope = if denom != 0.0 {
(n * sum_xy - sum_x * sum_y) / denom
} else {
0.0
};
let intercept = (sum_y - slope * sum_x) / n;
(slope, intercept)
}
fn rolling_linreg_apply<F>(prices: &[f64], timeperiod: usize, mut map: F) -> Vec<f64>
where
F: FnMut(f64, f64) -> f64,
{
let n = prices.len();
let mut result = vec![f64::NAN; n];
if timeperiod == 0 || n < timeperiod {
return result;
}
if prices.iter().any(|value| !value.is_finite()) {
for end in (timeperiod - 1)..n {
let window = &prices[(end + 1 - timeperiod)..=end];
let (slope, intercept) = linreg(window);
result[end] = map(slope, intercept);
}
return result;
}
let period = timeperiod as f64;
let last_x = (timeperiod - 1) as f64;
let sum_x = last_x * period / 2.0;
let sum_x2 = last_x * period * (2.0 * period - 1.0) / 6.0;
let denom = period * sum_x2 - sum_x * sum_x;
let mut sum_y = prices[..timeperiod].iter().sum::<f64>();
let mut sum_xy = prices[..timeperiod]
.iter()
.enumerate()
.map(|(idx, &value)| idx as f64 * value)
.sum::<f64>();
for end in (timeperiod - 1)..n {
let slope = if denom != 0.0 {
(period * sum_xy - sum_x * sum_y) / denom
} else {
0.0
};
let intercept = (sum_y - slope * sum_x) / period;
result[end] = map(slope, intercept);
if end + 1 < n {
let outgoing = prices[end + 1 - timeperiod];
let incoming = prices[end + 1];
let prev_sum_y = sum_y;
sum_y = prev_sum_y - outgoing + incoming;
sum_xy = sum_xy - (prev_sum_y - outgoing) + last_x * incoming;
}
}
result
}
// ---------------------------------------------------------------------------
// CCI / WILLR helpers (no external dep)
// ---------------------------------------------------------------------------
fn compute_cci(high: &[f64], low: &[f64], close: &[f64], timeperiod: usize) -> Vec<f64> {
let n = high.len();
let typical_price: Vec<f64> = high
.iter()
.zip(low.iter())
.zip(close.iter())
.map(|((&h, &l), &c)| (h + l + c) / 3.0)
.collect();
let mut result = vec![f64::NAN; n];
if timeperiod == 0 || n < timeperiod {
return result;
}
for end in (timeperiod - 1)..n {
let window = &typical_price[(end + 1 - timeperiod)..=end];
let mean = window.iter().sum::<f64>() / timeperiod as f64;
let mad = window
.iter()
.map(|&value| (value - mean).abs())
.sum::<f64>()
/ timeperiod as f64;
result[end] = if mad != 0.0 {
(typical_price[end] - mean) / (0.015 * mad)
} else {
0.0
};
}
result
}
fn compute_willr(high: &[f64], low: &[f64], close: &[f64], timeperiod: usize) -> Vec<f64> {
let n = high.len();
let mut result = vec![f64::NAN; n];
if timeperiod == 0 || n < timeperiod {
return result;
}
// Use simple sliding-window max/min
for end in (timeperiod - 1)..n {
let start = end + 1 - timeperiod;
let mut highest = f64::NEG_INFINITY;
let mut lowest = f64::INFINITY;
for i in start..=end {
if high[i] > highest {
highest = high[i];
}
if low[i] < lowest {
lowest = low[i];
}
}
let range = highest - lowest;
result[end] = if range != 0.0 {
-100.0 * (highest - close[end]) / range
} else {
-50.0
};
}
result
}
// ---------------------------------------------------------------------------
// batch_sma
// ---------------------------------------------------------------------------
/// Apply SMA to each column. Returns one output column per input column.
pub fn batch_sma(data: &[Vec<f64>], timeperiod: usize) -> Result<Vec<Vec<f64>>, String> {
if timeperiod == 0 {
return Err("timeperiod must be >= 1".into());
}
validate_columns(data)?;
Ok(data
.iter()
.map(|col| overlap::sma(col, timeperiod))
.collect())
}
// ---------------------------------------------------------------------------
// batch_ema
// ---------------------------------------------------------------------------
/// Apply EMA to each column.
pub fn batch_ema(data: &[Vec<f64>], timeperiod: usize) -> Result<Vec<Vec<f64>>, String> {
if timeperiod == 0 {
return Err("timeperiod must be >= 1".into());
}
validate_columns(data)?;
Ok(data
.iter()
.map(|col| overlap::ema(col, timeperiod))
.collect())
}
// ---------------------------------------------------------------------------
// batch_rsi
// ---------------------------------------------------------------------------
/// Apply RSI to each column.
pub fn batch_rsi(data: &[Vec<f64>], timeperiod: usize) -> Result<Vec<Vec<f64>>, String> {
if timeperiod == 0 {
return Err("timeperiod must be >= 1".into());
}
validate_columns(data)?;
Ok(data
.iter()
.map(|col| momentum::rsi(col, timeperiod))
.collect())
}
// ---------------------------------------------------------------------------
// batch_atr
// ---------------------------------------------------------------------------
/// Apply ATR to each set of (high, low, close) columns.
pub fn batch_atr(
high: &[Vec<f64>],
low: &[Vec<f64>],
close: &[Vec<f64>],
timeperiod: usize,
) -> Result<Vec<Vec<f64>>, String> {
if timeperiod == 0 {
return Err("timeperiod must be >= 1".into());
}
validate_hlc_columns(high, low, close)?;
Ok((0..high.len())
.map(|i| volatility::atr(&high[i], &low[i], &close[i], timeperiod))
.collect())
}
// ---------------------------------------------------------------------------
// batch_stoch
// ---------------------------------------------------------------------------
/// Apply Stochastic to each set of (high, low, close) columns.
/// Returns `(slowk_columns, slowd_columns)`.
pub fn batch_stoch(
high: &[Vec<f64>],
low: &[Vec<f64>],
close: &[Vec<f64>],
fastk_period: usize,
slowk_period: usize,
slowd_period: usize,
) -> Result<(Vec<Vec<f64>>, Vec<Vec<f64>>), String> {
validate_hlc_columns(high, low, close)?;
let mut all_k = Vec::with_capacity(high.len());
let mut all_d = Vec::with_capacity(high.len());
for i in 0..high.len() {
let (k, d) = momentum::stoch(
&high[i],
&low[i],
&close[i],
fastk_period,
slowk_period,
slowd_period,
);
all_k.push(k);
all_d.push(d);
}
Ok((all_k, all_d))
}
// ---------------------------------------------------------------------------
// batch_adx
// ---------------------------------------------------------------------------
/// Apply ADX to each set of (high, low, close) columns.
pub fn batch_adx(
high: &[Vec<f64>],
low: &[Vec<f64>],
close: &[Vec<f64>],
timeperiod: usize,
) -> Result<Vec<Vec<f64>>, String> {
if timeperiod == 0 {
return Err("timeperiod must be >= 1".into());
}
validate_hlc_columns(high, low, close)?;
Ok((0..high.len())
.map(|i| momentum::adx(&high[i], &low[i], &close[i], timeperiod))
.collect())
}
// ---------------------------------------------------------------------------
// run_close_indicators
// ---------------------------------------------------------------------------
fn validate_indicator_requests(names: &[String], timeperiods: &[usize]) -> Result<(), String> {
if names.len() != timeperiods.len() {
return Err(format!(
"names length ({}) must equal timeperiods length ({})",
names.len(),
timeperiods.len()
));
}
for (name, &tp) in names.iter().zip(timeperiods.iter()) {
if tp == 0 {
return Err(format!("{name}: timeperiod must be >= 1"));
}
}
Ok(())
}
fn compute_close_indicator(
name: &str,
close: &[f64],
timeperiod: usize,
) -> Result<Vec<f64>, String> {
match name {
"SMA" => Ok(overlap::sma(close, timeperiod)),
"EMA" => Ok(overlap::ema(close, timeperiod)),
"RSI" => Ok(momentum::rsi(close, timeperiod)),
"STDDEV" => Ok(statistic::stddev(close, timeperiod, 1.0)),
"VAR" => Ok(statistic::stddev(close, timeperiod, 1.0)
.into_iter()
.map(|v| if v.is_nan() { v } else { v * v })
.collect()),
"LINEARREG" => {
let last_x = (timeperiod - 1) as f64;
Ok(rolling_linreg_apply(close, timeperiod, |slope, intercept| {
intercept + slope * last_x
}))
}
"LINEARREG_SLOPE" => Ok(rolling_linreg_apply(close, timeperiod, |slope, _| slope)),
"LINEARREG_INTERCEPT" => {
Ok(rolling_linreg_apply(close, timeperiod, |_, intercept| {
intercept
}))
}
"LINEARREG_ANGLE" => Ok(rolling_linreg_apply(close, timeperiod, |slope, _| {
slope.atan() * 180.0 / std::f64::consts::PI
})),
"TSF" => {
let forecast_x = timeperiod as f64;
Ok(rolling_linreg_apply(close, timeperiod, |slope, intercept| {
intercept + slope * forecast_x
}))
}
_ => Err(format!(
"unsupported close indicator for grouped execution: {name}"
)),
}
}
/// Run multiple close-only indicators on the same series.
/// Returns `Vec<Result<Vec<f64>, String>>` — one result per (name, timeperiod) pair.
pub fn run_close_indicators(
close: &[f64],
names: &[String],
timeperiods: &[usize],
) -> Result<Vec<Vec<f64>>, String> {
validate_indicator_requests(names, timeperiods)?;
let mut results = Vec::with_capacity(names.len());
for (name, &tp) in names.iter().zip(timeperiods.iter()) {
results.push(compute_close_indicator(name, close, tp)?);
}
Ok(results)
}
// ---------------------------------------------------------------------------
// run_hlc_indicators
// ---------------------------------------------------------------------------
fn compute_hlc_indicator(
name: &str,
high: &[f64],
low: &[f64],
close: &[f64],
timeperiod: usize,
) -> Result<Vec<f64>, String> {
match name {
"ATR" => Ok(volatility::atr(high, low, close, timeperiod)),
"NATR" => {
let atr_vals = volatility::atr(high, low, close, timeperiod);
Ok(atr_vals
.into_iter()
.zip(close.iter())
.map(|(a, &c)| {
if a.is_nan() || c == 0.0 {
f64::NAN
} else {
(a / c) * 100.0
}
})
.collect())
}
"ADX" => Ok(momentum::adx(high, low, close, timeperiod)),
"ADXR" => Ok(momentum::adxr(high, low, close, timeperiod)),
"CCI" => Ok(compute_cci(high, low, close, timeperiod)),
"WILLR" => Ok(compute_willr(high, low, close, timeperiod)),
_ => Err(format!(
"unsupported HLC indicator for grouped execution: {name}"
)),
}
}
/// Run multiple HLC indicators on the same series.
pub fn run_hlc_indicators(
high: &[f64],
low: &[f64],
close: &[f64],
names: &[String],
timeperiods: &[usize],
) -> Result<Vec<Vec<f64>>, String> {
validate_indicator_requests(names, timeperiods)?;
if high.len() != low.len() || high.len() != close.len() {
return Err("high, low, and close must have equal length".into());
}
let mut results = Vec::with_capacity(names.len());
for (name, &tp) in names.iter().zip(timeperiods.iter()) {
results.push(compute_hlc_indicator(name, high, low, close, tp)?);
}
Ok(results)
}
// ---------------------------------------------------------------------------
// tests
// ---------------------------------------------------------------------------
#[cfg(test)]
mod tests {
use super::*;
fn close_data() -> Vec<f64> {
vec![
44.34, 44.09, 43.61, 44.33, 44.83, 45.10, 45.42, 45.84, 46.08, 45.89, 46.03, 45.61,
46.28, 46.28, 46.00, 46.03, 46.41, 46.22, 45.64,
]
}
fn hlc_data() -> (Vec<f64>, Vec<f64>, Vec<f64>) {
let close = close_data();
let high: Vec<f64> = close.iter().map(|c| c + 0.5).collect();
let low: Vec<f64> = close.iter().map(|c| c - 0.5).collect();
(high, low, close)
}
#[test]
fn test_batch_sma_basic() {
let col1 = vec![1.0, 2.0, 3.0, 4.0, 5.0];
let col2 = vec![10.0, 20.0, 30.0, 40.0, 50.0];
let data = vec![col1, col2];
let result = batch_sma(&data, 3).unwrap();
assert_eq!(result.len(), 2);
assert!(result[0][0].is_nan());
assert!(result[0][1].is_nan());
assert!((result[0][2] - 2.0).abs() < 1e-10);
assert!((result[1][2] - 20.0).abs() < 1e-10);
}
#[test]
fn test_batch_sma_zero_period() {
let data = vec![vec![1.0, 2.0]];
assert!(batch_sma(&data, 0).is_err());
}
#[test]
fn test_batch_ema_basic() {
let data = vec![vec![1.0, 2.0, 3.0, 4.0, 5.0]];
let result = batch_ema(&data, 3).unwrap();
assert_eq!(result.len(), 1);
assert!(result[0][0].is_nan());
}
#[test]
fn test_batch_rsi_basic() {
let data = vec![close_data()];
let result = batch_rsi(&data, 14).unwrap();
assert_eq!(result.len(), 1);
// First 14 values should be NaN
for i in 0..14 {
assert!(result[0][i].is_nan(), "index {i} should be NaN");
}
// Value at index 14 should be a valid RSI
let rsi_val = result[0][14];
assert!(!rsi_val.is_nan());
assert!(rsi_val >= 0.0 && rsi_val <= 100.0);
}
#[test]
fn test_batch_atr_basic() {
let (h, l, c) = hlc_data();
let high = vec![h];
let low = vec![l];
let close = vec![c];
let result = batch_atr(&high, &low, &close, 14).unwrap();
assert_eq!(result.len(), 1);
}
#[test]
fn test_batch_stoch_basic() {
let (h, l, c) = hlc_data();
let high = vec![h];
let low = vec![l];
let close = vec![c];
let (k, d) = batch_stoch(&high, &low, &close, 5, 3, 3).unwrap();
assert_eq!(k.len(), 1);
assert_eq!(d.len(), 1);
assert_eq!(k[0].len(), d[0].len());
}
#[test]
fn test_batch_adx_basic() {
let (h, l, c) = hlc_data();
let high = vec![h];
let low = vec![l];
let close = vec![c];
let result = batch_adx(&high, &low, &close, 14).unwrap();
assert_eq!(result.len(), 1);
}
#[test]
fn test_run_close_indicators_basic() {
let close = close_data();
let names = vec!["SMA".to_string(), "EMA".to_string()];
let timeperiods = vec![5, 5];
let result = run_close_indicators(&close, &names, &timeperiods).unwrap();
assert_eq!(result.len(), 2);
assert_eq!(result[0].len(), close.len());
assert_eq!(result[1].len(), close.len());
}
#[test]
fn test_run_close_indicators_mismatched_lengths() {
let close = close_data();
let names = vec!["SMA".to_string()];
let timeperiods = vec![5, 10]; // different length
assert!(run_close_indicators(&close, &names, &timeperiods).is_err());
}
#[test]
fn test_run_close_indicators_linreg_variants() {
let close = close_data();
let names = vec![
"LINEARREG".to_string(),
"LINEARREG_SLOPE".to_string(),
"LINEARREG_INTERCEPT".to_string(),
"LINEARREG_ANGLE".to_string(),
"TSF".to_string(),
];
let timeperiods = vec![5, 5, 5, 5, 5];
let result = run_close_indicators(&close, &names, &timeperiods).unwrap();
assert_eq!(result.len(), 5);
// First 4 values should be NaN for period=5
for series in &result {
for i in 0..4 {
assert!(series[i].is_nan());
}
assert!(!series[4].is_nan());
}
}
#[test]
fn test_run_hlc_indicators_basic() {
let (h, l, c) = hlc_data();
let names = vec!["ATR".to_string(), "CCI".to_string()];
let timeperiods = vec![14, 14];
let result = run_hlc_indicators(&h, &l, &c, &names, &timeperiods).unwrap();
assert_eq!(result.len(), 2);
}
#[test]
fn test_run_hlc_indicators_unsupported() {
let (h, l, c) = hlc_data();
let names = vec!["UNKNOWN".to_string()];
let timeperiods = vec![14];
assert!(run_hlc_indicators(&h, &l, &c, &names, &timeperiods).is_err());
}
#[test]
fn test_validate_hlc_mismatched_columns() {
let high = vec![vec![1.0, 2.0]];
let low = vec![vec![1.0, 2.0], vec![3.0, 4.0]]; // 2 cols vs 1
let close = vec![vec![1.0, 2.0]];
assert!(batch_atr(&high, &low, &close, 5).is_err());
}
#[test]
fn test_empty_data() {
let data: Vec<Vec<f64>> = vec![];
let result = batch_sma(&data, 3).unwrap();
assert!(result.is_empty());
}
#[test]
fn test_batch_multiple_columns() {
let data = vec![
vec![1.0, 2.0, 3.0, 4.0, 5.0],
vec![5.0, 4.0, 3.0, 2.0, 1.0],
vec![2.0, 4.0, 6.0, 8.0, 10.0],
];
let result = batch_sma(&data, 3).unwrap();
assert_eq!(result.len(), 3);
// col 0: sma(3) at index 2 = (1+2+3)/3 = 2.0
assert!((result[0][2] - 2.0).abs() < 1e-10);
// col 1: sma(3) at index 2 = (5+4+3)/3 = 4.0
assert!((result[1][2] - 4.0).abs() < 1e-10);
// col 2: sma(3) at index 2 = (2+4+6)/3 = 4.0
assert!((result[2][2] - 4.0).abs() < 1e-10);
}
}
+123
View File
@@ -0,0 +1,123 @@
//! Chunked / out-of-core execution helpers.
//!
//! - `trim_overlap` — remove the first N elements from a slice
//! - `stitch_chunks` — concatenate trimmed chunk results
//! - `make_chunk_ranges` — compute (start, end) index pairs for chunked processing
//! - `forward_fill_nan` — forward-fill NaN values
/// Remove the first `overlap` elements from a slice.
pub fn trim_overlap(chunk_out: &[f64], overlap: usize) -> Vec<f64> {
if overlap > chunk_out.len() {
return vec![];
}
chunk_out[overlap..].to_vec()
}
/// Concatenate a list of slices into a single Vec.
pub fn stitch_chunks(chunks: &[&[f64]]) -> Vec<f64> {
let mut out = Vec::new();
for &chunk in chunks {
out.extend_from_slice(chunk);
}
out
}
/// Compute (start, end) index pairs for chunked processing.
///
/// Returns a flat Vec of pairs: [start0, end0, start1, end1, ...].
/// `chunk_size` is the desired output bars per chunk, `overlap` is the warm-up prefix.
pub fn make_chunk_ranges(n: usize, chunk_size: usize, overlap: usize) -> Vec<i64> {
if chunk_size == 0 || n == 0 {
return vec![];
}
let mut ranges: Vec<i64> = Vec::new();
let mut start: usize = 0;
loop {
let end = (start + chunk_size + overlap).min(n);
ranges.push(start as i64);
ranges.push(end as i64);
if end >= n {
break;
}
start = end.saturating_sub(overlap);
}
ranges
}
/// Forward-fill NaN values in a 1-D array.
/// Leading NaN values are preserved until the first non-NaN value appears.
pub fn forward_fill_nan(values: &[f64]) -> Vec<f64> {
let mut out = Vec::with_capacity(values.len());
let mut last = f64::NAN;
for &value in values {
if value.is_nan() {
out.push(last);
} else {
last = value;
out.push(value);
}
}
out
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_trim_overlap() {
let data = vec![1.0, 2.0, 3.0, 4.0, 5.0];
let result = trim_overlap(&data, 2);
assert_eq!(result, vec![3.0, 4.0, 5.0]);
}
#[test]
fn test_trim_overlap_zero() {
let data = vec![1.0, 2.0, 3.0];
assert_eq!(trim_overlap(&data, 0), data);
}
#[test]
fn test_trim_overlap_exceeds() {
let data = vec![1.0, 2.0];
assert!(trim_overlap(&data, 5).is_empty());
}
#[test]
fn test_stitch_chunks() {
let a = vec![1.0, 2.0];
let b = vec![3.0, 4.0, 5.0];
let chunks: Vec<&[f64]> = vec![&a, &b];
let result = stitch_chunks(&chunks);
assert_eq!(result, vec![1.0, 2.0, 3.0, 4.0, 5.0]);
}
#[test]
fn test_make_chunk_ranges() {
let ranges = make_chunk_ranges(10, 4, 2);
// Expected: [0,6], [4,10]
assert_eq!(ranges.len() % 2, 0);
assert!(ranges.len() >= 4);
assert_eq!(ranges[0], 0);
}
#[test]
fn test_forward_fill_nan() {
let data = vec![f64::NAN, 1.0, f64::NAN, f64::NAN, 2.0, f64::NAN];
let result = forward_fill_nan(&data);
assert!(result[0].is_nan()); // leading NaN preserved
assert!((result[1] - 1.0).abs() < 1e-10);
assert!((result[2] - 1.0).abs() < 1e-10); // filled
assert!((result[3] - 1.0).abs() < 1e-10); // filled
assert!((result[4] - 2.0).abs() < 1e-10);
assert!((result[5] - 2.0).abs() < 1e-10); // filled
}
#[test]
fn test_empty() {
assert!(trim_overlap(&[], 0).is_empty());
assert!(stitch_chunks(&[]).is_empty());
assert!(make_chunk_ranges(0, 4, 2).is_empty());
assert!(forward_fill_nan(&[]).is_empty());
}
}
+295
View File
@@ -0,0 +1,295 @@
//! Commission, tax, and fee model for Indian and global markets.
//!
//! All `_rate` fields are fractions (0.001 = 0.1%).
//! All per-unit fields (`flat_per_order`, `per_lot`) are in base currency units (e.g., INR).
//! The model is self-contained: pass `trade_value`, `num_lots`, `is_buy` to get total cost.
#[cfg(feature = "serde")]
use serde::{Deserialize, Serialize};
/// Advanced commission and tax model.
///
/// # Fields (all public for direct construction)
/// - **Brokerage**: `flat_per_order`, `rate_of_value`, `per_lot`, `max_brokerage`
/// - **STT**: `stt_rate`, `stt_on_buy`, `stt_on_sell`
/// - **Levies**: `exchange_charges_rate`, `regulatory_charges_rate`, `gst_rate`, `stamp_duty_rate`
/// - **Sizing**: `lot_size`
///
/// # Indian market notes
/// - STT (Securities Transaction Tax) is applied on turnover (buy/sell legs vary by segment).
/// - Exchange charges and regulatory body charges are on turnover.
/// - GST (18%) applies on brokerage + exchange charges + regulatory body charges (not STT/stamp).
/// - Stamp duty is on buy-side value only.
#[derive(Clone, Debug, PartialEq)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub struct CommissionModel {
// --- Brokerage ---------------------------------------------------------
/// Fixed fee per order (e.g., ₹20 flat fee per order). 0.0 = none.
pub flat_per_order: f64,
/// Proportional brokerage as fraction of `trade_value` (e.g., 0.001 = 0.1%). 0.0 = none.
pub rate_of_value: f64,
/// Fixed fee per lot (e.g., ₹2 per lot). 0.0 = none.
pub per_lot: f64,
/// Brokerage cap in currency units. 0.0 = no cap.
/// Effective brokerage = min(flat + rate × value + per_lot × lots, max_brokerage).
pub max_brokerage: f64,
/// Bid-ask spread model in basis points. Half-spread is paid on each leg (entry and exit),
/// so total roundtrip cost = spread_bps in bps. 0.0 = no spread cost.
pub spread_bps: f64,
// --- Securities Transaction Tax (STT) ----------------------------------
/// STT rate as fraction of trade value. 0.0 = no STT.
pub stt_rate: f64,
/// Apply STT on the buy leg.
pub stt_on_buy: bool,
/// Apply STT on the sell leg.
pub stt_on_sell: bool,
// --- Exchange & Regulatory Levies --------------------------------------
/// Exchange transaction charges rate (fraction of trade value).
pub exchange_charges_rate: f64,
/// Regulatory body turnover charges rate (fraction of trade value). Typically ~0.000001.
pub regulatory_charges_rate: f64,
/// Indirect tax (GST) rate applied on (brokerage + exchange_charges + regulatory_charges).
/// Typically 0.18 in India.
pub gst_rate: f64,
/// Stamp duty rate on buy side only (fraction of trade value).
pub stamp_duty_rate: f64,
// --- Instrument Sizing ------------------------------------------------
/// Lot size for the instrument.
/// Equities: 1.0. Index futures/options: contract lot size (e.g., 25, 50, 75).
/// Used for per_lot cost: cost += per_lot × ceil(quantity / lot_size).
pub lot_size: f64,
// --- Short Selling ----------------------------------------------------
/// Annualised short borrow rate as a fraction (e.g. 0.03 = 3% p.a.).
/// Applied per bar to short positions. 0.0 = no borrow cost.
pub short_borrow_rate_annual: f64,
}
impl Default for CommissionModel {
fn default() -> Self {
Self {
flat_per_order: 0.0,
rate_of_value: 0.0,
per_lot: 0.0,
max_brokerage: 0.0,
spread_bps: 0.0,
stt_rate: 0.0,
stt_on_buy: false,
stt_on_sell: false,
exchange_charges_rate: 0.0,
regulatory_charges_rate: 0.0,
gst_rate: 0.0,
stamp_duty_rate: 0.0,
lot_size: 1.0,
short_borrow_rate_annual: 0.0,
}
}
}
impl CommissionModel {
// ------------------------------------------------------------------
// Core computation
// ------------------------------------------------------------------
/// Compute total transaction cost in **absolute currency units**.
///
/// # Parameters
/// - `trade_value`: price × quantity in base currency
/// - `num_lots`: number of lots transacted
/// - `is_buy`: true for buy (entry) leg, false for sell (exit) leg
pub fn total_cost(&self, trade_value: f64, num_lots: f64, is_buy: bool) -> f64 {
// Brokerage (optionally capped)
let raw_brokerage =
self.flat_per_order + self.rate_of_value * trade_value + self.per_lot * num_lots;
let brokerage = if self.max_brokerage > 0.0 {
raw_brokerage.min(self.max_brokerage)
} else {
raw_brokerage
};
// STT
let stt = if (is_buy && self.stt_on_buy) || (!is_buy && self.stt_on_sell) {
self.stt_rate * trade_value
} else {
0.0
};
let exchange = self.exchange_charges_rate * trade_value;
let regulatory = self.regulatory_charges_rate * trade_value;
// GST on brokerage + exchange + regulatory (NOT on STT or stamp duty)
let gst = self.gst_rate * (brokerage + exchange + regulatory);
// Stamp duty only on buy side
let stamp = if is_buy {
self.stamp_duty_rate * trade_value
} else {
0.0
};
// Bid-ask spread: half-spread paid on each leg
let spread_cost = self.spread_bps / 2.0 / 10_000.0 * trade_value;
brokerage + stt + exchange + regulatory + gst + stamp + spread_cost
}
/// Borrow cost per bar for a short position.
///
/// # Parameters
/// - `trade_value`: abs(price × quantity)
/// - `periods_per_year`: 252 for daily, 52 for weekly, etc.
pub fn short_borrow_cost(&self, trade_value: f64, periods_per_year: f64) -> f64 {
if self.short_borrow_rate_annual <= 0.0 || periods_per_year <= 0.0 {
return 0.0;
}
self.short_borrow_rate_annual / periods_per_year * trade_value
}
/// Compute cost as a **fraction of `initial_capital`** for use in normalised equity loops.
///
/// Returns 0.0 if `initial_capital` ≤ 0.
pub fn cost_fraction(
&self,
trade_value: f64,
num_lots: f64,
is_buy: bool,
initial_capital: f64,
) -> f64 {
if initial_capital <= 0.0 {
return 0.0;
}
self.total_cost(trade_value, num_lots, is_buy) / initial_capital
}
// ------------------------------------------------------------------
// Built-in Presets
// ------------------------------------------------------------------
/// Zero commission — useful for clean research/comparison runs.
pub fn zero() -> Self {
Self::default()
}
/// Indian equity **delivery** (long-term hold).
///
/// Brokerage: 0.1% (capped at ₹20), STT 0.1% both sides,
/// exchange charges, regulatory body charges, 18% GST, stamp duty.
pub fn equity_delivery_india() -> Self {
Self {
flat_per_order: 0.0,
rate_of_value: 0.001, // 0.1%
per_lot: 0.0,
max_brokerage: 20.0, // ₹20 cap
spread_bps: 0.0,
stt_rate: 0.001, // 0.1%
stt_on_buy: true,
stt_on_sell: true,
exchange_charges_rate: 0.0000297,
regulatory_charges_rate: 0.000001,
gst_rate: 0.18,
stamp_duty_rate: 0.00015,
lot_size: 1.0,
short_borrow_rate_annual: 0.0,
}
}
/// Indian equity **intraday** (same-day square-off).
///
/// Brokerage: 0.03% (capped at ₹20), STT 0.025% sell side only,
/// exchange charges, regulatory body charges, 18% GST, stamp duty on buy.
pub fn equity_intraday_india() -> Self {
Self {
flat_per_order: 0.0,
rate_of_value: 0.0003, // 0.03%
per_lot: 0.0,
max_brokerage: 20.0,
spread_bps: 0.0,
stt_rate: 0.00025, // 0.025%
stt_on_buy: false,
stt_on_sell: true,
exchange_charges_rate: 0.0000297,
regulatory_charges_rate: 0.000001,
gst_rate: 0.18,
stamp_duty_rate: 0.000003,
lot_size: 1.0,
short_borrow_rate_annual: 0.0,
}
}
/// Indian **index futures** (indicative rates per current regulations).
///
/// Flat ₹20 per order, STT 0.05% sell side only, exchange charges,
/// regulatory body charges, 18% GST, stamp duty on buy.
/// `lot_size` defaults to 25 — update as needed for the specific contract.
pub fn futures_india() -> Self {
Self {
flat_per_order: 20.0,
rate_of_value: 0.0,
per_lot: 0.0,
max_brokerage: 0.0,
spread_bps: 0.0,
stt_rate: 0.0005, // 0.05%
stt_on_buy: false,
stt_on_sell: true,
exchange_charges_rate: 0.0000019,
regulatory_charges_rate: 0.000001,
gst_rate: 0.18,
stamp_duty_rate: 0.00002,
lot_size: 25.0,
short_borrow_rate_annual: 0.0,
}
}
/// Indian **index options** (indicative rates per current regulations).
///
/// Flat ₹20 per order, STT 0.15% on premium sell side only, exchange charges,
/// regulatory body charges, 18% GST, stamp duty on buy.
/// `lot_size` defaults to 25 — update as needed for the specific contract.
pub fn options_india() -> Self {
Self {
flat_per_order: 20.0,
rate_of_value: 0.0,
per_lot: 0.0,
max_brokerage: 0.0,
spread_bps: 0.0,
stt_rate: 0.0015, // 0.15% on premium
stt_on_buy: false,
stt_on_sell: true,
exchange_charges_rate: 0.0000053,
regulatory_charges_rate: 0.000001,
gst_rate: 0.18,
stamp_duty_rate: 0.000003,
lot_size: 25.0,
short_borrow_rate_annual: 0.0,
}
}
/// Simple proportional model — e.g., `proportional(0.001)` = 0.1% both sides.
///
/// No taxes, no levies — suitable for non-Indian markets or simplified modelling.
pub fn proportional(rate: f64) -> Self {
Self {
rate_of_value: rate,
..Default::default()
}
}
// ------------------------------------------------------------------
// JSON serialization (requires "serde" feature)
// ------------------------------------------------------------------
/// Serialize to a pretty-printed JSON string.
#[cfg(feature = "serde")]
pub fn to_json(&self) -> Result<String, serde_json::Error> {
serde_json::to_string_pretty(self)
}
/// Deserialize from a JSON string.
#[cfg(feature = "serde")]
pub fn from_json(s: &str) -> Result<Self, serde_json::Error> {
serde_json::from_str(s)
}
}
+91
View File
@@ -0,0 +1,91 @@
//! Crypto and 24/7 market helpers.
//!
//! - `funding_cumulative_pnl` — cumulative PnL from periodic funding rate payments
//! - `continuous_bar_labels` — assign sequential integer labels based on fixed period size
//! - `mark_session_boundaries` — return indices where a new UTC day begins
/// Compute the cumulative PnL from funding rate payments.
///
/// `position_size` and `funding_rate` must have the same length.
/// PnL at period i = -position_size[i] * funding_rate[i] (longs pay when rate > 0).
pub fn funding_cumulative_pnl(position_size: &[f64], funding_rate: &[f64]) -> Vec<f64> {
let n = position_size.len();
let mut out = vec![0.0_f64; n];
let mut cumulative = 0.0_f64;
for i in 0..n {
cumulative += -position_size[i] * funding_rate[i];
out[i] = cumulative;
}
out
}
/// Assign a sequential integer label per bar based on a fixed-size period.
///
/// Bars 0..(period_bars-1) get label 0, bars period_bars..(2*period_bars-1) get label 1, etc.
/// `period_bars` must be >= 1.
pub fn continuous_bar_labels(n_bars: usize, period_bars: usize) -> Vec<i64> {
(0..n_bars).map(|i| (i / period_bars) as i64).collect()
}
/// Return bar indices where a new UTC day begins (based on nanosecond timestamps).
///
/// Bar 0 is always included as the first boundary.
pub fn mark_session_boundaries(timestamps_ns: &[i64]) -> Vec<i64> {
let n = timestamps_ns.len();
if n == 0 {
return vec![];
}
const NS_PER_DAY: i64 = 86_400_000_000_000;
let mut out = vec![0i64]; // bar 0 is always a boundary
let mut prev_day = timestamps_ns[0].div_euclid(NS_PER_DAY);
for (i, &t) in timestamps_ns.iter().enumerate().skip(1) {
let day = t.div_euclid(NS_PER_DAY);
if day != prev_day {
out.push(i as i64);
prev_day = day;
}
}
out
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_funding_cumulative_pnl() {
let pos = vec![100.0, 100.0, -50.0];
let rate = vec![0.001, -0.002, 0.001];
let result = funding_cumulative_pnl(&pos, &rate);
assert!((result[0] - (-0.1)).abs() < 1e-10);
assert!((result[1] - 0.1).abs() < 1e-10); // -0.1 + 0.2 = 0.1
assert!((result[2] - 0.15).abs() < 1e-10); // 0.1 + 0.05 = 0.15
}
#[test]
fn test_continuous_bar_labels() {
let labels = continuous_bar_labels(7, 3);
assert_eq!(labels, vec![0, 0, 0, 1, 1, 1, 2]);
}
#[test]
fn test_mark_session_boundaries() {
let ns_per_day: i64 = 86_400_000_000_000;
let ts = vec![
0, // day 0
ns_per_day / 2, // day 0
ns_per_day, // day 1
ns_per_day + ns_per_day / 2, // day 1
ns_per_day * 2, // day 2
];
let result = mark_session_boundaries(&ts);
assert_eq!(result, vec![0, 2, 4]);
}
#[test]
fn test_empty() {
assert!(funding_cumulative_pnl(&[], &[]).is_empty());
assert!(continuous_bar_labels(0, 1).is_empty());
assert!(mark_session_boundaries(&[]).is_empty());
}
}
+173
View File
@@ -0,0 +1,173 @@
//! Currency metadata and Indian number formatting.
/// Immutable currency descriptor.
///
/// Carries the currency code, symbol, decimal places, and whether to use
/// Indian lakh/crore grouping (1,23,45,678.00) instead of standard
/// Western grouping (1,234,567.89).
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct Currency {
/// IETF currency code, e.g. "INR", "USD".
pub code: &'static str,
/// Display symbol, e.g. "₹", "$".
pub symbol: &'static str,
/// Number of decimal places for formatting.
pub decimal_places: u8,
/// Use Indian lakh/crore digit grouping (true only for INR).
pub lakh_grouping: bool,
}
impl Currency {
pub const INR: Currency = Currency {
code: "INR",
symbol: "",
decimal_places: 2,
lakh_grouping: true,
};
pub const USD: Currency = Currency {
code: "USD",
symbol: "$",
decimal_places: 2,
lakh_grouping: false,
};
pub const EUR: Currency = Currency {
code: "EUR",
symbol: "",
decimal_places: 2,
lakh_grouping: false,
};
pub const GBP: Currency = Currency {
code: "GBP",
symbol: "£",
decimal_places: 2,
lakh_grouping: false,
};
pub const JPY: Currency = Currency {
code: "JPY",
symbol: "¥",
decimal_places: 0,
lakh_grouping: false,
};
pub const USDT: Currency = Currency {
code: "USDT",
symbol: "",
decimal_places: 2,
lakh_grouping: false,
};
/// Look up a currency by IETF code (case-insensitive).
/// Returns `None` if the code is not recognised.
pub fn from_code(code: &str) -> Option<&'static Currency> {
match code.to_ascii_uppercase().as_str() {
"INR" => Some(&Currency::INR),
"USD" => Some(&Currency::USD),
"EUR" => Some(&Currency::EUR),
"GBP" => Some(&Currency::GBP),
"JPY" => Some(&Currency::JPY),
"USDT" => Some(&Currency::USDT),
_ => None,
}
}
/// Format `amount` according to this currency's style.
///
/// - INR uses Indian lakh/crore grouping: `₹1,23,45,678.00`
/// - Others use standard Western grouping: `$1,234,567.89`
pub fn format(&self, amount: f64) -> String {
let neg = amount < 0.0;
let abs = amount.abs();
let integer_part = abs.floor() as u64;
let frac_part = abs - abs.floor();
let grouped = if self.lakh_grouping {
format_lakh(integer_part)
} else {
format_standard(integer_part)
};
let dp = self.decimal_places as usize;
let decimal_str = if dp > 0 {
let frac = (frac_part * 10f64.powi(dp as i32)).round() as u64;
format!(".{:0>width$}", frac, width = dp)
} else {
String::new()
};
let sign = if neg { "-" } else { "" };
format!("{}{}{}{}", sign, self.symbol, grouped, decimal_str)
}
}
/// Indian lakh/crore grouping: last 3 digits, then groups of 2 from the right.
/// e.g. 12345678 → "1,23,45,678"
fn format_lakh(n: u64) -> String {
let s = n.to_string();
if s.len() <= 3 {
return s;
}
let (rest, last3) = s.split_at(s.len() - 3);
let mut out = String::new();
let chars: Vec<char> = rest.chars().collect();
let first_len = chars.len() % 2;
if first_len > 0 {
out.push_str(&chars[..first_len].iter().collect::<String>());
}
let mut i = first_len;
while i < chars.len() {
if !out.is_empty() {
out.push(',');
}
out.push_str(&chars[i..i + 2].iter().collect::<String>());
i += 2;
}
if !out.is_empty() {
out.push(',');
}
out.push_str(last3);
out
}
/// Standard Western grouping: groups of 3 digits from the right.
/// e.g. 1234567 → "1,234,567"
fn format_standard(n: u64) -> String {
let s = n.to_string();
let mut out = String::new();
for (i, c) in s.chars().rev().enumerate() {
if i > 0 && i % 3 == 0 {
out.push(',');
}
out.push(c);
}
out.chars().rev().collect()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_inr_format() {
assert_eq!(Currency::INR.format(123456.78), "₹1,23,456.78");
assert_eq!(Currency::INR.format(10000000.0), "₹1,00,00,000.00");
assert_eq!(Currency::INR.format(100.0), "₹100.00");
assert_eq!(Currency::INR.format(-5000.0), "-₹5,000.00");
}
#[test]
fn test_usd_format() {
assert_eq!(Currency::USD.format(1234567.89), "$1,234,567.89");
assert_eq!(Currency::USD.format(0.5), "$0.50");
}
#[test]
fn test_jpy_format() {
assert_eq!(Currency::JPY.format(1000000.0), "¥1,000,000");
}
#[test]
fn test_from_code() {
assert_eq!(Currency::from_code("inr"), Some(&Currency::INR));
assert_eq!(Currency::from_code("USD"), Some(&Currency::USD));
assert_eq!(Currency::from_code("UNKNOWN"), None);
}
}
+370
View File
@@ -0,0 +1,370 @@
//! Cycle indicators — Hilbert Transform-based cycle analysis (Ehlers).
//!
//! Based on John Ehlers' Discrete Hilbert Transform as implemented in TA-Lib.
//! Reference: "Cybernetic Analysis for Stocks and Futures" by J.F. Ehlers
//!
//! All HT functions share a 63-bar lookback period.
use std::f64::consts::PI;
/// Number of leading bars that are set to NaN / zero.
pub const HT_LOOKBACK: usize = 63;
/// Shared output from the core Hilbert Transform computation.
pub struct HtCore {
pub trendline: Vec<f64>,
pub dc_period: Vec<f64>,
pub dc_phase: Vec<f64>,
pub inphase: Vec<f64>,
pub quadrature: Vec<f64>,
pub trend_mode: Vec<i32>,
}
/// Run the full Hilbert Transform pipeline on a slice of close prices.
pub fn compute_ht_core(prices: &[f64]) -> HtCore {
let n = prices.len();
let mut trendline = vec![f64::NAN; n];
let mut dc_period = vec![f64::NAN; n];
let mut dc_phase = vec![f64::NAN; n];
let mut inphase = vec![f64::NAN; n];
let mut quadrature = vec![f64::NAN; n];
let mut trend_mode = vec![0i32; n];
if n <= HT_LOOKBACK {
return HtCore {
trendline,
dc_period,
dc_phase,
inphase,
quadrature,
trend_mode,
};
}
// Step 1: Smooth the price series (4-bar weighted average)
let mut smooth = vec![0.0f64; n];
for i in 0..n {
smooth[i] = if i >= 3 {
(4.0 * prices[i] + 3.0 * prices[i - 1] + 2.0 * prices[i - 2] + prices[i - 3]) / 10.0
} else {
prices[i]
};
}
// Step 2: Full Hilbert Transform pipeline
let mut detrender = vec![0.0f64; n];
let mut q1 = vec![0.0f64; n];
let mut i1 = vec![0.0f64; n];
let mut ji = vec![0.0f64; n];
let mut jq = vec![0.0f64; n];
let mut i2 = vec![0.0f64; n];
let mut q2 = vec![0.0f64; n];
let mut re = vec![0.0f64; n];
let mut im = vec![0.0f64; n];
let mut period = vec![0.0f64; n];
let mut smooth_period = vec![0.0f64; n];
let mut phase = vec![0.0f64; n];
for i in 6..n {
let prev_period = period[i - 1];
// Alpha coefficient for HT filters depends on the current period estimate
let alpha = 0.075 * prev_period + 0.54;
// Discrete Hilbert Transform of smooth price (detrender)
detrender[i] = (0.0962 * smooth[i] + 0.5769 * smooth[i - 2]
- 0.5769 * smooth[i - 4]
- 0.0962 * smooth[i - 6])
* alpha;
// Q1: HT of detrender
if i >= 12 {
q1[i] = (0.0962 * detrender[i] + 0.5769 * detrender[i - 2]
- 0.5769 * detrender[i - 4]
- 0.0962 * detrender[i - 6])
* alpha;
}
// I1: delayed detrender
if i >= 9 {
i1[i] = detrender[i - 3];
}
// jI: HT of I1
if i >= 15 {
ji[i] = (0.0962 * i1[i] + 0.5769 * i1[i - 2] - 0.5769 * i1[i - 4] - 0.0962 * i1[i - 6])
* alpha;
}
// jQ: HT of Q1
if i >= 18 {
jq[i] = (0.0962 * q1[i] + 0.5769 * q1[i - 2] - 0.5769 * q1[i - 4] - 0.0962 * q1[i - 6])
* alpha;
}
// Phase components
let i2_raw = i1[i] - jq[i];
let q2_raw = q1[i] + ji[i];
// EMA smoothing of I2 and Q2
let i2_prev = i2[i - 1];
let q2_prev = q2[i - 1];
i2[i] = 0.2 * i2_raw + 0.8 * i2_prev;
q2[i] = 0.2 * q2_raw + 0.8 * q2_prev;
// Cross-product for period estimation
let re_raw = i2[i] * i2_prev + q2[i] * q2_prev;
let im_raw = i2[i] * q2_prev - q2[i] * i2_prev;
// EMA smoothing of Re and Im
re[i] = 0.2 * re_raw + 0.8 * re[i - 1];
im[i] = 0.2 * im_raw + 0.8 * im[i - 1];
// Compute period from cross-product of consecutive phasors.
let mut p = if re[i] != 0.0 && im[i] != 0.0 && re[i] > 0.0 {
2.0 * PI / (im[i] / re[i]).atan()
} else {
prev_period
};
// Clamp period relative to previous
if prev_period > 0.0 {
if p > 1.5 * prev_period {
p = 1.5 * prev_period;
}
if p < 0.67 * prev_period {
p = 0.67 * prev_period;
}
}
// Hard clamp to [6, 50] bars
p = p.clamp(6.0, 50.0);
// EMA smooth the period
period[i] = 0.2 * p + 0.8 * prev_period;
// Smooth the smoothed period once more
smooth_period[i] = 0.33 * period[i] + 0.67 * smooth_period[i - 1];
// Phase from I1 and Q1
phase[i] = if i1[i] != 0.0 {
q1[i].atan2(i1[i]) * 180.0 / PI
} else if q1[i] > 0.0 {
90.0
} else if q1[i] < 0.0 {
-90.0
} else {
0.0
};
// Write outputs once past lookback
if i >= HT_LOOKBACK {
dc_period[i] = smooth_period[i];
dc_phase[i] = phase[i];
inphase[i] = i1[i];
quadrature[i] = q1[i];
// Trend mode: cycle when SmoothPeriod >= 20, trend when < 20
trend_mode[i] = if smooth_period[i] < 20.0 { 1 } else { 0 };
}
}
// Trendline: average over the current dominant cycle period
for i in HT_LOOKBACK..n {
let sp = smooth_period[i];
let dc = (sp.round() as usize).max(1).min(i + 1);
let sum: f64 = (0..dc).map(|j| smooth[i - j]).sum();
trendline[i] = sum / dc as f64;
}
HtCore {
trendline,
dc_period,
dc_phase,
inphase,
quadrature,
trend_mode,
}
}
// ---------------------------------------------------------------------------
// Public indicator functions
// ---------------------------------------------------------------------------
/// Hilbert Transform Instantaneous Trendline (Ehlers).
/// Smooths price over the dominant cycle period.
pub fn ht_trendline(close: &[f64]) -> Vec<f64> {
compute_ht_core(close).trendline
}
/// Hilbert Transform Dominant Cycle Period in bars.
pub fn ht_dcperiod(close: &[f64]) -> Vec<f64> {
compute_ht_core(close).dc_period
}
/// Hilbert Transform Dominant Cycle Phase in degrees.
pub fn ht_dcphase(close: &[f64]) -> Vec<f64> {
compute_ht_core(close).dc_phase
}
/// Hilbert Transform Phasor components. Returns `(inphase, quadrature)`.
pub fn ht_phasor(close: &[f64]) -> (Vec<f64>, Vec<f64>) {
let core = compute_ht_core(close);
(core.inphase, core.quadrature)
}
/// Hilbert Transform SineWave. Returns `(sine, leadsine)` where leadsine
/// leads sine by 45 degrees.
pub fn ht_sine(close: &[f64]) -> (Vec<f64>, Vec<f64>) {
let n = close.len();
let core = compute_ht_core(close);
let mut sine = vec![f64::NAN; n];
let mut lead_sine = vec![f64::NAN; n];
for i in HT_LOOKBACK..n {
if !core.dc_phase[i].is_nan() {
let phase_rad = core.dc_phase[i] * PI / 180.0;
sine[i] = phase_rad.sin();
lead_sine[i] = (phase_rad + PI / 4.0).sin(); // 45-degree lead
}
}
(sine, lead_sine)
}
/// Hilbert Transform Trend vs Cycle Mode: 1 = trending, 0 = cycling.
pub fn ht_trendmode(close: &[f64]) -> Vec<i32> {
compute_ht_core(close).trend_mode
}
// ---------------------------------------------------------------------------
// Tests
// ---------------------------------------------------------------------------
#[cfg(test)]
mod tests {
use super::*;
/// Generate a simple sine wave for testing cycle detection.
fn sine_wave(n: usize, period: f64) -> Vec<f64> {
(0..n)
.map(|i| 100.0 + 10.0 * (2.0 * PI * i as f64 / period).sin())
.collect()
}
/// Flat price series for baseline testing.
fn flat_prices(n: usize) -> Vec<f64> {
vec![100.0; n]
}
#[test]
fn test_ht_trendline_length_and_lookback() {
let close = sine_wave(200, 20.0);
let result = ht_trendline(&close);
assert_eq!(result.len(), close.len());
// First HT_LOOKBACK values must be NaN
for v in &result[..HT_LOOKBACK] {
assert!(v.is_nan(), "expected NaN in lookback region");
}
// Values after lookback must be finite
for v in &result[HT_LOOKBACK..] {
assert!(v.is_finite(), "expected finite value after lookback");
}
}
#[test]
fn test_ht_dcperiod_length_and_lookback() {
let close = sine_wave(200, 20.0);
let result = ht_dcperiod(&close);
assert_eq!(result.len(), close.len());
for v in &result[..HT_LOOKBACK] {
assert!(v.is_nan());
}
// After lookback, period should be positive and finite
for v in &result[HT_LOOKBACK..] {
assert!(v.is_finite());
assert!(*v >= 6.0 && *v <= 50.0, "period {} out of [6,50]", v);
}
}
#[test]
fn test_ht_dcphase_length_and_lookback() {
let close = sine_wave(200, 20.0);
let result = ht_dcphase(&close);
assert_eq!(result.len(), close.len());
for v in &result[..HT_LOOKBACK] {
assert!(v.is_nan());
}
for v in &result[HT_LOOKBACK..] {
assert!(v.is_finite());
}
}
#[test]
fn test_ht_phasor_dual_output() {
let close = sine_wave(200, 20.0);
let (inp, quad) = ht_phasor(&close);
assert_eq!(inp.len(), close.len());
assert_eq!(quad.len(), close.len());
for v in &inp[..HT_LOOKBACK] {
assert!(v.is_nan());
}
for v in &quad[..HT_LOOKBACK] {
assert!(v.is_nan());
}
}
#[test]
fn test_ht_sine_dual_output() {
let close = sine_wave(200, 20.0);
let (s, ls) = ht_sine(&close);
assert_eq!(s.len(), close.len());
assert_eq!(ls.len(), close.len());
for v in &s[..HT_LOOKBACK] {
assert!(v.is_nan());
}
// Sine values should be in [-1, 1]
for v in &s[HT_LOOKBACK..] {
assert!(v.is_finite());
assert!(*v >= -1.0 && *v <= 1.0, "sine {} out of [-1,1]", v);
}
for v in &ls[HT_LOOKBACK..] {
assert!(v.is_finite());
assert!(*v >= -1.0 && *v <= 1.0, "leadsine {} out of [-1,1]", v);
}
}
#[test]
fn test_ht_trendmode_values() {
let close = sine_wave(200, 20.0);
let result = ht_trendmode(&close);
assert_eq!(result.len(), close.len());
// All values must be 0 or 1
for v in &result {
assert!(*v == 0 || *v == 1, "trend_mode {} not 0 or 1", v);
}
}
#[test]
fn test_short_input_all_nan() {
let close = vec![100.0; HT_LOOKBACK]; // exactly HT_LOOKBACK, not enough
let tl = ht_trendline(&close);
assert!(tl.iter().all(|v| v.is_nan()));
let dp = ht_dcperiod(&close);
assert!(dp.iter().all(|v| v.is_nan()));
}
#[test]
fn test_flat_prices_trendline_equals_price() {
let close = flat_prices(200);
let tl = ht_trendline(&close);
// For a flat price, trendline after lookback should be very close to the price
for v in &tl[HT_LOOKBACK..] {
assert!(
(v - 100.0).abs() < 1e-6,
"trendline {} diverged from flat price",
v
);
}
}
}
+977
View File
@@ -0,0 +1,977 @@
//! Extended indicators — pure Rust implementations (no PyO3, no numpy).
//!
//! These indicators are not part of TA-Lib and provide additional technical
//! analysis capabilities. All functions operate on `&[f64]` slices and return
//! `Vec<f64>` (or tuples thereof).
#![allow(clippy::too_many_arguments)]
use crate::math;
use crate::overlap;
// Note: we use a local compute_atr helper (seeds from bar 0) rather than
// crate::volatility::atr (which seeds from bar 1, TA-Lib style).
// ---------------------------------------------------------------------------
// Internal helpers
// ---------------------------------------------------------------------------
/// Compute ATR array using Wilder smoothing (same algorithm as in the PyO3
/// extended module — seeds from bar 0, not bar 1 like TA-Lib's `volatility::atr`).
fn compute_atr(high: &[f64], low: &[f64], close: &[f64], timeperiod: usize) -> Vec<f64> {
let n = high.len();
let mut result = vec![f64::NAN; n];
if n <= timeperiod {
return result;
}
// Seed: SMA of first `timeperiod` true range values
let mut seed_sum = high[0] - low[0]; // first TR has no prev_close
for i in 1..timeperiod {
let hl = high[i] - low[i];
let hc = (high[i] - close[i - 1]).abs();
let lc = (low[i] - close[i - 1]).abs();
seed_sum += hl.max(hc).max(lc);
}
let mut atr = seed_sum / timeperiod as f64;
result[timeperiod - 1] = atr;
let pf = (timeperiod - 1) as f64;
for i in timeperiod..n {
let hl = high[i] - low[i];
let hc = (high[i] - close[i - 1]).abs();
let lc = (low[i] - close[i - 1]).abs();
let tr = hl.max(hc).max(lc);
atr = (atr * pf + tr) / timeperiod as f64;
result[i] = atr;
}
result
}
// ---------------------------------------------------------------------------
// VWAP
// ---------------------------------------------------------------------------
/// Volume Weighted Average Price (cumulative or rolling).
///
/// # Arguments
/// * `high`, `low`, `close`, `volume` — equal-length price/volume slices.
/// * `timeperiod` — 0 for cumulative VWAP from bar 0; >= 1 for a rolling window.
///
/// # Returns
/// A `Vec<f64>` of VWAP values. For rolling mode the first `timeperiod - 1`
/// entries are `NaN`.
pub fn vwap(
high: &[f64],
low: &[f64],
close: &[f64],
volume: &[f64],
timeperiod: usize,
) -> Vec<f64> {
let n = high.len();
let mut result = vec![f64::NAN; n];
if timeperiod == 0 {
let mut cum_tpv = 0.0_f64;
let mut cum_vol = 0.0_f64;
for i in 0..n {
let tp = (high[i] + low[i] + close[i]) / 3.0;
cum_tpv += tp * volume[i];
cum_vol += volume[i];
result[i] = if cum_vol != 0.0 {
cum_tpv / cum_vol
} else {
f64::NAN
};
}
} else {
// Pre-compute cumulative sums for O(n) rolling window
let mut cum_tpv_arr = vec![0.0_f64; n];
let mut cum_vol_arr = vec![0.0_f64; n];
for i in 0..n {
let tp = (high[i] + low[i] + close[i]) / 3.0;
let tpv = tp * volume[i];
cum_tpv_arr[i] = tpv + if i > 0 { cum_tpv_arr[i - 1] } else { 0.0 };
cum_vol_arr[i] = volume[i] + if i > 0 { cum_vol_arr[i - 1] } else { 0.0 };
}
for i in (timeperiod - 1)..n {
let prev_tpv = if i >= timeperiod {
cum_tpv_arr[i - timeperiod]
} else {
0.0
};
let prev_vol = if i >= timeperiod {
cum_vol_arr[i - timeperiod]
} else {
0.0
};
let w_tpv = cum_tpv_arr[i] - prev_tpv;
let w_vol = cum_vol_arr[i] - prev_vol;
result[i] = if w_vol != 0.0 {
w_tpv / w_vol
} else {
f64::NAN
};
}
}
result
}
// ---------------------------------------------------------------------------
// VWMA
// ---------------------------------------------------------------------------
/// Volume Weighted Moving Average.
///
/// `VWMA = sum(close * volume, n) / sum(volume, n)`
///
/// # Arguments
/// * `close` — price series.
/// * `volume` — volume series (same length as `close`).
/// * `timeperiod` — rolling window size (>= 1).
///
/// # Returns
/// A `Vec<f64>` with `NaN` for the first `timeperiod - 1` entries.
pub fn vwma(close: &[f64], volume: &[f64], timeperiod: usize) -> Vec<f64> {
let n = close.len();
let mut result = vec![f64::NAN; n];
if timeperiod < 1 || n < timeperiod {
return result;
}
let mut cum_cv = vec![0.0_f64; n];
let mut cum_v = vec![0.0_f64; n];
for i in 0..n {
cum_cv[i] = close[i] * volume[i] + if i > 0 { cum_cv[i - 1] } else { 0.0 };
cum_v[i] = volume[i] + if i > 0 { cum_v[i - 1] } else { 0.0 };
}
for i in (timeperiod - 1)..n {
let prev_cv = if i >= timeperiod {
cum_cv[i - timeperiod]
} else {
0.0
};
let prev_v = if i >= timeperiod {
cum_v[i - timeperiod]
} else {
0.0
};
let w_cv = cum_cv[i] - prev_cv;
let w_v = cum_v[i] - prev_v;
result[i] = if w_v != 0.0 { w_cv / w_v } else { f64::NAN };
}
result
}
// ---------------------------------------------------------------------------
// SUPERTREND
// ---------------------------------------------------------------------------
/// ATR-based Supertrend indicator.
///
/// # Returns
/// `(supertrend_line, direction)` where direction values are:
/// * `1` = uptrend
/// * `-1` = downtrend
/// * `0` = warmup (first `timeperiod` bars)
pub fn supertrend(
high: &[f64],
low: &[f64],
close: &[f64],
timeperiod: usize,
multiplier: f64,
) -> (Vec<f64>, Vec<i8>) {
let n = high.len();
let mut supertrend_out = vec![f64::NAN; n];
let mut direction = vec![0_i8; n];
if timeperiod < 1 || n <= timeperiod {
return (supertrend_out, direction);
}
let atr = compute_atr(high, low, close, timeperiod);
let mut upper_band = vec![f64::NAN; n];
let mut lower_band = vec![f64::NAN; n];
let first_valid = timeperiod - 1;
if first_valid >= n || atr[first_valid].is_nan() {
return (supertrend_out, direction);
}
// Initialize band state at first valid ATR bar (compute basic bands inline)
{
let hl2 = (high[first_valid] + low[first_valid]) / 2.0;
upper_band[first_valid] = hl2 + multiplier * atr[first_valid];
lower_band[first_valid] = hl2 - multiplier * atr[first_valid];
}
for i in (first_valid + 1)..n {
if atr[i].is_nan() {
continue;
}
// Compute basic bands as scalars — no Vec allocation needed
let hl2 = (high[i] + low[i]) / 2.0;
let upper_basic = hl2 + multiplier * atr[i];
let lower_basic = hl2 - multiplier * atr[i];
// Adjust lower band
lower_band[i] = if lower_basic > lower_band[i - 1] || close[i - 1] < lower_band[i - 1]
{
lower_basic
} else {
lower_band[i - 1]
};
// Adjust upper band
upper_band[i] = if upper_basic < upper_band[i - 1] || close[i - 1] > upper_band[i - 1]
{
upper_basic
} else {
upper_band[i - 1]
};
// Direction and output only from index timeperiod (warmup = 0, NaN)
if i >= timeperiod {
let prev_dir = direction[i - 1];
direction[i] = if prev_dir == 0 {
if close[i] > upper_band[i] {
1
} else {
-1
}
} else if prev_dir == -1 {
if close[i] > upper_band[i] {
1
} else {
-1
}
} else if close[i] < lower_band[i] {
-1
} else {
1
};
supertrend_out[i] = if direction[i] == 1 {
lower_band[i]
} else {
upper_band[i]
};
}
}
(supertrend_out, direction)
}
// ---------------------------------------------------------------------------
// DONCHIAN
// ---------------------------------------------------------------------------
/// Donchian Channels — rolling highest high / lowest low.
///
/// # Returns
/// `(upper, middle, lower)` arrays.
pub fn donchian(
high: &[f64],
low: &[f64],
timeperiod: usize,
) -> (Vec<f64>, Vec<f64>, Vec<f64>) {
let n = high.len();
let mut upper = vec![f64::NAN; n];
let mut lower = vec![f64::NAN; n];
let mut middle = vec![f64::NAN; n];
if timeperiod < 1 || n < timeperiod {
return (upper, middle, lower);
}
let hh = math::sliding_max(high, timeperiod);
let ll = math::sliding_min(low, timeperiod);
for i in 0..n {
if !hh[i].is_nan() {
upper[i] = hh[i];
lower[i] = ll[i];
middle[i] = (upper[i] + lower[i]) / 2.0;
}
}
(upper, middle, lower)
}
// ---------------------------------------------------------------------------
// CHOPPINESS_INDEX
// ---------------------------------------------------------------------------
/// Choppiness Index — measures market choppiness vs trending.
///
/// Values near 100 indicate a choppy market; near 0 indicates trending.
/// The first `timeperiod` values are `NaN`.
pub fn choppiness_index(
high: &[f64],
low: &[f64],
close: &[f64],
timeperiod: usize,
) -> Vec<f64> {
let n = high.len();
let mut result = vec![f64::NAN; n];
if timeperiod < 1 || n <= timeperiod {
return result;
}
// ATR(1) = True Range per bar
let mut tr = vec![0.0_f64; n];
tr[0] = high[0] - low[0];
for i in 1..n {
let hl = high[i] - low[i];
let hc = (high[i] - close[i - 1]).abs();
let lc = (low[i] - close[i - 1]).abs();
tr[i] = hl.max(hc).max(lc);
}
// Cumulative TR for rolling sum
let mut cum_tr = vec![0.0_f64; n];
cum_tr[0] = tr[0];
for i in 1..n {
cum_tr[i] = cum_tr[i - 1] + tr[i];
}
let log_n = (timeperiod as f64).log10();
let hh = math::sliding_max(high, timeperiod);
let ll = math::sliding_min(low, timeperiod);
for i in (timeperiod)..n {
let prev_cum = cum_tr[i - timeperiod];
let sum_tr = cum_tr[i] - prev_cum;
let hl_range = hh[i] - ll[i];
if hl_range > 0.0 && log_n > 0.0 {
result[i] = 100.0 * (sum_tr / hl_range).log10() / log_n;
}
}
result
}
// ---------------------------------------------------------------------------
// KELTNER_CHANNELS
// ---------------------------------------------------------------------------
/// Keltner Channels — EMA +/- (multiplier x ATR).
///
/// # Returns
/// `(upper, middle, lower)` arrays.
pub fn keltner_channels(
high: &[f64],
low: &[f64],
close: &[f64],
timeperiod: usize,
atr_period: usize,
multiplier: f64,
) -> (Vec<f64>, Vec<f64>, Vec<f64>) {
let n = high.len();
if timeperiod < 1 || atr_period < 1 || n < timeperiod || n < atr_period {
let nan = vec![f64::NAN; n];
return (nan.clone(), nan.clone(), nan);
}
let middle = overlap::ema(close, timeperiod);
let atr = compute_atr(high, low, close, atr_period);
let mut upper = vec![f64::NAN; n];
let mut lower = vec![f64::NAN; n];
for i in 0..n {
if !middle[i].is_nan() && !atr[i].is_nan() {
let band = multiplier * atr[i];
upper[i] = middle[i] + band;
lower[i] = middle[i] - band;
}
}
(upper, middle, lower)
}
// ---------------------------------------------------------------------------
// HULL_MA
// ---------------------------------------------------------------------------
/// Hull Moving Average (HMA).
///
/// `HMA(n) = WMA(2 * WMA(n/2) - WMA(n), sqrt(n))`
pub fn hull_ma(close: &[f64], timeperiod: usize) -> Vec<f64> {
let n = close.len();
if timeperiod < 1 || n < timeperiod {
return vec![f64::NAN; n];
}
let half = (timeperiod / 2).max(1);
let sqrt_p = ((timeperiod as f64).sqrt().round() as usize).max(1);
let wma_full = overlap::wma(close, timeperiod);
let wma_half = overlap::wma(close, half);
// raw = 2 * wma_half - wma_full
let mut raw = vec![f64::NAN; n];
for i in 0..n {
if !wma_full[i].is_nan() && !wma_half[i].is_nan() {
raw[i] = 2.0 * wma_half[i] - wma_full[i];
}
}
// Find first valid index in raw
let first_valid = raw.iter().position(|x| !x.is_nan()).unwrap_or(n);
let mut hull = vec![f64::NAN; n];
if first_valid < n {
let raw_valid = &raw[first_valid..];
let hma_slice = overlap::wma(raw_valid, sqrt_p);
for (k, &v) in hma_slice.iter().enumerate() {
hull[first_valid + k] = v;
}
}
hull
}
// ---------------------------------------------------------------------------
// CHANDELIER_EXIT
// ---------------------------------------------------------------------------
/// Chandelier Exit — ATR-based trailing stop levels.
///
/// # Returns
/// `(long_exit, short_exit)` arrays.
pub fn chandelier_exit(
high: &[f64],
low: &[f64],
close: &[f64],
timeperiod: usize,
multiplier: f64,
) -> (Vec<f64>, Vec<f64>) {
let n = high.len();
if timeperiod < 1 || n < timeperiod {
return (vec![f64::NAN; n], vec![f64::NAN; n]);
}
let atr = compute_atr(high, low, close, timeperiod);
let highest_high = math::sliding_max(high, timeperiod);
let lowest_low = math::sliding_min(low, timeperiod);
let mut long_exit = vec![f64::NAN; n];
let mut short_exit = vec![f64::NAN; n];
for i in 0..n {
if !highest_high[i].is_nan() && !atr[i].is_nan() {
long_exit[i] = highest_high[i] - multiplier * atr[i];
short_exit[i] = lowest_low[i] + multiplier * atr[i];
}
}
(long_exit, short_exit)
}
// ---------------------------------------------------------------------------
// ICHIMOKU
// ---------------------------------------------------------------------------
/// Ichimoku Cloud (Ichimoku Kinko Hyo).
///
/// # Returns
/// `(tenkan, kijun, senkou_a, senkou_b, chikou)` arrays.
pub fn ichimoku(
high: &[f64],
low: &[f64],
close: &[f64],
tenkan_period: usize,
kijun_period: usize,
senkou_b_period: usize,
displacement: usize,
) -> (Vec<f64>, Vec<f64>, Vec<f64>, Vec<f64>, Vec<f64>) {
let n = high.len();
let nan = || vec![f64::NAN; n];
if tenkan_period < 1 || kijun_period < 1 || senkou_b_period < 1 {
return (nan(), nan(), nan(), nan(), nan());
}
// Helper: rolling (H+L)/2 via shared sliding_max / sliding_min
let midpoint_rolling = |period: usize| -> Vec<f64> {
let hh = math::sliding_max(high, period);
let ll = math::sliding_min(low, period);
let mut result = vec![f64::NAN; n];
for i in 0..n {
if !hh[i].is_nan() {
result[i] = (hh[i] + ll[i]) / 2.0;
}
}
result
};
let tenkan = midpoint_rolling(tenkan_period);
let kijun = midpoint_rolling(kijun_period);
let raw_b = midpoint_rolling(senkou_b_period);
// Senkou A: (tenkan + kijun) / 2 shifted back `displacement` bars
let mut senkou_a = vec![f64::NAN; n];
if n > displacement {
for i in displacement..n {
if !tenkan[i].is_nan() && !kijun[i].is_nan() {
senkou_a[i - displacement] = (tenkan[i] + kijun[i]) / 2.0;
}
}
}
// Senkou B: raw_b shifted back `displacement` bars
let mut senkou_b = vec![f64::NAN; n];
if n > displacement {
senkou_b[..n - displacement].copy_from_slice(&raw_b[displacement..]);
}
// Chikou: close shifted forward `displacement` bars
let mut chikou = vec![f64::NAN; n];
if n > displacement {
chikou[displacement..].copy_from_slice(&close[..n - displacement]);
}
(tenkan, kijun, senkou_a, senkou_b, chikou)
}
// ---------------------------------------------------------------------------
// PIVOT_POINTS
// ---------------------------------------------------------------------------
/// Pivot Points — support / resistance levels computed from the previous bar.
///
/// # Arguments
/// * `method` — `"classic"`, `"fibonacci"`, or `"camarilla"`. Returns all-NaN
/// vectors for unknown methods.
///
/// # Returns
/// `(pivot, r1, s1, r2, s2)` arrays. Index 0 is always `NaN` (no previous bar).
pub fn pivot_points(
high: &[f64],
low: &[f64],
close: &[f64],
method: &str,
) -> (Vec<f64>, Vec<f64>, Vec<f64>, Vec<f64>, Vec<f64>) {
let n = high.len();
let mut pivot = vec![f64::NAN; n];
let mut r1 = vec![f64::NAN; n];
let mut s1 = vec![f64::NAN; n];
let mut r2 = vec![f64::NAN; n];
let mut s2 = vec![f64::NAN; n];
let method_lower = method.to_lowercase();
if !matches!(method_lower.as_str(), "classic" | "fibonacci" | "camarilla") {
// Unknown method — return all NaN
return (pivot, r1, s1, r2, s2);
}
for i in 1..n {
let ph = high[i - 1];
let pl = low[i - 1];
let pc = close[i - 1];
let hl = ph - pl;
let p = (ph + pl + pc) / 3.0;
pivot[i] = p;
match method_lower.as_str() {
"classic" => {
r1[i] = 2.0 * p - pl;
s1[i] = 2.0 * p - ph;
r2[i] = p + hl;
s2[i] = p - hl;
}
"fibonacci" => {
r1[i] = p + 0.382 * hl;
s1[i] = p - 0.382 * hl;
r2[i] = p + 0.618 * hl;
s2[i] = p - 0.618 * hl;
}
"camarilla" => {
r1[i] = pc + 1.1 * hl / 12.0;
s1[i] = pc - 1.1 * hl / 12.0;
r2[i] = pc + 1.1 * hl / 6.0;
s2[i] = pc - 1.1 * hl / 6.0;
}
_ => unreachable!(),
}
}
(pivot, r1, s1, r2, s2)
}
// ---------------------------------------------------------------------------
// Tests
// ---------------------------------------------------------------------------
#[cfg(test)]
mod tests {
use super::*;
// Shared test data: 10-bar OHLCV
fn sample_ohlcv() -> (Vec<f64>, Vec<f64>, Vec<f64>, Vec<f64>) {
let high = vec![11.0, 12.0, 13.0, 14.0, 15.0, 14.5, 15.5, 16.0, 15.0, 14.0];
let low = vec![9.0, 10.0, 11.0, 12.0, 13.0, 12.5, 13.5, 14.0, 13.0, 12.0];
let close = vec![10.0, 11.0, 12.0, 13.0, 14.0, 13.5, 14.5, 15.0, 14.0, 13.0];
let volume = vec![
100.0, 150.0, 200.0, 250.0, 300.0, 200.0, 350.0, 400.0, 180.0, 220.0,
];
(high, low, close, volume)
}
// -----------------------------------------------------------------------
// VWAP tests
// -----------------------------------------------------------------------
#[test]
fn vwap_cumulative_basic() {
let (h, l, c, v) = sample_ohlcv();
let result = vwap(&h, &l, &c, &v, 0);
assert_eq!(result.len(), h.len());
// First bar: tp = (11+9+10)/3 = 10.0, tpv = 1000.0, vol = 100.0 => 10.0
assert!((result[0] - 10.0).abs() < 1e-10);
// All values should be non-NaN for cumulative
for val in &result {
assert!(!val.is_nan());
}
}
#[test]
fn vwap_empty_input() {
let result = vwap(&[], &[], &[], &[], 0);
assert!(result.is_empty());
}
#[test]
fn vwap_rolling_basic() {
let (h, l, c, v) = sample_ohlcv();
let result = vwap(&h, &l, &c, &v, 3);
assert_eq!(result.len(), h.len());
// First 2 values should be NaN
assert!(result[0].is_nan());
assert!(result[1].is_nan());
// From index 2 onward should be valid
assert!(!result[2].is_nan());
}
// -----------------------------------------------------------------------
// VWMA tests
// -----------------------------------------------------------------------
#[test]
fn vwma_basic() {
let (_, _, c, v) = sample_ohlcv();
let result = vwma(&c, &v, 3);
assert_eq!(result.len(), c.len());
assert!(result[0].is_nan());
assert!(result[1].is_nan());
// Index 2: sum(c*v, 0..3) / sum(v, 0..3) = (1000+1650+2400)/(100+150+200) = 5050/450
let expected = (10.0 * 100.0 + 11.0 * 150.0 + 12.0 * 200.0) / (100.0 + 150.0 + 200.0);
assert!((result[2] - expected).abs() < 1e-10);
}
#[test]
fn vwma_empty_input() {
let result = vwma(&[], &[], 3);
assert!(result.is_empty());
}
#[test]
fn vwma_period_larger_than_data() {
let result = vwma(&[1.0, 2.0], &[100.0, 200.0], 5);
assert_eq!(result.len(), 2);
assert!(result.iter().all(|v| v.is_nan()));
}
// -----------------------------------------------------------------------
// SUPERTREND tests
// -----------------------------------------------------------------------
#[test]
fn supertrend_basic() {
let (h, l, c, _) = sample_ohlcv();
let (st, dir) = supertrend(&h, &l, &c, 3, 2.0);
assert_eq!(st.len(), h.len());
assert_eq!(dir.len(), h.len());
// First 3 bars should be warmup (direction = 0, st = NaN)
for i in 0..3 {
assert_eq!(dir[i], 0);
assert!(st[i].is_nan());
}
// From bar 3 onward, direction should be 1 or -1
for i in 3..h.len() {
assert!(dir[i] == 1 || dir[i] == -1);
assert!(!st[i].is_nan());
}
}
#[test]
fn supertrend_empty_input() {
let (st, dir) = supertrend(&[], &[], &[], 3, 2.0);
assert!(st.is_empty());
assert!(dir.is_empty());
}
#[test]
fn supertrend_insufficient_data() {
let (st, dir) = supertrend(&[1.0, 2.0], &[0.5, 1.5], &[1.5, 1.8], 5, 2.0);
assert!(st.iter().all(|v| v.is_nan()));
assert!(dir.iter().all(|&d| d == 0));
}
// -----------------------------------------------------------------------
// DONCHIAN tests
// -----------------------------------------------------------------------
#[test]
fn donchian_basic() {
let (h, l, _, _) = sample_ohlcv();
let (upper, middle, lower) = donchian(&h, &l, 3);
assert_eq!(upper.len(), h.len());
// First 2 are NaN
assert!(upper[0].is_nan());
assert!(upper[1].is_nan());
// Index 2: max(11,12,13)=13, min(9,10,11)=9
assert!((upper[2] - 13.0).abs() < 1e-10);
assert!((lower[2] - 9.0).abs() < 1e-10);
assert!((middle[2] - 11.0).abs() < 1e-10);
}
#[test]
fn donchian_empty_input() {
let (u, m, l) = donchian(&[], &[], 3);
assert!(u.is_empty());
assert!(m.is_empty());
assert!(l.is_empty());
}
#[test]
fn donchian_period_1() {
let h = vec![5.0, 3.0, 7.0];
let l = vec![2.0, 1.0, 4.0];
let (upper, middle, lower) = donchian(&h, &l, 1);
// Every bar is its own window
assert!((upper[0] - 5.0).abs() < 1e-10);
assert!((lower[0] - 2.0).abs() < 1e-10);
assert!((middle[0] - 3.5).abs() < 1e-10);
}
// -----------------------------------------------------------------------
// CHOPPINESS_INDEX tests
// -----------------------------------------------------------------------
#[test]
fn choppiness_index_basic() {
let (h, l, c, _) = sample_ohlcv();
let result = choppiness_index(&h, &l, &c, 3);
assert_eq!(result.len(), h.len());
// First 3 values should be NaN (timeperiod=3, i+1 > 3 starts at i=3)
assert!(result[0].is_nan());
assert!(result[1].is_nan());
assert!(result[2].is_nan());
// Index 3 should have a valid value (i+1=4 > 3)
assert!(!result[3].is_nan());
// CI should be between 0 and 100
for val in result.iter().filter(|v| !v.is_nan()) {
assert!(*val >= 0.0 && *val <= 100.0);
}
}
#[test]
fn choppiness_index_empty_input() {
let result = choppiness_index(&[], &[], &[], 3);
assert!(result.is_empty());
}
// -----------------------------------------------------------------------
// KELTNER_CHANNELS tests
// -----------------------------------------------------------------------
#[test]
fn keltner_channels_basic() {
let (h, l, c, _) = sample_ohlcv();
let (upper, middle, lower) = keltner_channels(&h, &l, &c, 3, 3, 1.5);
assert_eq!(upper.len(), h.len());
// Where both EMA and ATR are valid, upper > middle > lower
for i in 0..h.len() {
if !upper[i].is_nan() && !lower[i].is_nan() {
assert!(upper[i] > middle[i]);
assert!(lower[i] < middle[i]);
}
}
}
#[test]
fn keltner_channels_empty_input() {
let (u, m, l) = keltner_channels(&[], &[], &[], 3, 3, 1.5);
assert!(u.is_empty());
assert!(m.is_empty());
assert!(l.is_empty());
}
// -----------------------------------------------------------------------
// HULL_MA tests
// -----------------------------------------------------------------------
#[test]
fn hull_ma_basic() {
let prices: Vec<f64> = (1..=20).map(|i| i as f64).collect();
let result = hull_ma(&prices, 4);
assert_eq!(result.len(), prices.len());
// Should have some NaN warmup, then valid values
let valid_count = result.iter().filter(|v| !v.is_nan()).count();
assert!(valid_count > 0);
}
#[test]
fn hull_ma_empty_input() {
let result = hull_ma(&[], 4);
assert!(result.is_empty());
}
#[test]
fn hull_ma_period_larger_than_data() {
let result = hull_ma(&[1.0, 2.0], 10);
assert!(result.iter().all(|v| v.is_nan()));
}
// -----------------------------------------------------------------------
// CHANDELIER_EXIT tests
// -----------------------------------------------------------------------
#[test]
fn chandelier_exit_basic() {
let (h, l, c, _) = sample_ohlcv();
let (long_exit, short_exit) = chandelier_exit(&h, &l, &c, 3, 2.0);
assert_eq!(long_exit.len(), h.len());
assert_eq!(short_exit.len(), h.len());
// Where valid, long_exit should be below highest high
for i in 0..h.len() {
if !long_exit[i].is_nan() {
// long_exit = highest_high - multiplier * atr, should be < max high
assert!(long_exit[i] < 20.0); // sanity
}
}
}
#[test]
fn chandelier_exit_empty_input() {
let (le, se) = chandelier_exit(&[], &[], &[], 3, 2.0);
assert!(le.is_empty());
assert!(se.is_empty());
}
// -----------------------------------------------------------------------
// ICHIMOKU tests
// -----------------------------------------------------------------------
#[test]
fn ichimoku_basic() {
// Use a larger dataset for ichimoku
let n = 60;
let high: Vec<f64> = (0..n).map(|i| 100.0 + i as f64 + 1.0).collect();
let low: Vec<f64> = (0..n).map(|i| 100.0 + i as f64 - 1.0).collect();
let close: Vec<f64> = (0..n).map(|i| 100.0 + i as f64).collect();
let (tenkan, kijun, senkou_a, senkou_b, chikou) =
ichimoku(&high, &low, &close, 9, 26, 52, 26);
assert_eq!(tenkan.len(), n);
assert_eq!(kijun.len(), n);
assert_eq!(senkou_a.len(), n);
assert_eq!(senkou_b.len(), n);
assert_eq!(chikou.len(), n);
// Tenkan: period 9, first valid at index 8
assert!(tenkan[7].is_nan());
assert!(!tenkan[8].is_nan());
// Kijun: period 26, first valid at index 25
assert!(kijun[24].is_nan());
assert!(!kijun[25].is_nan());
// Chikou: close shifted forward by 26 bars
assert!(chikou[25].is_nan());
assert!(!chikou[26].is_nan());
assert!((chikou[26] - close[0]).abs() < 1e-10);
}
#[test]
fn ichimoku_empty_input() {
let (t, k, sa, sb, ch) = ichimoku(&[], &[], &[], 9, 26, 52, 26);
assert!(t.is_empty());
assert!(k.is_empty());
assert!(sa.is_empty());
assert!(sb.is_empty());
assert!(ch.is_empty());
}
// -----------------------------------------------------------------------
// PIVOT_POINTS tests
// -----------------------------------------------------------------------
#[test]
fn pivot_points_classic() {
let h = vec![10.0, 12.0, 11.0];
let l = vec![8.0, 9.0, 8.5];
let c = vec![9.0, 11.0, 10.0];
let (pivot, r1, s1, r2, s2) = pivot_points(&h, &l, &c, "classic");
assert_eq!(pivot.len(), 3);
// Index 0 is NaN
assert!(pivot[0].is_nan());
// Index 1: prev bar H=10, L=8, C=9 => P=(10+8+9)/3=9.0
assert!((pivot[1] - 9.0).abs() < 1e-10);
// R1 = 2*P - L = 18 - 8 = 10
assert!((r1[1] - 10.0).abs() < 1e-10);
// S1 = 2*P - H = 18 - 10 = 8
assert!((s1[1] - 8.0).abs() < 1e-10);
// R2 = P + (H-L) = 9 + 2 = 11
assert!((r2[1] - 11.0).abs() < 1e-10);
// S2 = P - (H-L) = 9 - 2 = 7
assert!((s2[1] - 7.0).abs() < 1e-10);
}
#[test]
fn pivot_points_fibonacci() {
let h = vec![10.0, 12.0];
let l = vec![8.0, 9.0];
let c = vec![9.0, 11.0];
let (pivot, r1, s1, _, _) = pivot_points(&h, &l, &c, "fibonacci");
// Index 1: P = (10+8+9)/3 = 9.0, HL = 2
assert!((pivot[1] - 9.0).abs() < 1e-10);
assert!((r1[1] - (9.0 + 0.382 * 2.0)).abs() < 1e-10);
assert!((s1[1] - (9.0 - 0.382 * 2.0)).abs() < 1e-10);
}
#[test]
fn pivot_points_camarilla() {
let h = vec![10.0, 12.0];
let l = vec![8.0, 9.0];
let c = vec![9.0, 11.0];
let (pivot, r1, s1, _, _) = pivot_points(&h, &l, &c, "camarilla");
assert!((pivot[1] - 9.0).abs() < 1e-10);
// R1 = C + 1.1 * HL / 12 = 9 + 1.1*2/12
assert!((r1[1] - (9.0 + 1.1 * 2.0 / 12.0)).abs() < 1e-10);
assert!((s1[1] - (9.0 - 1.1 * 2.0 / 12.0)).abs() < 1e-10);
}
#[test]
fn pivot_points_unknown_method() {
let h = vec![10.0, 12.0];
let l = vec![8.0, 9.0];
let c = vec![9.0, 11.0];
let (pivot, r1, s1, r2, s2) = pivot_points(&h, &l, &c, "unknown");
assert!(pivot.iter().all(|v| v.is_nan()));
assert!(r1.iter().all(|v| v.is_nan()));
assert!(s1.iter().all(|v| v.is_nan()));
assert!(r2.iter().all(|v| v.is_nan()));
assert!(s2.iter().all(|v| v.is_nan()));
}
#[test]
fn pivot_points_empty_input() {
let (p, r1, s1, r2, s2) = pivot_points(&[], &[], &[], "classic");
assert!(p.is_empty());
assert!(r1.is_empty());
assert!(s1.is_empty());
assert!(r2.is_empty());
assert!(s2.is_empty());
}
}
+19
View File
@@ -26,11 +26,30 @@ assert!((sma[2] - 2.0).abs() < 1e-10);
```
*/
pub mod aggregation;
pub mod alerts;
pub mod attribution;
pub mod backtest;
pub mod batch;
pub mod chunked;
pub mod commission;
pub mod crypto;
pub mod currency;
pub mod cycle;
pub mod extended;
pub mod futures;
pub mod math;
pub mod math_ops;
pub mod momentum;
pub mod options;
pub mod overlap;
pub mod pattern;
pub mod portfolio;
pub mod price_transform;
pub mod regime;
pub mod resampling;
pub mod signals;
pub mod statistic;
pub mod streaming;
pub mod volatility;
pub mod volume;
+38 -9
View File
@@ -2,7 +2,14 @@
use std::collections::VecDeque;
/// Rolling sum over `timeperiod` bars.
/// Compute the rolling sum over `timeperiod` bars.
///
/// Returns a `Vec<f64>` of length `n`. The first `timeperiod - 1` values
/// are `NaN`. Uses an incremental algorithm (add new, subtract old) for O(n).
///
/// # Arguments
/// * `real` - Input series.
/// * `timeperiod` - Rolling window size (must be >= 1).
pub fn sum(real: &[f64], timeperiod: usize) -> Vec<f64> {
let n = real.len();
let mut result = vec![f64::NAN; n];
@@ -18,20 +25,38 @@ pub fn sum(real: &[f64], timeperiod: usize) -> Vec<f64> {
result
}
/// Rolling maximum over `timeperiod` bars — O(n) via monotonic deque.
/// Compute the rolling maximum over `timeperiod` bars.
///
/// Delegates to [`sliding_max`] for O(n) performance via a monotonic deque.
/// The first `timeperiod - 1` values are `NaN`.
///
/// # Arguments
/// * `real` - Input series.
/// * `timeperiod` - Rolling window size (must be >= 1).
pub fn max(real: &[f64], timeperiod: usize) -> Vec<f64> {
sliding_max(real, timeperiod)
}
/// Rolling minimum over `timeperiod` bars — O(n) via monotonic deque.
/// Compute the rolling minimum over `timeperiod` bars.
///
/// Delegates to [`sliding_min`] for O(n) performance via a monotonic deque.
/// The first `timeperiod - 1` values are `NaN`.
///
/// # Arguments
/// * `real` - Input series.
/// * `timeperiod` - Rolling window size (must be >= 1).
pub fn min(real: &[f64], timeperiod: usize) -> Vec<f64> {
sliding_min(real, timeperiod)
}
/// Sliding maximum over `timeperiod` bars O(n) via monotonic deque.
/// Compute the sliding maximum over `timeperiod` bars in O(n) time.
///
/// Equivalent to `max` but uses a monotonic deque for O(n) total time.
/// Leading `timeperiod - 1` values are NaN.
/// Uses a monotonic decreasing deque so each element is pushed/popped at
/// most once. The first `timeperiod - 1` values are `NaN`.
///
/// # Arguments
/// * `real` - Input series.
/// * `timeperiod` - Rolling window size (must be >= 1).
pub fn sliding_max(real: &[f64], timeperiod: usize) -> Vec<f64> {
let n = real.len();
let mut result = vec![f64::NAN; n];
@@ -56,10 +81,14 @@ pub fn sliding_max(real: &[f64], timeperiod: usize) -> Vec<f64> {
result
}
/// Sliding minimum over `timeperiod` bars O(n) via monotonic deque.
/// Compute the sliding minimum over `timeperiod` bars in O(n) time.
///
/// Equivalent to `min` but uses a monotonic deque for O(n) total time.
/// Leading `timeperiod - 1` values are NaN.
/// Uses a monotonic increasing deque so each element is pushed/popped at
/// most once. The first `timeperiod - 1` values are `NaN`.
///
/// # Arguments
/// * `real` - Input series.
/// * `timeperiod` - Rolling window size (must be >= 1).
pub fn sliding_min(real: &[f64], timeperiod: usize) -> Vec<f64> {
let n = real.len();
let mut result = vec![f64::NAN; n];
+154
View File
@@ -0,0 +1,154 @@
//! Rolling math operators — O(n) sliding window implementations.
//!
//! - `rolling_sum` — rolling sum over `timeperiod` bars (prefix-sum based)
//! - `rolling_max` — rolling maximum (O(n) monotonic deque)
//! - `rolling_min` — rolling minimum (O(n) monotonic deque)
//! - `rolling_maxindex` — index of rolling maximum
//! - `rolling_minindex` — index of rolling minimum
use std::collections::VecDeque;
/// Rolling sum over `timeperiod` bars using a prefix-sum array.
/// Leading `timeperiod - 1` values are NaN.
pub fn rolling_sum(real: &[f64], timeperiod: usize) -> Vec<f64> {
let n = real.len();
let mut result = vec![f64::NAN; n];
if timeperiod == 0 || n < timeperiod {
return result;
}
let mut cs = vec![0.0f64; n + 1];
for i in 0..n {
cs[i + 1] = cs[i] + real[i];
}
for i in (timeperiod - 1)..n {
result[i] = cs[i + 1] - cs[i + 1 - timeperiod];
}
result
}
/// Rolling maximum over `timeperiod` bars (O(n) monotonic deque).
/// Delegates to `math::sliding_max`.
pub fn rolling_max(real: &[f64], timeperiod: usize) -> Vec<f64> {
crate::math::sliding_max(real, timeperiod)
}
/// Rolling minimum over `timeperiod` bars (O(n) monotonic deque).
/// Delegates to `math::sliding_min`.
pub fn rolling_min(real: &[f64], timeperiod: usize) -> Vec<f64> {
crate::math::sliding_min(real, timeperiod)
}
/// Index of rolling maximum over `timeperiod` bars.
/// Returns 0-based index. During warmup the value is `-1`.
pub fn rolling_maxindex(real: &[f64], timeperiod: usize) -> Vec<i64> {
let n = real.len();
let mut result = vec![-1i64; n];
if timeperiod == 0 || n < timeperiod {
return result;
}
let mut dq: VecDeque<usize> = VecDeque::new();
for i in 0..n {
while dq.front().map(|&j| j + timeperiod <= i).unwrap_or(false) {
dq.pop_front();
}
while dq.back().map(|&j| real[j] <= real[i]).unwrap_or(false) {
dq.pop_back();
}
dq.push_back(i);
if i + 1 >= timeperiod {
result[i] = *dq.front().unwrap() as i64;
}
}
result
}
/// Index of rolling minimum over `timeperiod` bars.
/// Returns 0-based index. During warmup the value is `-1`.
pub fn rolling_minindex(real: &[f64], timeperiod: usize) -> Vec<i64> {
let n = real.len();
let mut result = vec![-1i64; n];
if timeperiod == 0 || n < timeperiod {
return result;
}
let mut dq: VecDeque<usize> = VecDeque::new();
for i in 0..n {
while dq.front().map(|&j| j + timeperiod <= i).unwrap_or(false) {
dq.pop_front();
}
while dq.back().map(|&j| real[j] >= real[i]).unwrap_or(false) {
dq.pop_back();
}
dq.push_back(i);
if i + 1 >= timeperiod {
result[i] = *dq.front().unwrap() as i64;
}
}
result
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_rolling_sum() {
let data = vec![1.0, 2.0, 3.0, 4.0, 5.0];
let result = rolling_sum(&data, 3);
assert!(result[0].is_nan());
assert!(result[1].is_nan());
assert!((result[2] - 6.0).abs() < 1e-10); // 1+2+3
assert!((result[3] - 9.0).abs() < 1e-10); // 2+3+4
assert!((result[4] - 12.0).abs() < 1e-10); // 3+4+5
}
#[test]
fn test_rolling_max() {
let data = vec![1.0, 3.0, 2.0, 5.0, 4.0];
let result = rolling_max(&data, 3);
assert!(result[0].is_nan());
assert!(result[1].is_nan());
assert!((result[2] - 3.0).abs() < 1e-10);
assert!((result[3] - 5.0).abs() < 1e-10);
assert!((result[4] - 5.0).abs() < 1e-10);
}
#[test]
fn test_rolling_min() {
let data = vec![5.0, 3.0, 4.0, 1.0, 2.0];
let result = rolling_min(&data, 3);
assert!(result[0].is_nan());
assert!(result[1].is_nan());
assert!((result[2] - 3.0).abs() < 1e-10);
assert!((result[3] - 1.0).abs() < 1e-10);
assert!((result[4] - 1.0).abs() < 1e-10);
}
#[test]
fn test_rolling_maxindex() {
let data = vec![1.0, 3.0, 2.0, 5.0, 4.0];
let result = rolling_maxindex(&data, 3);
assert_eq!(result[0], -1);
assert_eq!(result[1], -1);
assert_eq!(result[2], 1); // max(1,3,2) at index 1
assert_eq!(result[3], 3); // max(3,2,5) at index 3
assert_eq!(result[4], 3); // max(2,5,4) at index 3
}
#[test]
fn test_rolling_minindex() {
let data = vec![5.0, 3.0, 4.0, 1.0, 2.0];
let result = rolling_minindex(&data, 3);
assert_eq!(result[0], -1);
assert_eq!(result[1], -1);
assert_eq!(result[2], 1); // min(5,3,4) at index 1
assert_eq!(result[3], 3); // min(3,4,1) at index 3
assert_eq!(result[4], 3); // min(4,1,2) at index 3
}
#[test]
fn test_short_input() {
let data = vec![1.0, 2.0];
let result = rolling_sum(&data, 5);
assert!(result.iter().all(|v| v.is_nan()));
}
}
+175 -34
View File
@@ -1,11 +1,15 @@
//! Momentum indicators.
use crate::math::{sliding_max, sliding_min};
/// Relative Strength Index — TA-Lib compatible Wilder smoothing.
/// Compute the Relative Strength Index (RSI).
///
/// Seeds avg_gain/avg_loss with SMA of first `timeperiod` changes.
/// Uses branchless gain/loss split: `gain = diff.max(0.0)`, `loss = (-diff).max(0.0)`.
/// Returns values in the range `[0, 100]`. Uses Wilder's smoothing method
/// (TA-Lib compatible), seeding avg_gain/avg_loss with the SMA of the first
/// `timeperiod` price changes. The first `timeperiod` values are `NaN`.
///
/// # Arguments
/// * `close` - Price series.
/// * `timeperiod` - Lookback period (typically 14).
pub fn rsi(close: &[f64], timeperiod: usize) -> Vec<f64> {
let n = close.len();
let mut result = vec![f64::NAN; n];
@@ -46,7 +50,14 @@ pub fn rsi(close: &[f64], timeperiod: usize) -> Vec<f64> {
result
}
/// Momentum — `close[i] - close[i - timeperiod]`.
/// Compute the Momentum indicator: `close[i] - close[i - timeperiod]`.
///
/// Returns a `Vec<f64>` of length `n`. The first `timeperiod` values are `NaN`.
/// Positive values indicate upward price movement over the lookback window.
///
/// # Arguments
/// * `close` - Price series.
/// * `timeperiod` - Number of bars to look back (must be >= 1).
pub fn mom(close: &[f64], timeperiod: usize) -> Vec<f64> {
let n = close.len();
let mut result = vec![f64::NAN; n];
@@ -59,14 +70,21 @@ pub fn mom(close: &[f64], timeperiod: usize) -> Vec<f64> {
result
}
/// Stochastic Oscillator TA-Lib compatible.
/// Compute the Stochastic Oscillator (TA-Lib compatible).
///
/// Returns `(slowk, slowd)`.
/// - Fast %K[i] = 100 * (close[i] - min(low, fastk_period)) / (max(high, fastk_period) - min(low, fastk_period))
/// - Slow %K = SMA(fast %K, slowk_period)
/// - Slow %D = SMA(slow %K, slowd_period)
/// Returns `(slow_k, slow_d)`, both in the range `[0, 100]`.
/// - Fast %K = 100 * (close - lowest low) / (highest high - lowest low)
/// - Slow %K = SMA(fast %K, `slowk_period`)
/// - Slow %D = SMA(slow %K, `slowd_period`)
///
/// Uses O(n) sliding max/min via monotonic deques.
/// Uses O(n) sliding max/min via monotonic deques. Both outputs are
/// `NaN`-padded until slow %D becomes valid (TA-Lib convention).
///
/// # Arguments
/// * `high` / `low` / `close` - OHLC price series (same length).
/// * `fastk_period` - Lookback for highest high / lowest low.
/// * `slowk_period` - SMA period applied to fast %K.
/// * `slowd_period` - SMA period applied to slow %K.
pub fn stoch(
high: &[f64],
low: &[f64],
@@ -84,29 +102,38 @@ pub fn stoch(
return nan_pair();
}
let max_h = sliding_max(high, fastk_period);
let min_l = sliding_min(low, fastk_period);
let mut slowk = vec![f64::NAN; n];
let mut slowd = vec![f64::NAN; n];
// Fast %K is valid from index fastk_period-1 onward.
// Fused pass: compute fast %K inline with sliding max/min.
// For typical small windows (5-14), inline scan beats VecDeque overhead.
let fastk_start = fastk_period - 1;
let mut fastk_valid = vec![0.0; n - fastk_start];
let fk_len = n - fastk_start;
let mut fastk_valid = vec![0.0_f64; fk_len];
for i in fastk_start..n {
let range = max_h[i] - min_l[i];
// Inline sliding max(high) and min(low) over [i - fastk_period + 1 .. i].
let win_start = i + 1 - fastk_period;
let mut hh = high[win_start];
let mut ll = low[win_start];
for j in (win_start + 1)..=i {
let h = high[j];
let l = low[j];
if h > hh { hh = h; }
if l < ll { ll = l; }
}
let range = hh - ll;
fastk_valid[i - fastk_start] = if range != 0.0 {
100.0 * (close[i] - min_l[i]) / range
100.0 * (close[i] - ll) / range
} else {
0.0
};
}
// Slow %K = SMA(fastk_valid, slowk_period); write directly into `slowk` offset by `fastk_start`.
// Slow %K = SMA(fastk_valid, slowk_period).
crate::overlap::sma_into(&fastk_valid, slowk_period, &mut slowk, fastk_start);
// Slow %D = SMA(slowk, slowd_period).
// The valid part of slowk starts at `fastk_start + slowk_period - 1`.
let slowk_valid_start = fastk_start + slowk_period - 1;
let slowd_valid_start = slowk_valid_start + slowd_period - 1;
@@ -249,50 +276,164 @@ fn adx_inner(high: &[f64], low: &[f64], close: &[f64], period: usize) -> AdxInne
(b_pdm, b_mdm, b_pdi, b_mdi, b_dx, b_adx)
}
/// Plus Directional Movement (Wilder smoothed). Output length = n (bar 0 is NaN).
pub fn plus_dm(high: &[f64], low: &[f64], timeperiod: usize) -> Vec<f64> {
/// Compute all six ADX-family outputs in a single pass.
///
/// Returns `(plus_dm, minus_dm, plus_di, minus_di, dx, adx)`.
/// Use this when you need multiple ADX-family outputs to avoid redundant
/// computation. All values are in `[0, 100]` except DM which is unbounded.
/// Warmup: DI/DX valid from index `timeperiod`; ADX from `2 * timeperiod - 1`.
///
/// # Arguments
/// * `high` / `low` / `close` - OHLC price series (same length).
/// * `timeperiod` - Wilder smoothing period (typically 14).
pub fn adx_all(high: &[f64], low: &[f64], close: &[f64], timeperiod: usize) -> AdxInnerOutput {
adx_inner(high, low, close, timeperiod)
}
/// Internal helper for plus_dm and minus_dm that doesn't allocate dummy close prices.
/// Returns (plus_dm, minus_dm) smoothed with Wilder's method.
fn dm_only_inner(high: &[f64], low: &[f64], period: usize) -> (Vec<f64>, Vec<f64>) {
let n = high.len();
let closes = vec![0.0_f64; n];
let (pdm, _, _, _, _, _) = adx_inner(high, low, &closes, timeperiod);
let mut b_pdm = vec![f64::NAN; n];
let mut b_mdm = vec![f64::NAN; n];
if n < period || period < 1 || n < 2 {
return (b_pdm, b_mdm);
}
let m = n - 1;
let mut pdm = vec![0.0_f64; m];
let mut mdm = vec![0.0_f64; m];
for i in 0..m {
let j = i + 1;
let h_diff = high[j] - high[i];
let l_diff = low[i] - low[j];
pdm[i] = if h_diff > l_diff && h_diff > 0.0 {
h_diff
} else {
0.0
};
mdm[i] = if l_diff > h_diff && l_diff > 0.0 {
l_diff
} else {
0.0
};
}
if m < period {
return (b_pdm, b_mdm);
}
let mut pdm_s = pdm[..period].iter().sum::<f64>();
let mut mdm_s = mdm[..period].iter().sum::<f64>();
b_pdm[period] = pdm_s;
b_mdm[period] = mdm_s;
let decay = (period - 1) as f64 / period as f64;
for i in period..m {
pdm_s = pdm_s * decay + pdm[i];
mdm_s = mdm_s * decay + mdm[i];
b_pdm[i + 1] = pdm_s;
b_mdm[i + 1] = mdm_s;
}
(b_pdm, b_mdm)
}
/// Compute the Plus Directional Movement (+DM), Wilder smoothed.
///
/// Measures upward price movement. Returns a `Vec<f64>` of length `n`;
/// the first `timeperiod` values are `NaN`.
///
/// # Arguments
/// * `high` / `low` - High and low price series (same length).
/// * `timeperiod` - Wilder smoothing period.
pub fn plus_dm(high: &[f64], low: &[f64], timeperiod: usize) -> Vec<f64> {
let (pdm, _) = dm_only_inner(high, low, timeperiod);
pdm
}
/// Minus Directional Movement (Wilder smoothed). Output length = n (bar 0 is NaN).
/// Compute the Minus Directional Movement (-DM), Wilder smoothed.
///
/// Measures downward price movement. Returns a `Vec<f64>` of length `n`;
/// the first `timeperiod` values are `NaN`.
///
/// # Arguments
/// * `high` / `low` - High and low price series (same length).
/// * `timeperiod` - Wilder smoothing period.
pub fn minus_dm(high: &[f64], low: &[f64], timeperiod: usize) -> Vec<f64> {
let n = high.len();
let closes = vec![0.0_f64; n];
let (_, mdm, _, _, _, _) = adx_inner(high, low, &closes, timeperiod);
let (_, mdm) = dm_only_inner(high, low, timeperiod);
mdm
}
/// Plus Directional Indicator (Wilder smoothed). Output length = n.
/// Compute the Plus Directional Indicator (+DI), Wilder smoothed.
///
/// `+DI = 100 * smoothed(+DM) / smoothed(TR)`. Returns values in `[0, 100]`.
/// The first `timeperiod` values are `NaN`.
///
/// # Arguments
/// * `high` / `low` / `close` - OHLC price series (same length).
/// * `timeperiod` - Wilder smoothing period.
pub fn plus_di(high: &[f64], low: &[f64], close: &[f64], timeperiod: usize) -> Vec<f64> {
let (_, _, pdi, _, _, _) = adx_inner(high, low, close, timeperiod);
pdi
}
/// Minus Directional Indicator (Wilder smoothed). Output length = n.
/// Compute the Minus Directional Indicator (-DI), Wilder smoothed.
///
/// `-DI = 100 * smoothed(-DM) / smoothed(TR)`. Returns values in `[0, 100]`.
/// The first `timeperiod` values are `NaN`.
///
/// # Arguments
/// * `high` / `low` / `close` - OHLC price series (same length).
/// * `timeperiod` - Wilder smoothing period.
pub fn minus_di(high: &[f64], low: &[f64], close: &[f64], timeperiod: usize) -> Vec<f64> {
let (_, _, _, mdi, _, _) = adx_inner(high, low, close, timeperiod);
mdi
}
/// Directional Movement Index: 100 * |+DI DI| / (+DI + DI).
/// Compute the Directional Movement Index (DX).
///
/// `DX = 100 * |+DI - -DI| / (+DI + -DI)`. Returns values in `[0, 100]`.
/// The first `timeperiod` values are `NaN`.
///
/// # Arguments
/// * `high` / `low` / `close` - OHLC price series (same length).
/// * `timeperiod` - Wilder smoothing period.
pub fn dx(high: &[f64], low: &[f64], close: &[f64], timeperiod: usize) -> Vec<f64> {
let (_, _, _, _, dx_vals, _) = adx_inner(high, low, close, timeperiod);
dx_vals
}
/// Average Directional Movement Index (Wilder smoothing of DX).
/// Compute the Average Directional Movement Index (ADX).
///
/// ADX is Wilder's smoothing of DX, measuring trend strength regardless of
/// direction. Returns values in `[0, 100]`. The first `2 * timeperiod - 1`
/// values are `NaN` (DX warmup + ADX smoothing warmup).
///
/// # Arguments
/// * `high` / `low` / `close` - OHLC price series (same length).
/// * `timeperiod` - Wilder smoothing period (typically 14).
pub fn adx(high: &[f64], low: &[f64], close: &[f64], timeperiod: usize) -> Vec<f64> {
let (_, _, _, _, _, adx_vals) = adx_inner(high, low, close, timeperiod);
adx_vals
}
/// ADX Rating: (ADX[i] + ADX[i timeperiod]) / 2.
/// Compute the ADX Rating (ADXR).
///
/// `ADXR[i] = (ADX[i] + ADX[i - timeperiod]) / 2`. Smooths ADX further
/// by averaging current ADX with its value `timeperiod` bars ago.
/// Returns values in `[0, 100]`.
///
/// # Arguments
/// * `high` / `low` / `close` - OHLC price series (same length).
/// * `timeperiod` - Wilder smoothing period (typically 14).
pub fn adxr(high: &[f64], low: &[f64], close: &[f64], timeperiod: usize) -> Vec<f64> {
let n = high.len();
let adx_vals = adx(high, low, close, timeperiod);
// Reuse adx_all to compute ADX once, then derive ADXR from it
let (_, _, _, _, _, adx_vals) = adx_inner(high, low, close, timeperiod);
let mut result = vec![f64::NAN; n];
for i in timeperiod..n {
if !adx_vals[i].is_nan() && !adx_vals[i - timeperiod].is_nan() {
+252 -77
View File
@@ -3,7 +3,14 @@
//! All functions return a `Vec<f64>` of the same length as the input.
//! Leading values are `f64::NAN` for the warm-up period.
/// Simple Moving Average over `timeperiod` bars.
/// Compute the Simple Moving Average (SMA) over a rolling window.
///
/// Returns a `Vec<f64>` of the same length as `close`. The first
/// `timeperiod - 1` values are `NaN` (warmup period).
///
/// # Arguments
/// * `close` - Price series.
/// * `timeperiod` - Rolling window size (must be >= 1).
///
/// # Edge Cases
/// Returns all-NaN when `timeperiod < 1` or `close.len() < timeperiod`.
@@ -14,8 +21,17 @@ pub fn sma(close: &[f64], timeperiod: usize) -> Vec<f64> {
result
}
/// Simple Moving Average written directly into `dest` starting at `dest_offset`.
/// Leaves values before `dest_offset + timeperiod - 1` untouched (e.g. they can be NaN).
/// Write a Simple Moving Average directly into a pre-allocated buffer.
///
/// Values before `dest_offset + timeperiod - 1` are left untouched.
/// This avoids an intermediate allocation when composing indicators
/// (e.g., Stochastic slow %K and slow %D).
///
/// # Arguments
/// * `src` - Input price series.
/// * `timeperiod` - Rolling window size (must be >= 1).
/// * `dest` - Output buffer (must be at least `dest_offset + src.len()` long).
/// * `dest_offset` - Starting index in `dest` to write results.
pub fn sma_into(src: &[f64], timeperiod: usize, dest: &mut [f64], dest_offset: usize) {
let n = src.len();
if timeperiod < 1 || n < timeperiod {
@@ -66,7 +82,15 @@ pub fn sma_into(src: &[f64], timeperiod: usize, dest: &mut [f64], dest_offset: u
}
}
/// Exponential Moving Average — seeded with SMA of first `timeperiod` bars.
/// Compute the Exponential Moving Average (EMA).
///
/// The EMA is seeded with the SMA of the first `timeperiod` bars and uses
/// a smoothing factor of `k = 2 / (timeperiod + 1)`. Returns a `Vec<f64>`
/// of the same length as `close`; the first `timeperiod - 1` values are `NaN`.
///
/// # Arguments
/// * `close` - Price series.
/// * `timeperiod` - Lookback period (must be >= 1).
pub fn ema(close: &[f64], timeperiod: usize) -> Vec<f64> {
let n = close.len();
let mut result = vec![f64::NAN; n];
@@ -82,10 +106,15 @@ pub fn ema(close: &[f64], timeperiod: usize) -> Vec<f64> {
result
}
/// Weighted Moving Average — O(n) incremental algorithm using running weighted sum.
/// Compute the Weighted Moving Average (WMA).
///
/// Recurrence: `T[i] = T[i-1] + n*close[i] - S[i-1]`
/// where `S[i]` is the rolling sum over `timeperiod` bars.
/// Assigns linearly increasing weights (1, 2, ..., timeperiod) to the window.
/// Uses an O(n) incremental recurrence to avoid recomputing weights each bar.
/// Returns a `Vec<f64>` of length `n`; the first `timeperiod - 1` values are `NaN`.
///
/// # Arguments
/// * `close` - Price series.
/// * `timeperiod` - Rolling window size (must be >= 1).
pub fn wma(close: &[f64], timeperiod: usize) -> Vec<f64> {
let n = close.len();
let mut result = vec![f64::NAN; n];
@@ -157,10 +186,43 @@ pub fn wma(close: &[f64], timeperiod: usize) -> Vec<f64> {
result
}
/// Bollinger Bands returns `(upper, middle, lower)`.
/// Compute Bollinger Bands, returning `(upper, middle, lower)`.
///
/// Middle is SMA; bands are `± nbdev * stddev`.
/// Uses O(n) sliding `sum` and `sum_sq` windows for mean and variance.
/// The middle band is the SMA; upper and lower bands are offset by
/// `nbdevup` and `nbdevdn` standard deviations respectively. Uses
/// Welford's rolling algorithm for numerically stable variance in O(n).
///
/// # Arguments
/// * `close` - Price series.
/// * `timeperiod` - SMA / standard deviation window (must be >= 1).
/// * `nbdevup` - Number of standard deviations above the mean for the upper band.
/// * `nbdevdn` - Number of standard deviations below the mean for the lower band.
///
/// # Returns
/// `(upper, middle, lower)` -- each `Vec<f64>` of length `n`. The first
/// `timeperiod - 1` values in each vector are `NaN`.
///
/// ## Welford's rolling algorithm
///
/// We maintain `mean` and `m2` (sum of squared deviations from the current
/// mean) across a sliding window of size `N`. When a new value `x_new`
/// replaces an old value `x_old` (window size stays constant):
///
/// ```text
/// delta = x_new - x_old
/// old_mean = mean
/// mean += delta / N
/// m2 += delta * ((x_new - mean) + (x_old - old_mean))
///
/// variance = m2 / N // population variance
/// stddev = sqrt(variance)
/// ```
///
/// The initial window is seeded using the standard (non-rolling) Welford
/// incremental algorithm.
///
/// This avoids the catastrophic cancellation inherent in the naïve
/// `Σx²/N mean²` formula when values are large but close together.
pub fn bbands(
close: &[f64],
timeperiod: usize,
@@ -177,90 +239,135 @@ pub fn bbands(
let mut lower = vec![f64::NAN; n];
let p = timeperiod as f64;
// Seed sliding sums for the first window.
#[cfg(feature = "simd")]
let (mut sum, mut sum_sq) = {
use wide::f64x4;
let p_data = &close[..timeperiod];
let mut sum_simd = f64x4::splat(0.0);
let mut sq_simd = f64x4::splat(0.0);
let mut chunks = p_data.chunks_exact(4);
for chunk in &mut chunks {
let vals = f64x4::new([chunk[0], chunk[1], chunk[2], chunk[3]]);
sum_simd += vals;
sq_simd += vals * vals;
}
let s_arr = sum_simd.to_array();
let sq_arr = sq_simd.to_array();
let mut sum = s_arr[0] + s_arr[1] + s_arr[2] + s_arr[3];
let mut sum_sq = sq_arr[0] + sq_arr[1] + sq_arr[2] + sq_arr[3];
for &v in chunks.remainder() {
sum += v;
sum_sq += v * v;
}
(sum, sum_sq)
};
// --- Seed: build initial mean and m2 for the first window using
// Welford's incremental (non-rolling) algorithm. ---
let mut mean = 0.0_f64;
let mut m2 = 0.0_f64;
for (k, &x) in close[..timeperiod].iter().enumerate() {
let count = (k + 1) as f64;
let delta = x - mean;
mean += delta / count;
let delta2 = x - mean;
m2 += delta * delta2;
}
#[cfg(not(feature = "simd"))]
let (mut sum, mut sum_sq) = {
let s: f64 = close[..timeperiod].iter().sum();
let sq: f64 = close[..timeperiod].iter().map(|&x| x * x).sum();
(s, sq)
};
let mean = sum / p;
let var = (sum_sq / p - mean * mean).max(0.0);
let var = (m2 / p).max(0.0);
let std = var.sqrt();
middle[timeperiod - 1] = mean;
upper[timeperiod - 1] = mean + nbdevup * std;
lower[timeperiod - 1] = mean - nbdevdn * std;
// --- Rolling phase: slide the window one element at a time,
// removing the oldest value and adding the newest. ---
/// Inline helper: replace `x_old` with `x_new` in the Welford accumulator
/// (constant window size `p`), then write band values into the output slots.
///
/// Combined rolling Welford update (window size stays constant at N):
///
/// ```text
/// delta = x_new - x_old
/// old_mean = mean
/// mean += delta / N
/// m2 += delta * ((x_new - mean) + (x_old - old_mean))
/// ```
///
/// This is algebraically equivalent to removing `x_old` and adding `x_new`
/// in two separate Welford steps, but avoids the intermediate N-1 state.
#[inline(always)]
#[allow(clippy::too_many_arguments)]
fn welford_step(
x_old: f64,
x_new: f64,
mean: &mut f64,
m2: &mut f64,
p: f64,
nbdevup: f64,
nbdevdn: f64,
upper: &mut f64,
middle: &mut f64,
lower: &mut f64,
) {
let delta = x_new - x_old;
let old_mean = *mean;
*mean += delta / p;
// Update m2 using both old and new deviations.
*m2 += delta * ((x_new - *mean) + (x_old - old_mean));
// Clamp m2 to zero to guard against floating-point drift.
if *m2 < 0.0 {
*m2 = 0.0;
}
let var = *m2 / p;
let std = var.sqrt();
*middle = *mean;
*upper = *mean + nbdevup * std;
*lower = *mean - nbdevdn * std;
}
// Process two iterations at a time (loop unrolling) for throughput.
let mut i = timeperiod;
while i + 1 < n {
let old0 = close[i - timeperiod];
sum += close[i] - old0;
sum_sq += close[i] * close[i] - old0 * old0;
let mean = sum / p;
let var = (sum_sq / p - mean * mean).max(0.0);
let std = var.sqrt();
middle[i] = mean;
upper[i] = mean + nbdevup * std;
lower[i] = mean - nbdevdn * std;
let old1 = close[i + 1 - timeperiod];
sum += close[i + 1] - old1;
sum_sq += close[i + 1] * close[i + 1] - old1 * old1;
let mean1 = sum / p;
let var1 = (sum_sq / p - mean1 * mean1).max(0.0);
let std1 = var1.sqrt();
middle[i + 1] = mean1;
upper[i + 1] = mean1 + nbdevup * std1;
lower[i + 1] = mean1 - nbdevdn * std1;
welford_step(
close[i - timeperiod],
close[i],
&mut mean,
&mut m2,
p,
nbdevup,
nbdevdn,
&mut upper[i],
&mut middle[i],
&mut lower[i],
);
welford_step(
close[i + 1 - timeperiod],
close[i + 1],
&mut mean,
&mut m2,
p,
nbdevup,
nbdevdn,
&mut upper[i + 1],
&mut middle[i + 1],
&mut lower[i + 1],
);
i += 2;
}
if i < n {
let old = close[i - timeperiod];
sum += close[i] - old;
sum_sq += close[i] * close[i] - old * old;
let mean = sum / p;
let var = (sum_sq / p - mean * mean).max(0.0);
let std = var.sqrt();
middle[i] = mean;
upper[i] = mean + nbdevup * std;
lower[i] = mean - nbdevdn * std;
welford_step(
close[i - timeperiod],
close[i],
&mut mean,
&mut m2,
p,
nbdevup,
nbdevdn,
&mut upper[i],
&mut middle[i],
&mut lower[i],
);
}
(upper, middle, lower)
}
/// MACD — EMA(fastperiod) minus EMA(slowperiod), signal = EMA(macd, signalperiod).
/// Compute the Moving Average Convergence/Divergence (MACD).
///
/// Returns `(macd_line, signal_line, histogram)`, each of length `n`.
/// Leading values are `NaN` during warmup.
/// `fastperiod` must be less than `slowperiod`.
/// `MACD = EMA(close, fastperiod) - EMA(close, slowperiod)`.
/// The signal line is `EMA(macd, signalperiod)` and the histogram is
/// `macd - signal`. TA-Lib compatible: leading values are `NaN` up to
/// the point where all three outputs are valid.
///
/// Fast and slow EMAs are computed in a **single combined loop** to minimise
/// memory round-trips, then the signal EMA is computed in a second pass.
/// # Arguments
/// * `close` - Price series.
/// * `fastperiod` - Fast EMA period (must be < `slowperiod`).
/// * `slowperiod` - Slow EMA period.
/// * `signalperiod` - Signal line EMA period.
///
/// # Returns
/// `(macd_line, signal_line, histogram)` -- each `Vec<f64>` of length `n`.
pub fn macd(
close: &[f64],
fastperiod: usize,
@@ -383,6 +490,74 @@ mod tests {
assert!((lower[2] - 2.0).abs() < 1e-10);
}
#[test]
fn bbands_varying_prices() {
// Verify against hand-computed values for a small window.
let prices = vec![1.0, 2.0, 3.0, 4.0, 5.0];
let (upper, middle, lower) = bbands(&prices, 3, 2.0, 2.0);
// First two values should be NaN (warmup).
assert!(middle[0].is_nan());
assert!(middle[1].is_nan());
// Window [1,2,3]: mean = 2.0, pop_var = 2/3, std = sqrt(2/3)
let expected_mean = 2.0;
let expected_std = (2.0_f64 / 3.0).sqrt();
assert!((middle[2] - expected_mean).abs() < 1e-10);
assert!((upper[2] - (expected_mean + 2.0 * expected_std)).abs() < 1e-10);
assert!((lower[2] - (expected_mean - 2.0 * expected_std)).abs() < 1e-10);
// Window [2,3,4]: mean = 3.0, pop_var = 2/3, std = sqrt(2/3)
assert!((middle[3] - 3.0).abs() < 1e-10);
assert!((upper[3] - (3.0 + 2.0 * expected_std)).abs() < 1e-10);
// Window [3,4,5]: mean = 4.0, pop_var = 2/3, std = sqrt(2/3)
assert!((middle[4] - 4.0).abs() < 1e-10);
assert!((upper[4] - (4.0 + 2.0 * expected_std)).abs() < 1e-10);
}
#[test]
fn bbands_numerical_stability() {
// Large offset with tiny variation — this is where the naïve sum_sq
// formula suffers from catastrophic cancellation.
let base = 1e12;
let prices: Vec<f64> = (0..100).map(|i| base + (i as f64) * 0.01).collect();
let (upper, middle, lower) = bbands(&prices, 20, 2.0, 2.0);
// Check that middle band matches SMA.
for i in 19..100 {
let window = &prices[i - 19..=i];
let expected_mean: f64 = window.iter().sum::<f64>() / 20.0;
assert!(
(middle[i] - expected_mean).abs() < 1e-4,
"mean mismatch at {i}: got {} expected {}",
middle[i],
expected_mean,
);
// Bands should be above/below middle.
assert!(upper[i] >= middle[i]);
assert!(lower[i] <= middle[i]);
}
}
#[test]
fn bbands_edge_cases() {
// timeperiod == 1: every bar should have std = 0, bands == price.
let prices = vec![10.0, 20.0, 30.0];
let (upper, middle, lower) = bbands(&prices, 1, 2.0, 2.0);
for i in 0..3 {
assert!((middle[i] - prices[i]).abs() < 1e-10);
assert!((upper[i] - prices[i]).abs() < 1e-10);
assert!((lower[i] - prices[i]).abs() < 1e-10);
}
// Input shorter than timeperiod: all NaN.
let (u, m, l) = bbands(&[1.0, 2.0], 5, 2.0, 2.0);
assert!(u.iter().all(|v| v.is_nan()));
assert!(m.iter().all(|v| v.is_nan()));
assert!(l.iter().all(|v| v.is_nan()));
}
#[test]
fn macd_basic() {
// 40 bars of linearly increasing prices — MACD line should converge
File diff suppressed because it is too large Load Diff
+627
View File
@@ -0,0 +1,627 @@
//! Pure Rust portfolio analytics — no PyO3, no numpy, no ndarray.
//!
//! Functions:
//! - `portfolio_volatility` — sqrt(w' Σ w)
//! - `beta_full` — Cov/Var OLS beta
//! - `rolling_beta` — rolling beta with NaN warmup
//! - `drawdown_series` — per-bar drawdown + max drawdown
//! - `correlation_matrix` — pairwise Pearson correlation
//! - `relative_strength` — cumulative return ratio
//! - `spread` — A - hedge * B
//! - `ratio` — A / B (NaN for zero)
//! - `zscore_series` — rolling z-score, NaN warmup
//! - `compose_weighted` — weighted sum per row
// ---------------------------------------------------------------------------
// portfolio_volatility
// ---------------------------------------------------------------------------
/// Compute portfolio volatility: sqrt(w' Σ w).
///
/// `cov_matrix` is an n×n covariance matrix stored as a slice of row-Vecs.
/// `weights` has length n.
///
/// Panics if dimensions are inconsistent.
pub fn portfolio_volatility(cov_matrix: &[Vec<f64>], weights: &[f64]) -> f64 {
let n = weights.len();
assert!(
cov_matrix.len() == n,
"cov_matrix must have {} rows, got {}",
n,
cov_matrix.len()
);
let mut variance = 0.0_f64;
for i in 0..n {
assert!(
cov_matrix[i].len() == n,
"cov_matrix row {} must have length {}, got {}",
i,
n,
cov_matrix[i].len()
);
let mut row_sum = 0.0_f64;
for j in 0..n {
row_sum += weights[j] * cov_matrix[i][j];
}
variance += weights[i] * row_sum;
}
variance.max(0.0).sqrt()
}
// ---------------------------------------------------------------------------
// beta_full
// ---------------------------------------------------------------------------
/// Compute the full-sample OLS beta of `asset_returns` vs `benchmark_returns`.
///
/// Beta = Cov(asset, bench) / Var(bench).
///
/// Panics if lengths differ or are < 2, or if benchmark has zero variance.
pub fn beta_full(asset_returns: &[f64], benchmark_returns: &[f64]) -> f64 {
let n = asset_returns.len();
assert!(
n >= 2 && benchmark_returns.len() == n,
"asset_returns and benchmark_returns must have equal length >= 2"
);
let mean_a: f64 = asset_returns.iter().sum::<f64>() / n as f64;
let mean_b: f64 = benchmark_returns.iter().sum::<f64>() / n as f64;
let mut cov = 0.0_f64;
let mut var_b = 0.0_f64;
for i in 0..n {
let da = asset_returns[i] - mean_a;
let db = benchmark_returns[i] - mean_b;
cov += da * db;
var_b += db * db;
}
assert!(var_b != 0.0, "benchmark_returns has zero variance; cannot compute beta");
cov / var_b
}
// ---------------------------------------------------------------------------
// rolling_beta
// ---------------------------------------------------------------------------
/// Compute rolling beta of `asset` vs `benchmark` over a sliding `window`.
///
/// Returns a Vec of the same length as the inputs. The first `window - 1`
/// entries are NaN (warmup period). `window` must be >= 2.
pub fn rolling_beta(asset: &[f64], benchmark: &[f64], window: usize) -> Vec<f64> {
assert!(window >= 2, "window must be >= 2");
let n = asset.len();
assert!(
n > 0 && benchmark.len() == n,
"asset and benchmark must be non-empty and equal length"
);
let mut result = vec![f64::NAN; n];
for i in (window - 1)..n {
let start = i + 1 - window;
let a_win = &asset[start..=i];
let b_win = &benchmark[start..=i];
let mean_a: f64 = a_win.iter().sum::<f64>() / window as f64;
let mean_b: f64 = b_win.iter().sum::<f64>() / window as f64;
let mut cov = 0.0_f64;
let mut var_b = 0.0_f64;
for k in 0..window {
let da = a_win[k] - mean_a;
let db = b_win[k] - mean_b;
cov += da * db;
var_b += db * db;
}
result[i] = if var_b == 0.0 { f64::NAN } else { cov / var_b };
}
result
}
// ---------------------------------------------------------------------------
// drawdown_series
// ---------------------------------------------------------------------------
/// Compute the drawdown series and maximum drawdown for an equity/price series.
///
/// Drawdown at bar i = (equity[i] - running_max) / running_max (always <= 0).
///
/// Returns `(dd_array, max_dd)` where `max_dd` is the most negative drawdown.
///
/// Panics if `equity` is empty.
pub fn drawdown_series(equity: &[f64]) -> (Vec<f64>, f64) {
let n = equity.len();
assert!(n > 0, "equity must be non-empty");
let mut dd = vec![0.0_f64; n];
let mut peak = equity[0];
let mut max_dd = 0.0_f64;
for i in 0..n {
if equity[i] > peak {
peak = equity[i];
}
let d = if peak == 0.0 {
0.0
} else {
(equity[i] - peak) / peak
};
dd[i] = d;
if d < max_dd {
max_dd = d;
}
}
(dd, max_dd)
}
// ---------------------------------------------------------------------------
// correlation_matrix
// ---------------------------------------------------------------------------
/// Compute the pairwise Pearson correlation matrix.
///
/// `data` is a slice of column vectors — `data[j]` is the return series for
/// asset j, so `data[j][i]` is the return of asset j at bar i. All columns
/// must have the same length (>= 2).
///
/// Returns an n_assets × n_assets matrix stored as `Vec<Vec<f64>>`.
pub fn correlation_matrix(data: &[Vec<f64>]) -> Vec<Vec<f64>> {
let n_assets = data.len();
assert!(n_assets > 0, "data must contain at least one asset column");
let n_bars = data[0].len();
assert!(n_bars >= 2, "data must have at least 2 rows (bars)");
for j in 1..n_assets {
assert!(
data[j].len() == n_bars,
"all columns must have equal length; column 0 has {} but column {} has {}",
n_bars,
j,
data[j].len()
);
}
// Means
let mut means = vec![0.0_f64; n_assets];
for j in 0..n_assets {
means[j] = data[j].iter().sum::<f64>() / n_bars as f64;
}
// Standard deviations (population)
let mut stds = vec![0.0_f64; n_assets];
for j in 0..n_assets {
let var: f64 = data[j].iter().map(|&v| (v - means[j]).powi(2)).sum::<f64>() / n_bars as f64;
stds[j] = var.sqrt();
}
// Build correlation matrix (exploit symmetry: compute each pair once)
let mut result = vec![vec![0.0_f64; n_assets]; n_assets];
for j1 in 0..n_assets {
result[j1][j1] = 1.0;
for j2 in (j1 + 1)..n_assets {
let mut cov = 0.0_f64;
for i in 0..n_bars {
cov += (data[j1][i] - means[j1]) * (data[j2][i] - means[j2]);
}
cov /= n_bars as f64;
let denom = stds[j1] * stds[j2];
let corr = if denom == 0.0 { f64::NAN } else { cov / denom };
result[j1][j2] = corr;
result[j2][j1] = corr;
}
}
result
}
// ---------------------------------------------------------------------------
// relative_strength
// ---------------------------------------------------------------------------
/// Compute relative strength of an asset vs a benchmark.
///
/// result[i] = cumprod(1 + asset_returns[0..=i]) / cumprod(1 + benchmark_returns[0..=i])
///
/// Panics if lengths differ or are zero.
pub fn relative_strength(asset_returns: &[f64], benchmark_returns: &[f64]) -> Vec<f64> {
let n = asset_returns.len();
assert!(
n > 0 && benchmark_returns.len() == n,
"asset_returns and benchmark_returns must be non-empty and equal length"
);
let mut result = vec![0.0_f64; n];
let mut cum_a = 1.0_f64;
let mut cum_b = 1.0_f64;
for i in 0..n {
cum_a *= 1.0 + asset_returns[i];
cum_b *= 1.0 + benchmark_returns[i];
result[i] = if cum_b == 0.0 { f64::NAN } else { cum_a / cum_b };
}
result
}
// ---------------------------------------------------------------------------
// spread
// ---------------------------------------------------------------------------
/// Compute the spread between two series: a - hedge * b.
///
/// Panics if lengths differ or are zero.
pub fn spread(a: &[f64], b: &[f64], hedge: f64) -> Vec<f64> {
let n = a.len();
assert!(
n > 0 && b.len() == n,
"a and b must be non-empty and equal length"
);
a.iter().zip(b.iter()).map(|(&x, &y)| x - hedge * y).collect()
}
// ---------------------------------------------------------------------------
// ratio
// ---------------------------------------------------------------------------
/// Compute the ratio between two series: a / b.
///
/// Where b is 0, returns NaN.
///
/// Panics if lengths differ or are zero.
pub fn ratio(a: &[f64], b: &[f64]) -> Vec<f64> {
let n = a.len();
assert!(
n > 0 && b.len() == n,
"a and b must be non-empty and equal length"
);
a.iter()
.zip(b.iter())
.map(|(&x, &y)| if y == 0.0 { f64::NAN } else { x / y })
.collect()
}
// ---------------------------------------------------------------------------
// zscore_series
// ---------------------------------------------------------------------------
/// Compute the rolling Z-score of a 1-D series.
///
/// Z[i] = (x[i] - mean(window)) / std(window)
///
/// The first `window - 1` entries are NaN. `window` must be >= 2.
///
/// Panics if `x` is empty or `window < 2`.
pub fn zscore_series(x: &[f64], window: usize) -> Vec<f64> {
assert!(window >= 2, "window must be >= 2");
let n = x.len();
assert!(n > 0, "x must be non-empty");
let mut result = vec![f64::NAN; n];
for i in (window - 1)..n {
let win = &x[i + 1 - window..=i];
let mean: f64 = win.iter().sum::<f64>() / window as f64;
let var: f64 = win.iter().map(|v| (v - mean).powi(2)).sum::<f64>() / window as f64;
let std = var.sqrt();
result[i] = if std == 0.0 {
f64::NAN
} else {
(x[i] - mean) / std
};
}
result
}
// ---------------------------------------------------------------------------
// compose_weighted
// ---------------------------------------------------------------------------
/// Weighted combination of multiple signal columns.
///
/// `data` is a slice of column vectors — `data[j]` is one signal column.
/// `weights` has one entry per column.
///
/// Returns a Vec of length n_bars where each entry is the weighted sum across
/// columns for that bar.
///
/// Panics if weights length != number of columns, or columns have unequal lengths.
pub fn compose_weighted(data: &[Vec<f64>], weights: &[f64]) -> Vec<f64> {
let n_sigs = data.len();
assert!(
weights.len() == n_sigs,
"weights length ({}) must equal number of signal columns ({})",
weights.len(),
n_sigs
);
if n_sigs == 0 {
return vec![];
}
let n_bars = data[0].len();
for j in 1..n_sigs {
assert!(
data[j].len() == n_bars,
"all columns must have equal length"
);
}
let mut result = vec![0.0_f64; n_bars];
for i in 0..n_bars {
let mut s = 0.0_f64;
for j in 0..n_sigs {
s += data[j][i] * weights[j];
}
result[i] = s;
}
result
}
// ---------------------------------------------------------------------------
// Tests
// ---------------------------------------------------------------------------
#[cfg(test)]
mod tests {
use super::*;
const EPS: f64 = 1e-10;
fn approx_eq(a: f64, b: f64) -> bool {
(a - b).abs() < EPS
}
// -- portfolio_volatility -------------------------------------------------
#[test]
fn test_portfolio_volatility_identity_cov() {
// Identity covariance, equal weights => sqrt(sum(w_i^2))
let cov = vec![
vec![1.0, 0.0],
vec![0.0, 1.0],
];
let w = vec![0.5, 0.5];
let vol = portfolio_volatility(&cov, &w);
// w' I w = 0.25 + 0.25 = 0.5, sqrt = 0.7071...
assert!(approx_eq(vol, (0.5_f64).sqrt()));
}
#[test]
fn test_portfolio_volatility_single_asset() {
let cov = vec![vec![0.04]];
let w = vec![1.0];
assert!(approx_eq(portfolio_volatility(&cov, &w), 0.2));
}
#[test]
fn test_portfolio_volatility_correlated() {
// Fully correlated: cov = [[0.04, 0.04], [0.04, 0.04]]
let cov = vec![
vec![0.04, 0.04],
vec![0.04, 0.04],
];
let w = vec![0.5, 0.5];
// w' Σ w = 0.04, sqrt = 0.2
let vol = portfolio_volatility(&cov, &w);
assert!(approx_eq(vol, 0.2));
}
// -- beta_full ------------------------------------------------------------
#[test]
fn test_beta_full_same_series() {
let r = vec![0.01, -0.02, 0.03, -0.01, 0.02];
assert!(approx_eq(beta_full(&r, &r), 1.0));
}
#[test]
fn test_beta_full_double() {
let bench = vec![0.01, -0.02, 0.03, -0.01, 0.02];
let asset: Vec<f64> = bench.iter().map(|x| x * 2.0).collect();
assert!(approx_eq(beta_full(&asset, &bench), 2.0));
}
#[test]
#[should_panic]
fn test_beta_full_zero_variance() {
let a = vec![0.01, 0.02];
let b = vec![0.05, 0.05]; // zero variance
beta_full(&a, &b);
}
// -- rolling_beta ---------------------------------------------------------
#[test]
fn test_rolling_beta_warmup_nan() {
let a = vec![0.01, -0.02, 0.03, -0.01, 0.02];
let b = vec![0.01, -0.02, 0.03, -0.01, 0.02];
let rb = rolling_beta(&a, &b, 3);
assert_eq!(rb.len(), 5);
assert!(rb[0].is_nan());
assert!(rb[1].is_nan());
// From index 2 onward, beta of identical series = 1.0
assert!(approx_eq(rb[2], 1.0));
assert!(approx_eq(rb[3], 1.0));
assert!(approx_eq(rb[4], 1.0));
}
#[test]
fn test_rolling_beta_double() {
let bench = vec![0.01, -0.02, 0.03, -0.01, 0.02];
let asset: Vec<f64> = bench.iter().map(|x| x * 3.0).collect();
let rb = rolling_beta(&asset, &bench, 3);
for i in 2..5 {
assert!(approx_eq(rb[i], 3.0));
}
}
// -- drawdown_series ------------------------------------------------------
#[test]
fn test_drawdown_series_monotonic_up() {
let eq = vec![100.0, 110.0, 120.0, 130.0];
let (dd, max_dd) = drawdown_series(&eq);
for &d in &dd {
assert!(approx_eq(d, 0.0));
}
assert!(approx_eq(max_dd, 0.0));
}
#[test]
fn test_drawdown_series_with_dip() {
let eq = vec![100.0, 120.0, 90.0, 110.0];
let (dd, max_dd) = drawdown_series(&eq);
assert!(approx_eq(dd[0], 0.0));
assert!(approx_eq(dd[1], 0.0));
// dd[2] = (90 - 120) / 120 = -0.25
assert!(approx_eq(dd[2], -0.25));
// dd[3] = (110 - 120) / 120 = -1/12
assert!((dd[3] - (-1.0 / 12.0)).abs() < EPS);
assert!(approx_eq(max_dd, -0.25));
}
// -- correlation_matrix ---------------------------------------------------
#[test]
fn test_correlation_matrix_identical() {
let col = vec![0.01, -0.02, 0.03, -0.01, 0.02];
let data = vec![col.clone(), col.clone()];
let cm = correlation_matrix(&data);
assert_eq!(cm.len(), 2);
assert!(approx_eq(cm[0][0], 1.0));
assert!(approx_eq(cm[1][1], 1.0));
assert!(approx_eq(cm[0][1], 1.0));
assert!(approx_eq(cm[1][0], 1.0));
}
#[test]
fn test_correlation_matrix_negatively_correlated() {
let col_a = vec![1.0, 2.0, 3.0, 4.0, 5.0];
let col_b: Vec<f64> = col_a.iter().map(|x| -x).collect();
let data = vec![col_a, col_b];
let cm = correlation_matrix(&data);
assert!(approx_eq(cm[0][1], -1.0));
assert!(approx_eq(cm[1][0], -1.0));
}
#[test]
fn test_correlation_matrix_single_asset() {
let data = vec![vec![1.0, 2.0, 3.0]];
let cm = correlation_matrix(&data);
assert_eq!(cm.len(), 1);
assert!(approx_eq(cm[0][0], 1.0));
}
// -- relative_strength ----------------------------------------------------
#[test]
fn test_relative_strength_equal() {
let r = vec![0.01, -0.02, 0.03];
let rs = relative_strength(&r, &r);
for &v in &rs {
assert!(approx_eq(v, 1.0));
}
}
#[test]
fn test_relative_strength_outperformance() {
let a = vec![0.10, 0.10];
let b = vec![0.05, 0.05];
let rs = relative_strength(&a, &b);
// rs[0] = 1.10 / 1.05
assert!((rs[0] - 1.10 / 1.05).abs() < EPS);
// rs[1] = 1.21 / 1.1025
assert!((rs[1] - 1.21 / 1.1025).abs() < EPS);
}
// -- spread ---------------------------------------------------------------
#[test]
fn test_spread_basic() {
let a = vec![10.0, 20.0, 30.0];
let b = vec![5.0, 10.0, 15.0];
let s = spread(&a, &b, 2.0);
assert!(approx_eq(s[0], 0.0));
assert!(approx_eq(s[1], 0.0));
assert!(approx_eq(s[2], 0.0));
}
#[test]
fn test_spread_hedge_one() {
let a = vec![10.0, 20.0];
let b = vec![3.0, 7.0];
let s = spread(&a, &b, 1.0);
assert!(approx_eq(s[0], 7.0));
assert!(approx_eq(s[1], 13.0));
}
// -- ratio ----------------------------------------------------------------
#[test]
fn test_ratio_basic() {
let a = vec![10.0, 20.0, 30.0];
let b = vec![5.0, 10.0, 15.0];
let r = ratio(&a, &b);
assert!(approx_eq(r[0], 2.0));
assert!(approx_eq(r[1], 2.0));
assert!(approx_eq(r[2], 2.0));
}
#[test]
fn test_ratio_zero_denominator() {
let a = vec![10.0, 20.0];
let b = vec![0.0, 5.0];
let r = ratio(&a, &b);
assert!(r[0].is_nan());
assert!(approx_eq(r[1], 4.0));
}
// -- zscore_series --------------------------------------------------------
#[test]
fn test_zscore_warmup_nan() {
let x = vec![1.0, 2.0, 3.0, 4.0, 5.0];
let z = zscore_series(&x, 3);
assert!(z[0].is_nan());
assert!(z[1].is_nan());
assert!(!z[2].is_nan());
assert!(!z[3].is_nan());
assert!(!z[4].is_nan());
}
#[test]
fn test_zscore_constant_window() {
// All same values in window => std = 0 => NaN
let x = vec![5.0, 5.0, 5.0, 5.0];
let z = zscore_series(&x, 3);
assert!(z[2].is_nan());
assert!(z[3].is_nan());
}
#[test]
fn test_zscore_known_value() {
// Window [1, 2, 3]: mean=2, pop_std = sqrt(2/3) ~0.8165
// z = (3 - 2) / sqrt(2/3) = sqrt(3/2) ~ 1.2247
let x = vec![1.0, 2.0, 3.0];
let z = zscore_series(&x, 3);
let expected = (3.0_f64 / 2.0).sqrt();
assert!((z[2] - expected).abs() < EPS);
}
// -- compose_weighted -----------------------------------------------------
#[test]
fn test_compose_weighted_basic() {
let data = vec![
vec![1.0, 2.0, 3.0],
vec![4.0, 5.0, 6.0],
];
let weights = vec![0.3, 0.7];
let cw = compose_weighted(&data, &weights);
// bar 0: 1*0.3 + 4*0.7 = 3.1
assert!(approx_eq(cw[0], 3.1));
// bar 1: 2*0.3 + 5*0.7 = 4.1
assert!(approx_eq(cw[1], 4.1));
// bar 2: 3*0.3 + 6*0.7 = 5.1
assert!(approx_eq(cw[2], 5.1));
}
#[test]
fn test_compose_weighted_single_column() {
let data = vec![vec![10.0, 20.0]];
let weights = vec![2.0];
let cw = compose_weighted(&data, &weights);
assert!(approx_eq(cw[0], 20.0));
assert!(approx_eq(cw[1], 40.0));
}
#[test]
fn test_compose_weighted_empty() {
let data: Vec<Vec<f64>> = vec![];
let weights: Vec<f64> = vec![];
let cw = compose_weighted(&data, &weights);
assert!(cw.is_empty());
}
}
@@ -0,0 +1,89 @@
//! Price transformations — synthesize OHLC arrays into single price arrays.
/// Average Price: (open + high + low + close) / 4.
pub fn avgprice(open: &[f64], high: &[f64], low: &[f64], close: &[f64]) -> Vec<f64> {
open.iter()
.zip(high.iter())
.zip(low.iter())
.zip(close.iter())
.map(|(((&o, &h), &l), &c)| (o + h + l + c) / 4.0)
.collect()
}
/// Median Price: (high + low) / 2.
pub fn medprice(high: &[f64], low: &[f64]) -> Vec<f64> {
high.iter()
.zip(low.iter())
.map(|(&h, &l)| (h + l) / 2.0)
.collect()
}
/// Typical Price: (high + low + close) / 3.
pub fn typprice(high: &[f64], low: &[f64], close: &[f64]) -> Vec<f64> {
high.iter()
.zip(low.iter())
.zip(close.iter())
.map(|((&h, &l), &c)| (h + l + c) / 3.0)
.collect()
}
/// Weighted Close Price: (high + low + close * 2) / 4.
pub fn wclprice(high: &[f64], low: &[f64], close: &[f64]) -> Vec<f64> {
high.iter()
.zip(low.iter())
.zip(close.iter())
.map(|((&h, &l), &c)| (h + l + c * 2.0) / 4.0)
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_avgprice() {
let o = vec![1.0, 2.0, 3.0];
let h = vec![4.0, 5.0, 6.0];
let l = vec![0.5, 1.5, 2.5];
let c = vec![2.5, 3.5, 4.5];
let result = avgprice(&o, &h, &l, &c);
assert_eq!(result.len(), 3);
assert!((result[0] - 2.0).abs() < 1e-10); // (1+4+0.5+2.5)/4 = 2.0
}
#[test]
fn test_medprice() {
let h = vec![10.0, 20.0];
let l = vec![6.0, 12.0];
let result = medprice(&h, &l);
assert!((result[0] - 8.0).abs() < 1e-10);
assert!((result[1] - 16.0).abs() < 1e-10);
}
#[test]
fn test_typprice() {
let h = vec![10.0];
let l = vec![6.0];
let c = vec![8.0];
let result = typprice(&h, &l, &c);
assert!((result[0] - 8.0).abs() < 1e-10); // (10+6+8)/3 = 8.0
}
#[test]
fn test_wclprice() {
let h = vec![10.0];
let l = vec![6.0];
let c = vec![8.0];
let result = wclprice(&h, &l, &c);
assert!((result[0] - 8.0).abs() < 1e-10); // (10+6+16)/4 = 8.0
}
#[test]
fn test_empty_inputs() {
let empty: Vec<f64> = vec![];
assert!(avgprice(&empty, &empty, &empty, &empty).is_empty());
assert!(medprice(&empty, &empty).is_empty());
assert!(typprice(&empty, &empty, &empty).is_empty());
assert!(wclprice(&empty, &empty, &empty).is_empty());
}
}
+171
View File
@@ -0,0 +1,171 @@
//! Regime detection and structural breaks.
//!
//! - `regime_adx` — label trend (1) vs range (0) using ADX threshold
//! - `regime_combined` — combine ADX + ATR-ratio for robust regime labelling
//! - `detect_breaks_cusum` — CUSUM-based structural break detection
//! - `rolling_variance_break` — variance ratio break detection
/// Label each bar as trend (1) or range (0) based on ADX level.
///
/// Returns `Vec<i8>`: `1` = trend (ADX > threshold), `0` = range, `-1` = NaN/warmup.
pub fn regime_adx(adx: &[f64], threshold: f64) -> Vec<i8> {
adx.iter()
.map(|&v| {
if v.is_nan() {
-1i8
} else if v > threshold {
1i8
} else {
0i8
}
})
.collect()
}
/// Label each bar as trend (1) or range (0) using ADX + ATR-ratio rule.
///
/// A bar is trending when: `adx[i] > adx_threshold` AND `atr[i] / close[i] > atr_pct_threshold`.
///
/// Returns `Vec<i8>`: `1` = trend, `0` = range, `-1` = NaN.
pub fn regime_combined(
adx: &[f64],
atr: &[f64],
close: &[f64],
adx_threshold: f64,
atr_pct_threshold: f64,
) -> Vec<i8> {
let n = adx.len();
(0..n)
.map(|i| {
let av = adx[i];
let rv = atr[i];
let cv = close[i];
if av.is_nan() || rv.is_nan() || cv.is_nan() || cv == 0.0 {
-1i8
} else if av > adx_threshold && (rv / cv) > atr_pct_threshold {
1i8
} else {
0i8
}
})
.collect()
}
/// Detect structural breaks using a CUSUM (cumulative sum) approach.
///
/// `window` must be >= 2. Returns `Vec<i8>`: `1` at break bars, `0` elsewhere.
pub fn detect_breaks_cusum(
series: &[f64],
window: usize,
threshold: f64,
slack: f64,
) -> Vec<i8> {
let n = series.len();
let mut out = vec![0i8; n];
if n < window || window < 2 {
return out;
}
let mut cusum_pos = 0.0_f64;
let mut cusum_neg = 0.0_f64;
for i in window..n {
let slice = &series[(i - window)..i];
let mean: f64 = slice.iter().sum::<f64>() / window as f64;
let var: f64 =
slice.iter().map(|&v| (v - mean) * (v - mean)).sum::<f64>() / (window - 1) as f64;
let std = var.sqrt();
if std == 0.0 || std.is_nan() || series[i].is_nan() {
continue;
}
let z = (series[i] - mean) / std;
cusum_pos = (cusum_pos + z - slack).max(0.0);
cusum_neg = (cusum_neg - z - slack).max(0.0);
if cusum_pos > threshold || cusum_neg > threshold {
out[i] = 1;
cusum_pos = 0.0;
cusum_neg = 0.0;
}
}
out
}
/// Detect volatility regime breaks using rolling variance ratio.
///
/// `short_window` must be >= 2, `long_window` must be > `short_window`.
/// Returns `Vec<i8>`: `1` at break bars, `0` elsewhere.
pub fn rolling_variance_break(
series: &[f64],
short_window: usize,
long_window: usize,
threshold: f64,
) -> Vec<i8> {
let n = series.len();
let mut out = vec![0i8; n];
if n < long_window || short_window < 2 || long_window <= short_window {
return out;
}
let variance = |slice: &[f64]| -> f64 {
let k = slice.len();
let mean: f64 = slice.iter().sum::<f64>() / k as f64;
slice.iter().map(|&v| (v - mean) * (v - mean)).sum::<f64>() / (k - 1) as f64
};
for i in long_window..n {
let long_slice = &series[(i - long_window)..i];
let short_slice = &series[(i - short_window)..i];
let long_var = variance(long_slice);
let short_var = variance(short_slice);
if long_var == 0.0 || long_var.is_nan() || short_var.is_nan() {
continue;
}
if short_var / long_var > threshold {
out[i] = 1;
}
}
out
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_regime_adx_basic() {
let adx = vec![f64::NAN, 20.0, 30.0, 10.0, 50.0];
let result = regime_adx(&adx, 25.0);
assert_eq!(result, vec![-1, 0, 1, 0, 1]);
}
#[test]
fn test_regime_combined() {
let adx = vec![30.0, 30.0, 10.0];
let atr = vec![1.0, 0.001, 1.0];
let close = vec![100.0, 100.0, 100.0];
let result = regime_combined(&adx, &atr, &close, 25.0, 0.005);
assert_eq!(result[0], 1); // ADX>25 and ATR/close=0.01>0.005
assert_eq!(result[1], 0); // ATR/close=0.00001 < 0.005
assert_eq!(result[2], 0); // ADX<25
}
#[test]
fn test_detect_breaks_cusum_short_input() {
let series = vec![1.0, 2.0];
let result = detect_breaks_cusum(&series, 5, 3.0, 0.5);
assert!(result.iter().all(|&v| v == 0));
}
#[test]
fn test_rolling_variance_break_short_input() {
let series = vec![1.0, 2.0, 3.0];
let result = rolling_variance_break(&series, 2, 5, 2.0);
assert!(result.iter().all(|&v| v == 0));
}
#[test]
fn test_empty() {
assert!(regime_adx(&[], 25.0).is_empty());
assert!(regime_combined(&[], &[], &[], 25.0, 0.005).is_empty());
assert!(detect_breaks_cusum(&[], 2, 3.0, 0.5).is_empty());
assert!(rolling_variance_break(&[], 2, 5, 2.0).is_empty());
}
}
+278
View File
@@ -0,0 +1,278 @@
//! Resampling — OHLCV resampling and multi-timeframe helpers, pure Rust.
//!
//! # Functions
//! - `volume_bars` — Aggregate OHLCV bars into bars of fixed volume size.
//! - `ohlcv_agg` — Aggregate OHLCV bars given contiguous integer group labels.
/// OHLCV 5-tuple return type alias.
type Ohlcv5 = (Vec<f64>, Vec<f64>, Vec<f64>, Vec<f64>, Vec<f64>);
// ---------------------------------------------------------------------------
// volume_bars
// ---------------------------------------------------------------------------
/// Aggregate OHLCV data into volume bars of a fixed volume threshold.
///
/// Each output bar accumulates input bars until `volume_threshold` units of
/// volume have been consumed. The resulting bar has:
/// - open = first open of the group
/// - high = max high of the group
/// - low = min low of the group
/// - close = last close of the group
/// - volume = sum of volumes (approximately `volume_threshold`)
///
/// Returns `(open, high, low, close, volume)`.
///
/// # Panics
/// Panics if arrays are empty, have unequal lengths, or `volume_threshold <= 0`.
pub fn volume_bars(
open: &[f64],
high: &[f64],
low: &[f64],
close: &[f64],
volume: &[f64],
volume_threshold: f64,
) -> Ohlcv5 {
assert!(volume_threshold > 0.0, "volume_threshold must be > 0");
let n = open.len();
assert!(n > 0, "input arrays must be non-empty");
assert!(
high.len() == n && low.len() == n && close.len() == n && volume.len() == n,
"all input arrays must have equal length"
);
let mut out_open: Vec<f64> = Vec::new();
let mut out_high: Vec<f64> = Vec::new();
let mut out_low: Vec<f64> = Vec::new();
let mut out_close: Vec<f64> = Vec::new();
let mut out_vol: Vec<f64> = Vec::new();
let mut bar_open = open[0];
let mut bar_high = high[0];
let mut bar_low = low[0];
let mut bar_close = close[0];
let mut bar_vol = volume[0];
for i in 1..n {
bar_high = bar_high.max(high[i]);
bar_low = bar_low.min(low[i]);
bar_close = close[i];
bar_vol += volume[i];
if bar_vol >= volume_threshold {
out_open.push(bar_open);
out_high.push(bar_high);
out_low.push(bar_low);
out_close.push(bar_close);
out_vol.push(bar_vol);
// Start new bar
if i + 1 < n {
bar_open = open[i + 1];
bar_high = high[i + 1];
bar_low = low[i + 1];
bar_close = close[i + 1];
bar_vol = volume[i + 1];
}
}
}
// Push any remaining partial bar
if bar_vol > 0.0 && out_vol.last().is_none_or(|&last| last != bar_vol) {
out_open.push(bar_open);
out_high.push(bar_high);
out_low.push(bar_low);
out_close.push(bar_close);
out_vol.push(bar_vol);
}
(out_open, out_high, out_low, out_close, out_vol)
}
// ---------------------------------------------------------------------------
// ohlcv_agg
// ---------------------------------------------------------------------------
/// Aggregate OHLCV bars by integer group labels.
///
/// Groups consecutive bars with the same label and computes:
/// - open = first open of the group
/// - high = max high of the group
/// - low = min low of the group
/// - close = last close of the group
/// - volume = sum of volumes
///
/// `labels` must be non-decreasing (groups are contiguous).
///
/// Returns `(open, high, low, close, volume)`.
///
/// # Panics
/// Panics if arrays are empty or have unequal lengths.
pub fn ohlcv_agg(
open: &[f64],
high: &[f64],
low: &[f64],
close: &[f64],
volume: &[f64],
labels: &[i64],
) -> Ohlcv5 {
let n = open.len();
assert!(n > 0, "input arrays must be non-empty");
assert!(
high.len() == n
&& low.len() == n
&& close.len() == n
&& volume.len() == n
&& labels.len() == n,
"all input arrays must have equal length"
);
let mut out_open: Vec<f64> = Vec::new();
let mut out_high: Vec<f64> = Vec::new();
let mut out_low: Vec<f64> = Vec::new();
let mut out_close: Vec<f64> = Vec::new();
let mut out_vol: Vec<f64> = Vec::new();
let mut cur_label = labels[0];
let mut bar_open = open[0];
let mut bar_high = high[0];
let mut bar_low = low[0];
let mut bar_close = close[0];
let mut bar_vol = volume[0];
for i in 1..n {
if labels[i] != cur_label {
out_open.push(bar_open);
out_high.push(bar_high);
out_low.push(bar_low);
out_close.push(bar_close);
out_vol.push(bar_vol);
cur_label = labels[i];
bar_open = open[i];
bar_high = high[i];
bar_low = low[i];
bar_close = close[i];
bar_vol = volume[i];
} else {
bar_high = bar_high.max(high[i]);
bar_low = bar_low.min(low[i]);
bar_close = close[i];
bar_vol += volume[i];
}
}
out_open.push(bar_open);
out_high.push(bar_high);
out_low.push(bar_low);
out_close.push(bar_close);
out_vol.push(bar_vol);
(out_open, out_high, out_low, out_close, out_vol)
}
// ---------------------------------------------------------------------------
// Tests
// ---------------------------------------------------------------------------
#[cfg(test)]
mod tests {
use super::*;
// -- volume_bars ---------------------------------------------------------
#[test]
fn test_volume_bars_basic() {
let o = [100.0, 101.0, 102.0, 103.0, 104.0];
let h = [105.0, 106.0, 107.0, 108.0, 109.0];
let l = [95.0, 96.0, 97.0, 98.0, 99.0];
let c = [101.0, 102.0, 103.0, 104.0, 105.0];
let v = [50.0, 60.0, 40.0, 70.0, 30.0];
// threshold 100: first bar covers indices 0..2 (vol=110>=100)
let (ro, rh, rl, rc, rv) = volume_bars(&o, &h, &l, &c, &v, 100.0);
assert!(rv.len() >= 2);
// First bar: vol = 50+60 = 110
assert!((rv[0] - 110.0).abs() < 1e-10);
assert!((ro[0] - 100.0).abs() < 1e-10);
assert!((rh[0] - 106.0).abs() < 1e-10);
assert!((rl[0] - 95.0).abs() < 1e-10);
assert!((rc[0] - 102.0).abs() < 1e-10);
}
#[test]
fn test_volume_bars_single_element() {
let (ro, rh, rl, rc, rv) =
volume_bars(&[10.0], &[12.0], &[8.0], &[11.0], &[50.0], 100.0);
assert_eq!(rv.len(), 1);
assert!((rv[0] - 50.0).abs() < 1e-10);
assert!((ro[0] - 10.0).abs() < 1e-10);
}
#[test]
#[should_panic(expected = "volume_threshold must be > 0")]
fn test_volume_bars_zero_threshold() {
volume_bars(&[1.0], &[1.0], &[1.0], &[1.0], &[1.0], 0.0);
}
#[test]
#[should_panic(expected = "input arrays must be non-empty")]
fn test_volume_bars_empty() {
volume_bars(&[], &[], &[], &[], &[], 100.0);
}
// -- ohlcv_agg -----------------------------------------------------------
#[test]
fn test_ohlcv_agg_basic() {
let o = [100.0, 101.0, 102.0, 103.0];
let h = [105.0, 106.0, 108.0, 109.0];
let l = [95.0, 96.0, 97.0, 98.0];
let c = [101.0, 102.0, 103.0, 104.0];
let v = [10.0, 20.0, 30.0, 40.0];
let labels: [i64; 4] = [0, 0, 1, 1];
let (ro, rh, rl, rc, rv) = ohlcv_agg(&o, &h, &l, &c, &v, &labels);
assert_eq!(ro.len(), 2);
// Group 0: open=100, high=max(105,106)=106, low=min(95,96)=95, close=102, vol=30
assert!((ro[0] - 100.0).abs() < 1e-10);
assert!((rh[0] - 106.0).abs() < 1e-10);
assert!((rl[0] - 95.0).abs() < 1e-10);
assert!((rc[0] - 102.0).abs() < 1e-10);
assert!((rv[0] - 30.0).abs() < 1e-10);
// Group 1: open=102, high=max(108,109)=109, low=min(97,98)=97, close=104, vol=70
assert!((ro[1] - 102.0).abs() < 1e-10);
assert!((rh[1] - 109.0).abs() < 1e-10);
assert!((rl[1] - 97.0).abs() < 1e-10);
assert!((rc[1] - 104.0).abs() < 1e-10);
assert!((rv[1] - 70.0).abs() < 1e-10);
}
#[test]
fn test_ohlcv_agg_single_group() {
let o = [100.0, 101.0];
let h = [105.0, 106.0];
let l = [95.0, 96.0];
let c = [101.0, 102.0];
let v = [10.0, 20.0];
let labels: [i64; 2] = [0, 0];
let (ro, rh, rl, rc, rv) = ohlcv_agg(&o, &h, &l, &c, &v, &labels);
assert_eq!(ro.len(), 1);
assert!((rv[0] - 30.0).abs() < 1e-10);
}
#[test]
fn test_ohlcv_agg_each_bar_own_group() {
let o = [100.0, 101.0, 102.0];
let h = [105.0, 106.0, 107.0];
let l = [95.0, 96.0, 97.0];
let c = [101.0, 102.0, 103.0];
let v = [10.0, 20.0, 30.0];
let labels: [i64; 3] = [0, 1, 2];
let (ro, _rh, _rl, _rc, rv) = ohlcv_agg(&o, &h, &l, &c, &v, &labels);
assert_eq!(ro.len(), 3);
assert!((rv[0] - 10.0).abs() < 1e-10);
assert!((rv[1] - 20.0).abs() < 1e-10);
assert!((rv[2] - 30.0).abs() < 1e-10);
}
#[test]
#[should_panic(expected = "input arrays must be non-empty")]
fn test_ohlcv_agg_empty() {
ohlcv_agg(&[], &[], &[], &[], &[], &[]);
}
}
+131
View File
@@ -0,0 +1,131 @@
//! Signal processing helpers.
//!
//! - `rank_values` — fractional rank of a slice (1-based, ties averaged)
//! - `compose_rank` — rank-based composite scores for a 2-D signal matrix
//! - `top_n_indices` — indices of the N largest values
//! - `bottom_n_indices` — indices of the N smallest values
/// Compute fractional rank of each element (1-based, ascending).
/// Ties receive the average of their rank positions.
pub fn rank_values(x: &[f64]) -> Vec<f64> {
let n = x.len();
let mut order: Vec<usize> = (0..n).collect();
order.sort_by(|&a, &b| x[a].partial_cmp(&x[b]).unwrap_or(std::cmp::Ordering::Equal));
let mut ranks = vec![0.0_f64; n];
let mut i = 0;
while i < n {
let val = x[order[i]];
let mut j = i + 1;
while j < n && x[order[j]] == val {
j += 1;
}
let avg_rank = (i + 1 + j) as f64 / 2.0;
for k in i..j {
ranks[order[k]] = avg_rank;
}
i = j;
}
ranks
}
/// Compute rank-based composite scores for a 2-D signal matrix.
///
/// Each column is ranked independently, and the per-row ranks are summed.
/// `signals` is a slice of columns, each column being a `&[f64]` of the same length.
pub fn compose_rank(signals: &[&[f64]]) -> Vec<f64> {
if signals.is_empty() {
return vec![];
}
let n_bars = signals[0].len();
let mut scores = vec![0.0_f64; n_bars];
for &column in signals {
let ranks = rank_values(column);
for (bar_idx, rank) in ranks.into_iter().enumerate() {
scores[bar_idx] += rank;
}
}
scores
}
/// Return the indices of the N largest values in `x` (descending by value).
pub fn top_n_indices(x: &[f64], n: usize) -> Vec<i64> {
let len = x.len();
let k = n.min(len);
let mut order: Vec<usize> = (0..len).collect();
order.sort_by(|&a, &b| x[b].partial_cmp(&x[a]).unwrap_or(std::cmp::Ordering::Equal));
order[..k].iter().map(|&i| i as i64).collect()
}
/// Return the indices of the N smallest values in `x` (ascending by value).
pub fn bottom_n_indices(x: &[f64], n: usize) -> Vec<i64> {
let len = x.len();
let k = n.min(len);
let mut order: Vec<usize> = (0..len).collect();
order.sort_by(|&a, &b| x[a].partial_cmp(&x[b]).unwrap_or(std::cmp::Ordering::Equal));
order[..k].iter().map(|&i| i as i64).collect()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_rank_values() {
let x = vec![3.0, 1.0, 2.0];
let ranks = rank_values(&x);
assert!((ranks[0] - 3.0).abs() < 1e-10); // 3.0 is largest → rank 3
assert!((ranks[1] - 1.0).abs() < 1e-10); // 1.0 is smallest → rank 1
assert!((ranks[2] - 2.0).abs() < 1e-10); // 2.0 is middle → rank 2
}
#[test]
fn test_rank_values_ties() {
let x = vec![1.0, 2.0, 2.0, 4.0];
let ranks = rank_values(&x);
assert!((ranks[0] - 1.0).abs() < 1e-10);
assert!((ranks[1] - 2.5).abs() < 1e-10); // tied → average
assert!((ranks[2] - 2.5).abs() < 1e-10);
assert!((ranks[3] - 4.0).abs() < 1e-10);
}
#[test]
fn test_compose_rank() {
let col1 = vec![3.0, 1.0, 2.0];
let col2 = vec![1.0, 3.0, 2.0];
let signals: Vec<&[f64]> = vec![&col1, &col2];
let scores = compose_rank(&signals);
// Row 0: rank(3.0)=3 + rank(1.0)=1 = 4
// Row 1: rank(1.0)=1 + rank(3.0)=3 = 4
// Row 2: rank(2.0)=2 + rank(2.0)=2 = 4
assert!((scores[0] - 4.0).abs() < 1e-10);
assert!((scores[1] - 4.0).abs() < 1e-10);
assert!((scores[2] - 4.0).abs() < 1e-10);
}
#[test]
fn test_top_n_indices() {
let x = vec![10.0, 50.0, 30.0, 20.0, 40.0];
let result = top_n_indices(&x, 3);
assert_eq!(result.len(), 3);
assert_eq!(result[0], 1); // 50.0
assert_eq!(result[1], 4); // 40.0
assert_eq!(result[2], 2); // 30.0
}
#[test]
fn test_bottom_n_indices() {
let x = vec![10.0, 50.0, 30.0, 20.0, 40.0];
let result = bottom_n_indices(&x, 2);
assert_eq!(result.len(), 2);
assert_eq!(result[0], 0); // 10.0
assert_eq!(result[1], 3); // 20.0
}
#[test]
fn test_top_n_exceeds_len() {
let x = vec![1.0, 2.0];
let result = top_n_indices(&x, 5);
assert_eq!(result.len(), 2);
}
}
+9 -1
View File
@@ -1,6 +1,14 @@
//! Statistic functions.
/// Standard deviation — population (`ddof = 0`).
/// Compute the rolling population standard deviation, scaled by `nbdev`.
///
/// Uses population variance (`ddof = 0`). Returns `nbdev * stddev` for
/// each window. The first `timeperiod - 1` values are `NaN`.
///
/// # Arguments
/// * `real` - Input series.
/// * `timeperiod` - Rolling window size (must be >= 1).
/// * `nbdev` - Multiplier applied to the standard deviation (use 1.0 for raw stddev).
pub fn stddev(real: &[f64], timeperiod: usize, nbdev: f64) -> Vec<f64> {
let n = real.len();
let mut result = vec![f64::NAN; n];
+947
View File
@@ -0,0 +1,947 @@
//! Streaming / Incremental Indicators — bar-by-bar stateful structs.
//!
//! Pure Rust implementations with no PyO3 dependency. Each struct:
//! - Accepts one value per call to `update()`.
//! - Returns `NaN` (or a NaN tuple) during the warm-up window.
//! - Exposes a `reset()` method to restart from scratch.
//! - Has a `period()` accessor (where applicable).
use std::collections::VecDeque;
// ---------------------------------------------------------------------------
// Error type
// ---------------------------------------------------------------------------
/// Validation error for streaming indicator parameters.
#[derive(Debug, Clone)]
pub struct StreamingError(pub String);
impl std::fmt::Display for StreamingError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{}", self.0)
}
}
impl std::error::Error for StreamingError {}
fn validate_timeperiod(value: usize, name: &str, minimum: usize) -> Result<(), StreamingError> {
if value < minimum {
return Err(StreamingError(format!(
"{} must be >= {}, got {}",
name, minimum, value
)));
}
Ok(())
}
// ---------------------------------------------------------------------------
// Internal helper: EMA state (used inside composite classes)
// ---------------------------------------------------------------------------
/// SMA-seeded EMA state machine. Not exposed directly — used by
/// `StreamingEMA`, `StreamingMACD`, etc.
pub(crate) struct EmaState {
period: usize,
alpha: f64,
ema: f64,
seed_buf: Vec<f64>,
seeded: bool,
}
impl EmaState {
pub fn new(period: usize) -> Self {
Self {
period,
alpha: 2.0 / (period as f64 + 1.0),
ema: 0.0,
seed_buf: Vec::with_capacity(period),
seeded: false,
}
}
pub fn update(&mut self, value: f64) -> f64 {
if !self.seeded {
self.seed_buf.push(value);
if self.seed_buf.len() < self.period {
return f64::NAN;
}
let seed = self.seed_buf.iter().sum::<f64>() / self.period as f64;
self.ema = seed;
self.seeded = true;
return seed;
}
self.ema += self.alpha * (value - self.ema);
self.ema
}
pub fn reset(&mut self) {
self.ema = 0.0;
self.seed_buf.clear();
self.seeded = false;
}
pub fn period(&self) -> usize {
self.period
}
}
// ---------------------------------------------------------------------------
// Internal helper: ATR state (Wilder smoothing)
// ---------------------------------------------------------------------------
/// Wilder-smoothed ATR state machine. Used by `StreamingATR` and
/// `StreamingSupertrend`.
pub(crate) struct AtrState {
period: usize,
prev_close: f64,
tr_buf: Vec<f64>,
atr: f64,
seeded: bool,
has_prev: bool,
}
impl AtrState {
pub fn new(period: usize) -> Self {
Self {
period,
prev_close: 0.0,
tr_buf: Vec::with_capacity(period),
atr: 0.0,
seeded: false,
has_prev: false,
}
}
pub fn update(&mut self, high: f64, low: f64, close: f64) -> f64 {
let tr = if self.has_prev {
let hl = high - low;
let hc = (high - self.prev_close).abs();
let lc = (low - self.prev_close).abs();
hl.max(hc).max(lc)
} else {
high - low
};
self.prev_close = close;
self.has_prev = true;
if !self.seeded {
self.tr_buf.push(tr);
if self.tr_buf.len() < self.period {
return f64::NAN;
}
let seed = self.tr_buf.iter().sum::<f64>() / self.period as f64;
self.atr = seed;
self.seeded = true;
return f64::NAN; // first `period` bars (including this one) return NaN
}
let pf = (self.period - 1) as f64;
self.atr = (self.atr * pf + tr) / self.period as f64;
self.atr
}
pub fn reset(&mut self) {
self.prev_close = 0.0;
self.has_prev = false;
self.tr_buf.clear();
self.atr = 0.0;
self.seeded = false;
}
pub fn period(&self) -> usize {
self.period
}
}
// ---------------------------------------------------------------------------
// StreamingSMA
// ---------------------------------------------------------------------------
/// Simple Moving Average — O(1) per update via running sum.
///
/// Returns NaN during the first `period - 1` bars.
pub struct StreamingSMA {
period: usize,
buf: VecDeque<f64>,
running_sum: f64,
count: usize,
}
impl StreamingSMA {
pub fn new(period: usize) -> Result<Self, StreamingError> {
validate_timeperiod(period, "period", 1)?;
Ok(Self {
period,
buf: VecDeque::with_capacity(period + 1),
running_sum: 0.0,
count: 0,
})
}
/// Add a new bar and return the current SMA (NaN during warmup).
pub fn update(&mut self, value: f64) -> f64 {
if self.buf.len() == self.period {
if let Some(old) = self.buf.pop_front() {
self.running_sum -= old;
}
}
self.buf.push_back(value);
self.running_sum += value;
self.count += 1;
if self.count < self.period {
f64::NAN
} else {
self.running_sum / self.period as f64
}
}
/// Reset state to initial condition.
pub fn reset(&mut self) {
self.buf.clear();
self.running_sum = 0.0;
self.count = 0;
}
pub fn period(&self) -> usize {
self.period
}
}
// ---------------------------------------------------------------------------
// StreamingEMA
// ---------------------------------------------------------------------------
/// Exponential Moving Average with SMA seeding.
///
/// Uses a simple SMA for the first `period` bars to seed the EMA, then
/// switches to the standard EMA formula (alpha = 2 / (period + 1)).
/// Returns NaN during the warmup window.
pub struct StreamingEMA {
inner: EmaState,
}
impl StreamingEMA {
pub fn new(period: usize) -> Result<Self, StreamingError> {
validate_timeperiod(period, "period", 1)?;
Ok(Self {
inner: EmaState::new(period),
})
}
/// Add a new bar and return the current EMA (NaN during warmup).
pub fn update(&mut self, value: f64) -> f64 {
self.inner.update(value)
}
pub fn reset(&mut self) {
self.inner.reset();
}
pub fn period(&self) -> usize {
self.inner.period()
}
}
// ---------------------------------------------------------------------------
// StreamingRSI
// ---------------------------------------------------------------------------
/// Relative Strength Index with TA-Lib-compatible Wilder seeding.
///
/// Returns NaN during the first `period` bars.
pub struct StreamingRSI {
period: usize,
prev: f64,
has_prev: bool,
gains: Vec<f64>,
losses: Vec<f64>,
avg_gain: f64,
avg_loss: f64,
seeded: bool,
}
impl StreamingRSI {
pub fn new(period: usize) -> Result<Self, StreamingError> {
validate_timeperiod(period, "period", 1)?;
Ok(Self {
period,
prev: 0.0,
has_prev: false,
gains: Vec::with_capacity(period),
losses: Vec::with_capacity(period),
avg_gain: 0.0,
avg_loss: 0.0,
seeded: false,
})
}
/// Add a new close and return RSI in [0, 100] (NaN during warmup).
pub fn update(&mut self, value: f64) -> f64 {
if !self.has_prev {
self.prev = value;
self.has_prev = true;
return f64::NAN;
}
let delta = value - self.prev;
self.prev = value;
let gain = if delta > 0.0 { delta } else { 0.0 };
let loss = if delta < 0.0 { -delta } else { 0.0 };
if !self.seeded {
self.gains.push(gain);
self.losses.push(loss);
if self.gains.len() < self.period {
return f64::NAN;
}
self.avg_gain = self.gains.iter().sum::<f64>() / self.period as f64;
self.avg_loss = self.losses.iter().sum::<f64>() / self.period as f64;
self.seeded = true;
} else {
let pf = (self.period - 1) as f64;
self.avg_gain = (self.avg_gain * pf + gain) / self.period as f64;
self.avg_loss = (self.avg_loss * pf + loss) / self.period as f64;
}
if self.avg_loss == 0.0 {
return 100.0;
}
let rs = self.avg_gain / self.avg_loss;
100.0 - 100.0 / (1.0 + rs)
}
pub fn reset(&mut self) {
self.prev = 0.0;
self.has_prev = false;
self.gains.clear();
self.losses.clear();
self.avg_gain = 0.0;
self.avg_loss = 0.0;
self.seeded = false;
}
pub fn period(&self) -> usize {
self.period
}
}
// ---------------------------------------------------------------------------
// StreamingATR
// ---------------------------------------------------------------------------
/// Average True Range with TA-Lib-compatible Wilder seeding.
///
/// Accepts (high, low, close) per bar.
/// Returns NaN during the first `period` bars.
pub struct StreamingATR {
inner: AtrState,
}
impl StreamingATR {
pub fn new(period: usize) -> Result<Self, StreamingError> {
validate_timeperiod(period, "period", 1)?;
Ok(Self {
inner: AtrState::new(period),
})
}
/// Add a new bar (high, low, close) and return ATR (NaN during warmup).
pub fn update(&mut self, high: f64, low: f64, close: f64) -> f64 {
self.inner.update(high, low, close)
}
pub fn reset(&mut self) {
self.inner.reset();
}
pub fn period(&self) -> usize {
self.inner.period()
}
}
// ---------------------------------------------------------------------------
// StreamingBBands
// ---------------------------------------------------------------------------
/// Bollinger Bands — streaming variant using Welford's online algorithm.
///
/// Returns (upper, middle, lower).
/// NaN tuple during the warmup window.
pub struct StreamingBBands {
period: usize,
nbdevup: f64,
nbdevdn: f64,
buf: VecDeque<f64>,
mean: f64,
m2: f64,
}
impl StreamingBBands {
pub fn new(period: usize, nbdevup: f64, nbdevdn: f64) -> Result<Self, StreamingError> {
validate_timeperiod(period, "period", 2)?;
Ok(Self {
period,
nbdevup,
nbdevdn,
buf: VecDeque::with_capacity(period + 1),
mean: 0.0,
m2: 0.0,
})
}
/// Add a new bar; return (upper, middle, lower). NaN tuple during warmup.
pub fn update(&mut self, value: f64) -> (f64, f64, f64) {
let n = self.buf.len();
if n == self.period {
let x_old = self.buf.pop_front().unwrap();
let count = self.period as f64;
let delta_old = x_old - self.mean;
self.mean -= delta_old / (count - 1.0);
let delta2_old = x_old - self.mean;
self.m2 -= delta_old * delta2_old;
}
self.buf.push_back(value);
let count = self.buf.len() as f64;
let delta_new = value - self.mean;
self.mean += delta_new / count;
let delta2_new = value - self.mean;
self.m2 += delta_new * delta2_new;
if self.m2 < 0.0 {
self.m2 = 0.0;
}
if self.buf.len() < self.period {
return (f64::NAN, f64::NAN, f64::NAN);
}
let variance = self.m2 / (count - 1.0);
let std = variance.sqrt();
(
self.mean + self.nbdevup * std,
self.mean,
self.mean - self.nbdevdn * std,
)
}
pub fn reset(&mut self) {
self.buf.clear();
self.mean = 0.0;
self.m2 = 0.0;
}
pub fn period(&self) -> usize {
self.period
}
}
// ---------------------------------------------------------------------------
// StreamingMACD
// ---------------------------------------------------------------------------
/// MACD — fast EMA, slow EMA, signal EMA.
///
/// Returns (macd_line, signal_line, histogram).
/// NaN values during warmup.
pub struct StreamingMACD {
fast: EmaState,
slow: EmaState,
signal: EmaState,
}
impl StreamingMACD {
pub fn new(
fastperiod: usize,
slowperiod: usize,
signalperiod: usize,
) -> Result<Self, StreamingError> {
validate_timeperiod(fastperiod, "fastperiod", 1)?;
validate_timeperiod(slowperiod, "slowperiod", 1)?;
validate_timeperiod(signalperiod, "signalperiod", 1)?;
if fastperiod >= slowperiod {
return Err(StreamingError(
"fastperiod must be < slowperiod".to_string(),
));
}
Ok(Self {
fast: EmaState::new(fastperiod),
slow: EmaState::new(slowperiod),
signal: EmaState::new(signalperiod),
})
}
/// Add a new close; return (macd_line, signal_line, histogram).
pub fn update(&mut self, value: f64) -> (f64, f64, f64) {
let fast_val = self.fast.update(value);
let slow_val = self.slow.update(value);
if slow_val.is_nan() {
return (f64::NAN, f64::NAN, f64::NAN);
}
let macd = fast_val - slow_val;
let signal = self.signal.update(macd);
if signal.is_nan() {
return (macd, f64::NAN, f64::NAN);
}
(macd, signal, macd - signal)
}
pub fn reset(&mut self) {
self.fast.reset();
self.slow.reset();
self.signal.reset();
}
pub fn fast_period(&self) -> usize {
self.fast.period()
}
pub fn slow_period(&self) -> usize {
self.slow.period()
}
pub fn signal_period(&self) -> usize {
self.signal.period()
}
}
// ---------------------------------------------------------------------------
// StreamingStoch
// ---------------------------------------------------------------------------
/// Slow Stochastic (SMA-smoothed).
///
/// Returns (slowk, slowd).
/// NaN tuple during warmup.
pub struct StreamingStoch {
fastk_period: usize,
slowk_period: usize,
slowd_period: usize,
high_buf: VecDeque<f64>,
low_buf: VecDeque<f64>,
fastk_buf: VecDeque<f64>,
slowk_buf: VecDeque<f64>,
}
impl StreamingStoch {
pub fn new(
fastk_period: usize,
slowk_period: usize,
slowd_period: usize,
) -> Result<Self, StreamingError> {
validate_timeperiod(fastk_period, "fastk_period", 1)?;
validate_timeperiod(slowk_period, "slowk_period", 1)?;
validate_timeperiod(slowd_period, "slowd_period", 1)?;
Ok(Self {
fastk_period,
slowk_period,
slowd_period,
high_buf: VecDeque::with_capacity(fastk_period + 1),
low_buf: VecDeque::with_capacity(fastk_period + 1),
fastk_buf: VecDeque::with_capacity(slowk_period + 1),
slowk_buf: VecDeque::with_capacity(slowd_period + 1),
})
}
/// Add a new bar (high, low, close); return (slowk, slowd).
pub fn update(&mut self, high: f64, low: f64, close: f64) -> (f64, f64) {
if self.high_buf.len() == self.fastk_period {
self.high_buf.pop_front();
self.low_buf.pop_front();
}
self.high_buf.push_back(high);
self.low_buf.push_back(low);
if self.high_buf.len() < self.fastk_period {
return (f64::NAN, f64::NAN);
}
let max_h = self
.high_buf
.iter()
.cloned()
.fold(f64::NEG_INFINITY, f64::max);
let min_l = self.low_buf.iter().cloned().fold(f64::INFINITY, f64::min);
let fastk = if max_h != min_l {
100.0 * (close - min_l) / (max_h - min_l)
} else {
0.0
};
if self.fastk_buf.len() == self.slowk_period {
self.fastk_buf.pop_front();
}
self.fastk_buf.push_back(fastk);
if self.fastk_buf.len() < self.slowk_period {
return (f64::NAN, f64::NAN);
}
let slowk = self.fastk_buf.iter().sum::<f64>() / self.slowk_period as f64;
if self.slowk_buf.len() == self.slowd_period {
self.slowk_buf.pop_front();
}
self.slowk_buf.push_back(slowk);
if self.slowk_buf.len() < self.slowd_period {
return (slowk, f64::NAN);
}
let slowd = self.slowk_buf.iter().sum::<f64>() / self.slowd_period as f64;
(slowk, slowd)
}
pub fn reset(&mut self) {
self.high_buf.clear();
self.low_buf.clear();
self.fastk_buf.clear();
self.slowk_buf.clear();
}
pub fn period(&self) -> usize {
self.fastk_period
}
}
// ---------------------------------------------------------------------------
// StreamingVWAP
// ---------------------------------------------------------------------------
/// Cumulative Volume Weighted Average Price.
///
/// Resets automatically when `reset()` is called (e.g. at session open).
/// Accepts (high, low, close, volume) per bar.
#[derive(Default)]
pub struct StreamingVWAP {
cum_tpv: f64,
cum_vol: f64,
}
impl StreamingVWAP {
pub fn new() -> Self {
Self {
cum_tpv: 0.0,
cum_vol: 0.0,
}
}
/// Add a new bar (high, low, close, volume) and return cumulative VWAP.
pub fn update(&mut self, high: f64, low: f64, close: f64, volume: f64) -> f64 {
let tp = (high + low + close) / 3.0;
self.cum_tpv += tp * volume;
self.cum_vol += volume;
if self.cum_vol == 0.0 {
f64::NAN
} else {
self.cum_tpv / self.cum_vol
}
}
/// Reset for a new session.
pub fn reset(&mut self) {
self.cum_tpv = 0.0;
self.cum_vol = 0.0;
}
}
// ---------------------------------------------------------------------------
// StreamingSupertrend
// ---------------------------------------------------------------------------
/// ATR-based Supertrend — streaming variant.
///
/// Accepts (high, low, close) per bar.
/// Returns (supertrend_line, direction).
/// direction: 1 = uptrend, -1 = downtrend, 0 = warmup.
pub struct StreamingSupertrend {
period: usize,
multiplier: f64,
atr: AtrState,
upper_band: f64,
lower_band: f64,
has_bands: bool,
direction: i8,
prev_close: f64,
has_prev: bool,
}
impl StreamingSupertrend {
pub fn new(period: usize, multiplier: f64) -> Result<Self, StreamingError> {
validate_timeperiod(period, "period", 1)?;
Ok(Self {
period,
multiplier,
atr: AtrState::new(period),
upper_band: 0.0,
lower_band: 0.0,
has_bands: false,
direction: 0,
prev_close: 0.0,
has_prev: false,
})
}
/// Add a new bar (high, low, close); return (supertrend_line, direction).
pub fn update(&mut self, high: f64, low: f64, close: f64) -> (f64, i8) {
let atr = self.atr.update(high, low, close);
if atr.is_nan() {
self.prev_close = close;
self.has_prev = true;
return (f64::NAN, 0);
}
let hl2 = (high + low) / 2.0;
let upper_basic = hl2 + self.multiplier * atr;
let lower_basic = hl2 - self.multiplier * atr;
if !self.has_bands {
self.upper_band = upper_basic;
self.lower_band = lower_basic;
self.has_bands = true;
self.direction = -1;
self.prev_close = close;
self.has_prev = true;
return (self.upper_band, self.direction);
}
let prev_close = self.prev_close;
let new_lower = if lower_basic > self.lower_band || prev_close < self.lower_band {
lower_basic
} else {
self.lower_band
};
let new_upper = if upper_basic < self.upper_band || prev_close > self.upper_band {
upper_basic
} else {
self.upper_band
};
self.lower_band = new_lower;
self.upper_band = new_upper;
self.direction = if self.direction == -1 {
if close > new_upper {
1
} else {
-1
}
} else if close < new_lower {
-1
} else {
1
};
self.prev_close = close;
let line = if self.direction == 1 {
new_lower
} else {
new_upper
};
(line, self.direction)
}
pub fn reset(&mut self) {
self.atr.reset();
self.upper_band = 0.0;
self.lower_band = 0.0;
self.has_bands = false;
self.direction = 0;
self.prev_close = 0.0;
self.has_prev = false;
}
pub fn period(&self) -> usize {
self.period
}
}
// ---------------------------------------------------------------------------
// Tests
// ---------------------------------------------------------------------------
#[cfg(test)]
mod tests {
use super::*;
/// Helper: compare two f64 values, treating NaN == NaN as true.
fn approx_eq(a: f64, b: f64, tol: f64) -> bool {
if a.is_nan() && b.is_nan() {
return true;
}
(a - b).abs() < tol
}
#[test]
fn test_sma_basic() {
let mut sma = StreamingSMA::new(3).unwrap();
assert!(sma.update(1.0).is_nan());
assert!(sma.update(2.0).is_nan());
let v = sma.update(3.0);
assert!(approx_eq(v, 2.0, 1e-10));
let v = sma.update(4.0);
assert!(approx_eq(v, 3.0, 1e-10));
let v = sma.update(5.0);
assert!(approx_eq(v, 4.0, 1e-10));
assert_eq!(sma.period(), 3);
}
#[test]
fn test_sma_reset() {
let mut sma = StreamingSMA::new(2).unwrap();
sma.update(10.0);
sma.update(20.0);
sma.reset();
assert!(sma.update(5.0).is_nan());
let v = sma.update(7.0);
assert!(approx_eq(v, 6.0, 1e-10));
}
#[test]
fn test_ema_warmup_and_decay() {
let mut ema = StreamingEMA::new(3).unwrap();
assert!(ema.update(2.0).is_nan());
assert!(ema.update(4.0).is_nan());
// Third bar: SMA seed = (2+4+6)/3 = 4.0
let v = ema.update(6.0);
assert!(approx_eq(v, 4.0, 1e-10));
// Fourth bar: alpha = 0.5, ema = 4.0 + 0.5*(8.0-4.0) = 6.0
let v = ema.update(8.0);
assert!(approx_eq(v, 6.0, 1e-10));
}
#[test]
fn test_rsi_warmup() {
let mut rsi = StreamingRSI::new(3).unwrap();
// First bar: no prev
assert!(rsi.update(44.0).is_nan());
// Bars 2-4: collecting gains/losses
assert!(rsi.update(44.5).is_nan());
assert!(rsi.update(43.5).is_nan());
// Bar 5: seeded
let v = rsi.update(44.5);
assert!(!v.is_nan());
assert!(v >= 0.0 && v <= 100.0);
}
#[test]
fn test_atr_warmup() {
let mut atr = StreamingATR::new(3).unwrap();
// First 3 bars return NaN (period = 3, seed happens on bar 3 but still NaN)
assert!(atr.update(10.0, 9.0, 9.5).is_nan());
assert!(atr.update(11.0, 9.5, 10.5).is_nan());
assert!(atr.update(10.5, 9.0, 9.5).is_nan());
// Bar 4: first real value
let v = atr.update(11.0, 10.0, 10.5);
assert!(!v.is_nan());
assert!(v > 0.0);
}
#[test]
fn test_bbands_warmup() {
let mut bb = StreamingBBands::new(3, 2.0, 2.0).unwrap();
let (u, m, l) = bb.update(10.0);
assert!(u.is_nan() && m.is_nan() && l.is_nan());
let (u, m, l) = bb.update(11.0);
assert!(u.is_nan() && m.is_nan() && l.is_nan());
let (u, m, l) = bb.update(12.0);
assert!(!u.is_nan() && !m.is_nan() && !l.is_nan());
assert!(approx_eq(m, 11.0, 1e-10));
assert!(u > m && l < m);
}
#[test]
fn test_macd_basic() {
let mut macd = StreamingMACD::new(3, 5, 2).unwrap();
// Feed enough bars for the slow (5) to seed
for i in 0..4 {
let (m, s, h) = macd.update(100.0 + i as f64);
assert!(m.is_nan());
}
// Bar 5: slow seeds
let (m, s, _h) = macd.update(104.0);
assert!(!m.is_nan());
}
#[test]
fn test_macd_fast_ge_slow_rejected() {
assert!(StreamingMACD::new(5, 3, 2).is_err());
assert!(StreamingMACD::new(5, 5, 2).is_err());
}
#[test]
fn test_stoch_basic() {
let mut stoch = StreamingStoch::new(3, 2, 2).unwrap();
// Need fastk_period bars, then slowk_period, then slowd_period
let (k, d) = stoch.update(10.0, 8.0, 9.0);
assert!(k.is_nan() && d.is_nan());
let (k, d) = stoch.update(11.0, 9.0, 10.0);
assert!(k.is_nan() && d.is_nan());
// Bar 3: fastk ready, collecting slowk
let (k, d) = stoch.update(12.0, 10.0, 11.0);
assert!(k.is_nan());
// Bar 4
let (k, d) = stoch.update(13.0, 11.0, 12.0);
assert!(!k.is_nan());
}
#[test]
fn test_vwap_basic() {
let mut vwap = StreamingVWAP::new();
let v = vwap.update(10.0, 8.0, 9.0, 100.0);
// tp = (10+8+9)/3 = 9.0, vwap = 9.0*100/100 = 9.0
assert!(approx_eq(v, 9.0, 1e-10));
let v = vwap.update(12.0, 10.0, 11.0, 200.0);
// tp2 = 11.0, cum_tpv = 900+2200=3100, cum_vol=300, vwap=10.333..
assert!(approx_eq(v, 3100.0 / 300.0, 1e-10));
}
#[test]
fn test_vwap_zero_volume() {
let mut vwap = StreamingVWAP::new();
let v = vwap.update(10.0, 8.0, 9.0, 0.0);
assert!(v.is_nan());
}
#[test]
fn test_supertrend_warmup() {
let mut st = StreamingSupertrend::new(3, 2.0).unwrap();
let (line, dir) = st.update(10.0, 9.0, 9.5);
assert!(line.is_nan() && dir == 0);
let (line, dir) = st.update(11.0, 9.5, 10.5);
assert!(line.is_nan() && dir == 0);
let (line, dir) = st.update(10.5, 9.0, 9.5);
assert!(line.is_nan() && dir == 0);
// Bar 4: first real value
let (line, dir) = st.update(11.0, 10.0, 10.5);
assert!(!line.is_nan());
assert!(dir == 1 || dir == -1);
}
#[test]
fn test_streaming_sma_matches_batch() {
// Compare streaming SMA against a simple batch computation
let data = vec![1.0, 3.0, 5.0, 7.0, 9.0, 11.0, 13.0];
let period = 3;
let mut sma = StreamingSMA::new(period).unwrap();
let streaming: Vec<f64> = data.iter().map(|&v| sma.update(v)).collect();
// Batch SMA
for i in 0..data.len() {
if i + 1 < period {
assert!(streaming[i].is_nan(), "bar {} should be NaN", i);
} else {
let batch: f64 =
data[i + 1 - period..=i].iter().sum::<f64>() / period as f64;
assert!(
approx_eq(streaming[i], batch, 1e-10),
"bar {}: streaming={} batch={}",
i,
streaming[i],
batch
);
}
}
}
}
+17 -5
View File
@@ -1,10 +1,15 @@
//! Volatility indicators.
/// Average True Range Wilder smoothed (TA-Lib compatible).
/// Compute the Average True Range (ATR), Wilder smoothed (TA-Lib compatible).
///
/// Seeds ATR with SMA of TR[1..=timeperiod] (bar 0 is skipped, matching TA-Lib).
/// First valid output is at index `timeperiod`; indices 0..timeperiod are NaN.
/// TR is computed on-the-fly (no separate tr Vec allocation).
/// ATR measures market volatility by smoothing the True Range with Wilder's
/// method. Seeded with the SMA of `TR[1..=timeperiod]` (bar 0 is skipped,
/// matching TA-Lib). Returns non-negative values; the first `timeperiod`
/// indices are `NaN`.
///
/// # Arguments
/// * `high` / `low` / `close` - OHLC price series (same length).
/// * `timeperiod` - Smoothing period (typically 14).
pub fn atr(high: &[f64], low: &[f64], close: &[f64], timeperiod: usize) -> Vec<f64> {
let n = high.len();
let mut result = vec![f64::NAN; n];
@@ -33,7 +38,14 @@ pub fn atr(high: &[f64], low: &[f64], close: &[f64], timeperiod: usize) -> Vec<f
result
}
/// True Range — max(H-L, |H-Cprev|, |L-Cprev|).
/// Compute the True Range for each bar.
///
/// `TR = max(H - L, |H - C_prev|, |L - C_prev|)`. For bar 0, TR is
/// simply `H - L` (no previous close available). Returns non-negative
/// values for every bar (no `NaN` warmup).
///
/// # Arguments
/// * `high` / `low` / `close` - OHLC price series (same length).
pub fn trange(high: &[f64], low: &[f64], close: &[f64]) -> Vec<f64> {
let n = high.len();
let mut result = vec![f64::NAN; n];
+91 -5
View File
@@ -1,6 +1,14 @@
//! Volume indicators.
/// On-Balance Volume.
/// Compute On-Balance Volume (OBV).
///
/// OBV is a cumulative indicator that adds volume on up-close bars and
/// subtracts volume on down-close bars. Unchanged closes contribute zero.
/// Returns a `Vec<f64>` of length `n` with no `NaN` values.
///
/// # Arguments
/// * `close` - Price series.
/// * `volume` - Volume series (same length as `close`).
pub fn obv(close: &[f64], volume: &[f64]) -> Vec<f64> {
let n = close.len();
let mut result = vec![0.0_f64; n];
@@ -21,11 +29,17 @@ pub fn obv(close: &[f64], volume: &[f64]) -> Vec<f64> {
result
}
/// Money Flow Index — O(n) sliding-window implementation without per-bar allocation.
/// Compute the Money Flow Index (MFI).
///
/// MFI = 100 - 100 / (1 + positive_flow / negative_flow) over `timeperiod` bars.
/// typical_price = (high + low + close) / 3; raw_money_flow = typical_price * volume.
/// Leading `timeperiod` values are NaN.
/// MFI is a volume-weighted RSI, returning values in `[0, 100]`.
/// `typical_price = (H + L + C) / 3`; money flow is positive when
/// typical price rises, negative when it falls. The first `timeperiod`
/// values are `NaN`.
///
/// # Arguments
/// * `high` / `low` / `close` - OHLC price series (same length).
/// * `volume` - Volume series (same length).
/// * `timeperiod` - Lookback window (typically 14).
pub fn mfi(
high: &[f64],
low: &[f64],
@@ -78,6 +92,51 @@ pub fn mfi(
result
}
/// Chaikin Accumulation/Distribution Line.
///
/// Cumulates `(close - low - (high - close)) / (high - low) * volume`.
pub fn ad(high: &[f64], low: &[f64], close: &[f64], volume: &[f64]) -> Vec<f64> {
let n = high.len();
let mut result = vec![0.0_f64; n];
let mut ad_val = 0.0_f64;
for i in 0..n {
let hl = high[i] - low[i];
let clv = if hl != 0.0 {
((close[i] - low[i]) - (high[i] - close[i])) / hl
} else {
0.0
};
ad_val += clv * volume[i];
result[i] = ad_val;
}
result
}
/// Chaikin A/D Oscillator: fast EMA of AD minus slow EMA of AD.
///
/// Uses the core EMA implementation from `overlap::ema`.
pub fn adosc(
high: &[f64],
low: &[f64],
close: &[f64],
volume: &[f64],
fastperiod: usize,
slowperiod: usize,
) -> Vec<f64> {
let n = high.len();
let ad_vals = ad(high, low, close, volume);
let fast_ema = crate::overlap::ema(&ad_vals, fastperiod);
let slow_ema = crate::overlap::ema(&ad_vals, slowperiod);
let warmup = slowperiod - 1;
let mut result = vec![f64::NAN; n];
for i in warmup..n {
if !fast_ema[i].is_nan() && !slow_ema[i].is_nan() {
result[i] = fast_ema[i] - slow_ema[i];
}
}
result
}
#[cfg(test)]
mod tests {
use super::*;
@@ -92,6 +151,33 @@ mod tests {
assert!((result[2] - 600.0).abs() < 1e-10);
}
#[test]
fn ad_basic() {
let h = vec![10.0, 12.0, 11.0];
let l = vec![8.0, 9.0, 9.0];
let c = vec![9.0, 11.0, 10.0];
let v = vec![1000.0, 2000.0, 1500.0];
let result = ad(&h, &l, &c, &v);
assert_eq!(result.len(), 3);
// CLV[0] = ((9-8) - (10-9)) / (10-8) = (1 - 1) / 2 = 0
assert!((result[0] - 0.0).abs() < 1e-10);
}
#[test]
fn adosc_basic() {
let n = 30;
let h: Vec<f64> = (1..=n).map(|i| i as f64 + 1.0).collect();
let l: Vec<f64> = (1..=n).map(|i| i as f64 - 1.0).collect();
let c: Vec<f64> = (1..=n).map(|i| i as f64).collect();
let v: Vec<f64> = vec![1000.0; n];
let result = adosc(&h, &l, &c, &v, 3, 10);
assert_eq!(result.len(), n);
// Warmup period should be NaN
for i in 0..9 {
assert!(result[i].is_nan());
}
}
#[test]
fn mfi_range() {
let n = 50;
+73
View File
@@ -10,6 +10,11 @@ a Python technical analysis library.
* - Area
- Status
- What it is
* - Backtesting engine
- Adjacent
- Vectorized Rust backtester: OHLCV fill, stop-loss/TP, 23 performance
metrics, trade extraction, parallel Monte Carlo, walk-forward analysis,
and multi-asset portfolio simulation. See :ref:`backtesting-engine`.
* - Derivatives analytics
- Adjacent
- Options pricing, Greeks, implied volatility helpers, futures basis,
@@ -36,6 +41,74 @@ a Python technical analysis library.
- Registry and plugin packaging model for custom indicators. See
:doc:`plugins`.
.. _backtesting-engine:
Backtesting Engine
------------------
``ferro_ta.analysis.backtest`` ships a production-grade backtesting engine
backed entirely by Rust hot-path functions.
**Core API:**
.. code-block:: python
from ferro_ta.analysis.backtest import BacktestEngine, monte_carlo, walk_forward
result = (
BacktestEngine()
.with_commission(0.001)
.with_slippage(5.0) # basis points
.with_ohlcv(high=high, low=low, open_=open_)
.with_stop_loss(0.02)
.with_take_profit(0.04)
.run(close, "sma_crossover")
)
print(result.metrics["sharpe"]) # one of 23 metrics
print(result.trades) # pandas DataFrame
print(result.drawdown_series.min()) # max drawdown
mc = monte_carlo(result, n_sims=1000) # parallel bootstrap
wf = walk_forward(close, "rsi", param_grid=[{"timeperiod": t} for t in [10,14,20]],
train_bars=500, test_bars=100)
**Available Rust primitives** (``ferro_ta._ferro_ta``):
- ``backtest_core`` — close-only, vectorized, commission + slippage
- ``backtest_ohlcv_core`` — fill at open, intrabar stop-loss / take-profit
- ``compute_performance_metrics`` — 23 metrics in one pass (Sharpe, Sortino,
Calmar, CAGR, Omega, Ulcer, win rate, profit factor, tail ratio, etc.)
- ``extract_trades_ohlcv`` — 9 parallel arrays (entry/exit bar, MAE, MFE, …)
- ``backtest_multi_asset_core`` — N-asset parallel backtest via Rayon
- ``monte_carlo_bootstrap`` — parallel block bootstrap, returns (n_sims, n_bars)
- ``walk_forward_indices`` — anchored/rolling fold index generator
- ``kelly_fraction`` / ``half_kelly_fraction``
**Speed vs competitors** (100k bars, SMA crossover, Apple M-series):
.. list-table::
:header-rows: 1
* - Library
- Time
- vs ferro-ta
* - ferro-ta ``backtest_core``
- 0.29 ms
- —
* - NumPy vectorized
- 0.46 ms
- 1.6× slower
* - vectorbt
- 2.9 ms
- 10× slower
* - backtesting.py
- 320 ms
- 1,100× slower
* - backtrader
- ~520 ms (10k bars)
- >15,000× slower
How to read the project
-----------------------
+6255 -4994
View File
File diff suppressed because it is too large Load Diff
+83
View File
@@ -13,9 +13,92 @@ The authoritative benchmark workflow lives in ``benchmarks/``:
- Cross-library speed suite: ``benchmarks/test_speed.py``
- Cross-library accuracy suite: ``benchmarks/test_accuracy.py``
- TA-Lib head-to-head script: ``benchmarks/bench_vs_talib.py``
- Backtesting engine benchmark: ``benchmarks/bench_backtest.py``
- Table generation from benchmark JSON: ``benchmarks/benchmark_table.py``
- Perf-contract artifact bundle: ``benchmarks/run_perf_contract.py``
Backtesting engine — competitor comparison
------------------------------------------
Measured on Apple M-series, Python 3.13, Rust 1.91, using an SMA(20/50)
crossover strategy with 0.1% commission and 5 bps slippage. Median of 5 runs.
.. list-table:: Speed vs backtesting libraries (signal → equity curve)
:header-rows: 1
* - Library
- 1k bars
- 10k bars
- 100k bars
- vs ferro-ta core (100k)
* - **ferro-ta** ``backtest_core``
- 0.004 ms
- 0.033 ms
- 0.286 ms
- —
* - **ferro-ta** ``backtest_ohlcv_core``
- 0.004 ms
- 0.037 ms
- 0.332 ms
- ~same
* - NumPy vectorized (manual)
- 0.013 ms
- 0.042 ms
- 0.459 ms
- 1.6× slower
* - vectorbt 0.28
- 1.32 ms
- 1.31 ms
- 2.90 ms
- **10× slower**
* - backtesting.py
- 10.5 ms
- 42.3 ms
- 319.6 ms
- **1,117× slower**
* - backtrader 1.9
- 53.9 ms
- 518 ms
- n/a (skipped)
- **>15,000× slower**
Accuracy: ferro-ta positions and bar-returns are **bit-exact** against the NumPy
reference implementation (max per-bar equity diff = 0.00e+00 with zero
commission/slippage).
Additional ferro-ta capabilities not present in the libraries above:
.. list-table::
:header-rows: 1
* - Capability
- ferro-ta result
- NumPy baseline
- Speedup
* - Monte Carlo 1,000 sims (100k bars)
- 50 ms (parallel Rayon + LCG)
- 612 ms (Python loop)
- **12×**
* - 23 performance metrics, single call (100k bars)
- 2.8 ms
- 0.36 ms (2 metrics only)
- 0.12 ms / metric
* - Multi-asset 100 assets (100k bars)
- 43 ms parallel / 88 ms serial
- —
- 2× parallel speedup
* - Walk-forward fold indices (100k bars)
- 0.3 µs
- —
- —
Reproduce the backtest benchmark:
.. code-block:: bash
python benchmarks/bench_backtest.py --sizes 10000 100000 \
--json benchmarks/artifacts/latest/bench_backtest_results.json
Latest checked-in TA-Lib artifact
---------------------------------
+204 -1
View File
@@ -1,7 +1,210 @@
Release Notes
=============
These docs track package version ``1.0.6``.
These docs track package version ``1.2.0``.
1.2.0-audit (2026-03-28)
------------------------
**Comprehensive audit: 90 findings addressed**
*Code quality & correctness*
- **Welford's algorithm for BBANDS**: replaced naive ``sum_sq/N - mean^2`` variance
with numerically stable Welford's rolling algorithm in both batch and streaming BBANDS.
Fixes catastrophic cancellation for large-valued series (e.g., prices near 1e12).
- **FFI boundary safety**: ``transpose_to_series_major()`` in ``batch/mod.rs`` now
returns ``PyResult`` instead of using ``expect()``. Remaining ``as_slice().expect()``
calls in ``allow_threads`` closures are documented with SAFETY comments (structurally
infallible after C-contiguous transpose).
- **Clippy clean**: resolved all clippy warnings — complex type in ``adx_all`` extracted
to ``AdxAllResult`` type alias; ``welford_step`` helper annotated with
``#[allow(clippy::too_many_arguments)]``.
*Performance*
- **``target-cpu=native``**: new ``.cargo/config.toml`` enables native CPU instruction
set (AVX2, NEON, etc.) for all non-WASM targets. CI can override via ``RUSTFLAGS``.
*Testing*
- **Streaming unit tests**: 37 new tests in ``tests/unit/streaming/test_streaming.py``
covering ``StreamingSMA``, ``StreamingEMA``, ``StreamingRSI`` — batch parity, warmup
NaN behavior, reset, edge cases, and large dataset numerical stability.
- **Edge case tests**: 31 new tests in ``tests/unit/test_edge_cases.py`` — empty arrays,
single elements, all-NaN input, NaN propagation, extreme values (1e300, 1e-300),
constant series, period boundary conditions, OHLCV edge cases, and dtype coercion
(float32, int64).
- **Property-based tests**: expanded Hypothesis tests for EMA, BBANDS, MACD, ATR, WMA,
and OBV with algebraic invariants (upper >= middle >= lower, histogram == macd - signal,
ATR non-negative, etc.).
- **Pandas/polars integration tests**: new ``test_dataframe_integration.py`` verifying
transparent ``pd.Series`` and ``polars.Series`` support across SMA, EMA, RSI, BBANDS,
MACD, and end-to-end DataFrame workflows.
- **Fuzzing**: expanded from 2 to 9 fuzz targets — added EMA, BBANDS, MACD, ATR, STOCH,
MFI, and WMA with output invariant assertions.
- **Test helpers**: new ``tests/unit/helpers.py`` consolidating duplicated assertion
patterns (``nan_count``, ``finite``, ``assert_nan_warmup``, ``assert_output_length``,
``assert_range``, ``make_ohlcv``).
*Documentation*
- **README benchmarks**: updated to match actual artifact data — MFI 3.25x, WMA 2.20x,
BBANDS 1.97x, SMA 1.93x; corrected win count from 6 to 7 at 100k bars.
- **Rust doc comments**: added comprehensive ``///`` documentation to all public functions
in ``ferro_ta_core`` — overlap (SMA, EMA, WMA, BBANDS, MACD), momentum (RSI, STOCH,
ADX family), volatility (ATR, TRANGE), volume (OBV, MFI), statistic (STDDEV), and math
(sum, max, min, sliding_max, sliding_min).
*Linting*
- **Ruff clean**: fixed import sorting, unused imports, trailing whitespace, and
formatting across all Python files.
- **cargo fmt**: all Rust code formatted.
1.2.0 (2026-03-28)
------------------
**Phase 1 — Simulation fidelity**
- **Bid-ask spread model**: new ``CommissionModel.spread_bps`` field (basis points).
Half-spread is deducted per leg (entry and exit), modelling real market microstructure costs.
- **Breakeven stop**: new ``backtest_ohlcv_core`` parameter ``breakeven_pct`` and
``BacktestEngine.with_breakeven_stop(pct)``. Once profit reaches ``pct``, the
effective stop-loss is moved to the entry price, guaranteeing at worst a breakeven exit.
- **Bracket order priority**: when both stop-loss and take-profit are breached on the
same bar, the level closer to the bar's open price fires first (previously SL always won).
**Phase 2 — Portfolio & risk**
- **Short borrow cost**: new ``CommissionModel.short_borrow_rate_annual`` field.
Accrued per bar for short positions at the specified annualised rate.
- **Leverage / margin modeling**: new ``BacktestEngine.with_leverage(margin_ratio, margin_call_pct)``.
Tracks margin usage and triggers a margin-call force-close when equity falls below
``margin_call_pct × initial_margin``.
- **Loss circuit breakers**: new ``BacktestEngine.with_loss_limits(daily, total)``.
Halts all trading when a per-bar loss or total drawdown threshold is breached.
- **Portfolio constraints**: new ``BacktestEngine.with_portfolio_constraints(max_asset_weight,
max_gross_exposure, max_net_exposure)`` for multi-asset backtests.
**Phase 3 — Data & UX**
- **Bar aggregation** (``ferro_ta.analysis.resample``): ``resample_ohlcv()``, ``align_to_coarse()``,
``resample_ohlcv_labels()`` — pure-NumPy OHLCV resampling from any fine TF to any coarser TF.
- **Multi-timeframe engine** (``ferro_ta.analysis.multitf``): ``MultiTimeframeEngine`` — compute
strategy signals on coarser bars and execute on finer bars, with automatic signal alignment.
- **Dividend/split adjustment** (``ferro_ta.analysis.adjust``): ``adjust_ohlcv()``,
``adjust_for_splits()``, ``adjust_for_dividends()`` — backward-adjusted price series for
equity/index strategies.
- **Visualization** (``ferro_ta.analysis.plot``): ``plot_backtest()`` — interactive Plotly chart
with equity curve, drawdown panel, position panel, trade markers, and optional benchmark overlay.
**Phase 4 — Differentiation**
- **Regime detection** (``ferro_ta.analysis.regime``): ``detect_volatility_regime()``,
``detect_trend_regime()``, ``detect_combined_regime()``, ``RegimeFilter`` — pure-NumPy
6-state market regime labeling and signal filtering; no external ML dependencies.
- **Portfolio optimization** (``ferro_ta.analysis.optimize``): ``PortfolioOptimizer``,
``mean_variance_optimize()``, ``risk_parity_optimize()``, ``max_sharpe_optimize()``
minimum-variance, risk-parity, and maximum-Sharpe portfolios via SLSQP (requires scipy).
- **Paper trading bridge** (``ferro_ta.analysis.live``): ``PaperTrader`` — event-driven
bar-by-bar simulator matching ``backtest_ohlcv_core`` logic exactly; supports streaming
data, live state inspection, and seamless strategy migration from backtesting to live.
1.1.0 (2026-03-27)
------------------
**Advanced commission and fee model (Indian market support)**
- New ``CommissionModel`` class (pure Rust in ``ferro_ta_core``, exposed via
PyO3 and WASM) replaces the broken flat ``commission_per_trade`` scalar. The
old code subtracted an absolute currency amount from a 1.0-normalised equity
curve — equivalent to a 2 000 % error on a ₹1 lakh account. The new model
correctly converts every charge to a fraction of ``initial_capital`` before
deducting it from the equity curve.
- ``CommissionModel`` supports: proportional brokerage (``rate_of_value``),
flat per-order fee (``flat_per_order``), per-lot fee (``per_lot``), brokerage
cap (``max_brokerage``), Securities Transaction Tax (``stt_rate`` with
configurable buy/sell sides), exchange transaction charges, SEBI regulatory
charges, 18 % GST on brokerage + exchange + regulatory levies, and stamp duty
on buy leg only.
- Built-in presets: ``CommissionModel.equity_delivery_india()``,
``CommissionModel.equity_intraday_india()``,
``CommissionModel.futures_india()``, ``CommissionModel.options_india()``,
``CommissionModel.proportional(rate)``, ``CommissionModel.zero()``.
- JSON persistence: ``model.to_json()`` / ``CommissionModel.from_json(s)``,
``model.save(path)`` / ``CommissionModel.load(path)``.
- ``BacktestEngine.with_commission_model(model)`` — pass a full
``CommissionModel``; old ``with_commission(rate)`` kept as a shim.
- New ``initial_capital`` parameter (default ₹1,00,000) on both
``backtest_core`` and ``backtest_ohlcv_core``; also exposed as
``BacktestEngine.with_initial_capital(capital)``.
**Currency system — INR default with lakh/crore formatting**
- New ``Currency`` immutable descriptor in the Python layer with constants
``INR``, ``USD``, ``EUR``, ``GBP``, ``JPY``, ``USDT``.
- ``INR`` is the default currency for ``BacktestEngine``; change via
``engine.with_currency("USD")`` or ``engine.with_currency(EUR)``.
- ``currency.format(amount)`` produces Indian lakh/crore grouping for INR
(e.g. ``₹1,23,45,678.00``) and standard Western grouping for other
currencies.
- Module-level helper ``format_currency(amount, currency=INR)``.
- ``AdvancedBacktestResult`` gains ``currency``, ``initial_capital``, and
``equity_abs`` (absolute currency equity curve) slots.
- ``summary()`` now includes ``initial_capital``, ``final_capital``,
``absolute_pnl``, and ``currency`` keys.
- ``AdvancedBacktestResult.__repr__`` shows the final capital in the correct
currency symbol (e.g. ``final=₹1,23,450.00``).
- Trade log gains a ``pnl_abs`` column (PnL in absolute currency units).
- ``to_equity_dataframe()`` now includes an ``equity_abs`` column.
**Trailing stop loss**
- ``backtest_ohlcv_core`` (and ``BacktestEngine.with_trailing_stop(pct)``)
now supports a trailing stop implemented intrabar in Rust: the high-water
mark is updated each bar; the position is exited at
``trail_high × (1 pct)`` when ``low[i]`` crosses below it (long trades),
or ``trail_low × (1 + pct)`` for short trades.
**Benchmark comparison metrics**
- ``compute_performance_metrics`` accepts an optional ``benchmark_returns``
array. When provided, ``summary()`` includes: ``benchmark_total_return``,
``benchmark_cagr``, ``benchmark_annualized_vol``, ``benchmark_sharpe``,
``alpha`` (active return), ``beta``, ``tracking_error``, and
``information_ratio``.
- ``BacktestEngine.with_benchmark(close_array)`` — pass benchmark close prices.
**Volatility-target position sizing**
- New ``"volatility_target"`` method for ``with_position_sizing()``:
``engine.with_position_sizing("volatility_target", target_vol=0.15, vol_window=20)``.
Signals are pre-scaled in Python by ``clip(target_vol / rolling_annualised_vol, 0, 3)``
before the Rust core call, keeping the hot loop unchanged.
**Backtesting engine v2 — full feature set**
- ``BacktestEngine`` now supports true two-pass Kelly / half-Kelly position
sizing: a unit-signal pass computes win statistics, then the core engine
re-runs with signals scaled by the Kelly fraction.
- Added ``fixed_fractional`` position sizing method:
``engine.with_position_sizing("fixed_fractional", fraction=0.5)``.
- New ``StreamingBacktest`` Rust class for bar-by-bar incremental backtesting
(no bulk arrays needed); exposes ``.on_bar()``, ``.summary()``, ``.reset()``.
- ``AdvancedBacktestResult.to_equity_dataframe(freq)`` — returns equity,
returns, and drawdown as a ``pd.DataFrame`` with a synthetic DatetimeIndex.
- ``AdvancedBacktestResult.summary()`` — concise dict of the 9 most commonly
cited metrics plus ``n_trades``.
**Core indicator speedup**
- ADX-family indicators (``adx_all`` public API): all six series (PDM, MDM,
+DI, -DI, DX, ADX) can now be computed from a single TR/PDM/MDM pass via
``ferro_ta.adx_all()``, eliminating the 6× redundant computation that
occurred when callers fetched each series independently.
- ``adxr`` now reuses a single ``adx_inner`` call internally (was calling
``adx()`` which re-ran the inner loop).
1.0.6 (2026-03-24)
------------------
+1
View File
@@ -59,6 +59,7 @@ Core library:
Adjacent and experimental tooling:
- **Backtesting engine** — OHLCV fill, 23 metrics, Monte Carlo, walk-forward, multi-asset — see :doc:`adjacent_tooling`
- Derivatives analytics — see :doc:`derivatives`
- Agentic workflow and LangChain tool wrappers — see `Agentic guide <https://github.com/pratikbhadane24/ferro-ta/blob/main/docs/agentic.md>`_
- MCP server for MCP-compatible clients — see `MCP guide <https://github.com/pratikbhadane24/ferro-ta/blob/main/docs/mcp.md>`_
+29 -7
View File
@@ -234,6 +234,29 @@ wrapper with validation and `_to_f64`; all computation runs in the extension.
and `benchmarks/profile_runtime_hotspots.py` record timings with git/runtime
metadata so you can compare apples to apples across machines and commits.
## Backtesting Performance
ferro-ta's backtesting engine is the fastest in the Python ecosystem for
vectorized single- and multi-asset scenarios.
| Library | 100k bars | vs ferro-ta |
|---------|-----------|-------------|
| ferro-ta `backtest_core` | **0.29 ms** | — |
| ferro-ta `backtest_ohlcv_core` | **0.33 ms** | ~same |
| NumPy vectorized | 0.46 ms | 1.6× slower |
| vectorbt | 2.90 ms | 10× slower |
| backtesting.py | 319 ms | 1,117× slower |
| backtrader | ~50,000 ms (est.) | >15,000× slower |
Additional capabilities measured at 100k bars:
| Capability | Time |
|---|---|
| Monte Carlo 1,000 sims (parallel) | 50 ms — 12× faster than NumPy loop |
| 23 performance metrics | 2.8 ms (0.12 ms/metric) |
| Multi-asset 100 symbols, parallel | 43 ms — 2× vs serial |
| Walk-forward index generation | 0.3 µs |
## Benchmark Tooling
The benchmark suite now includes a small set of machine-readable scripts for
@@ -241,6 +264,7 @@ performance work beyond the full pytest benchmark table:
- `python benchmarks/bench_batch.py --json batch_benchmark.json`
- `python benchmarks/bench_streaming.py --json streaming_benchmark.json`
- `python benchmarks/bench_backtest.py --json bench_backtest_results.json`
- `python benchmarks/profile_runtime_hotspots.py --json runtime_hotspots.json`
- `python benchmarks/bench_simd.py --json simd_benchmark.json`
- `python benchmarks/run_perf_contract.py --output-dir benchmarks/artifacts/latest`
@@ -301,13 +325,11 @@ for history and commits.
Maintainer-facing list of slower paths and optional improvements. Update as
bottlenecks are fixed or deferred.
**Backtest** (`python/ferro_ta/backtest.py`):
- Equity with commission uses an O(n) Python loop (lines 374380). Could
vectorize (e.g. cumsum of commission events) or move to a small Rust helper.
- When both slippage and commission are used, `position_changed` is computed
twice; compute once and reuse.
- Built-in strategies do redundant `np.asarray(..., dtype=np.float64)` if
callers already pass contiguous float64; minor.
**Backtest** (`python/ferro_ta/analysis/backtest.py`):
- Core signal→equity loop is fully in Rust (`backtest_core`, `backtest_ohlcv_core`).
- Commission and slippage applied inside Rust; no Python loop on the hot path.
- `compute_performance_metrics` computes all 23 metrics in a single Rust pass.
- Monte Carlo runs in parallel Rayon threads with LCG seeding (GIL released).
**Batch** (`python/ferro_ta/batch.py`):
- `batch_apply` runs a Python loop over columns (one Python call per column).
+74 -1
View File
@@ -59,10 +59,83 @@ Module status
* - ``ferro_ta.analysis.*``
- Adjacent tooling
- Useful analytics helpers, but not the primary product story.
* - ``ferro_ta.analysis.resample``
- Supported (v1.2.0)
- ``resample_ohlcv()``, ``align_to_coarse()``, ``resample_ohlcv_labels()`` — pure-NumPy
OHLCV bar aggregation across timeframes.
* - ``ferro_ta.analysis.multitf``
- Supported (v1.2.0)
- ``MultiTimeframeEngine`` — multi-timeframe signal generation with automatic alignment.
* - ``ferro_ta.analysis.adjust``
- Supported (v1.2.0)
- ``adjust_ohlcv()``, ``adjust_for_splits()``, ``adjust_for_dividends()`` — backward-adjusted
price series for equity/index strategies.
* - ``ferro_ta.analysis.plot``
- Supported (v1.2.0)
- ``plot_backtest()`` — interactive Plotly backtest visualization (requires plotly).
* - ``ferro_ta.analysis.regime``
- Supported (v1.2.0)
- ``detect_volatility_regime()``, ``detect_trend_regime()``, ``detect_combined_regime()``,
``RegimeFilter`` — pure-NumPy 6-state market regime labeling; no ML dependencies.
* - ``ferro_ta.analysis.optimize``
- Supported (v1.2.0)
- ``PortfolioOptimizer``, ``mean_variance_optimize()``, ``risk_parity_optimize()``,
``max_sharpe_optimize()`` — portfolio optimization via SLSQP (requires scipy).
* - ``ferro_ta.analysis.live``
- Supported (v1.2.0)
- ``PaperTrader`` — event-driven paper trading bridge matching backtest logic exactly.
* - MCP, WASM, GPU, plugin, and agent-oriented tooling
- Experimental or adjacent
- Evaluate these independently from the core indicator library.
Backtesting engine features
---------------------------
.. list-table::
:header-rows: 1
* - Feature
- Status
- Notes
* - Flat/proportional commission
- Supported
- Via ``CommissionModel`` presets and ``BacktestEngine.with_commission_model()``.
* - Bid-ask spread model (``spread_bps``)
- Supported (v1.2.0)
- New ``CommissionModel.spread_bps`` field; half-spread deducted per leg.
* - Short borrow cost (``short_borrow_rate_annual``)
- Supported (v1.2.0)
- New ``CommissionModel.short_borrow_rate_annual`` field; accrued per bar for short positions.
* - Trailing stop loss
- Supported
- ``BacktestEngine.with_trailing_stop(pct)`` — intrabar high-water mark tracking.
* - Breakeven stop (``breakeven_pct``)
- Supported (v1.2.0)
- ``BacktestEngine.with_breakeven_stop(pct)`` — moves stop to entry once profit reaches ``pct``.
* - Bracket order priority
- Supported (v1.2.0)
- When both SL and TP are breached on the same bar, the level closer to open fires first.
* - Leverage / margin modeling
- Supported (v1.2.0)
- ``BacktestEngine.with_leverage(margin_ratio, margin_call_pct)`` — tracks margin and
triggers force-close on margin call.
* - Loss circuit breakers
- Supported (v1.2.0)
- ``BacktestEngine.with_loss_limits(daily, total)`` — halts trading on drawdown breach.
* - Portfolio constraints
- Supported (v1.2.0)
- ``BacktestEngine.with_portfolio_constraints(max_asset_weight, max_gross_exposure,
max_net_exposure)`` for multi-asset backtests.
* - Volatility-target position sizing
- Supported
- ``BacktestEngine.with_position_sizing("volatility_target", ...)``.
* - Walk-forward / Monte Carlo
- Supported
- Available via ``BacktestEngine`` higher-level methods.
* - Benchmark comparison
- Supported
- ``BacktestEngine.with_benchmark(close_array)`` — alpha, beta, information ratio.
Supported Python versions
-------------------------
@@ -107,7 +180,7 @@ For source builds, packaging details, and platform notes, see
Release status
--------------
These docs track package version ``1.0.6``.
These docs track package version ``1.2.0``.
- Release notes by version: :doc:`changelog`
- Canonical project changelog: `CHANGELOG.md <https://github.com/pratikbhadane24/ferro-ta/blob/main/CHANGELOG.md>`_
+49
View File
@@ -28,5 +28,54 @@ test = false
doc = false
bench = false
[[bin]]
name = "fuzz_ema"
path = "fuzz_targets/fuzz_ema.rs"
test = false
doc = false
bench = false
[[bin]]
name = "fuzz_bbands"
path = "fuzz_targets/fuzz_bbands.rs"
test = false
doc = false
bench = false
[[bin]]
name = "fuzz_macd"
path = "fuzz_targets/fuzz_macd.rs"
test = false
doc = false
bench = false
[[bin]]
name = "fuzz_atr"
path = "fuzz_targets/fuzz_atr.rs"
test = false
doc = false
bench = false
[[bin]]
name = "fuzz_stoch"
path = "fuzz_targets/fuzz_stoch.rs"
test = false
doc = false
bench = false
[[bin]]
name = "fuzz_mfi"
path = "fuzz_targets/fuzz_mfi.rs"
test = false
doc = false
bench = false
[[bin]]
name = "fuzz_wma"
path = "fuzz_targets/fuzz_wma.rs"
test = false
doc = false
bench = false
[profile.release]
debug = 1
+48
View File
@@ -0,0 +1,48 @@
/*!
Fuzz target for `ferro_ta_core::volatility::atr`.
Verifies that ATR never panics, output length matches input, and all
finite values are non-negative (ATR is always >= 0).
*/
#![no_main]
use libfuzzer_sys::fuzz_target;
use ferro_ta_core::volatility;
fuzz_target!(|data: &[u8]| {
if data.len() < 2 {
return;
}
let timeperiod = ((data[0] as usize) % 64) + 1;
// Need 3 f64s per bar (high, low, close)
let float_bytes = &data[1..];
let n_floats = float_bytes.len() / 8;
let n_bars = n_floats / 3;
if n_bars == 0 {
return;
}
let all_floats: Vec<f64> = (0..n_bars * 3)
.map(|i| {
let chunk: [u8; 8] = float_bytes[i * 8..(i + 1) * 8].try_into().unwrap();
f64::from_le_bytes(chunk)
})
.collect();
let high = &all_floats[..n_bars];
let low = &all_floats[n_bars..n_bars * 2];
let close = &all_floats[n_bars * 2..n_bars * 3];
let result = volatility::atr(high, low, close, timeperiod);
assert_eq!(result.len(), high.len(), "ATR output length mismatch");
// ATR values should be non-negative when finite
for (i, &v) in result.iter().enumerate() {
if v.is_finite() {
assert!(v >= 0.0, "ATR result[{i}] = {v} is negative");
}
}
});
+60
View File
@@ -0,0 +1,60 @@
/*!
Fuzz target for `ferro_ta_core::overlap::bbands`.
Verifies that BBANDS never panics and that the three output vectors
(upper, middle, lower) always have the same length as the input.
When finite, upper >= middle >= lower must hold.
*/
#![no_main]
use libfuzzer_sys::fuzz_target;
use ferro_ta_core::overlap;
fuzz_target!(|data: &[u8]| {
if data.len() < 3 {
return;
}
let timeperiod = ((data[0] as usize) % 64) + 1;
// Use second byte for deviation multipliers (1.0 - 4.0 range)
let nbdevup = 1.0 + (data[1] as f64 / 255.0) * 3.0;
let nbdevdn = 1.0 + (data[2] as f64 / 255.0) * 3.0;
let float_bytes = &data[3..];
let n_floats = float_bytes.len() / 8;
if n_floats == 0 {
return;
}
let close: Vec<f64> = (0..n_floats)
.map(|i| {
let chunk: [u8; 8] = float_bytes[i * 8..(i + 1) * 8].try_into().unwrap();
f64::from_le_bytes(chunk)
})
.collect();
let (upper, middle, lower) = overlap::bbands(&close, timeperiod, nbdevup, nbdevdn);
assert_eq!(upper.len(), close.len(), "BBANDS upper length mismatch");
assert_eq!(middle.len(), close.len(), "BBANDS middle length mismatch");
assert_eq!(lower.len(), close.len(), "BBANDS lower length mismatch");
// When all three are finite, upper >= middle >= lower
for i in 0..close.len() {
if upper[i].is_finite() && middle[i].is_finite() && lower[i].is_finite() {
assert!(
upper[i] >= middle[i],
"BBANDS upper[{i}] ({}) < middle[{i}] ({})",
upper[i],
middle[i]
);
assert!(
middle[i] >= lower[i],
"BBANDS middle[{i}] ({}) < lower[{i}] ({})",
middle[i],
lower[i]
);
}
}
});
+35
View File
@@ -0,0 +1,35 @@
/*!
Fuzz target for `ferro_ta_core::overlap::ema`.
Verifies that EMA never panics for any input and that the output length
always matches the input length.
*/
#![no_main]
use libfuzzer_sys::fuzz_target;
use ferro_ta_core::overlap;
fuzz_target!(|data: &[u8]| {
if data.len() < 2 {
return;
}
let timeperiod = ((data[0] as usize) % 64) + 1;
let float_bytes = &data[1..];
let n_floats = float_bytes.len() / 8;
if n_floats == 0 {
return;
}
let close: Vec<f64> = (0..n_floats)
.map(|i| {
let chunk: [u8; 8] = float_bytes[i * 8..(i + 1) * 8].try_into().unwrap();
f64::from_le_bytes(chunk)
})
.collect();
let result = overlap::ema(&close, timeperiod);
assert_eq!(result.len(), close.len(), "EMA output length mismatch");
});
+41
View File
@@ -0,0 +1,41 @@
/*!
Fuzz target for `ferro_ta_core::overlap::macd`.
Verifies that MACD never panics and that all three output vectors
(macd, signal, histogram) match the input length.
*/
#![no_main]
use libfuzzer_sys::fuzz_target;
use ferro_ta_core::overlap;
fuzz_target!(|data: &[u8]| {
if data.len() < 4 {
return;
}
// Extract periods from first 3 bytes (1-64 range each)
let fastperiod = ((data[0] as usize) % 32) + 1;
let slowperiod = ((data[1] as usize) % 32) + fastperiod + 1; // slow > fast
let signalperiod = ((data[2] as usize) % 32) + 1;
let float_bytes = &data[3..];
let n_floats = float_bytes.len() / 8;
if n_floats == 0 {
return;
}
let close: Vec<f64> = (0..n_floats)
.map(|i| {
let chunk: [u8; 8] = float_bytes[i * 8..(i + 1) * 8].try_into().unwrap();
f64::from_le_bytes(chunk)
})
.collect();
let (macd, signal, hist) = overlap::macd(&close, fastperiod, slowperiod, signalperiod);
assert_eq!(macd.len(), close.len(), "MACD line length mismatch");
assert_eq!(signal.len(), close.len(), "MACD signal length mismatch");
assert_eq!(hist.len(), close.len(), "MACD histogram length mismatch");
});
+51
View File
@@ -0,0 +1,51 @@
/*!
Fuzz target for `ferro_ta_core::volume::mfi`.
Verifies that MFI never panics, output length matches input, and finite
values lie in [0, 100].
*/
#![no_main]
use libfuzzer_sys::fuzz_target;
use ferro_ta_core::volume;
fuzz_target!(|data: &[u8]| {
if data.len() < 2 {
return;
}
let timeperiod = ((data[0] as usize) % 64) + 1;
// Need 4 f64s per bar (high, low, close, volume)
let float_bytes = &data[1..];
let n_floats = float_bytes.len() / 8;
let n_bars = n_floats / 4;
if n_bars == 0 {
return;
}
let all_floats: Vec<f64> = (0..n_bars * 4)
.map(|i| {
let chunk: [u8; 8] = float_bytes[i * 8..(i + 1) * 8].try_into().unwrap();
f64::from_le_bytes(chunk)
})
.collect();
let high = &all_floats[..n_bars];
let low = &all_floats[n_bars..n_bars * 2];
let close = &all_floats[n_bars * 2..n_bars * 3];
let vol = &all_floats[n_bars * 3..n_bars * 4];
let result = volume::mfi(high, low, close, vol, timeperiod);
assert_eq!(result.len(), high.len(), "MFI output length mismatch");
for (i, &v) in result.iter().enumerate() {
if v.is_finite() {
assert!(
v >= 0.0 && v <= 100.0,
"MFI result[{i}] = {v} is out of [0, 100]"
);
}
}
});
+63
View File
@@ -0,0 +1,63 @@
/*!
Fuzz target for `ferro_ta_core::momentum::stoch`.
Verifies that STOCH never panics, output lengths match, and finite
values lie in [0, 100].
*/
#![no_main]
use libfuzzer_sys::fuzz_target;
use ferro_ta_core::momentum;
fuzz_target!(|data: &[u8]| {
if data.len() < 4 {
return;
}
let fastk_period = ((data[0] as usize) % 32) + 1;
let slowk_period = ((data[1] as usize) % 16) + 1;
let slowd_period = ((data[2] as usize) % 16) + 1;
// Need 3 f64s per bar (high, low, close)
let float_bytes = &data[3..];
let n_floats = float_bytes.len() / 8;
let n_bars = n_floats / 3;
if n_bars == 0 {
return;
}
let all_floats: Vec<f64> = (0..n_bars * 3)
.map(|i| {
let chunk: [u8; 8] = float_bytes[i * 8..(i + 1) * 8].try_into().unwrap();
f64::from_le_bytes(chunk)
})
.collect();
let high = &all_floats[..n_bars];
let low = &all_floats[n_bars..n_bars * 2];
let close = &all_floats[n_bars * 2..n_bars * 3];
let (slowk, slowd) = momentum::stoch(high, low, close, fastk_period, slowk_period, slowd_period);
assert_eq!(slowk.len(), high.len(), "STOCH slowk length mismatch");
assert_eq!(slowd.len(), high.len(), "STOCH slowd length mismatch");
// Finite values should be in [0, 100]
for (i, &v) in slowk.iter().enumerate() {
if v.is_finite() {
assert!(
v >= 0.0 && v <= 100.0,
"STOCH slowk[{i}] = {v} is out of [0, 100]"
);
}
}
for (i, &v) in slowd.iter().enumerate() {
if v.is_finite() {
assert!(
v >= 0.0 && v <= 100.0,
"STOCH slowd[{i}] = {v} is out of [0, 100]"
);
}
}
});
+35
View File
@@ -0,0 +1,35 @@
/*!
Fuzz target for `ferro_ta_core::overlap::wma`.
Verifies that WMA never panics and that the output length always
matches the input length.
*/
#![no_main]
use libfuzzer_sys::fuzz_target;
use ferro_ta_core::overlap;
fuzz_target!(|data: &[u8]| {
if data.len() < 2 {
return;
}
let timeperiod = ((data[0] as usize) % 64) + 1;
let float_bytes = &data[1..];
let n_floats = float_bytes.len() / 8;
if n_floats == 0 {
return;
}
let close: Vec<f64> = (0..n_floats)
.map(|i| {
let chunk: [u8; 8] = float_bytes[i * 8..(i + 1) * 8].try_into().unwrap();
f64::from_le_bytes(chunk)
})
.collect();
let result = overlap::wma(&close, timeperiod);
assert_eq!(result.len(), close.len(), "WMA output length mismatch");
});
+19 -1
View File
@@ -4,7 +4,7 @@ build-backend = "maturin"
[project]
name = "ferro-ta"
version = "1.0.6"
version = "1.2.0"
description = "Rust-powered Python technical analysis library with a TA-Lib-compatible API"
readme = "README.md"
license = { text = "MIT" }
@@ -43,6 +43,10 @@ comparison = [
"pandas-ta>=0.3; python_version >= '3.12'",
"ta>=0.10",
"pandas>=1.0",
"vectorbt>=0.28",
"backtrader>=1.9",
"backtesting>=0.6",
"quantstats>=0.0.81",
]
gpu = ["torch>=2.0"]
options = []
@@ -50,6 +54,7 @@ mcp = ["mcp>=1.0"]
all = ["pandas>=1.0", "polars>=0.19", "pytest>=7.0", "pytest-benchmark>=4.0"]
dev = [
"pytest>=7.0",
"pytest-cov>=4.0",
"hypothesis>=6.0",
"pandas>=1.0",
"polars>=0.19",
@@ -72,6 +77,19 @@ Repository = "https://github.com/pratikbhadane24/ferro-ta"
[tool.pytest.ini_options]
testpaths = ["tests/unit", "tests/integration"]
[tool.coverage.run]
source = ["ferro_ta"]
omit = ["*/_ferro_ta*", "*/mcp/*"]
[tool.coverage.report]
show_missing = true
fail_under = 65
exclude_lines = [
"pragma: no cover",
"if TYPE_CHECKING:",
"if __name__ ==",
]
[tool.maturin]
python-source = "python"
module-name = "ferro_ta._ferro_ta"
+6 -1
View File
@@ -74,7 +74,12 @@ def _to_f64(data: ArrayLike) -> np.ndarray:
data = np.array(data.to_list(), dtype=np.float64) # type: ignore[union-attr]
arr = np.ascontiguousarray(data, dtype=np.float64)
if arr.ndim != 1:
raise ValueError("Input must be a 1-D array or list of prices.")
from ferro_ta.core.exceptions import FerroTAInputError
raise FerroTAInputError(
f"Input must be a 1-D array or list of prices, got {arr.ndim}-D array.",
suggestion="Flatten your array with .ravel() or pass a 1-D Series/list.",
)
return arr
+39
View File
@@ -1,6 +1,7 @@
"""
ferro_ta.analysis Portfolio analytics, strategy analysis, and financial modelling.
Sub-modules
-----------
* :mod:`ferro_ta.analysis.portfolio` Portfolio and multi-asset analytics
@@ -15,9 +16,47 @@ Sub-modules
* :mod:`ferro_ta.analysis.futures` Futures basis, curve, roll, and synthetic analytics
* :mod:`ferro_ta.analysis.options_strategy` Typed derivatives strategy schemas
* :mod:`ferro_ta.analysis.derivatives_payoff` Multi-leg payoff and Greeks aggregation
* :mod:`ferro_ta.analysis.resample` OHLCV bar aggregation utilities
* :mod:`ferro_ta.analysis.multitf` Multi-timeframe signal utilities
* :mod:`ferro_ta.analysis.adjust` Corporate action price adjustment utilities
* :mod:`ferro_ta.analysis.plot` Plotly-based backtest visualization
Example usage::
from ferro_ta.analysis.portfolio import portfolio_returns
from ferro_ta.analysis.backtest import backtest
from ferro_ta.analysis.resample import resample_ohlcv, align_to_coarse, resample_ohlcv_labels
from ferro_ta.analysis.multitf import MultiTimeframeEngine
from ferro_ta.analysis.adjust import adjust_ohlcv, adjust_for_splits, adjust_for_dividends
from ferro_ta.analysis.plot import plot_backtest
"""
import importlib as _importlib
_LAZY_IMPORTS: dict[str, tuple[str, str]] = {
"detect_volatility_regime": (
"ferro_ta.analysis.regime",
"detect_volatility_regime",
),
"detect_trend_regime": ("ferro_ta.analysis.regime", "detect_trend_regime"),
"detect_combined_regime": ("ferro_ta.analysis.regime", "detect_combined_regime"),
"RegimeFilter": ("ferro_ta.analysis.regime", "RegimeFilter"),
"PortfolioOptimizer": ("ferro_ta.analysis.optimize", "PortfolioOptimizer"),
"mean_variance_optimize": ("ferro_ta.analysis.optimize", "mean_variance_optimize"),
"risk_parity_optimize": ("ferro_ta.analysis.optimize", "risk_parity_optimize"),
"max_sharpe_optimize": ("ferro_ta.analysis.optimize", "max_sharpe_optimize"),
"PaperTrader": ("ferro_ta.analysis.live", "PaperTrader"),
"BarResult": ("ferro_ta.analysis.live", "BarResult"),
"TradeRecord": ("ferro_ta.analysis.live", "TradeRecord"),
}
def __getattr__(name: str):
"""Lazy imports for heavy sub-modules to avoid startup cost."""
if name in _LAZY_IMPORTS:
module_path, attr = _LAZY_IMPORTS[name]
mod = _importlib.import_module(module_path)
obj = getattr(mod, attr)
globals()[name] = obj # cache so subsequent access skips __getattr__
return obj
raise AttributeError(f"module 'ferro_ta.analysis' has no attribute {name!r}")
+194
View File
@@ -0,0 +1,194 @@
"""
Corporate action price adjustment utilities.
adjust_for_splits(close, split_factors, split_indices)
Apply split adjustments to a close price series (backward-adjusted).
adjust_for_dividends(close, dividends, ex_dates)
Apply dividend adjustments to a close price series (backward-adjusted).
adjust_ohlcv(open_, high, low, close, volume, split_factors=None, split_indices=None,
dividends=None, ex_date_indices=None)
Apply both split and dividend adjustments to a full OHLCV dataset.
Returns (adj_open, adj_high, adj_low, adj_close, adj_volume).
"""
from typing import Optional
import numpy as np
from numpy.typing import ArrayLike, NDArray
__all__ = ["adjust_for_splits", "adjust_for_dividends", "adjust_ohlcv"]
def adjust_for_splits(
close: ArrayLike,
split_factors: ArrayLike, # e.g. [2.0, 3.0] means 2-for-1 then 3-for-1
split_indices: ArrayLike, # bar indices of each split (must be sorted ascending)
) -> NDArray:
"""Backward-adjust close prices for stock splits.
All prices BEFORE a split are divided by the split factor.
e.g. a 2-for-1 split at bar 100: prices[0:100] are halved.
Parameters
----------
close : array-like
Raw close prices.
split_factors : array-like
Split factor for each split event (e.g. 2.0 for a 2-for-1 split).
split_indices : array-like
Bar index of each split event (0-based, must be sorted ascending).
Returns
-------
NDArray of adjusted close prices.
"""
c = np.asarray(close, dtype=np.float64).copy()
factors = np.asarray(split_factors, dtype=np.float64)
indices = np.asarray(split_indices, dtype=np.intp)
# Process splits in chronological order; apply backward adjustment
# (all bars before the split are divided by the factor)
for idx, factor in zip(indices, factors):
if factor <= 0:
raise ValueError(f"split_factor must be > 0, got {factor}")
c[:idx] /= factor
return c
def adjust_for_dividends(
close: ArrayLike,
dividends: ArrayLike, # dividend amount per ex-date
ex_date_indices: ArrayLike, # bar indices of ex-dividend dates
) -> NDArray:
"""Backward-adjust close prices for cash dividends (proportional method).
Adjustment factor at ex-date i = (close[i-1] - dividend) / close[i-1].
All bars before ex-date are multiplied by the cumulative adjustment.
Parameters
----------
close : array-like
Raw close prices.
dividends : array-like
Dividend amount (in currency units) at each ex-dividend date.
ex_date_indices : array-like
Bar index of each ex-dividend date (0-based, sorted ascending).
Returns
-------
NDArray of adjusted close prices.
"""
c = np.asarray(close, dtype=np.float64).copy()
divs = np.asarray(dividends, dtype=np.float64)
indices = np.asarray(ex_date_indices, dtype=np.intp)
# Process in chronological order
for idx, div in zip(indices, divs):
if idx == 0:
# No prior bar; skip adjustment (nothing to adjust)
continue
prev_close = c[idx - 1]
if prev_close <= 0:
continue
adj_factor = (prev_close - div) / prev_close
if adj_factor <= 0:
continue
# All prices before ex-date are multiplied by adj_factor
c[:idx] *= adj_factor
return c
def adjust_ohlcv(
open_: ArrayLike,
high: ArrayLike,
low: ArrayLike,
close: ArrayLike,
volume: ArrayLike,
split_factors: Optional[ArrayLike] = None,
split_indices: Optional[ArrayLike] = None,
dividends: Optional[ArrayLike] = None,
ex_date_indices: Optional[ArrayLike] = None,
) -> tuple[NDArray, NDArray, NDArray, NDArray, NDArray]:
"""Apply split and dividend adjustments to full OHLCV data.
Price arrays are multiplied by cumulative adjustment factor.
Volume is divided by split factors (shares outstanding adjust inversely).
Returns (adj_open, adj_high, adj_low, adj_close, adj_volume).
Parameters
----------
open_, high, low, close : array-like
Raw OHLCV price arrays.
volume : array-like
Raw volume array.
split_factors : array-like, optional
Split factors for each split event.
split_indices : array-like, optional
Bar indices of split events (required if split_factors provided).
dividends : array-like, optional
Dividend amounts for each ex-date.
ex_date_indices : array-like, optional
Bar indices of ex-dividend dates (required if dividends provided).
Returns
-------
(adj_open, adj_high, adj_low, adj_close, adj_volume)
"""
o = np.asarray(open_, dtype=np.float64).copy()
h = np.asarray(high, dtype=np.float64).copy()
low_arr = np.asarray(low, dtype=np.float64).copy()
c = np.asarray(close, dtype=np.float64).copy()
v = np.asarray(volume, dtype=np.float64).copy()
n = len(c)
# Build a per-bar cumulative adjustment factor for prices (starts at 1.0)
price_adj = np.ones(n, dtype=np.float64)
# Separate inverse adjustment for volume (splits only)
vol_adj = np.ones(n, dtype=np.float64)
# -----------------------------------------------------------------------
# Apply split adjustments
# -----------------------------------------------------------------------
if split_factors is not None and split_indices is not None:
sf = np.asarray(split_factors, dtype=np.float64)
si = np.asarray(split_indices, dtype=np.intp)
for idx, factor in zip(si, sf):
if factor <= 0:
raise ValueError(f"split_factor must be > 0, got {factor}")
# Prices before split are divided by factor
price_adj[:idx] /= factor
# Volume before split is multiplied by factor (more shares pre-split)
vol_adj[:idx] *= factor
# -----------------------------------------------------------------------
# Apply dividend adjustments (prices only)
# -----------------------------------------------------------------------
if dividends is not None and ex_date_indices is not None:
divs = np.asarray(dividends, dtype=np.float64)
ei = np.asarray(ex_date_indices, dtype=np.intp)
# We need the split-adjusted close at (idx-1) for each dividend event.
# Instead of recomputing the full array each iteration, read the single
# element we need: c[idx-1] * price_adj[idx-1].
for idx, div in zip(ei, divs):
if idx == 0:
continue
prev_close = c[idx - 1] * price_adj[idx - 1]
if prev_close <= 0:
continue
adj_factor = (prev_close - div) / prev_close
if adj_factor <= 0:
continue
price_adj[:idx] *= adj_factor
adj_open = o * price_adj
adj_high = h * price_adj
adj_low = low_arr * price_adj
adj_close = c * price_adj
adj_volume = v * vol_adj
return adj_open, adj_high, adj_low, adj_close, adj_volume
File diff suppressed because it is too large Load Diff
+544
View File
@@ -0,0 +1,544 @@
"""
Paper trading bridge event-driven bar-by-bar simulation.
PaperTrader
Simulates live order execution using the same logic as the backtester,
but processes one bar at a time. Maintains live state (position, equity, trades).
Usage:
from ferro_ta.analysis.live import PaperTrader
trader = PaperTrader(initial_capital=100_000)
for bar in streaming_bars:
signal = my_strategy(bar)
result = trader.on_bar(
open_=bar.open, high=bar.high, low=bar.low, close=bar.close,
signal=signal
)
if result.filled:
print(f"Order filled at {result.fill_price}")
"""
from __future__ import annotations
import math
from dataclasses import dataclass
from typing import Optional
@dataclass
class BarResult:
"""Result of processing one bar through PaperTrader."""
bar_index: int
filled: bool # whether an order was executed this bar
fill_price: float # NaN if no fill
position: float # position after this bar
equity: float # equity after this bar (normalized, initial = 1.0)
equity_abs: float # absolute equity in currency units
pnl_bar: float # P&L this bar as fraction of initial capital
regime: Optional[int] = None # regime label if regime detection is enabled
@dataclass
class TradeRecord:
"""Record of a completed round-trip trade."""
entry_bar: int
exit_bar: int
entry_price: float
exit_price: float
position: float # +1 long, -1 short
pnl_pct: float # P&L as fraction of initial capital
pnl_abs: float # P&L in currency units
class PaperTrader:
"""Event-driven paper trading simulator.
Processes bars one at a time, maintaining live state.
Supports stop-loss, take-profit, trailing stop, and breakeven stop.
Parameters
----------
initial_capital : float
Starting capital in base currency.
stop_loss_pct : float
Stop-loss distance from entry (fraction). 0 = disabled.
take_profit_pct : float
Take-profit distance from entry (fraction). 0 = disabled.
trailing_stop_pct : float
Trailing stop distance (fraction). 0 = disabled.
breakeven_pct : float
Move stop to breakeven when this profit is reached. 0 = disabled.
slippage_bps : float
Slippage in basis points per fill.
commission_model : optional CommissionModel
Full commission model. None = zero commission.
"""
def __init__(
self,
initial_capital: float = 100_000.0,
stop_loss_pct: float = 0.0,
take_profit_pct: float = 0.0,
trailing_stop_pct: float = 0.0,
breakeven_pct: float = 0.0,
slippage_bps: float = 0.0,
commission_model=None,
) -> None:
self.initial_capital = float(initial_capital)
self.stop_loss_pct = float(stop_loss_pct)
self.take_profit_pct = float(take_profit_pct)
self.trailing_stop_pct = float(trailing_stop_pct)
self.breakeven_pct = float(breakeven_pct)
self.slippage_bps = float(slippage_bps)
self.commission_model = commission_model
# Live state
self._position: float = 0.0
self._entry_price: float = float("nan")
self._equity: float = 1.0 # normalized
self._prev_close: float = float("nan")
self._bar_index: int = 0
self._trail_high: float = float("nan")
self._trail_low: float = float("nan")
self._breakeven_activated: bool = False
self._breakeven_stop: float = float("nan")
self._trades: list[TradeRecord] = []
self._equity_history: list[float] = []
# One-bar-lag signal state
self._pending_signal: float = 0.0
self._first_bar: bool = True
def _close_position(self) -> None:
"""Reset all trade-tracking state to flat (mirrors Rust OhlcvState.close_position)."""
self._position = 0.0
self._entry_price = float("nan")
self._trail_high = float("nan")
self._trail_low = float("nan")
self._breakeven_activated = False
self._breakeven_stop = float("nan")
def _commission_cost(self, fill_price: float, pos_size: float) -> float:
"""Compute commission cost as fraction of initial capital."""
if self.commission_model is None:
return 0.0
try:
trade_value = abs(pos_size) * fill_price * self.initial_capital
if hasattr(self.commission_model, "cost_fraction"):
return self.commission_model.cost_fraction(
trade_value, 1.0, pos_size > 0, self.initial_capital
)
except Exception:
pass
return 0.0
def on_bar(
self,
open_: float,
high: float,
low: float,
close: float,
signal: float,
) -> BarResult:
"""Process one bar and return a BarResult.
signal : float
Desired position (+1, -1, or 0). Applied next bar (standard bar-by-bar logic).
For this bar, the signal from the PREVIOUS bar is acted upon.
"""
nan = float("nan")
slip = self.slippage_bps / 10_000.0
bar_idx = self._bar_index
self._bar_index += 1
# On the very first bar: record signal, no action (no prev signal yet)
if self._first_bar:
self._pending_signal = signal
self._first_bar = False
self._prev_close = close
self._equity_history.append(self._equity)
return BarResult(
bar_index=bar_idx,
filled=False,
fill_price=nan,
position=self._position,
equity=self._equity,
equity_abs=self._equity * self.initial_capital,
pnl_bar=0.0,
)
# The signal to act on this bar is from the previous call
desired_pos = (
self._pending_signal if not math.isnan(self._pending_signal) else 0.0
)
# Store current bar's signal for next bar
self._pending_signal = signal
prev_close = self._prev_close
self._prev_close = close
strategy_return = 0.0
fill_price_this_bar = nan
filled = False
forced_close = False
# ---- Update trailing stop water marks ----
if self.trailing_stop_pct > 0.0:
if self._position > 0.0 and not math.isnan(self._trail_high):
self._trail_high = max(self._trail_high, high)
if self._position < 0.0 and not math.isnan(self._trail_low):
self._trail_low = min(self._trail_low, low)
close_ret = (close - prev_close) / prev_close if prev_close != 0.0 else 0.0
# ---- Trailing stop check ----
if (
self.trailing_stop_pct > 0.0
and self._position != 0.0
and not math.isnan(self._entry_price)
):
if self._position > 0.0 and not math.isnan(self._trail_high):
trail_stop = self._trail_high * (1.0 - self.trailing_stop_pct)
if low <= trail_stop:
stop_ret = (
(trail_stop - prev_close) / prev_close
if prev_close != 0.0
else -self.trailing_stop_pct
)
comm = self._commission_cost(trail_stop, self._position)
strategy_return = self._position * stop_ret - slip - comm
fill_price_this_bar = trail_stop
filled = True
self._record_trade(bar_idx, trail_stop)
self._close_position()
forced_close = True
elif self._position < 0.0 and not math.isnan(self._trail_low):
trail_stop = self._trail_low * (1.0 + self.trailing_stop_pct)
if high >= trail_stop:
stop_ret = (
(trail_stop - prev_close) / prev_close
if prev_close != 0.0
else self.trailing_stop_pct
)
comm = self._commission_cost(trail_stop, self._position)
strategy_return = self._position * stop_ret - slip - comm
fill_price_this_bar = trail_stop
filled = True
self._record_trade(bar_idx, trail_stop)
self._close_position()
forced_close = True
# ---- Breakeven stop activation ----
if (
self.breakeven_pct > 0.0
and self._position != 0.0
and not math.isnan(self._entry_price)
and not self._breakeven_activated
):
if self._position > 0.0 and high >= self._entry_price * (
1.0 + self.breakeven_pct
):
self._breakeven_activated = True
self._breakeven_stop = self._entry_price
elif self._position < 0.0 and low <= self._entry_price * (
1.0 - self.breakeven_pct
):
self._breakeven_activated = True
self._breakeven_stop = self._entry_price
# ---- SL/TP combined bracket check ----
if (
not forced_close
and self._position != 0.0
and not math.isnan(self._entry_price)
):
entry = self._entry_price
has_stop = self._breakeven_activated or self.stop_loss_pct > 0.0
stop_long = (
self._breakeven_stop
if self._breakeven_activated
else entry * (1.0 - self.stop_loss_pct)
)
stop_short = (
self._breakeven_stop
if self._breakeven_activated
else entry * (1.0 + self.stop_loss_pct)
)
has_tp = self.take_profit_pct > 0.0
tp_long = entry * (1.0 + self.take_profit_pct)
tp_short = entry * (1.0 - self.take_profit_pct)
if self._position > 0.0:
sl_triggered = has_stop and low <= stop_long
tp_triggered = has_tp and high >= tp_long
if sl_triggered and tp_triggered:
sl_dist = abs(open_ - stop_long)
tp_dist = abs(tp_long - open_)
if sl_dist <= tp_dist:
# SL first
sr = (
(stop_long - prev_close) / prev_close
if prev_close != 0.0
else -self.stop_loss_pct
)
comm = self._commission_cost(stop_long, self._position)
strategy_return = self._position * sr - slip - comm
fill_price_this_bar = stop_long
else:
sr = (
(tp_long - prev_close) / prev_close
if prev_close != 0.0
else self.take_profit_pct
)
comm = self._commission_cost(tp_long, self._position)
strategy_return = self._position * sr - slip - comm
fill_price_this_bar = tp_long
filled = True
self._record_trade(bar_idx, fill_price_this_bar)
self._close_position()
forced_close = True
elif sl_triggered:
sr = (
(stop_long - prev_close) / prev_close
if prev_close != 0.0
else -self.stop_loss_pct
)
comm = self._commission_cost(stop_long, self._position)
strategy_return = self._position * sr - slip - comm
fill_price_this_bar = stop_long
filled = True
self._record_trade(bar_idx, stop_long)
self._close_position()
forced_close = True
elif tp_triggered:
sr = (
(tp_long - prev_close) / prev_close
if prev_close != 0.0
else self.take_profit_pct
)
comm = self._commission_cost(tp_long, self._position)
strategy_return = self._position * sr - slip - comm
fill_price_this_bar = tp_long
filled = True
self._record_trade(bar_idx, tp_long)
self._close_position()
forced_close = True
elif self._position < 0.0:
sl_triggered = has_stop and high >= stop_short
tp_triggered = has_tp and low <= tp_short
if sl_triggered and tp_triggered:
sl_dist = abs(stop_short - open_)
tp_dist = abs(open_ - tp_short)
if sl_dist <= tp_dist:
sr = (
(stop_short - prev_close) / prev_close
if prev_close != 0.0
else self.stop_loss_pct
)
comm = self._commission_cost(stop_short, self._position)
strategy_return = self._position * sr - slip - comm
fill_price_this_bar = stop_short
else:
sr = (
(tp_short - prev_close) / prev_close
if prev_close != 0.0
else -self.take_profit_pct
)
comm = self._commission_cost(tp_short, self._position)
strategy_return = self._position * sr - slip - comm
fill_price_this_bar = tp_short
filled = True
self._record_trade(bar_idx, fill_price_this_bar)
self._close_position()
forced_close = True
elif sl_triggered:
sr = (
(stop_short - prev_close) / prev_close
if prev_close != 0.0
else self.stop_loss_pct
)
comm = self._commission_cost(stop_short, self._position)
strategy_return = self._position * sr - slip - comm
fill_price_this_bar = stop_short
filled = True
self._record_trade(bar_idx, stop_short)
self._close_position()
forced_close = True
elif tp_triggered:
sr = (
(tp_short - prev_close) / prev_close
if prev_close != 0.0
else -self.take_profit_pct
)
comm = self._commission_cost(tp_short, self._position)
strategy_return = self._position * sr - slip - comm
fill_price_this_bar = tp_short
filled = True
self._record_trade(bar_idx, tp_short)
self._close_position()
forced_close = True
# ---- Normal signal execution ----
if not forced_close:
pos_changed = abs(desired_pos - self._position) > 1e-12
# Fill at open (market_open mode, same as Rust default)
base_fill = open_
if desired_pos > self._position:
actual_fill = base_fill * (1.0 + slip)
elif desired_pos < self._position:
actual_fill = base_fill * (1.0 - slip)
else:
actual_fill = base_fill
if pos_changed:
fill_price_this_bar = actual_fill
filled = True
old_pos = self._position
if desired_pos != 0.0 and old_pos == 0.0:
r = (
desired_pos * (close - actual_fill) / actual_fill
if actual_fill != 0.0
else 0.0
)
comm = self._commission_cost(actual_fill, desired_pos)
strategy_return = r - comm
self._set_entry(bar_idx, actual_fill, desired_pos)
elif desired_pos == 0.0:
r = (
old_pos * (actual_fill - prev_close) / prev_close
if prev_close != 0.0
else 0.0
)
comm = self._commission_cost(actual_fill, old_pos)
strategy_return = r - comm
self._record_trade(bar_idx, actual_fill)
self._close_position()
else:
exit_r = (
old_pos * (actual_fill - prev_close) / prev_close
if prev_close != 0.0
else 0.0
)
entry_r = (
desired_pos * (close - actual_fill) / actual_fill
if actual_fill != 0.0
else 0.0
)
exit_comm = self._commission_cost(actual_fill, old_pos)
entry_comm = self._commission_cost(actual_fill, desired_pos)
strategy_return = exit_r + entry_r - exit_comm - entry_comm
if old_pos != 0.0:
self._record_trade(bar_idx, actual_fill)
self._set_entry(bar_idx, actual_fill, desired_pos)
self._position = desired_pos
else:
# Hold: full bar return (close-to-close on existing position)
strategy_return = self._position * close_ret
# Update equity
prev_equity = self._equity
self._equity = self._equity * (1.0 + strategy_return)
pnl_bar = self._equity - prev_equity
self._equity_history.append(self._equity)
return BarResult(
bar_index=bar_idx,
filled=filled,
fill_price=fill_price_this_bar,
position=self._position,
equity=self._equity,
equity_abs=self._equity * self.initial_capital,
pnl_bar=pnl_bar,
)
def _record_trade(self, exit_bar: int, exit_price: float) -> None:
"""Record a completed round-trip trade."""
if math.isnan(self._entry_price):
return
entry_price = self._entry_price
pos = self._position
# P&L = position * (exit - entry) / entry as fraction
if entry_price != 0.0:
pnl_pct = pos * (exit_price - entry_price) / entry_price
else:
pnl_pct = 0.0
pnl_abs = pnl_pct * self.initial_capital
self._trades.append(
TradeRecord(
entry_bar=getattr(self, "_trade_entry_bar", 0),
exit_bar=exit_bar,
entry_price=entry_price,
exit_price=exit_price,
position=pos,
pnl_pct=pnl_pct,
pnl_abs=pnl_abs,
)
)
def _set_entry(self, bar_idx: int, fill_price: float, pos: float) -> None:
"""Set entry state — call after position changes to new non-zero position."""
self._entry_price = fill_price
self._trade_entry_bar = bar_idx
self._trail_high = fill_price if pos > 0.0 else float("nan")
self._trail_low = fill_price if pos < 0.0 else float("nan")
self._breakeven_activated = False
self._breakeven_stop = float("nan")
@property
def position(self) -> float:
"""Current open position."""
return self._position
@property
def equity(self) -> float:
"""Current normalized equity."""
return self._equity
@property
def equity_abs(self) -> float:
"""Current absolute equity in base currency."""
return self._equity * self.initial_capital
@property
def trades(self) -> list[TradeRecord]:
"""List of completed trades."""
return list(self._trades)
@property
def equity_curve(self) -> list[float]:
"""Equity history (normalized)."""
return list(self._equity_history)
def reset(self) -> None:
"""Reset all state to initial values."""
self._position = 0.0
self._entry_price = float("nan")
self._equity = 1.0
self._prev_close = float("nan")
self._bar_index = 0
self._trail_high = float("nan")
self._trail_low = float("nan")
self._breakeven_activated = False
self._breakeven_stop = float("nan")
self._trades = []
self._equity_history = []
self._pending_signal = 0.0
self._first_bar = True
+185
View File
@@ -0,0 +1,185 @@
"""
Multi-timeframe signal utilities.
MultiTimeframeEngine wraps BacktestEngine with a higher-timeframe signal computation step.
Usage:
from ferro_ta.analysis.multitf import MultiTimeframeEngine
result = (
MultiTimeframeEngine(factor=4) # 4 fine bars per coarse bar
.with_htf_strategy("rsi_30_70") # strategy runs on coarse bars
.with_ohlcv(high=h, low=l, open_=o)
.with_stop_loss(0.02)
.run(close_fine)
)
"""
from __future__ import annotations
import numpy as np
from numpy.typing import ArrayLike
from ferro_ta.analysis.backtest import AdvancedBacktestResult, BacktestEngine
from ferro_ta.analysis.resample import align_to_coarse, resample_ohlcv
__all__ = ["MultiTimeframeEngine"]
class MultiTimeframeEngine:
"""Backtests using signals computed on a higher timeframe (coarser bars).
Parameters
----------
factor : int
Number of fine-resolution bars per coarse bar.
"""
def __init__(self, factor: int) -> None:
if factor < 1:
raise ValueError(f"factor must be >= 1, got {factor}")
self._factor = factor
self._htf_strategy = "rsi_30_70"
self._inner = BacktestEngine()
# Store OHLCV separately so we can resample them
self._high: np.ndarray | None = None
self._low: np.ndarray | None = None
self._open: np.ndarray | None = None
def with_htf_strategy(self, strategy) -> MultiTimeframeEngine:
"""Set the strategy function or name used on coarse bars."""
self._htf_strategy = strategy
return self
def with_ohlcv(self, *, high, low, open_) -> MultiTimeframeEngine:
"""Store OHLCV data for resampling and pass to inner engine after resampling."""
self._high = np.asarray(high, dtype=np.float64)
self._low = np.asarray(low, dtype=np.float64)
self._open = np.asarray(open_, dtype=np.float64)
return self
def with_stop_loss(self, pct: float) -> MultiTimeframeEngine:
self._inner.with_stop_loss(pct)
return self
def with_take_profit(self, pct: float) -> MultiTimeframeEngine:
self._inner.with_take_profit(pct)
return self
def with_trailing_stop(self, pct: float) -> MultiTimeframeEngine:
self._inner.with_trailing_stop(pct)
return self
def with_commission(self, rate: float) -> MultiTimeframeEngine:
self._inner.with_commission(rate)
return self
def with_commission_model(self, model) -> MultiTimeframeEngine:
self._inner.with_commission_model(model)
return self
def with_slippage(self, bps: float) -> MultiTimeframeEngine:
self._inner.with_slippage(bps)
return self
def with_initial_capital(self, capital: float) -> MultiTimeframeEngine:
self._inner.with_initial_capital(capital)
return self
def with_fill_mode(self, mode: str) -> MultiTimeframeEngine:
self._inner.with_fill_mode(mode)
return self
def with_leverage(
self, margin_ratio: float, margin_call_pct: float = 0.5
) -> MultiTimeframeEngine:
self._inner.with_leverage(margin_ratio, margin_call_pct)
return self
def with_loss_limits(
self, daily: float = 0.0, total: float = 0.0
) -> MultiTimeframeEngine:
self._inner.with_loss_limits(daily, total)
return self
def run(
self, close_fine: ArrayLike, **htf_strategy_kwargs
) -> AdvancedBacktestResult:
"""Run multi-timeframe backtest.
1. Resample close_fine (and stored OHLCV) to coarse bars
2. Run htf_strategy on coarse close to get coarse signals
3. Align coarse signals back to fine resolution (repeat each coarse signal `factor` times)
4. Run BacktestEngine on fine bars with aligned signals
Parameters
----------
close_fine : array-like
Fine-resolution close prices.
**htf_strategy_kwargs
Extra keyword arguments passed to the HTF strategy.
Returns
-------
AdvancedBacktestResult
"""
c_fine = np.asarray(close_fine, dtype=np.float64)
n_fine = len(c_fine)
factor = self._factor
# ------------------------------------------------------------------
# 1. Resample close to coarse resolution
# ------------------------------------------------------------------
# Build dummy OHLCV if OHLCV not provided
if self._high is not None and self._low is not None and self._open is not None:
coarse_o, coarse_h, coarse_l, coarse_c, _ = resample_ohlcv(
self._open,
self._high,
self._low,
c_fine,
np.ones(n_fine), # volume placeholder
factor,
)
else:
coarse_o, coarse_h, coarse_l, coarse_c, _ = resample_ohlcv(
c_fine,
c_fine,
c_fine,
c_fine,
np.ones(n_fine),
factor,
)
# ------------------------------------------------------------------
# 2. Compute coarse-bar signals via htf_strategy
# ------------------------------------------------------------------
from ferro_ta.analysis.backtest import _resolve_strategy
strategy_fn = _resolve_strategy(self._htf_strategy)
# Ensure the coarse close array is C-contiguous (required by Rust kernels)
coarse_c = np.ascontiguousarray(coarse_c, dtype=np.float64)
coarse_signals = np.asarray(
strategy_fn(coarse_c, **htf_strategy_kwargs), dtype=np.float64
)
# ------------------------------------------------------------------
# 3. Align coarse signals back to fine resolution
# ------------------------------------------------------------------
aligned_signals = align_to_coarse(coarse_signals, factor, n_fine)
# ------------------------------------------------------------------
# 4. Set up OHLCV on inner engine if provided and run
# ------------------------------------------------------------------
if self._high is not None and self._low is not None and self._open is not None:
self._inner.with_ohlcv(
high=self._high,
low=self._low,
open_=self._open,
)
# Use a passthrough lambda so the already-computed aligned_signals are used
return self._inner.run(
c_fine,
strategy=lambda c, **kw: aligned_signals,
)
+318
View File
@@ -0,0 +1,318 @@
"""
Portfolio optimization utilities.
mean_variance_optimize(returns, target_return=None, allow_short=False)
Minimum-variance portfolio (or target-return portfolio on efficient frontier).
Uses scipy.optimize.minimize with SLSQP.
Returns weight array summing to 1.
risk_parity_optimize(returns, risk_budget=None)
Equal risk contribution portfolio (or custom risk budget).
Each asset contributes equally to total portfolio volatility.
Returns weight array summing to 1.
max_sharpe_optimize(returns, risk_free_rate=0.0)
Maximize Sharpe ratio portfolio.
Returns weight array.
PortfolioOptimizer
Fluent builder that wraps the above functions and integrates with
BacktestEngine for portfolio-level signal generation.
"""
from __future__ import annotations
from typing import Optional
import numpy as np
from numpy.typing import ArrayLike, NDArray
def mean_variance_optimize(
returns: ArrayLike,
target_return: Optional[float] = None,
allow_short: bool = False,
risk_free_rate: float = 0.0,
) -> NDArray:
"""Compute minimum variance (or target return) portfolio weights.
Parameters
----------
returns : (T, N) array of asset returns
target_return : float or None
If None, return minimum-variance portfolio.
If float, return minimum-variance portfolio with this expected return.
allow_short : bool
If False, weights are constrained to [0, 1].
risk_free_rate : float
Not used directly here (kept for API symmetry with max_sharpe).
Returns
-------
weights : (N,) array summing to 1.0
"""
try:
from scipy.optimize import minimize
except ImportError:
raise ImportError(
"scipy is required for portfolio optimization: pip install scipy"
)
r = np.asarray(returns, dtype=np.float64)
if r.ndim == 1:
r = r[:, np.newaxis]
n_assets = r.shape[1]
if n_assets == 1:
return np.array([1.0])
mu = r.mean(axis=0)
cov = np.cov(r, rowvar=False)
# Regularize to handle near-singular covariance matrices
cov += 1e-8 * np.eye(n_assets)
# Objective: minimize portfolio variance w^T @ cov @ w
def portfolio_variance(w: np.ndarray) -> float:
return float(w @ cov @ w)
def portfolio_variance_grad(w: np.ndarray) -> np.ndarray:
return 2.0 * cov @ w
# Constraints: weights sum to 1
constraints = [{"type": "eq", "fun": lambda w: np.sum(w) - 1.0}]
# Optional target return constraint
if target_return is not None:
constraints.append(
{"type": "eq", "fun": lambda w, mu=mu, tr=target_return: float(w @ mu) - tr}
)
# Bounds
bounds = None if allow_short else [(0.0, 1.0)] * n_assets
# Initial guess: equal weights
w0 = np.ones(n_assets) / n_assets
result = minimize(
portfolio_variance,
w0,
jac=portfolio_variance_grad,
method="SLSQP",
bounds=bounds,
constraints=constraints,
options={"ftol": 1e-12, "maxiter": 1000},
)
weights = result.x
# Normalize to ensure exact sum=1 (numerical noise)
weights = weights / weights.sum()
if not allow_short:
weights = np.maximum(weights, 0.0)
s = weights.sum()
if s > 0:
weights /= s
return weights
def risk_parity_optimize(
returns: ArrayLike,
risk_budget: Optional[ArrayLike] = None,
) -> NDArray:
"""Compute risk parity weights (equal risk contribution).
Parameters
----------
returns : (T, N) array of asset returns
risk_budget : (N,) array or None
Target risk contribution per asset (normalized internally). None = equal.
Returns
-------
weights : (N,) array summing to 1.0
"""
try:
from scipy.optimize import minimize
except ImportError:
raise ImportError(
"scipy is required for portfolio optimization: pip install scipy"
)
r = np.asarray(returns, dtype=np.float64)
if r.ndim == 1:
r = r[:, np.newaxis]
n_assets = r.shape[1]
if n_assets == 1:
return np.array([1.0])
cov = np.cov(r, rowvar=False)
cov += 1e-8 * np.eye(n_assets)
if risk_budget is None:
budget = np.ones(n_assets) / n_assets
else:
budget = np.asarray(risk_budget, dtype=np.float64)
budget = budget / budget.sum()
def risk_contribution(w: np.ndarray) -> np.ndarray:
"""Return marginal risk contribution of each asset."""
sigma = np.sqrt(w @ cov @ w)
if sigma < 1e-12:
return np.zeros(n_assets)
mrc = cov @ w / sigma
return w * mrc
def objective(w: np.ndarray) -> float:
"""Minimize squared deviation from target risk budget."""
rc = risk_contribution(w)
total_rc = rc.sum()
if total_rc < 1e-12:
return float(np.sum((rc - budget) ** 2))
rc_normalized = rc / total_rc
return float(np.sum((rc_normalized - budget) ** 2))
constraints = [{"type": "eq", "fun": lambda w: np.sum(w) - 1.0}]
bounds = [(1e-6, 1.0)] * n_assets # risk parity requires positive weights
w0 = np.ones(n_assets) / n_assets
result = minimize(
objective,
w0,
method="SLSQP",
bounds=bounds,
constraints=constraints,
options={"ftol": 1e-12, "maxiter": 2000},
)
weights = result.x
weights = np.maximum(weights, 0.0)
s = weights.sum()
if s > 0:
weights /= s
return weights
def max_sharpe_optimize(
returns: ArrayLike,
risk_free_rate: float = 0.0,
allow_short: bool = False,
) -> NDArray:
"""Compute maximum Sharpe ratio portfolio weights.
Returns
-------
weights : (N,) array summing to 1.0
"""
try:
from scipy.optimize import minimize
except ImportError:
raise ImportError(
"scipy is required for portfolio optimization: pip install scipy"
)
r = np.asarray(returns, dtype=np.float64)
if r.ndim == 1:
r = r[:, np.newaxis]
n_assets = r.shape[1]
if n_assets == 1:
return np.array([1.0])
mu = r.mean(axis=0)
cov = np.cov(r, rowvar=False)
cov += 1e-8 * np.eye(n_assets)
# Maximize Sharpe = minimize negative Sharpe
def neg_sharpe(w: np.ndarray) -> float:
port_return = float(w @ mu)
port_vol = float(np.sqrt(w @ cov @ w))
if port_vol < 1e-12:
return 0.0
return -(port_return - risk_free_rate) / port_vol
constraints = [{"type": "eq", "fun": lambda w: np.sum(w) - 1.0}]
bounds = None if allow_short else [(0.0, 1.0)] * n_assets
w0 = np.ones(n_assets) / n_assets
result = minimize(
neg_sharpe,
w0,
method="SLSQP",
bounds=bounds,
constraints=constraints,
options={"ftol": 1e-12, "maxiter": 1000},
)
weights = result.x
weights = weights / weights.sum()
if not allow_short:
weights = np.maximum(weights, 0.0)
s = weights.sum()
if s > 0:
weights /= s
return weights
class PortfolioOptimizer:
"""Fluent interface for portfolio weight optimization.
Example
-------
weights = (
PortfolioOptimizer()
.with_method("risk_parity")
.with_lookback(252)
.optimize(returns_matrix)
)
"""
def __init__(self) -> None:
self._method: str = "min_variance"
self._lookback: Optional[int] = None
self._allow_short: bool = False
self._risk_free_rate: float = 0.0
self._target_return: Optional[float] = None
self._risk_budget: Optional[NDArray] = None
def with_method(self, method: str) -> PortfolioOptimizer:
"""Method: 'min_variance', 'risk_parity', 'max_sharpe'."""
valid = ("min_variance", "risk_parity", "max_sharpe")
if method not in valid:
raise ValueError(f"method must be one of {valid}")
self._method = method
return self
def with_lookback(self, n_bars: int) -> PortfolioOptimizer:
"""Use only the last n_bars for covariance estimation."""
self._lookback = int(n_bars)
return self
def with_short_selling(self, allow: bool = True) -> PortfolioOptimizer:
self._allow_short = allow
return self
def with_risk_free_rate(self, rate: float) -> PortfolioOptimizer:
self._risk_free_rate = float(rate)
return self
def with_target_return(self, target: float) -> PortfolioOptimizer:
self._target_return = float(target)
return self
def with_risk_budget(self, budget: ArrayLike) -> PortfolioOptimizer:
self._risk_budget = np.asarray(budget, dtype=np.float64)
return self
def optimize(self, returns: ArrayLike) -> NDArray:
"""Run optimization and return weight array."""
r = np.asarray(returns, dtype=np.float64)
if self._lookback is not None:
r = r[-self._lookback :]
if self._method == "min_variance":
return mean_variance_optimize(
r, self._target_return, self._allow_short, self._risk_free_rate
)
elif self._method == "risk_parity":
return risk_parity_optimize(r, self._risk_budget)
else:
return max_sharpe_optimize(r, self._risk_free_rate, self._allow_short)
+277
View File
@@ -0,0 +1,277 @@
"""
Visualization utilities for backtest results.
plot_backtest(result, *, title="Backtest", show=True, return_fig=False)
Generate an interactive Plotly chart with:
- Top panel: equity curve (normalized to 1.0)
- Middle panel: drawdown series (negative values, shaded red)
- Bottom panel: position/signal over time
Optional trade markers: entry (green triangle up) and exit (red triangle down) on equity curve.
Requires plotly -- raises ImportError with install hint if not available.
"""
from __future__ import annotations
from typing import TYPE_CHECKING
if TYPE_CHECKING:
pass
__all__ = ["plot_backtest"]
def plot_backtest(
result, # AdvancedBacktestResult
*,
title: str = "Backtest",
show: bool = True,
return_fig: bool = False,
benchmark: bool = True,
):
"""Plot equity curve, drawdown, and positions.
Parameters
----------
result : AdvancedBacktestResult
Backtest result object with equity, drawdown_series, positions, and trades.
title : str
Chart title.
show : bool
Call fig.show() if True.
return_fig : bool
Return the plotly Figure object.
benchmark : bool
Overlay benchmark equity curve if result has benchmark returns.
Returns
-------
plotly.graph_objects.Figure if return_fig=True, else None.
Raises
------
ImportError
If plotly is not installed.
"""
try:
import plotly.graph_objects as go
from plotly.subplots import make_subplots
except ImportError:
raise ImportError(
"plotly is required for visualization. Install with: pip install plotly"
)
import numpy as np
# ------------------------------------------------------------------
# Extract result fields
# ------------------------------------------------------------------
equity = np.asarray(result.equity, dtype=np.float64)
n = len(equity)
bars = np.arange(n)
# Drawdown: prefer pre-computed drawdown_series, else compute from equity
if hasattr(result, "drawdown_series") and result.drawdown_series is not None:
drawdown = np.asarray(result.drawdown_series, dtype=np.float64)
else:
cum_max = np.maximum.accumulate(equity)
drawdown = np.where(cum_max > 0, equity / cum_max - 1.0, 0.0)
positions = (
np.asarray(result.positions, dtype=np.float64)
if hasattr(result, "positions")
else np.zeros(n)
)
# Trades (may be empty or None)
trades = getattr(result, "trades", None)
# Benchmark equity (optional)
benchmark_equity = None
if (
benchmark
and hasattr(result, "benchmark_equity")
and result.benchmark_equity is not None
):
benchmark_equity = np.asarray(result.benchmark_equity, dtype=np.float64)
# ------------------------------------------------------------------
# Build 3-panel subplot
# ------------------------------------------------------------------
fig = make_subplots(
rows=3,
cols=1,
shared_xaxes=True,
row_heights=[0.5, 0.25, 0.25],
vertical_spacing=0.04,
subplot_titles=("Equity Curve", "Drawdown", "Positions"),
)
# ---- Panel 1: Equity curve ----------------------------------------
fig.add_trace(
go.Scatter(
x=bars,
y=equity,
name="Strategy",
line=dict(color="#00d4ff", width=1.5),
hovertemplate="Bar %{x}<br>Equity: %{y:.4f}<extra></extra>",
),
row=1,
col=1,
)
# Benchmark overlay
if benchmark_equity is not None:
fig.add_trace(
go.Scatter(
x=bars[: len(benchmark_equity)],
y=benchmark_equity,
name="Benchmark",
line=dict(color="#f0a500", width=1.2, dash="dot"),
hovertemplate="Bar %{x}<br>Benchmark: %{y:.4f}<extra></extra>",
),
row=1,
col=1,
)
# Trade markers
if trades is not None and hasattr(trades, "__len__") and len(trades) > 0:
# trades may be a pd.DataFrame or a list of dicts
try:
# pandas DataFrame path
entry_bars = trades["entry_bar"].values
exit_bars = trades["exit_bar"].values
except (TypeError, KeyError, AttributeError):
# list-of-dicts path
try:
entry_bars = np.array([t["entry_bar"] for t in trades])
exit_bars = np.array([t["exit_bar"] for t in trades])
except (KeyError, TypeError):
entry_bars = np.array([])
exit_bars = np.array([])
if len(entry_bars) > 0:
# Clip indices to equity length
entry_bars = np.clip(entry_bars.astype(int), 0, n - 1)
exit_bars = np.clip(exit_bars.astype(int), 0, n - 1)
fig.add_trace(
go.Scatter(
x=entry_bars,
y=equity[entry_bars],
mode="markers",
name="Entry",
marker=dict(
symbol="triangle-up",
size=10,
color="lime",
line=dict(color="darkgreen", width=1),
),
hovertemplate="Entry Bar %{x}<br>Equity: %{y:.4f}<extra></extra>",
),
row=1,
col=1,
)
fig.add_trace(
go.Scatter(
x=exit_bars,
y=equity[exit_bars],
mode="markers",
name="Exit",
marker=dict(
symbol="triangle-down",
size=10,
color="red",
line=dict(color="darkred", width=1),
),
hovertemplate="Exit Bar %{x}<br>Equity: %{y:.4f}<extra></extra>",
),
row=1,
col=1,
)
# ---- Panel 2: Drawdown -------------------------------------------
fig.add_trace(
go.Scatter(
x=bars,
y=drawdown,
name="Drawdown",
fill="tozeroy",
fillcolor="rgba(220, 50, 50, 0.25)",
line=dict(color="rgba(220, 50, 50, 0.8)", width=1.0),
hovertemplate="Bar %{x}<br>Drawdown: %{y:.2%}<extra></extra>",
),
row=2,
col=1,
)
# ---- Panel 3: Positions ------------------------------------------
fig.add_trace(
go.Scatter(
x=bars,
y=positions,
name="Position",
fill="tozeroy",
fillcolor="rgba(0, 150, 255, 0.2)",
line=dict(color="rgba(0, 150, 255, 0.7)", width=1.0),
hovertemplate="Bar %{x}<br>Position: %{y:.2f}<extra></extra>",
),
row=3,
col=1,
)
# ------------------------------------------------------------------
# Styling: dark theme + ferro-ta branding
# ------------------------------------------------------------------
metrics = getattr(result, "metrics", {})
sharpe_str = f"Sharpe: {metrics.get('sharpe', float('nan')):.2f}" if metrics else ""
dd_str = (
f"Max DD: {metrics.get('max_drawdown', float('nan')):.1%}" if metrics else ""
)
subtitle = " | ".join(filter(None, [sharpe_str, dd_str]))
fig.update_layout(
title=dict(
text=f"<b>{title}</b>" + (f"<br><sub>{subtitle}</sub>" if subtitle else ""),
font=dict(size=18, color="#e0e0e0"),
),
template="plotly_dark",
paper_bgcolor="#0e1117",
plot_bgcolor="#0e1117",
font=dict(color="#b0b8c1", size=11),
legend=dict(
orientation="h",
yanchor="bottom",
y=1.01,
xanchor="right",
x=1,
bgcolor="rgba(0,0,0,0)",
),
hovermode="x unified",
height=700,
margin=dict(l=60, r=40, t=80, b=40),
)
# Axis styling
axis_style = dict(
gridcolor="rgba(255,255,255,0.07)",
zerolinecolor="rgba(255,255,255,0.15)",
tickfont=dict(size=10),
)
fig.update_xaxes(**axis_style)
fig.update_yaxes(**axis_style)
# Y-axis labels
fig.update_yaxes(title_text="Equity (norm.)", row=1, col=1)
fig.update_yaxes(title_text="Drawdown", tickformat=".1%", row=2, col=1)
fig.update_yaxes(title_text="Position", row=3, col=1)
fig.update_xaxes(title_text="Bar", row=3, col=1)
# ------------------------------------------------------------------
if show:
fig.show()
if return_fig:
return fig
return None
+258
View File
@@ -277,6 +277,264 @@ def regime(
raise ValueError(f"Unknown regime method '{method}'. Use 'adx' or 'combined'.")
# ---------------------------------------------------------------------------
# Phase 4: Volatility/Trend regime detection (pure NumPy)
# ---------------------------------------------------------------------------
try:
from ferro_ta._ferro_ta import sma as _rust_sma
except ImportError:
_rust_sma = None
def _rolling_sma_pure(arr: np.ndarray, window: int) -> np.ndarray:
"""Rolling SMA — delegates to the Rust SMA when available."""
if _rust_sma is not None:
return np.asarray(_rust_sma(arr, window), dtype=np.float64)
# Fallback: O(n) rolling SMA using cumsum
n = len(arr)
out = np.full(n, np.nan)
if window > n:
return out
cs = np.cumsum(arr)
out[window - 1] = cs[window - 1] / window
if window < n:
out[window:] = (cs[window:] - cs[: n - window]) / window
return out
def _rolling_std_pure(arr: np.ndarray, window: int) -> np.ndarray:
"""O(n) rolling std using cumsum-of-squares on the valid (non-NaN) portion.
Handles leading NaN values (e.g., log returns where arr[0] is NaN).
NaN is returned for warm-up bars.
"""
n = len(arr)
out = np.full(n, np.nan)
if window < 2 or window > n:
return out
# Find the first non-NaN index
first_valid = 0
while first_valid < n and np.isnan(arr[first_valid]):
first_valid += 1
if first_valid >= n:
return out # all NaN
# Work on the valid slice
valid_slice = arr[first_valid:]
m = len(valid_slice)
if window > m:
return out
cs = np.cumsum(valid_slice)
cs2 = np.cumsum(valid_slice**2)
n_windows = m - window + 1
s = np.empty(n_windows)
s2 = np.empty(n_windows)
s[0] = cs[window - 1]
s2[0] = cs2[window - 1]
if n_windows > 1:
s[1:] = cs[window:] - cs[: m - window]
s2[1:] = cs2[window:] - cs2[: m - window]
mean = s / window
var = np.maximum(s2 / window - mean**2, 0.0)
stds = np.sqrt(var)
# Place back into output (first result is at index first_valid + window - 1)
start_out = first_valid + window - 1
out[start_out : start_out + n_windows] = stds
return out
def detect_volatility_regime(
close: ArrayLike,
window: int = 20,
n_regimes: int = 3,
) -> NDArray:
"""Label bars by rolling volatility percentile bucket (0 = lowest vol regime).
Uses rolling standard deviation of log returns. NaN for warm-up bars
(returned as -1 in the integer output).
Parameters
----------
close : array-like
Close price series.
window : int
Rolling window for std computation (default 20).
n_regimes : int
Number of volatility regimes (default 3: low/mid/high = 0/1/2).
Returns
-------
NDArray[int64]
Integer array where each element is in {-1, 0, ..., n_regimes-1}.
-1 indicates NaN (warm-up) bars.
"""
c = np.asarray(close, dtype=np.float64)
n = len(c)
out = np.full(n, -1, dtype=np.int64)
log_ret = np.full(n, np.nan)
with np.errstate(divide="ignore", invalid="ignore"):
log_ret[1:] = np.log(c[1:] / c[:-1])
rolling_vol = _rolling_std_pure(log_ret, window)
valid = ~np.isnan(rolling_vol)
if not np.any(valid):
return out
vol_vals = rolling_vol[valid]
pcts = [100.0 * k / n_regimes for k in range(1, n_regimes)]
boundaries = np.percentile(vol_vals, pcts) if pcts else np.array([])
labels = np.digitize(vol_vals, boundaries).astype(np.int64)
out[valid] = labels
return out
def detect_trend_regime(
close: ArrayLike,
fast: int = 50,
slow: int = 200,
) -> NDArray:
"""Label bars: 1=bull (fast SMA > slow SMA), -1=bear, 0=sideways/NaN warmup.
Parameters
----------
close : array-like
Close price series.
fast : int
Fast SMA period (default 50).
slow : int
Slow SMA period (default 200).
Returns
-------
NDArray[int64]
Integer array with values in {-1, 0, 1}.
0 for warm-up bars where either SMA is NaN.
"""
c = np.asarray(close, dtype=np.float64)
n = len(c)
out = np.zeros(n, dtype=np.int64)
fast_sma = _rolling_sma_pure(c, fast)
slow_sma = _rolling_sma_pure(c, slow)
valid = ~np.isnan(fast_sma) & ~np.isnan(slow_sma)
out[valid & (fast_sma > slow_sma)] = 1
out[valid & (fast_sma < slow_sma)] = -1
return out
def detect_combined_regime(
close: ArrayLike,
vol_window: int = 20,
fast: int = 50,
slow: int = 200,
) -> NDArray:
"""Combine trend + vol into 6-state integer regime label.
States: 0=bull+low-vol, 1=bull+mid-vol, 2=bull+high-vol,
3=bear+low-vol, 4=bear+mid-vol, 5=bear+high-vol.
NaN bars (warm-up or sideways) -1.
Parameters
----------
close : array-like
Close price series.
vol_window : int
Rolling window for volatility regime detection.
fast, slow : int
SMA periods for trend regime detection.
Returns
-------
NDArray[int64]
Integer array with values in {-1, 0, 1, 2, 3, 4, 5}.
"""
c = np.asarray(close, dtype=np.float64)
n = len(c)
out = np.full(n, -1, dtype=np.int64)
trend = detect_trend_regime(c, fast=fast, slow=slow)
vol = detect_volatility_regime(c, window=vol_window, n_regimes=3)
bull_valid = (trend == 1) & (vol >= 0)
bear_valid = (trend == -1) & (vol >= 0)
out[bull_valid] = vol[bull_valid] # 0, 1, or 2
out[bear_valid] = 3 + vol[bear_valid] # 3, 4, or 5
return out
class RegimeFilter:
"""Filter trading signals to only fire in allowed market regimes.
Parameters
----------
allowed_regimes : list[int]
Which regime labels to trade in. Signals in other regimes are zeroed out.
vol_window : int
Rolling window for volatility regime detection.
fast, slow : int
SMA periods for trend regime detection.
"""
def __init__(
self,
allowed_regimes: list[int],
vol_window: int = 20,
fast: int = 50,
slow: int = 200,
) -> None:
self.allowed_regimes = list(allowed_regimes)
self._allowed_regimes_arr = np.array(allowed_regimes, dtype=np.int64)
self.vol_window = int(vol_window)
self.fast = int(fast)
self.slow = int(slow)
def filter(self, signals: ArrayLike, close: ArrayLike) -> NDArray:
"""Zero out signals where regime is not in allowed_regimes.
Parameters
----------
signals : array-like
Signal array (+1, -1, 0, or NaN).
close : array-like
Close price series (same length as signals).
Returns
-------
NDArray[float64]
Filtered signal array signals in disallowed regimes are set to 0.
"""
s = np.asarray(signals, dtype=np.float64).copy()
regimes = detect_combined_regime(
close,
vol_window=self.vol_window,
fast=self.fast,
slow=self.slow,
)
in_allowed = np.isin(regimes, self._allowed_regimes_arr)
s[~in_allowed] = 0.0
return s
# ---------------------------------------------------------------------------
# (original structural_breaks below)
# ---------------------------------------------------------------------------
def structural_breaks(
series: ArrayLike,
method: str = "cusum",
+139
View File
@@ -0,0 +1,139 @@
"""
OHLCV bar aggregation utilities.
resample_ohlcv(open, high, low, close, volume, factor)
Aggregate every `factor` bars into one OHLCV bar.
open = first bar's open
high = max of highs
low = min of lows
close = last bar's close
volume = sum of volumes
resample_ohlcv_labels(n_bars, factor)
Return an integer label array of length n_bars where label[i] = i // factor.
Useful for aligning fine-bar signals with coarse-bar indicators.
align_to_coarse(coarse_values, factor, n_fine_bars)
Broadcast a coarse-bar array back to fine-bar length by repeating each value `factor` times.
Handles the case where n_fine_bars % factor != 0 (last group may be partial).
"""
import numpy as np
from numpy.typing import ArrayLike, NDArray
__all__ = ["resample_ohlcv", "resample_ohlcv_labels", "align_to_coarse"]
def resample_ohlcv(
open_: ArrayLike,
high: ArrayLike,
low: ArrayLike,
close: ArrayLike,
volume: ArrayLike,
factor: int,
) -> tuple[NDArray, NDArray, NDArray, NDArray, NDArray]:
"""Aggregate fine-bar OHLCV into coarser bars.
Parameters
----------
open_ : array-like
Fine-bar open prices.
high : array-like
Fine-bar high prices.
low : array-like
Fine-bar low prices.
close : array-like
Fine-bar close prices.
volume : array-like
Fine-bar volume.
factor : int
Number of fine bars per coarse bar (e.g. 5 for 1-min -> 5-min).
Returns
-------
(open, high, low, close, volume) arrays of length ceil(n / factor).
Only complete groups are returned if n % factor != 0, trailing bars are dropped.
"""
if factor < 1:
raise ValueError(f"factor must be >= 1, got {factor}")
o = np.asarray(open_, dtype=np.float64)
h = np.asarray(high, dtype=np.float64)
low_arr = np.asarray(low, dtype=np.float64)
c = np.asarray(close, dtype=np.float64)
v = np.asarray(volume, dtype=np.float64)
n = len(o)
n_complete = (n // factor) * factor # truncate to complete bars
o = o[:n_complete].reshape(-1, factor)
h = h[:n_complete].reshape(-1, factor)
low_arr = low_arr[:n_complete].reshape(-1, factor)
c = c[:n_complete].reshape(-1, factor)
v = v[:n_complete].reshape(-1, factor)
return (
o[:, 0], # open = first bar's open
h.max(axis=1), # high = max of highs
low_arr.min(axis=1), # low = min of lows
c[:, -1], # close = last bar's close
v.sum(axis=1), # volume = sum of volumes
)
def resample_ohlcv_labels(n_bars: int, factor: int) -> NDArray:
"""Return coarse-bar index for each fine bar (i // factor).
Parameters
----------
n_bars : int
Number of fine-resolution bars.
factor : int
Number of fine bars per coarse bar.
Returns
-------
NDArray of int64, shape (n_bars,), where label[i] = i // factor.
"""
if factor < 1:
raise ValueError(f"factor must be >= 1, got {factor}")
return np.arange(n_bars, dtype=np.int64) // factor
def align_to_coarse(coarse_values: ArrayLike, factor: int, n_fine_bars: int) -> NDArray:
"""Broadcast coarse-bar array back to fine-bar resolution.
Each coarse value is repeated `factor` times. If n_fine_bars % factor != 0,
the last coarse value covers the partial group at the end.
Parameters
----------
coarse_values : array-like
Values at coarse resolution, shape (n_coarse,).
factor : int
Number of fine bars per coarse bar.
n_fine_bars : int
Total number of fine bars to produce.
Returns
-------
NDArray of shape (n_fine_bars,).
"""
if factor < 1:
raise ValueError(f"factor must be >= 1, got {factor}")
coarse = np.asarray(coarse_values, dtype=np.float64)
n_coarse = len(coarse)
# Build the full repeated array (may be longer than n_fine_bars if partial group exists)
repeated = np.repeat(coarse, factor)
# If repeated is shorter than n_fine_bars (shouldn't happen with correct n_coarse,
# but handle defensively), pad with last value
if len(repeated) < n_fine_bars:
pad = np.full(
n_fine_bars - len(repeated), coarse[-1] if n_coarse > 0 else np.nan
)
repeated = np.concatenate([repeated, pad])
return repeated[:n_fine_bars]
+10 -5
View File
@@ -39,6 +39,7 @@ from ferro_ta._ferro_ta import aggregate_tick_bars as _rust_tick_bars
from ferro_ta._ferro_ta import aggregate_time_bars as _rust_time_bars
from ferro_ta._ferro_ta import aggregate_volume_bars_ticks as _rust_volume_bars_ticks
from ferro_ta._utils import _to_f64
from ferro_ta.core.exceptions import FerroTAValueError
__all__ = [
"aggregate_ticks",
@@ -61,23 +62,25 @@ def _parse_rule(rule: str) -> tuple[str, float]:
"""
parts = rule.split(":", 1)
if len(parts) != 2:
raise ValueError(
raise FerroTAValueError(
f"Invalid rule format: {rule!r}. "
"Expected 'time:<seconds>', 'volume:<threshold>', or 'tick:<n>'."
)
bar_type = parts[0].lower().strip()
if bar_type not in ("time", "volume", "tick"):
raise ValueError(
raise FerroTAValueError(
f"Unknown bar type {bar_type!r}. Supported types: 'time', 'volume', 'tick'."
)
try:
param = float(parts[1].strip())
except ValueError as exc:
raise ValueError(
raise FerroTAValueError(
f"Cannot parse parameter {parts[1]!r} as a number in rule {rule!r}."
) from exc
if param <= 0:
raise ValueError(f"Rule parameter must be > 0, got {param} in rule {rule!r}.")
raise FerroTAValueError(
f"Rule parameter must be > 0, got {param} in rule {rule!r}."
)
return bar_type, param
@@ -171,7 +174,9 @@ def aggregate_ticks(
extra = None
else: # time
if ts_arr is None:
raise ValueError("Time bars require a timestamp column in the tick data.")
raise FerroTAValueError(
"Time bars require a timestamp column in the tick data."
)
period_secs = int(param)
labels = (ts_arr // period_secs).astype(np.int64)
ro, rh, rl, rc, rv, lbl = _rust_time_bars(price_arr, size_arr, labels)
+46 -17
View File
@@ -66,6 +66,7 @@ from ferro_ta._ferro_ta import (
vwma as _rust_vwma,
)
from ferro_ta._utils import _to_f64
from ferro_ta.core.exceptions import FerroTAValueError, _normalize_rust_error
def VWAP(
@@ -103,14 +104,15 @@ def VWAP(
Implemented in Rust for maximum performance.
"""
if timeperiod < 0:
from ferro_ta.core.exceptions import FerroTAValueError
raise FerroTAValueError("timeperiod must be >= 0 for VWAP")
h = _to_f64(high)
lo = _to_f64(low)
c = _to_f64(close)
v = _to_f64(volume)
return np.asarray(_rust_vwap(h, lo, c, v, timeperiod))
try:
return np.asarray(_rust_vwap(h, lo, c, v, timeperiod))
except ValueError as e:
_normalize_rust_error(e)
def SUPERTREND(
@@ -164,7 +166,10 @@ def SUPERTREND(
h = _to_f64(high)
lo = _to_f64(low)
c = _to_f64(close)
st, d = _rust_supertrend(h, lo, c, timeperiod, multiplier)
try:
st, d = _rust_supertrend(h, lo, c, timeperiod, multiplier)
except ValueError as e:
_normalize_rust_error(e)
return np.asarray(st), np.asarray(d)
@@ -205,9 +210,12 @@ def ICHIMOKU(
h = _to_f64(high)
lo = _to_f64(low)
c = _to_f64(close)
t, k, sa, sb, ch = _rust_ichimoku(
h, lo, c, tenkan_period, kijun_period, senkou_b_period, displacement
)
try:
t, k, sa, sb, ch = _rust_ichimoku(
h, lo, c, tenkan_period, kijun_period, senkou_b_period, displacement
)
except ValueError as e:
_normalize_rust_error(e)
return (
np.asarray(t),
np.asarray(k),
@@ -241,7 +249,10 @@ def DONCHIAN(
"""
h = _to_f64(high)
lo = _to_f64(low)
upper, middle, lower = _rust_donchian(h, lo, timeperiod)
try:
upper, middle, lower = _rust_donchian(h, lo, timeperiod)
except ValueError as e:
_normalize_rust_error(e)
return np.asarray(upper), np.asarray(middle), np.asarray(lower)
@@ -279,13 +290,16 @@ def PIVOT_POINTS(
"""
valid_methods = {"classic", "fibonacci", "camarilla"}
if method.lower() not in valid_methods:
raise ValueError(
raise FerroTAValueError(
f"Unknown pivot method '{method}'. Use 'classic', 'fibonacci', or 'camarilla'."
)
h = _to_f64(high)
lo = _to_f64(low)
c = _to_f64(close)
pivot, r1, s1, r2, s2 = _rust_pivot_points(h, lo, c, method)
try:
pivot, r1, s1, r2, s2 = _rust_pivot_points(h, lo, c, method)
except ValueError as e:
_normalize_rust_error(e)
return (
np.asarray(pivot),
np.asarray(r1),
@@ -328,9 +342,12 @@ def KELTNER_CHANNELS(
h = _to_f64(high)
lo = _to_f64(low)
c = _to_f64(close)
upper, middle, lower = _rust_keltner_channels(
h, lo, c, timeperiod, atr_period, multiplier
)
try:
upper, middle, lower = _rust_keltner_channels(
h, lo, c, timeperiod, atr_period, multiplier
)
except ValueError as e:
_normalize_rust_error(e)
return np.asarray(upper), np.asarray(middle), np.asarray(lower)
@@ -358,7 +375,10 @@ def HULL_MA(
Implemented in Rust all WMA computations are in-process.
"""
c = _to_f64(close)
return np.asarray(_rust_hull_ma(c, timeperiod))
try:
return np.asarray(_rust_hull_ma(c, timeperiod))
except ValueError as e:
_normalize_rust_error(e)
def CHANDELIER_EXIT(
@@ -391,7 +411,10 @@ def CHANDELIER_EXIT(
h = _to_f64(high)
lo = _to_f64(low)
c = _to_f64(close)
long_exit, short_exit = _rust_chandelier_exit(h, lo, c, timeperiod, multiplier)
try:
long_exit, short_exit = _rust_chandelier_exit(h, lo, c, timeperiod, multiplier)
except ValueError as e:
_normalize_rust_error(e)
return np.asarray(long_exit), np.asarray(short_exit)
@@ -419,7 +442,10 @@ def VWMA(
"""
c = _to_f64(close)
v = _to_f64(volume)
return np.asarray(_rust_vwma(c, v, timeperiod))
try:
return np.asarray(_rust_vwma(c, v, timeperiod))
except ValueError as e:
_normalize_rust_error(e)
def CHOPPINESS_INDEX(
@@ -452,7 +478,10 @@ def CHOPPINESS_INDEX(
h = _to_f64(high)
lo = _to_f64(low)
c = _to_f64(close)
return np.asarray(_rust_choppiness_index(h, lo, c, timeperiod))
try:
return np.asarray(_rust_choppiness_index(h, lo, c, timeperiod))
except ValueError as e:
_normalize_rust_error(e)
__all__ = [
+1 -1
View File
@@ -60,7 +60,7 @@ CARRIERS = [
VersionCarrier(
"cargo_core_dep",
ROOT / "Cargo.toml",
r'(ferro_ta_core = \{ path = "crates/ferro_ta_core", version = ")([^"]+)(" \})',
r'(ferro_ta_core = \{ path = "crates/ferro_ta_core", version = ")([^"]+)("[^}]*\})',
r"\g<1>{version}\g<3>",
),
VersionCarrier(
+20 -192
View File
@@ -1,18 +1,9 @@
//! Tick / Trade Aggregation Pipeline — Rust implementations.
//!
//! Aggregates raw tick/trade data into OHLCV bars:
//! - **time bars** — fixed duration buckets (label-based via Python timestamps)
//! - **volume bars** — fixed volume threshold per bar
//! - **tick bars** — fixed number of ticks per bar
//!
//! The Python layer (ferro_ta.aggregation) provides the timestamp bucketing
//! for time bars; this module handles the compute-intensive OHLCV accumulation.
//! Tick/trade aggregation (thin PyO3 wrapper over ferro_ta_core::aggregation).
use numpy::{IntoPyArray, PyArray1, PyReadonlyArray1};
use pyo3::exceptions::PyValueError;
use pyo3::prelude::*;
/// Return type for functions that return five OHLCV 1-D arrays.
type Ohlcv5<'py> = (
Bound<'py, PyArray1<f64>>,
Bound<'py, PyArray1<f64>>,
@@ -21,7 +12,6 @@ type Ohlcv5<'py> = (
Bound<'py, PyArray1<f64>>,
);
/// Return type for time bars: five OHLCV arrays plus labels.
type Ohlcv5AndLabels<'py> = (
Bound<'py, PyArray1<f64>>,
Bound<'py, PyArray1<f64>>,
@@ -31,21 +21,7 @@ type Ohlcv5AndLabels<'py> = (
Bound<'py, PyArray1<i64>>,
);
// ---------------------------------------------------------------------------
// aggregate_tick_bars
// ---------------------------------------------------------------------------
/// Aggregate tick/trade data into tick bars (every N ticks become one bar).
///
/// Parameters
/// ----------
/// price, size : 1-D float64 arrays (equal length, one entry per trade/tick)
/// ticks_per_bar : int — number of ticks per bar (must be >= 1)
///
/// Returns
/// -------
/// Tuple of five 1-D arrays: (open, high, low, close, volume)
/// where volume = sum of sizes in each bar.
#[pyfunction]
#[pyo3(signature = (price, size, ticks_per_bar))]
pub fn aggregate_tick_bars<'py>(
@@ -65,58 +41,17 @@ pub fn aggregate_tick_bars<'py>(
"price and size must be non-empty and equal length",
));
}
let n_bars = n.div_ceil(ticks_per_bar);
let mut out_open = Vec::with_capacity(n_bars);
let mut out_high = Vec::with_capacity(n_bars);
let mut out_low = Vec::with_capacity(n_bars);
let mut out_close = Vec::with_capacity(n_bars);
let mut out_vol = Vec::with_capacity(n_bars);
let mut i = 0;
while i < n {
let end = (i + ticks_per_bar).min(n);
let bar_p = &p[i..end];
let bar_s = &s[i..end];
let bar_open = bar_p[0];
let bar_high = bar_p.iter().cloned().fold(f64::NEG_INFINITY, f64::max);
let bar_low = bar_p.iter().cloned().fold(f64::INFINITY, f64::min);
let bar_close = *bar_p.last().unwrap();
let bar_vol: f64 = bar_s.iter().sum();
out_open.push(bar_open);
out_high.push(bar_high);
out_low.push(bar_low);
out_close.push(bar_close);
out_vol.push(bar_vol);
i = end;
}
let (ro, rh, rl, rc, rv) = ferro_ta_core::aggregation::aggregate_tick_bars(p, s, ticks_per_bar);
Ok((
out_open.into_pyarray(py),
out_high.into_pyarray(py),
out_low.into_pyarray(py),
out_close.into_pyarray(py),
out_vol.into_pyarray(py),
ro.into_pyarray(py),
rh.into_pyarray(py),
rl.into_pyarray(py),
rc.into_pyarray(py),
rv.into_pyarray(py),
))
}
// ---------------------------------------------------------------------------
// aggregate_volume_bars_ticks
// ---------------------------------------------------------------------------
/// Aggregate tick data into volume bars (fixed volume threshold).
///
/// Accumulates ticks until cumulative size >= `volume_threshold`, then emits
/// a bar.
///
/// Parameters
/// ----------
/// price, size : 1-D float64 arrays (equal length)
/// volume_threshold : float — cumulative size threshold per bar (must be > 0)
///
/// Returns
/// -------
/// Tuple of five 1-D arrays: (open, high, low, close, volume)
#[pyfunction]
#[pyo3(signature = (price, size, volume_threshold))]
pub fn aggregate_volume_bars_ticks<'py>(
@@ -136,78 +71,17 @@ pub fn aggregate_volume_bars_ticks<'py>(
"price and size must be non-empty and equal length",
));
}
let mut out_open: Vec<f64> = Vec::new();
let mut out_high: Vec<f64> = Vec::new();
let mut out_low: Vec<f64> = Vec::new();
let mut out_close: Vec<f64> = Vec::new();
let mut out_vol: Vec<f64> = Vec::new();
let mut bar_open = p[0];
let mut bar_high = p[0];
let mut bar_low = p[0];
let mut bar_close = p[0];
let mut bar_vol = s[0];
for i in 1..n {
bar_high = bar_high.max(p[i]);
bar_low = bar_low.min(p[i]);
bar_close = p[i];
bar_vol += s[i];
if bar_vol >= volume_threshold {
out_open.push(bar_open);
out_high.push(bar_high);
out_low.push(bar_low);
out_close.push(bar_close);
out_vol.push(bar_vol);
if i + 1 < n {
bar_open = p[i + 1];
bar_high = p[i + 1];
bar_low = p[i + 1];
bar_close = p[i + 1];
bar_vol = s[i + 1];
} else {
bar_vol = 0.0;
}
}
}
// Push remaining partial bar
if bar_vol > 0.0 {
out_open.push(bar_open);
out_high.push(bar_high);
out_low.push(bar_low);
out_close.push(bar_close);
out_vol.push(bar_vol);
}
let (ro, rh, rl, rc, rv) = ferro_ta_core::aggregation::aggregate_volume_bars_ticks(p, s, volume_threshold);
Ok((
out_open.into_pyarray(py),
out_high.into_pyarray(py),
out_low.into_pyarray(py),
out_close.into_pyarray(py),
out_vol.into_pyarray(py),
ro.into_pyarray(py),
rh.into_pyarray(py),
rl.into_pyarray(py),
rc.into_pyarray(py),
rv.into_pyarray(py),
))
}
// ---------------------------------------------------------------------------
// aggregate_time_bars
// ---------------------------------------------------------------------------
/// Aggregate tick data into time bars using pre-computed integer bucket labels.
///
/// Each tick is assigned a `label` (e.g. unix_ts // period_secs). Ticks with
/// the same label are accumulated into one bar. Labels must be non-decreasing.
///
/// Parameters
/// ----------
/// price, size : 1-D float64 arrays
/// labels : 1-D int64 array — bucket label per tick (non-decreasing)
///
/// Returns
/// -------
/// Tuple of five 1-D arrays: (open, high, low, close, volume)
/// and a 1-D int64 array of unique labels (one per bar).
#[pyfunction]
#[pyo3(signature = (price, size, labels))]
pub fn aggregate_time_bars<'py>(
@@ -225,63 +99,17 @@ pub fn aggregate_time_bars<'py>(
"price, size, and labels must be non-empty and equal length",
));
}
let mut out_open: Vec<f64> = Vec::new();
let mut out_high: Vec<f64> = Vec::new();
let mut out_low: Vec<f64> = Vec::new();
let mut out_close: Vec<f64> = Vec::new();
let mut out_vol: Vec<f64> = Vec::new();
let mut out_labels: Vec<i64> = Vec::new();
let mut cur_label = lbl[0];
let mut bar_open = p[0];
let mut bar_high = p[0];
let mut bar_low = p[0];
let mut bar_close = p[0];
let mut bar_vol = s[0];
for i in 1..n {
if lbl[i] != cur_label {
out_open.push(bar_open);
out_high.push(bar_high);
out_low.push(bar_low);
out_close.push(bar_close);
out_vol.push(bar_vol);
out_labels.push(cur_label);
cur_label = lbl[i];
bar_open = p[i];
bar_high = p[i];
bar_low = p[i];
bar_close = p[i];
bar_vol = s[i];
} else {
bar_high = bar_high.max(p[i]);
bar_low = bar_low.min(p[i]);
bar_close = p[i];
bar_vol += s[i];
}
}
out_open.push(bar_open);
out_high.push(bar_high);
out_low.push(bar_low);
out_close.push(bar_close);
out_vol.push(bar_vol);
out_labels.push(cur_label);
let (ro, rh, rl, rc, rv, rlbl) = ferro_ta_core::aggregation::aggregate_time_bars(p, s, lbl);
Ok((
out_open.into_pyarray(py),
out_high.into_pyarray(py),
out_low.into_pyarray(py),
out_close.into_pyarray(py),
out_vol.into_pyarray(py),
out_labels.into_pyarray(py),
ro.into_pyarray(py),
rh.into_pyarray(py),
rl.into_pyarray(py),
rc.into_pyarray(py),
rv.into_pyarray(py),
rlbl.into_pyarray(py),
))
}
// ---------------------------------------------------------------------------
// Register
// ---------------------------------------------------------------------------
pub fn register(m: &Bound<'_, PyModule>) -> PyResult<()> {
m.add_function(wrap_pyfunction!(aggregate_tick_bars, m)?)?;
m.add_function(wrap_pyfunction!(aggregate_volume_bars_ticks, m)?)?;
+13 -110
View File
@@ -1,40 +1,20 @@
//! Alerts — condition evaluation helpers.
//!
//! These Rust functions evaluate conditions over price/indicator series and
//! return boolean or integer arrays indicating where conditions fire. They
//! are designed to be called once per batch (backtest) or per bar (live) and
//! return the full history of firings.
//!
//! Functions
//! ---------
//! - `check_threshold` — fires when a series crosses above/below a level
//! - `check_cross` — fires when *fast* crosses above or below *slow*
//! - `collect_alert_bars` — returns indices of bars where a bool mask is True
//! Alerts — condition evaluation helpers (thin PyO3 wrapper over ferro_ta_core::alerts).
use numpy::{IntoPyArray, PyArray1, PyReadonlyArray1};
use pyo3::exceptions::PyValueError;
use pyo3::prelude::*;
// ---------------------------------------------------------------------------
// check_threshold
// ---------------------------------------------------------------------------
/// Fire an alert when *series* crosses a threshold level.
///
/// Parameters
/// ----------
/// series : 1-D float64 array — indicator values (e.g. RSI)
/// series : 1-D float64 array
/// level : float — threshold value
/// direction : int
/// ``1`` → fire when series crosses **above** *level* (value goes from
/// ≤ level to > level).
/// ``-1`` → fire when series crosses **below** *level* (value goes from
/// ≥ level to < level).
/// direction : int — ``1`` (cross above) or ``-1`` (cross below)
///
/// Returns
/// -------
/// 1-D int8 array — 1 at the bar where the crossing occurs, 0 elsewhere.
/// Element 0 is always 0 (no crossing possible without a prior bar).
/// 1-D int8 array — 1 at crossing bars, 0 elsewhere.
#[pyfunction]
pub fn check_threshold<'py>(
py: Python<'py>,
@@ -48,50 +28,15 @@ pub fn check_threshold<'py>(
));
}
let s = series.as_slice()?;
let n = s.len();
let mut out = vec![0i8; n];
if n < 2 {
return Ok(out.into_pyarray(py));
}
for i in 1..n {
let prev = s[i - 1];
let curr = s[i];
if prev.is_nan() || curr.is_nan() {
continue;
}
if direction == 1 {
// cross above: was at or below level, now above
if prev <= level && curr > level {
out[i] = 1;
}
} else {
// cross below: was at or above level, now below
if prev >= level && curr < level {
out[i] = 1;
}
}
}
Ok(out.into_pyarray(py))
let result = ferro_ta_core::alerts::check_threshold(s, level, direction);
Ok(result.into_pyarray(py))
}
// ---------------------------------------------------------------------------
// check_cross
// ---------------------------------------------------------------------------
/// Detect cross-over / cross-under events between two series.
///
/// Parameters
/// ----------
/// fast : 1-D float64 array — the "fast" series (e.g. short SMA)
/// slow : 1-D float64 array — the "slow" series (e.g. long SMA)
///
/// Returns
/// -------
/// 1-D int8 array:
/// ``1`` at bars where *fast* crosses **above** *slow* (bullish cross)
/// ``-1`` at bars where *fast* crosses **below** *slow* (bearish cross)
/// ``0`` elsewhere
/// Element 0 is always 0.
/// 1-D int8 array: ``1`` = bullish, ``-1`` = bearish, ``0`` = none.
#[pyfunction]
pub fn check_cross<'py>(
py: Python<'py>,
@@ -100,68 +45,26 @@ pub fn check_cross<'py>(
) -> PyResult<Bound<'py, PyArray1<i8>>> {
let f = fast.as_slice()?;
let s = slow.as_slice()?;
let n = f.len();
if n != s.len() {
if f.len() != s.len() {
return Err(PyValueError::new_err(
"fast and slow must have the same length",
));
}
let mut out = vec![0i8; n];
if n < 2 {
return Ok(out.into_pyarray(py));
}
for i in 1..n {
let fp = f[i - 1];
let fc = f[i];
let sp = s[i - 1];
let sc = s[i];
if fp.is_nan() || fc.is_nan() || sp.is_nan() || sc.is_nan() {
continue;
}
// Bullish: fast was below slow, now above
if fp <= sp && fc > sc {
out[i] = 1;
}
// Bearish: fast was above slow, now below
else if fp >= sp && fc < sc {
out[i] = -1;
}
}
Ok(out.into_pyarray(py))
let result = ferro_ta_core::alerts::check_cross(f, s);
Ok(result.into_pyarray(py))
}
// ---------------------------------------------------------------------------
// collect_alert_bars
// ---------------------------------------------------------------------------
/// Collect bar indices where *mask* is non-zero (i.e. condition fired).
///
/// Parameters
/// ----------
/// mask : 1-D int8 array (output of ``check_threshold`` or ``check_cross``)
///
/// Returns
/// -------
/// 1-D int64 array — indices of fired bars (ascending order)
/// Collect bar indices where *mask* is non-zero.
#[pyfunction]
pub fn collect_alert_bars<'py>(
py: Python<'py>,
mask: PyReadonlyArray1<'py, i8>,
) -> PyResult<Bound<'py, PyArray1<i64>>> {
let m = mask.as_slice()?;
let indices: Vec<i64> = m
.iter()
.enumerate()
.filter(|(_, &v)| v != 0)
.map(|(i, _)| i as i64)
.collect();
Ok(indices.into_pyarray(py))
let result = ferro_ta_core::alerts::collect_alert_bars(m);
Ok(result.into_pyarray(py))
}
// ---------------------------------------------------------------------------
// Register
// ---------------------------------------------------------------------------
pub fn register(m: &Bound<'_, PyModule>) -> PyResult<()> {
m.add_function(wrap_pyfunction!(check_threshold, m)?)?;
m.add_function(wrap_pyfunction!(check_cross, m)?)?;
+10 -183
View File
@@ -1,41 +1,12 @@
//! Performance attribution and trade analysis.
//!
//! Functions
//! ---------
//! - `trade_stats` — compute win rate, avg win/loss, hold time,
//! profit factor from a list of trade PnLs and hold durations.
//! - `monthly_contribution` — group bar returns by month index and sum, for
//! time-based performance attribution.
//! - `signal_attribution` — given signal labels per bar and bar returns,
//! compute the PnL contribution of each signal.
//! Performance attribution (thin PyO3 wrapper over ferro_ta_core::attribution).
use numpy::{IntoPyArray, PyArray1, PyReadonlyArray1};
use pyo3::exceptions::PyValueError;
use pyo3::prelude::*;
use std::collections::HashMap;
// ---------------------------------------------------------------------------
// trade_stats
// ---------------------------------------------------------------------------
use crate::validation;
/// Compute trade-level statistics from trade PnL and hold durations.
///
/// Parameters
/// ----------
/// pnl : 1-D float64 array — per-trade profit/loss (positive = win)
/// hold_bars : 1-D float64 array — hold duration in bars for each trade
/// (same length as *pnl*)
///
/// Returns
/// -------
/// tuple of 5 floats:
/// ``(win_rate, avg_win, avg_loss, profit_factor, avg_hold_bars)``
///
/// - **win_rate** : fraction of trades with PnL > 0
/// - **avg_win** : mean PnL of winning trades (or 0 if none)
/// - **avg_loss** : mean PnL of losing trades (negative; or 0 if none)
/// - **profit_factor** : gross profit / |gross loss| (inf if no losses)
/// - **avg_hold_bars** : mean hold duration across all trades
#[pyfunction]
pub fn trade_stats(
pnl: PyReadonlyArray1<'_, f64>,
@@ -47,68 +18,11 @@ pub fn trade_stats(
if n == 0 {
return Err(PyValueError::new_err("pnl must be non-empty"));
}
if n != h.len() {
return Err(PyValueError::new_err(
"pnl and hold_bars must have the same length",
));
}
let mut wins: Vec<f64> = Vec::new();
let mut losses: Vec<f64> = Vec::new();
for &v in p.iter() {
if v > 0.0 {
wins.push(v);
} else if v < 0.0 {
losses.push(v);
}
}
let win_rate = wins.len() as f64 / n as f64;
let avg_win = if wins.is_empty() {
0.0
} else {
wins.iter().sum::<f64>() / wins.len() as f64
};
let avg_loss = if losses.is_empty() {
0.0
} else {
losses.iter().sum::<f64>() / losses.len() as f64
};
let gross_profit: f64 = wins.iter().sum();
let gross_loss: f64 = losses.iter().map(|v| v.abs()).sum();
let profit_factor = if gross_loss == 0.0 {
f64::INFINITY
} else {
gross_profit / gross_loss
};
let avg_hold = h.iter().sum::<f64>() / n as f64;
Ok((win_rate, avg_win, avg_loss, profit_factor, avg_hold))
validation::validate_equal_length(&[(n, "pnl"), (h.len(), "hold_bars")])?;
Ok(ferro_ta_core::attribution::trade_stats(p, h))
}
// ---------------------------------------------------------------------------
// monthly_contribution
// ---------------------------------------------------------------------------
/// Group per-bar returns by month index and sum each month's contribution.
///
/// The ``month_index`` array assigns each bar to a month bucket (0-based
/// integer, e.g. 0 = January year 1, 1 = February year 1, …). The function
/// returns the **unique sorted month indices** and the corresponding
/// **total return** for each month.
///
/// Parameters
/// ----------
/// bar_returns : 1-D float64 array — per-bar strategy returns
/// month_index : 1-D int64 array — month bucket for each bar (same length)
///
/// Returns
/// -------
/// tuple ``(months, contributions)``:
/// - ``months`` : 1-D int64 array — sorted unique month indices
/// - ``contributions`` : 1-D float64 array — summed return per month
#[pyfunction]
#[allow(clippy::type_complexity)]
pub fn monthly_contribution<'py>(
@@ -119,48 +33,12 @@ pub fn monthly_contribution<'py>(
let ret = bar_returns.as_slice()?;
let mi = month_index.as_slice()?;
let n = ret.len();
if n != mi.len() {
return Err(PyValueError::new_err(
"bar_returns and month_index must have the same length",
));
}
// Accumulate contributions by month
let mut map: HashMap<i64, f64> = HashMap::new();
for i in 0..n {
if !ret[i].is_nan() {
*map.entry(mi[i]).or_insert(0.0) += ret[i];
}
}
// Sort by month index
let mut months: Vec<i64> = map.keys().copied().collect();
months.sort_unstable();
let contributions: Vec<f64> = months.iter().map(|m| map[m]).collect();
validation::validate_equal_length(&[(n, "bar_returns"), (mi.len(), "month_index")])?;
let (months, contributions) = ferro_ta_core::attribution::monthly_contribution(ret, mi);
Ok((months.into_pyarray(py), contributions.into_pyarray(py)))
}
// ---------------------------------------------------------------------------
// signal_attribution
// ---------------------------------------------------------------------------
/// Attribute per-bar returns to each signal label.
///
/// Each bar has a *signal_label* (integer) indicating which signal or rule
/// triggered the trade. ``-1`` means "no signal / flat". The function sums
/// bar returns per signal label.
///
/// Parameters
/// ----------
/// bar_returns : 1-D float64 array — per-bar strategy returns
/// signal_labels : 1-D int64 array — signal label per bar (same length)
///
/// Returns
/// -------
/// tuple ``(labels, contributions)``:
/// - ``labels`` : 1-D int64 array — sorted unique signal labels
/// - ``contributions`` : 1-D float64 array — summed return per label
#[pyfunction]
#[allow(clippy::type_complexity)]
pub fn signal_attribution<'py>(
@@ -171,33 +49,12 @@ pub fn signal_attribution<'py>(
let ret = bar_returns.as_slice()?;
let lbl = signal_labels.as_slice()?;
let n = ret.len();
if n != lbl.len() {
return Err(PyValueError::new_err(
"bar_returns and signal_labels must have the same length",
));
}
let mut map: HashMap<i64, f64> = HashMap::new();
for i in 0..n {
if !ret[i].is_nan() {
*map.entry(lbl[i]).or_insert(0.0) += ret[i];
}
}
let mut labels: Vec<i64> = map.keys().copied().collect();
labels.sort_unstable();
let contributions: Vec<f64> = labels.iter().map(|l| map[l]).collect();
validation::validate_equal_length(&[(n, "bar_returns"), (lbl.len(), "signal_labels")])?;
let (labels, contributions) = ferro_ta_core::attribution::signal_attribution(ret, lbl);
Ok((labels.into_pyarray(py), contributions.into_pyarray(py)))
}
// ---------------------------------------------------------------------------
// extract_trades
// ---------------------------------------------------------------------------
/// Extract trade-level pnl and hold durations from positions and strategy returns.
///
/// A trade is a maximal contiguous run of non-zero position values.
#[pyfunction]
#[allow(clippy::type_complexity)]
pub fn extract_trades<'py>(
@@ -208,41 +65,11 @@ pub fn extract_trades<'py>(
let pos = positions.as_slice()?;
let ret = strategy_returns.as_slice()?;
let n = pos.len();
if n != ret.len() {
return Err(PyValueError::new_err(
"positions and strategy_returns must have the same length",
));
}
let mut pnl = Vec::<f64>::new();
let mut hold = Vec::<f64>::new();
let mut i = 0usize;
while i < n {
if pos[i] == 0.0 {
i += 1;
continue;
}
let mut j = i + 1;
while j < n && pos[j] == pos[i] {
j += 1;
}
let mut trade_pnl = 0.0_f64;
for v in ret.iter().take(j).skip(i) {
trade_pnl += *v;
}
pnl.push(trade_pnl);
hold.push((j - i) as f64);
i = j;
}
validation::validate_equal_length(&[(n, "positions"), (ret.len(), "strategy_returns")])?;
let (pnl, hold) = ferro_ta_core::attribution::extract_trades(pos, ret);
Ok((pnl.into_pyarray(py), hold.into_pyarray(py)))
}
// ---------------------------------------------------------------------------
// Register
// ---------------------------------------------------------------------------
pub fn register(m: &Bound<'_, PyModule>) -> PyResult<()> {
m.add_function(wrap_pyfunction!(trade_stats, m)?)?;
m.add_function(wrap_pyfunction!(monthly_contribution, m)?)?;
+294
View File
@@ -0,0 +1,294 @@
//! PyO3 wrapper around `ferro_ta_core::commission::CommissionModel`.
//!
//! Exposes all fields as Python properties, provides static preset constructors,
//! and supports JSON persistence (save/load).
use ferro_ta_core::commission::CommissionModel as CoreModel;
use pyo3::exceptions::PyValueError;
use pyo3::prelude::*;
use std::fs;
/// Advanced commission and tax model for Indian and global markets.
///
/// All `_rate` fields are fractions (e.g. 0.001 = 0.1%).
/// Per-unit fields (`flat_per_order`, `per_lot`) are in base currency units (e.g. INR).
///
/// ## Example
/// ```python
/// from ferro_ta._ferro_ta import CommissionModel
///
/// # Use a built-in preset
/// m = CommissionModel.equity_delivery_india()
/// cost = m.total_cost(100_000.0, 1.0, True)
/// print(f"Buy cost: ₹{cost:.2f}")
///
/// # Save and reload
/// m.save("/tmp/my_commission.json")
/// m2 = CommissionModel.load("/tmp/my_commission.json")
/// ```
#[pyclass(module = "ferro_ta._ferro_ta", name = "CommissionModel")]
#[derive(Clone, Default)]
pub struct PyCommissionModel {
pub(crate) inner: CoreModel,
}
#[pymethods]
impl PyCommissionModel {
/// Create a zero-commission model (all fields = 0, lot_size = 1).
#[new]
pub fn new() -> Self {
Self::default()
}
// ---- Brokerage fields -----------------------------------------------
#[getter]
pub fn flat_per_order(&self) -> f64 {
self.inner.flat_per_order
}
#[setter]
pub fn set_flat_per_order(&mut self, v: f64) {
self.inner.flat_per_order = v;
}
#[getter]
pub fn rate_of_value(&self) -> f64 {
self.inner.rate_of_value
}
#[setter]
pub fn set_rate_of_value(&mut self, v: f64) {
self.inner.rate_of_value = v;
}
#[getter]
pub fn per_lot(&self) -> f64 {
self.inner.per_lot
}
#[setter]
pub fn set_per_lot(&mut self, v: f64) {
self.inner.per_lot = v;
}
#[getter]
pub fn max_brokerage(&self) -> f64 {
self.inner.max_brokerage
}
#[setter]
pub fn set_max_brokerage(&mut self, v: f64) {
self.inner.max_brokerage = v;
}
#[getter]
pub fn spread_bps(&self) -> f64 {
self.inner.spread_bps
}
#[setter]
pub fn set_spread_bps(&mut self, v: f64) {
self.inner.spread_bps = v;
}
// ---- STT fields -----------------------------------------------------
#[getter]
pub fn stt_rate(&self) -> f64 {
self.inner.stt_rate
}
#[setter]
pub fn set_stt_rate(&mut self, v: f64) {
self.inner.stt_rate = v;
}
#[getter]
pub fn stt_on_buy(&self) -> bool {
self.inner.stt_on_buy
}
#[setter]
pub fn set_stt_on_buy(&mut self, v: bool) {
self.inner.stt_on_buy = v;
}
#[getter]
pub fn stt_on_sell(&self) -> bool {
self.inner.stt_on_sell
}
#[setter]
pub fn set_stt_on_sell(&mut self, v: bool) {
self.inner.stt_on_sell = v;
}
// ---- Exchange / regulatory fields -----------------------------------
#[getter]
pub fn exchange_charges_rate(&self) -> f64 {
self.inner.exchange_charges_rate
}
#[setter]
pub fn set_exchange_charges_rate(&mut self, v: f64) {
self.inner.exchange_charges_rate = v;
}
#[getter]
pub fn regulatory_charges_rate(&self) -> f64 {
self.inner.regulatory_charges_rate
}
#[setter]
pub fn set_regulatory_charges_rate(&mut self, v: f64) {
self.inner.regulatory_charges_rate = v;
}
#[getter]
pub fn gst_rate(&self) -> f64 {
self.inner.gst_rate
}
#[setter]
pub fn set_gst_rate(&mut self, v: f64) {
self.inner.gst_rate = v;
}
#[getter]
pub fn stamp_duty_rate(&self) -> f64 {
self.inner.stamp_duty_rate
}
#[setter]
pub fn set_stamp_duty_rate(&mut self, v: f64) {
self.inner.stamp_duty_rate = v;
}
#[getter]
pub fn lot_size(&self) -> f64 {
self.inner.lot_size
}
#[setter]
pub fn set_lot_size(&mut self, v: f64) {
self.inner.lot_size = v;
}
#[getter]
pub fn short_borrow_rate_annual(&self) -> f64 {
self.inner.short_borrow_rate_annual
}
#[setter]
pub fn set_short_borrow_rate_annual(&mut self, v: f64) {
self.inner.short_borrow_rate_annual = v;
}
// ---- Compute --------------------------------------------------------
/// Total transaction cost in absolute currency units.
///
/// Args:
/// trade_value: price × quantity in base currency
/// num_lots: number of lots transacted
/// is_buy: True for buy (entry) leg, False for sell (exit) leg
pub fn total_cost(&self, trade_value: f64, num_lots: f64, is_buy: bool) -> f64 {
self.inner.total_cost(trade_value, num_lots, is_buy)
}
/// Cost as fraction of `initial_capital` (for normalised equity loops).
///
/// Returns 0.0 if `initial_capital` ≤ 0.
pub fn cost_fraction(
&self,
trade_value: f64,
num_lots: f64,
is_buy: bool,
initial_capital: f64,
) -> f64 {
self.inner
.cost_fraction(trade_value, num_lots, is_buy, initial_capital)
}
// ---- Presets (static constructors) ----------------------------------
/// Zero-commission model (all fields = 0).
#[staticmethod]
pub fn zero() -> Self {
Self {
inner: CoreModel::zero(),
}
}
/// Indian equity delivery preset (0.1% brokerage capped ₹20, STT both sides, full levies).
#[staticmethod]
pub fn equity_delivery_india() -> Self {
Self {
inner: CoreModel::equity_delivery_india(),
}
}
/// Indian equity intraday preset (0.03% brokerage capped ₹20, STT sell only, full levies).
#[staticmethod]
pub fn equity_intraday_india() -> Self {
Self {
inner: CoreModel::equity_intraday_india(),
}
}
/// Indian index futures preset (₹20 flat, STT sell only, lot_size=25).
#[staticmethod]
pub fn futures_india() -> Self {
Self {
inner: CoreModel::futures_india(),
}
}
/// Indian index options preset (₹20 flat, STT on premium sell side, lot_size=25).
#[staticmethod]
pub fn options_india() -> Self {
Self {
inner: CoreModel::options_india(),
}
}
/// Simple proportional model — `rate` fraction applied both ways, no taxes.
#[staticmethod]
pub fn proportional(rate: f64) -> Self {
Self {
inner: CoreModel::proportional(rate),
}
}
// ---- JSON persistence -----------------------------------------------
/// Serialize this model to a JSON string.
pub fn to_json(&self) -> PyResult<String> {
self.inner
.to_json()
.map_err(|e| PyValueError::new_err(e.to_string()))
}
/// Deserialize a `CommissionModel` from a JSON string.
#[staticmethod]
pub fn from_json(s: &str) -> PyResult<Self> {
CoreModel::from_json(s)
.map(|inner| Self { inner })
.map_err(|e| PyValueError::new_err(e.to_string()))
}
/// Save this model to a JSON file at `path`.
pub fn save(&self, path: &str) -> PyResult<()> {
let json = self.to_json()?;
fs::write(path, json).map_err(|e| PyValueError::new_err(e.to_string()))
}
/// Load a `CommissionModel` from a JSON file at `path`.
#[staticmethod]
pub fn load(path: &str) -> PyResult<Self> {
let s = fs::read_to_string(path).map_err(|e| PyValueError::new_err(e.to_string()))?;
Self::from_json(&s)
}
fn __repr__(&self) -> String {
format!(
"CommissionModel(flat={}, rate_pct={:.4}%, stt={:.4}%, lot_size={})",
self.inner.flat_per_order,
self.inner.rate_of_value * 100.0,
self.inner.stt_rate * 100.0,
self.inner.lot_size,
)
}
fn __eq__(&self, other: &Self) -> bool {
self.inner == other.inner
}
}
+133
View File
@@ -0,0 +1,133 @@
//! PyO3 wrapper around `ferro_ta_core::currency::Currency`.
use ferro_ta_core::currency::Currency as CoreCurrency;
use pyo3::exceptions::PyValueError;
use pyo3::prelude::*;
/// Immutable currency descriptor with formatting support.
///
/// ## Example
/// ```python
/// from ferro_ta._ferro_ta import Currency
///
/// inr = Currency.INR()
/// print(inr.format(123456.78)) # ₹1,23,456.78
///
/// usd = Currency.from_code("USD")
/// print(usd.format(1234567.89)) # $1,234,567.89
/// ```
#[pyclass(name = "Currency", module = "ferro_ta._ferro_ta", frozen)]
#[derive(Clone)]
pub struct PyCurrency {
pub(crate) inner: &'static CoreCurrency,
}
#[pymethods]
impl PyCurrency {
/// Format *amount* according to this currency's style.
pub fn format(&self, amount: f64) -> String {
self.inner.format(amount)
}
#[getter]
pub fn code(&self) -> &str {
self.inner.code
}
#[getter]
pub fn symbol(&self) -> &str {
self.inner.symbol
}
#[getter]
pub fn decimal_places(&self) -> u8 {
self.inner.decimal_places
}
#[getter]
pub fn lakh_grouping(&self) -> bool {
self.inner.lakh_grouping
}
// ---- Static constructors (presets) ----
#[staticmethod]
pub fn from_code(code: &str) -> PyResult<Self> {
CoreCurrency::from_code(code)
.map(|c| PyCurrency { inner: c })
.ok_or_else(|| {
PyValueError::new_err(format!(
"Unknown currency code '{code}'. Supported: INR, USD, EUR, GBP, JPY, USDT"
))
})
}
/// Indian Rupee.
#[staticmethod]
#[allow(non_snake_case)]
pub fn INR() -> Self {
PyCurrency {
inner: &CoreCurrency::INR,
}
}
/// US Dollar.
#[staticmethod]
#[allow(non_snake_case)]
pub fn USD() -> Self {
PyCurrency {
inner: &CoreCurrency::USD,
}
}
/// Euro.
#[staticmethod]
#[allow(non_snake_case)]
pub fn EUR() -> Self {
PyCurrency {
inner: &CoreCurrency::EUR,
}
}
/// British Pound.
#[staticmethod]
#[allow(non_snake_case)]
pub fn GBP() -> Self {
PyCurrency {
inner: &CoreCurrency::GBP,
}
}
/// Japanese Yen.
#[staticmethod]
#[allow(non_snake_case)]
pub fn JPY() -> Self {
PyCurrency {
inner: &CoreCurrency::JPY,
}
}
/// Tether USD.
#[staticmethod]
#[allow(non_snake_case)]
pub fn USDT() -> Self {
PyCurrency {
inner: &CoreCurrency::USDT,
}
}
fn __repr__(&self) -> String {
format!("Currency({:?})", self.inner.code)
}
fn __eq__(&self, other: &Self) -> bool {
self.inner.code == other.inner.code
}
fn __hash__(&self) -> u64 {
use std::hash::{Hash, Hasher};
let mut hasher = std::collections::hash_map::DefaultHasher::new();
self.inner.code.hash(&mut hasher);
hasher.finish()
}
}
+714 -149
View File
@@ -1,34 +1,126 @@
//! Rust-backed strategy signal generation and backtest core.
//!
//! These functions move the hot loops from Python into Rust while preserving
//! the public Python behavior.
//! Thin PyO3 wrappers delegating to `ferro_ta_core::backtest`.
use crate::validation;
use numpy::{IntoPyArray, PyArray1, PyReadonlyArray1};
pub mod commission;
pub mod currency;
use commission::PyCommissionModel;
use currency::PyCurrency;
use ferro_ta_core::backtest as core_bt;
use ndarray::Array2;
use numpy::{IntoPyArray, PyArray1, PyArray2, PyReadonlyArray1, PyReadonlyArray2};
use pyo3::exceptions::PyValueError;
use pyo3::prelude::*;
use pyo3::types::PyDict;
use rayon::prelude::*;
fn nan_to_num_with_numpy_defaults(v: f64) -> f64 {
if v.is_nan() {
0.0
} else if v.is_infinite() {
if v.is_sign_positive() {
f64::MAX
} else {
-f64::MAX
use crate::validation;
// ---------------------------------------------------------------------------
// BacktestConfig pyclass wrapping core struct
// ---------------------------------------------------------------------------
#[pyclass(name = "BacktestConfig")]
#[derive(Clone)]
pub struct BacktestConfig {
#[pyo3(get, set)]
pub fill_mode: String,
#[pyo3(get, set)]
pub stop_loss_pct: f64,
#[pyo3(get, set)]
pub take_profit_pct: f64,
#[pyo3(get, set)]
pub trailing_stop_pct: f64,
#[pyo3(get, set)]
pub slippage_bps: f64,
#[pyo3(get, set)]
pub initial_capital: f64,
#[pyo3(get, set)]
pub commission_per_trade: f64,
#[pyo3(get, set)]
pub max_hold_bars: usize,
#[pyo3(get, set)]
pub slippage_pct_range: f64,
#[pyo3(get, set)]
pub breakeven_pct: f64,
#[pyo3(get, set)]
pub periods_per_year: f64,
#[pyo3(get, set)]
pub margin_ratio: f64,
#[pyo3(get, set)]
pub margin_call_pct: f64,
#[pyo3(get, set)]
pub daily_loss_limit: f64,
#[pyo3(get, set)]
pub total_loss_limit: f64,
#[pyo3(get, set)]
pub commission: Option<PyCommissionModel>,
}
#[pymethods]
impl BacktestConfig {
#[new]
#[pyo3(signature = (
fill_mode = "market_open",
stop_loss_pct = 0.0,
take_profit_pct = 0.0,
trailing_stop_pct = 0.0,
slippage_bps = 0.0,
initial_capital = 100_000.0,
commission_per_trade = 0.0,
max_hold_bars = 0,
slippage_pct_range = 0.0,
breakeven_pct = 0.0,
periods_per_year = 252.0,
margin_ratio = 0.0,
margin_call_pct = 0.5,
daily_loss_limit = 0.0,
total_loss_limit = 0.0,
commission = None,
))]
#[allow(clippy::too_many_arguments)]
pub fn new(
fill_mode: &str,
stop_loss_pct: f64,
take_profit_pct: f64,
trailing_stop_pct: f64,
slippage_bps: f64,
initial_capital: f64,
commission_per_trade: f64,
max_hold_bars: usize,
slippage_pct_range: f64,
breakeven_pct: f64,
periods_per_year: f64,
margin_ratio: f64,
margin_call_pct: f64,
daily_loss_limit: f64,
total_loss_limit: f64,
commission: Option<PyCommissionModel>,
) -> Self {
BacktestConfig {
fill_mode: fill_mode.to_string(),
stop_loss_pct,
take_profit_pct,
trailing_stop_pct,
slippage_bps,
initial_capital,
commission_per_trade,
max_hold_bars,
slippage_pct_range,
breakeven_pct,
periods_per_year,
margin_ratio,
margin_call_pct,
daily_loss_limit,
total_loss_limit,
commission,
}
} else {
v
}
}
// ---------------------------------------------------------------------------
// Strategy signal helpers
// Signal generators
// ---------------------------------------------------------------------------
/// RSI threshold strategy:
/// +1 when RSI <= oversold, -1 when RSI >= overbought, 0 otherwise.
/// Warm-up bars are NaN.
#[pyfunction]
#[pyo3(signature = (close, timeperiod = 14, oversold = 30.0, overbought = 70.0))]
pub fn rsi_threshold_signals<'py>(
@@ -40,26 +132,10 @@ pub fn rsi_threshold_signals<'py>(
) -> PyResult<Bound<'py, PyArray1<f64>>> {
validation::validate_timeperiod(timeperiod, "timeperiod", 1)?;
let prices = close.as_slice()?;
let rsi = ferro_ta_core::momentum::rsi(prices, timeperiod);
let out: Vec<f64> = rsi
.iter()
.map(|&v| {
if v.is_nan() {
f64::NAN
} else if v <= oversold {
1.0
} else if v >= overbought {
-1.0
} else {
0.0
}
})
.collect();
let out = core_bt::rsi_threshold_signals(prices, timeperiod, oversold, overbought);
Ok(out.into_pyarray(py))
}
/// SMA crossover strategy:
/// +1 when fast SMA > slow SMA, -1 otherwise. Warm-up bars are NaN.
#[pyfunction]
#[pyo3(signature = (close, fast = 10, slow = 30))]
pub fn sma_crossover_signals<'py>(
@@ -70,32 +146,12 @@ pub fn sma_crossover_signals<'py>(
) -> PyResult<Bound<'py, PyArray1<f64>>> {
validation::validate_timeperiod(fast, "fast", 1)?;
validation::validate_timeperiod(slow, "slow", 1)?;
if fast >= slow {
return Err(PyValueError::new_err(format!(
"fast ({fast}) must be less than slow ({slow})"
)));
}
let prices = close.as_slice()?;
let sma_fast = ferro_ta_core::overlap::sma(prices, fast);
let sma_slow = ferro_ta_core::overlap::sma(prices, slow);
let out: Vec<f64> = sma_fast
.iter()
.zip(sma_slow.iter())
.map(|(&f, &s)| {
if f.is_nan() || s.is_nan() {
f64::NAN
} else if f > s {
1.0
} else {
-1.0
}
})
.collect();
let out = core_bt::sma_crossover_signals(prices, fast, slow)
.map_err(|e| PyValueError::new_err(e))?;
Ok(out.into_pyarray(py))
}
/// MACD crossover strategy:
/// +1 when MACD line > signal line, -1 otherwise. Warm-up bars are NaN.
#[pyfunction]
#[pyo3(signature = (close, fastperiod = 12, slowperiod = 26, signalperiod = 9))]
pub fn macd_crossover_signals<'py>(
@@ -108,137 +164,646 @@ pub fn macd_crossover_signals<'py>(
validation::validate_timeperiod(fastperiod, "fastperiod", 1)?;
validation::validate_timeperiod(slowperiod, "slowperiod", 1)?;
validation::validate_timeperiod(signalperiod, "signalperiod", 1)?;
if fastperiod >= slowperiod {
return Err(PyValueError::new_err(format!(
"fastperiod ({fastperiod}) must be less than slowperiod ({slowperiod})"
)));
}
let prices = close.as_slice()?;
let (macd_line, signal_line, _) =
ferro_ta_core::overlap::macd(prices, fastperiod, slowperiod, signalperiod);
let out: Vec<f64> = macd_line
.iter()
.zip(signal_line.iter())
.map(|(&m, &s)| {
if m.is_nan() || s.is_nan() {
f64::NAN
} else if m > s {
1.0
} else {
-1.0
}
})
.collect();
let out = core_bt::macd_crossover_signals(prices, fastperiod, slowperiod, signalperiod)
.map_err(|e| PyValueError::new_err(e))?;
Ok(out.into_pyarray(py))
}
// ---------------------------------------------------------------------------
// Backtest core
// Backtest core (close-only)
// ---------------------------------------------------------------------------
/// Backtest core loop over close prices and strategy signals.
///
/// Returns `(positions, bar_returns, strategy_returns, equity)`.
#[pyfunction]
#[pyo3(signature = (close, signals, commission_per_trade = 0.0, slippage_bps = 0.0))]
#[pyo3(signature = (
close, signals,
commission = None,
slippage_bps = 0.0,
initial_capital = 100_000.0,
commission_per_trade = 0.0,
))]
#[allow(clippy::type_complexity)]
pub fn backtest_core<'py>(
py: Python<'py>,
close: PyReadonlyArray1<'py, f64>,
signals: PyReadonlyArray1<'py, f64>,
commission_per_trade: f64,
commission: Option<PyRef<'py, PyCommissionModel>>,
slippage_bps: f64,
initial_capital: f64,
commission_per_trade: f64,
) -> PyResult<(
Bound<'py, PyArray1<f64>>,
Bound<'py, PyArray1<f64>>,
Bound<'py, PyArray1<f64>>,
Bound<'py, PyArray1<f64>>,
)> {
let c = close.as_slice()?;
let s = signals.as_slice()?;
validation::validate_equal_length(&[(c.len(), "close"), (s.len(), "signals")])?;
let cm = commission.as_ref().map(|c| &c.inner);
let result = core_bt::backtest_core(c, s, cm, slippage_bps, initial_capital, commission_per_trade)
.map_err(|e| PyValueError::new_err(e))?;
Ok((
result.positions.into_pyarray(py),
result.bar_returns.into_pyarray(py),
result.strategy_returns.into_pyarray(py),
result.equity.into_pyarray(py),
))
}
// ---------------------------------------------------------------------------
// OHLCV backtest
// ---------------------------------------------------------------------------
#[pyfunction]
#[pyo3(signature = (
open, high, low, close, signals,
fill_mode = "market_open",
stop_loss_pct = 0.0,
take_profit_pct = 0.0,
trailing_stop_pct = 0.0,
commission = None,
slippage_bps = 0.0,
initial_capital = 100_000.0,
commission_per_trade = 0.0,
limit_prices = None,
max_hold_bars = 0,
slippage_pct_range = 0.0,
breakeven_pct = 0.0,
periods_per_year = 252.0,
margin_ratio = 0.0,
margin_call_pct = 0.5,
daily_loss_limit = 0.0,
total_loss_limit = 0.0,
))]
#[allow(clippy::too_many_arguments, clippy::type_complexity)]
pub fn backtest_ohlcv_core<'py>(
py: Python<'py>,
open: PyReadonlyArray1<'py, f64>,
high: PyReadonlyArray1<'py, f64>,
low: PyReadonlyArray1<'py, f64>,
close: PyReadonlyArray1<'py, f64>,
signals: PyReadonlyArray1<'py, f64>,
fill_mode: &str,
stop_loss_pct: f64,
take_profit_pct: f64,
trailing_stop_pct: f64,
commission: Option<PyRef<'py, PyCommissionModel>>,
slippage_bps: f64,
initial_capital: f64,
commission_per_trade: f64,
limit_prices: Option<PyReadonlyArray1<'py, f64>>,
max_hold_bars: usize,
slippage_pct_range: f64,
breakeven_pct: f64,
periods_per_year: f64,
margin_ratio: f64,
margin_call_pct: f64,
daily_loss_limit: f64,
total_loss_limit: f64,
) -> PyResult<(
Bound<'py, PyArray1<f64>>,
Bound<'py, PyArray1<f64>>,
Bound<'py, PyArray1<f64>>,
Bound<'py, PyArray1<f64>>,
Bound<'py, PyArray1<f64>>,
)> {
let o = open.as_slice()?;
let h = high.as_slice()?;
let l = low.as_slice()?;
let c = close.as_slice()?;
let s = signals.as_slice()?;
let n = c.len();
validation::validate_equal_length(&[(n, "close"), (s.len(), "signals")])?;
let mut positions = vec![0.0_f64; n];
if n > 1 {
for i in 1..n {
positions[i] = nan_to_num_with_numpy_defaults(s[i - 1]);
}
validation::validate_equal_length(&[
(n, "close"),
(o.len(), "open"),
(h.len(), "high"),
(l.len(), "low"),
(s.len(), "signals"),
])?;
let config = core_bt::BacktestConfig {
fill_mode: fill_mode.to_string(),
stop_loss_pct,
take_profit_pct,
trailing_stop_pct,
slippage_bps,
initial_capital,
commission_per_trade,
max_hold_bars,
slippage_pct_range,
breakeven_pct,
periods_per_year,
margin_ratio,
margin_call_pct,
daily_loss_limit,
total_loss_limit,
commission: commission.as_ref().map(|c| c.inner.clone()),
};
let lp_opt: Option<&[f64]> = limit_prices.as_ref().and_then(|lp| lp.as_slice().ok());
let result = core_bt::backtest_ohlcv_core(o, h, l, c, s, &config, lp_opt)
.map_err(|e| PyValueError::new_err(e))?;
Ok((
result.positions.into_pyarray(py),
result.fill_prices.into_pyarray(py),
result.bar_returns.into_pyarray(py),
result.strategy_returns.into_pyarray(py),
result.equity.into_pyarray(py),
))
}
// ---------------------------------------------------------------------------
// Performance metrics
// ---------------------------------------------------------------------------
#[pyfunction]
#[pyo3(signature = (strategy_returns, equity, periods_per_year = 252.0, risk_free_rate = 0.0, benchmark_returns = None))]
pub fn compute_performance_metrics<'py>(
py: Python<'py>,
strategy_returns: PyReadonlyArray1<'py, f64>,
equity: PyReadonlyArray1<'py, f64>,
periods_per_year: f64,
risk_free_rate: f64,
benchmark_returns: Option<PyReadonlyArray1<'py, f64>>,
) -> PyResult<Bound<'py, PyDict>> {
let r = strategy_returns.as_slice()?;
let eq = equity.as_slice()?;
let br = benchmark_returns.as_ref().and_then(|b| b.as_slice().ok());
let metrics = core_bt::compute_performance_metrics(r, eq, periods_per_year, risk_free_rate, br)
.map_err(|e| PyValueError::new_err(e))?;
let dict = PyDict::new(py);
dict.set_item("total_return", metrics.total_return)?;
dict.set_item("cagr", metrics.cagr)?;
dict.set_item("annualized_vol", metrics.annualized_vol)?;
dict.set_item("sharpe", metrics.sharpe)?;
dict.set_item("sortino", metrics.sortino)?;
dict.set_item("calmar", metrics.calmar)?;
dict.set_item("max_drawdown", metrics.max_drawdown)?;
dict.set_item("avg_drawdown", metrics.avg_drawdown)?;
dict.set_item("max_drawdown_duration_bars", metrics.max_drawdown_duration_bars as i64)?;
dict.set_item("avg_drawdown_duration_bars", metrics.avg_drawdown_duration_bars)?;
dict.set_item("ulcer_index", metrics.ulcer_index)?;
dict.set_item("omega_ratio", metrics.omega_ratio)?;
dict.set_item("win_rate", metrics.win_rate)?;
dict.set_item("profit_factor", metrics.profit_factor)?;
dict.set_item("r_expectancy", metrics.r_expectancy)?;
dict.set_item("avg_win", metrics.avg_win)?;
dict.set_item("avg_loss", metrics.avg_loss)?;
dict.set_item("tail_ratio", metrics.tail_ratio)?;
dict.set_item("skewness", metrics.skewness)?;
dict.set_item("kurtosis", metrics.kurtosis)?;
dict.set_item("best_bar", metrics.best_bar)?;
dict.set_item("worst_bar", metrics.worst_bar)?;
dict.set_item("n_trades", metrics.n_trades as i64)?;
dict.set_item("n_position_changes", metrics.n_position_changes as i64)?;
if let Some(v) = metrics.benchmark_total_return {
dict.set_item("benchmark_total_return", v)?;
}
if let Some(v) = metrics.benchmark_cagr {
dict.set_item("benchmark_cagr", v)?;
}
if let Some(v) = metrics.benchmark_annualized_vol {
dict.set_item("benchmark_annualized_vol", v)?;
}
if let Some(v) = metrics.benchmark_sharpe {
dict.set_item("benchmark_sharpe", v)?;
}
if let Some(v) = metrics.alpha {
dict.set_item("alpha", v)?;
}
if let Some(v) = metrics.beta {
dict.set_item("beta", v)?;
}
if let Some(v) = metrics.tracking_error {
dict.set_item("tracking_error", v)?;
}
if let Some(v) = metrics.information_ratio {
dict.set_item("information_ratio", v)?;
}
let mut bar_returns = vec![0.0_f64; n];
for i in 1..n {
bar_returns[i] = (c[i] - c[i - 1]) / c[i - 1];
}
Ok(dict)
}
let mut strategy_returns = vec![0.0_f64; n];
for i in 0..n {
strategy_returns[i] = positions[i] * bar_returns[i];
}
// ---------------------------------------------------------------------------
// Trade extraction
// ---------------------------------------------------------------------------
let mut position_changed = vec![false; n];
for i in 1..n {
position_changed[i] = positions[i] != positions[i - 1];
}
#[pyfunction]
#[allow(clippy::type_complexity)]
pub fn extract_trades_ohlcv<'py>(
py: Python<'py>,
positions: PyReadonlyArray1<'py, f64>,
fill_prices: PyReadonlyArray1<'py, f64>,
high: PyReadonlyArray1<'py, f64>,
low: PyReadonlyArray1<'py, f64>,
) -> PyResult<(
Bound<'py, PyArray1<i64>>,
Bound<'py, PyArray1<i64>>,
Bound<'py, PyArray1<f64>>,
Bound<'py, PyArray1<f64>>,
Bound<'py, PyArray1<f64>>,
Bound<'py, PyArray1<f64>>,
Bound<'py, PyArray1<i64>>,
Bound<'py, PyArray1<f64>>,
Bound<'py, PyArray1<f64>>,
)> {
let pos = positions.as_slice()?;
let fp = fill_prices.as_slice()?;
let h = high.as_slice()?;
let l = low.as_slice()?;
if slippage_bps > 0.0 {
let slip = slippage_bps / 10_000.0;
for i in 0..n {
if position_changed[i] {
strategy_returns[i] -= slip;
}
}
}
validation::validate_equal_length(&[
(pos.len(), "positions"),
(fp.len(), "fill_prices"),
(h.len(), "high"),
(l.len(), "low"),
])?;
let mut equity = vec![1.0_f64; n];
if n > 0 {
if commission_per_trade <= 0.0 {
let mut gross = 1.0_f64;
for i in 0..n {
gross *= 1.0 + strategy_returns[i];
equity[i] = gross;
}
} else {
let mut gross_equity = vec![1.0_f64; n];
let mut gross = 1.0_f64;
for i in 0..n {
gross *= 1.0 + strategy_returns[i];
gross_equity[i] = gross;
}
let trades = core_bt::extract_trades_ohlcv(pos, fp, h, l)
.map_err(|e| PyValueError::new_err(e))?;
if gross_equity.contains(&0.0) {
equity[0] = 1.0;
for i in 1..n {
equity[i] = equity[i - 1] * (1.0 + strategy_returns[i]);
if position_changed[i] {
equity[i] -= commission_per_trade;
}
}
} else {
let mut discounted_commissions = 0.0_f64;
for i in 0..n {
if position_changed[i] {
discounted_commissions += commission_per_trade / gross_equity[i];
}
equity[i] = gross_equity[i] * (1.0 - discounted_commissions);
}
}
}
let mut entry_bars: Vec<i64> = Vec::with_capacity(trades.len());
let mut exit_bars: Vec<i64> = Vec::with_capacity(trades.len());
let mut directions: Vec<f64> = Vec::with_capacity(trades.len());
let mut entry_prices: Vec<f64> = Vec::with_capacity(trades.len());
let mut exit_prices: Vec<f64> = Vec::with_capacity(trades.len());
let mut pnl_pcts: Vec<f64> = Vec::with_capacity(trades.len());
let mut duration_bars_vec: Vec<i64> = Vec::with_capacity(trades.len());
let mut maes: Vec<f64> = Vec::with_capacity(trades.len());
let mut mfes: Vec<f64> = Vec::with_capacity(trades.len());
for t in &trades {
entry_bars.push(t.entry_bar);
exit_bars.push(t.exit_bar);
directions.push(t.direction);
entry_prices.push(t.entry_price);
exit_prices.push(t.exit_price);
pnl_pcts.push(t.pnl_pct);
duration_bars_vec.push(t.duration_bars);
maes.push(t.mae);
mfes.push(t.mfe);
}
Ok((
positions.into_pyarray(py),
bar_returns.into_pyarray(py),
strategy_returns.into_pyarray(py),
equity.into_pyarray(py),
entry_bars.into_pyarray(py),
exit_bars.into_pyarray(py),
directions.into_pyarray(py),
entry_prices.into_pyarray(py),
exit_prices.into_pyarray(py),
pnl_pcts.into_pyarray(py),
duration_bars_vec.into_pyarray(py),
maes.into_pyarray(py),
mfes.into_pyarray(py),
))
}
// ---------------------------------------------------------------------------
// Multi-asset backtest
// ---------------------------------------------------------------------------
#[pyfunction]
#[pyo3(signature = (
close_2d, weights_2d,
commission_per_trade = 0.0,
slippage_bps = 0.0,
parallel = true,
max_asset_weight = 1.0,
max_gross_exposure = 0.0,
max_net_exposure = 0.0,
))]
#[allow(clippy::too_many_arguments, clippy::type_complexity)]
pub fn backtest_multi_asset_core<'py>(
py: Python<'py>,
close_2d: PyReadonlyArray2<'py, f64>,
weights_2d: PyReadonlyArray2<'py, f64>,
commission_per_trade: f64,
slippage_bps: f64,
parallel: bool,
max_asset_weight: f64,
max_gross_exposure: f64,
max_net_exposure: f64,
) -> PyResult<(
Bound<'py, PyArray2<f64>>,
Bound<'py, PyArray1<f64>>,
Bound<'py, PyArray1<f64>>,
)> {
let c_arr = close_2d.as_array();
let w_arr = weights_2d.as_array();
let (n_bars, n_assets) = c_arr.dim();
if w_arr.dim() != (n_bars, n_assets) {
return Err(PyValueError::new_err(format!(
"weights_2d shape {:?} must match close_2d shape {:?}",
w_arr.dim(),
c_arr.dim()
)));
}
// Transpose to (n_assets, n_bars) for the core function
let mut close_cm: Vec<Vec<f64>> = vec![vec![0.0; n_bars]; n_assets];
let mut weights_cm: Vec<Vec<f64>> = vec![vec![0.0; n_bars]; n_assets];
for j in 0..n_assets {
for i in 0..n_bars {
close_cm[j][i] = c_arr[[i, j]];
weights_cm[j][i] = w_arr[[i, j]];
}
}
// For parallel execution, use rayon directly on the core's single_asset_backtest.
// Apply portfolio constraints first via the core function's logic.
// Apply constraints
if max_asset_weight != 1.0 || max_gross_exposure > 0.0 || max_net_exposure > 0.0 {
for i in 0..n_bars {
if max_asset_weight < f64::INFINITY && max_asset_weight > 0.0 {
for j in 0..n_assets {
let w = weights_cm[j][i];
if w.abs() > max_asset_weight {
weights_cm[j][i] = w.signum() * max_asset_weight;
}
}
}
if max_gross_exposure > 0.0 {
let gross: f64 = (0..n_assets).map(|j| weights_cm[j][i].abs()).sum();
if gross > max_gross_exposure {
let scale = max_gross_exposure / gross;
for j in 0..n_assets {
weights_cm[j][i] *= scale;
}
}
}
if max_net_exposure > 0.0 {
let net: f64 = (0..n_assets).map(|j| weights_cm[j][i]).sum();
if net.abs() > max_net_exposure {
let excess = net - net.signum() * max_net_exposure;
let adj_per_asset = excess / n_assets as f64;
for j in 0..n_assets {
weights_cm[j][i] -= adj_per_asset;
}
}
}
}
}
// Run per-asset backtests (parallel or serial)
let asset_strategy_returns: Vec<Vec<f64>> = py.allow_threads(|| {
let run_asset = |j: usize| -> Vec<f64> {
let (_, strat_rets, _) = core_bt::single_asset_backtest(
&close_cm[j],
&weights_cm[j],
commission_per_trade,
slippage_bps,
);
strat_rets
};
if parallel {
(0..n_assets).into_par_iter().map(run_asset).collect()
} else {
(0..n_assets).map(run_asset).collect()
}
});
// Assemble asset_returns 2D array (n_bars, n_assets)
let mut asset_ret_arr = Array2::<f64>::zeros((n_bars, n_assets));
for j in 0..n_assets {
for i in 0..n_bars {
asset_ret_arr[[i, j]] = asset_strategy_returns[j][i];
}
}
// Portfolio returns
let mut portfolio_returns = vec![0.0_f64; n_bars];
for i in 0..n_bars {
let mut s = 0.0_f64;
for j in 0..n_assets {
s += asset_ret_arr[[i, j]];
}
portfolio_returns[i] = s;
}
// Portfolio equity
let mut portfolio_equity = vec![1.0_f64; n_bars];
let mut cum = 1.0_f64;
for i in 0..n_bars {
cum *= 1.0 + portfolio_returns[i];
portfolio_equity[i] = cum;
}
Ok((
asset_ret_arr.into_pyarray(py),
portfolio_returns.into_pyarray(py),
portfolio_equity.into_pyarray(py),
))
}
// ---------------------------------------------------------------------------
// Monte Carlo bootstrap
// ---------------------------------------------------------------------------
#[pyfunction]
#[pyo3(signature = (strategy_returns, n_sims = 1000, seed = 42, block_size = 1))]
pub fn monte_carlo_bootstrap<'py>(
py: Python<'py>,
strategy_returns: PyReadonlyArray1<'py, f64>,
n_sims: usize,
seed: u64,
block_size: usize,
) -> PyResult<Bound<'py, PyArray2<f64>>> {
let r = strategy_returns.as_slice()?;
let n = r.len();
// Use rayon for parallel Monte Carlo (preserving the original parallel behavior)
if n < 2 {
return Err(PyValueError::new_err(
"strategy_returns must have at least 2 elements",
));
}
if n_sims == 0 {
return Err(PyValueError::new_err("n_sims must be >= 1"));
}
let bsize = block_size.max(1).min(n);
let mut result = Array2::<f64>::zeros((n_sims, n));
py.allow_threads(|| {
result
.as_slice_mut()
.unwrap()
.par_chunks_mut(n)
.enumerate()
.for_each(|(sim_idx, row)| {
let mut state = seed
.wrapping_mul(6_364_136_223_846_793_005_u64)
.wrapping_add((sim_idx as u64).wrapping_mul(2_862_933_555_777_941_757_u64));
core_bt::lcg_next(&mut state);
core_bt::lcg_next(&mut state);
if bsize == 1 {
for dst in row.iter_mut() {
*dst = r[core_bt::lcg_index(&mut state, n)];
}
} else {
let mut filled = 0_usize;
while filled < n {
let start = core_bt::lcg_index(&mut state, n);
let take = bsize.min(n - filled);
for k in 0..take {
row[filled + k] = r[(start + k) % n];
}
filled += take;
}
}
let mut cum = 1.0_f64;
for elem in row.iter_mut().take(n) {
cum *= 1.0 + *elem;
*elem = cum;
}
});
});
Ok(result.into_pyarray(py))
}
// ---------------------------------------------------------------------------
// Walk-forward indices
// ---------------------------------------------------------------------------
#[pyfunction]
#[pyo3(signature = (n_bars, train_bars, test_bars, anchored = false, step_bars = 0))]
pub fn walk_forward_indices<'py>(
py: Python<'py>,
n_bars: usize,
train_bars: usize,
test_bars: usize,
anchored: bool,
step_bars: usize,
) -> PyResult<Bound<'py, PyArray2<i64>>> {
let folds = core_bt::walk_forward_indices(n_bars, train_bars, test_bars, anchored, step_bars)
.map_err(|e| PyValueError::new_err(e))?;
let n_folds = folds.len();
let mut arr = Array2::<i64>::zeros((n_folds, 4));
for (i, fold) in folds.iter().enumerate() {
for j in 0..4 {
arr[[i, j]] = fold[j];
}
}
Ok(arr.into_pyarray(py))
}
// ---------------------------------------------------------------------------
// Kelly criterion
// ---------------------------------------------------------------------------
#[pyfunction]
pub fn kelly_fraction(win_rate: f64, avg_win: f64, avg_loss: f64) -> PyResult<f64> {
core_bt::kelly_fraction(win_rate, avg_win, avg_loss).map_err(|e| PyValueError::new_err(e))
}
#[pyfunction]
pub fn half_kelly_fraction(win_rate: f64, avg_win: f64, avg_loss: f64) -> PyResult<f64> {
core_bt::half_kelly_fraction(win_rate, avg_win, avg_loss).map_err(|e| PyValueError::new_err(e))
}
// ---------------------------------------------------------------------------
// StreamingBacktest
// ---------------------------------------------------------------------------
#[pyclass(name = "StreamingBacktest")]
pub struct StreamingBacktest {
inner: core_bt::StreamingBacktest,
}
#[pymethods]
impl StreamingBacktest {
#[new]
#[pyo3(signature = (commission_per_trade=0.0, slippage_bps=0.0))]
pub fn new(commission_per_trade: f64, slippage_bps: f64) -> Self {
StreamingBacktest {
inner: core_bt::StreamingBacktest::new(commission_per_trade, slippage_bps),
}
}
pub fn on_bar<'py>(
&mut self,
py: Python<'py>,
close: f64,
signal: f64,
) -> PyResult<Bound<'py, PyDict>> {
let result = self.inner.on_bar(close, signal);
let d = PyDict::new(py);
d.set_item("position", result.position)?;
d.set_item("bar_return", result.bar_return)?;
d.set_item("equity", result.equity)?;
d.set_item("n_trades", result.n_trades)?;
Ok(d)
}
#[getter]
pub fn equity(&self) -> f64 {
self.inner.equity
}
#[getter]
pub fn position(&self) -> f64 {
self.inner.position
}
#[getter]
pub fn n_trades(&self) -> usize {
self.inner.n_trades
}
pub fn summary<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyDict>> {
let s = self.inner.summary();
let d = PyDict::new(py);
d.set_item("equity", s.equity)?;
d.set_item("n_trades", s.n_trades)?;
d.set_item("total_commission", s.total_commission)?;
d.set_item("win_rate", s.win_rate)?;
d.set_item("avg_win", s.avg_win)?;
d.set_item("avg_loss", s.avg_loss)?;
d.set_item("kelly_fraction", s.kelly_fraction)?;
Ok(d)
}
pub fn reset(&mut self) {
self.inner.reset();
}
}
// ---------------------------------------------------------------------------
// Register
// ---------------------------------------------------------------------------
pub fn register(m: &Bound<'_, PyModule>) -> PyResult<()> {
m.add_function(wrap_pyfunction!(rsi_threshold_signals, m)?)?;
m.add_function(wrap_pyfunction!(sma_crossover_signals, m)?)?;
m.add_function(wrap_pyfunction!(macd_crossover_signals, m)?)?;
m.add_function(wrap_pyfunction!(backtest_core, m)?)?;
m.add_function(wrap_pyfunction!(backtest_ohlcv_core, m)?)?;
m.add_function(wrap_pyfunction!(compute_performance_metrics, m)?)?;
m.add_function(wrap_pyfunction!(extract_trades_ohlcv, m)?)?;
m.add_function(wrap_pyfunction!(backtest_multi_asset_core, m)?)?;
m.add_function(wrap_pyfunction!(monte_carlo_bootstrap, m)?)?;
m.add_function(wrap_pyfunction!(walk_forward_indices, m)?)?;
m.add_function(wrap_pyfunction!(kelly_fraction, m)?)?;
m.add_function(wrap_pyfunction!(half_kelly_fraction, m)?)?;
m.add_class::<BacktestConfig>()?;
m.add_class::<StreamingBacktest>()?;
m.add_class::<PyCommissionModel>()?;
m.add_class::<PyCurrency>()?;
Ok(())
}
+159 -396
View File
@@ -8,19 +8,57 @@
//! [Rayon](https://docs.rs/rayon) after releasing the GIL. For small inputs
//! the sequential path (`parallel = false`) may be faster due to thread-pool
//! overhead.
//!
//! All indicator logic lives in `ferro_ta_core::batch`. This module is a thin
//! PyO3 wrapper that converts numpy ↔ Rust types and optionally adds Rayon
//! parallelism.
use ndarray::{Array2, ArrayView2};
use ndarray::Array2;
use numpy::{IntoPyArray, PyArray1, PyArray2, PyReadonlyArray1, PyReadonlyArray2};
use pyo3::exceptions::PyValueError;
use pyo3::prelude::*;
use rayon::prelude::*;
use ta::indicators::{Maximum, Minimum};
use ta::Next;
fn transpose_to_series_major(data: ArrayView2<'_, f64>) -> Array2<f64> {
let (n_samples, n_series) = data.dim();
Array2::from_shape_vec((n_series, n_samples), data.t().iter().copied().collect())
.expect("shape matches transposed data")
// ---------------------------------------------------------------------------
// numpy ↔ Vec<Vec<f64>> helpers
// ---------------------------------------------------------------------------
/// Convert a numpy (n_samples, n_series) array into `Vec<Vec<f64>>` where
/// `result[j]` is column j (one time-series of length n_samples).
fn numpy2d_to_columns(arr: &ndarray::ArrayView2<'_, f64>) -> Vec<Vec<f64>> {
let (_n_samples, n_series) = arr.dim();
(0..n_series)
.map(|j| arr.column(j).to_vec())
.collect()
}
/// Convert `Vec<Vec<f64>>` (columns) back into a numpy (n_samples, n_series) array.
fn columns_to_numpy2d<'py>(
py: Python<'py>,
n_samples: usize,
columns: Vec<Vec<f64>>,
) -> Bound<'py, PyArray2<f64>> {
let n_series = columns.len();
let mut result = Array2::<f64>::from_elem((n_samples, n_series), f64::NAN);
for (j, col) in columns.into_iter().enumerate() {
for (i, val) in col.into_iter().enumerate() {
result[[i, j]] = val;
}
}
result.into_pyarray(py)
}
/// Convert a pair of column-vectors into a pair of numpy 2-D arrays.
fn column_pair_to_numpy2d<'py>(
py: Python<'py>,
n_samples: usize,
cols_a: Vec<Vec<f64>>,
cols_b: Vec<Vec<f64>>,
) -> (Bound<'py, PyArray2<f64>>, Bound<'py, PyArray2<f64>>) {
(
columns_to_numpy2d(py, n_samples, cols_a),
columns_to_numpy2d(py, n_samples, cols_b),
)
}
fn validate_same_shape(
@@ -38,245 +76,41 @@ fn validate_same_shape(
}
}
fn finish_single_output<'py>(
py: Python<'py>,
n_samples: usize,
n_series: usize,
col_results: Vec<Vec<f64>>,
) -> Bound<'py, PyArray2<f64>> {
let mut result = Array2::<f64>::from_elem((n_samples, n_series), f64::NAN);
for (series_idx, values) in col_results.into_iter().enumerate() {
debug_assert_eq!(values.len(), n_samples);
for (sample_idx, value) in values.into_iter().enumerate() {
result[[sample_idx, series_idx]] = value;
}
}
result.into_pyarray(py)
fn map_core_err(err: String) -> PyErr {
PyValueError::new_err(err)
}
fn finish_pair_output<'py>(
py: Python<'py>,
n_samples: usize,
n_series: usize,
col_results: Vec<(Vec<f64>, Vec<f64>)>,
) -> (Bound<'py, PyArray2<f64>>, Bound<'py, PyArray2<f64>>) {
let mut result_k = Array2::<f64>::from_elem((n_samples, n_series), f64::NAN);
let mut result_d = Array2::<f64>::from_elem((n_samples, n_series), f64::NAN);
// ---------------------------------------------------------------------------
// Parallel-aware unary batch helper
// ---------------------------------------------------------------------------
for (series_idx, (k_values, d_values)) in col_results.into_iter().enumerate() {
debug_assert_eq!(k_values.len(), n_samples);
debug_assert_eq!(d_values.len(), n_samples);
for (sample_idx, value) in k_values.into_iter().enumerate() {
result_k[[sample_idx, series_idx]] = value;
}
for (sample_idx, value) in d_values.into_iter().enumerate() {
result_d[[sample_idx, series_idx]] = value;
}
}
(result_k.into_pyarray(py), result_d.into_pyarray(py))
}
fn run_unary_batch<'py, F>(
/// Run a unary batch function. When `parallel` is true, split column extraction
/// across Rayon threads and process in parallel; otherwise delegate sequentially
/// to `ferro_ta_core::batch`.
fn run_unary_batch_par<'py, F>(
py: Python<'py>,
data: PyReadonlyArray2<'py, f64>,
parallel: bool,
process_col: F,
) -> Bound<'py, PyArray2<f64>>
per_col: F,
) -> PyResult<Bound<'py, PyArray2<f64>>>
where
F: Fn(&[f64]) -> Vec<f64> + Sync,
{
let arr = data.as_array();
let (n_samples, n_series) = arr.dim();
let series_major = transpose_to_series_major(arr);
let (n_samples, _n_series) = arr.dim();
let columns = numpy2d_to_columns(&arr);
let col_results: Vec<Vec<f64>> = py.allow_threads(|| {
let run = |series_idx: usize| {
let column_row = series_major.row(series_idx);
let column = column_row
.as_slice()
.expect("series-major rows are contiguous");
process_col(column)
};
if parallel {
(0..n_series).into_par_iter().map(run).collect()
columns.par_iter().map(|col| per_col(col)).collect()
} else {
(0..n_series).map(run).collect()
columns.iter().map(|col| per_col(col)).collect()
}
});
finish_single_output(py, n_samples, n_series, col_results)
Ok(columns_to_numpy2d(py, n_samples, col_results))
}
fn validate_indicator_requests(names: &[String], timeperiods: &[usize]) -> PyResult<()> {
if names.len() != timeperiods.len() {
return Err(PyValueError::new_err(format!(
"names length ({}) must equal timeperiods length ({})",
names.len(),
timeperiods.len()
)));
}
for (name, &timeperiod) in names.iter().zip(timeperiods.iter()) {
if timeperiod == 0 {
return Err(PyValueError::new_err(format!(
"{name}: timeperiod must be >= 1"
)));
}
}
Ok(())
}
fn compute_cci(high: &[f64], low: &[f64], close: &[f64], timeperiod: usize) -> Vec<f64> {
let n = high.len();
let typical_price: Vec<f64> = high
.iter()
.zip(low.iter())
.zip(close.iter())
.map(|((&h, &l), &c)| (h + l + c) / 3.0)
.collect();
let mut result = vec![f64::NAN; n];
for end in (timeperiod - 1)..n {
let window = &typical_price[(end + 1 - timeperiod)..=end];
let mean = window.iter().sum::<f64>() / timeperiod as f64;
let mad = window
.iter()
.map(|&value| (value - mean).abs())
.sum::<f64>()
/ timeperiod as f64;
result[end] = if mad != 0.0 {
(typical_price[end] - mean) / (0.015 * mad)
} else {
0.0
};
}
result
}
fn compute_willr(
high: &[f64],
low: &[f64],
close: &[f64],
timeperiod: usize,
) -> PyResult<Vec<f64>> {
let n = high.len();
let mut result = vec![f64::NAN; n];
let mut max_ind =
Maximum::new(timeperiod).map_err(|err| PyValueError::new_err(err.to_string()))?;
let mut min_ind =
Minimum::new(timeperiod).map_err(|err| PyValueError::new_err(err.to_string()))?;
for (idx, ((&high_value, &low_value), &close_value)) in
high.iter().zip(low.iter()).zip(close.iter()).enumerate()
{
let highest = max_ind.next(high_value);
let lowest = min_ind.next(low_value);
if idx + 1 >= timeperiod {
let range = highest - lowest;
result[idx] = if range != 0.0 {
-100.0 * (highest - close_value) / range
} else {
-50.0
};
}
}
Ok(result)
}
fn compute_close_indicator(name: &str, close: &[f64], timeperiod: usize) -> PyResult<Vec<f64>> {
match name {
"SMA" => Ok(ferro_ta_core::overlap::sma(close, timeperiod)),
"EMA" => Ok(ferro_ta_core::overlap::ema(close, timeperiod)),
"RSI" => Ok(ferro_ta_core::momentum::rsi(close, timeperiod)),
"STDDEV" => Ok(ferro_ta_core::statistic::stddev(close, timeperiod, 1.0)),
"VAR" => Ok(ferro_ta_core::statistic::stddev(close, timeperiod, 1.0)
.into_iter()
.map(|value| if value.is_nan() { value } else { value * value })
.collect()),
"LINEARREG" => {
use crate::statistic::common::rolling_linreg_apply;
let last_x = (timeperiod - 1) as f64;
Ok(rolling_linreg_apply(
close,
timeperiod,
|slope: f64, intercept: f64| intercept + slope * last_x,
))
}
"LINEARREG_SLOPE" => {
use crate::statistic::common::rolling_linreg_apply;
Ok(rolling_linreg_apply(
close,
timeperiod,
|slope: f64, _: f64| slope,
))
}
"LINEARREG_INTERCEPT" => {
use crate::statistic::common::rolling_linreg_apply;
Ok(rolling_linreg_apply(
close,
timeperiod,
|_: f64, intercept: f64| intercept,
))
}
"LINEARREG_ANGLE" => {
use crate::statistic::common::rolling_linreg_apply;
Ok(rolling_linreg_apply(
close,
timeperiod,
|slope: f64, _: f64| slope.atan() * 180.0 / std::f64::consts::PI,
))
}
"TSF" => {
use crate::statistic::common::rolling_linreg_apply;
let forecast_x = timeperiod as f64;
Ok(rolling_linreg_apply(
close,
timeperiod,
|slope: f64, intercept: f64| intercept + slope * forecast_x,
))
}
_ => Err(PyValueError::new_err(format!(
"unsupported close indicator for grouped execution: {name}"
))),
}
}
fn compute_hlc_indicator(
name: &str,
high: &[f64],
low: &[f64],
close: &[f64],
timeperiod: usize,
) -> PyResult<Vec<f64>> {
match name {
"ATR" => Ok(ferro_ta_core::volatility::atr(high, low, close, timeperiod)),
"NATR" => {
let atr = ferro_ta_core::volatility::atr(high, low, close, timeperiod);
Ok(atr
.into_iter()
.zip(close.iter())
.map(|(atr_value, &close_value)| {
if atr_value.is_nan() || close_value == 0.0 {
f64::NAN
} else {
(atr_value / close_value) * 100.0
}
})
.collect())
}
"ADX" => Ok(ferro_ta_core::momentum::adx(high, low, close, timeperiod)),
"ADXR" => Ok(ferro_ta_core::momentum::adxr(high, low, close, timeperiod)),
"CCI" => Ok(compute_cci(high, low, close, timeperiod)),
"WILLR" => compute_willr(high, low, close, timeperiod),
_ => Err(PyValueError::new_err(format!(
"unsupported HLC indicator for grouped execution: {name}"
))),
}
}
type IndicatorArrayList = Vec<Py<PyArray1<f64>>>;
// ---------------------------------------------------------------------------
// batch_sma
// ---------------------------------------------------------------------------
@@ -309,9 +143,9 @@ pub fn batch_sma<'py>(
log::debug!(
"batch_sma: timeperiod={timeperiod}, shape=({n_samples}, {n_series}), parallel={parallel}"
);
Ok(run_unary_batch(py, data, parallel, |col| {
run_unary_batch_par(py, data, parallel, |col| {
ferro_ta_core::overlap::sma(col, timeperiod)
}))
})
}
// ---------------------------------------------------------------------------
@@ -319,17 +153,6 @@ pub fn batch_sma<'py>(
// ---------------------------------------------------------------------------
/// Batch Exponential Moving Average — applies EMA to every column.
///
/// Parameters
/// ----------
/// data : numpy array, shape (n_samples, n_series), dtype float64
/// timeperiod : int
/// parallel : bool, default True
/// When True, columns are processed in parallel via Rayon (GIL released).
///
/// Returns
/// -------
/// numpy array, shape (n_samples, n_series), dtype float64
#[pyfunction]
#[pyo3(signature = (data, timeperiod = 30, parallel = true))]
pub fn batch_ema<'py>(
@@ -345,9 +168,9 @@ pub fn batch_ema<'py>(
log::debug!(
"batch_ema: timeperiod={timeperiod}, shape=({n_samples}, {n_series}), parallel={parallel}"
);
Ok(run_unary_batch(py, data, parallel, |col| {
run_unary_batch_par(py, data, parallel, |col| {
ferro_ta_core::overlap::ema(col, timeperiod)
}))
})
}
// ---------------------------------------------------------------------------
@@ -355,18 +178,6 @@ pub fn batch_ema<'py>(
// ---------------------------------------------------------------------------
/// Batch RSI — applies RSI (Wilder seeding) to every column.
///
/// Parameters
/// ----------
/// data : numpy array, shape (n_samples, n_series), dtype float64
/// timeperiod : int
/// parallel : bool, default True
/// When True, columns are processed in parallel via Rayon (GIL released).
///
/// Returns
/// -------
/// numpy array, shape (n_samples, n_series), dtype float64
/// Values in [0, 100]; NaN during warmup.
#[pyfunction]
#[pyo3(signature = (data, timeperiod = 14, parallel = true))]
pub fn batch_rsi<'py>(
@@ -382,49 +193,9 @@ pub fn batch_rsi<'py>(
log::debug!(
"batch_rsi: timeperiod={timeperiod}, shape=({n_samples}, {n_series}), parallel={parallel}"
);
let period_f = timeperiod as f64;
Ok(run_unary_batch(py, data, parallel, |col| {
let mut col_result = vec![f64::NAN; n_samples];
if n_samples <= timeperiod {
return col_result;
}
let mut avg_gain = 0.0_f64;
let mut avg_loss = 0.0_f64;
for i in 1..=timeperiod {
let delta = col[i] - col[i - 1];
if delta > 0.0 {
avg_gain += delta;
} else {
avg_loss += -delta;
}
}
avg_gain /= period_f;
avg_loss /= period_f;
let rs = if avg_loss == 0.0 {
f64::MAX
} else {
avg_gain / avg_loss
};
col_result[timeperiod] = 100.0 - 100.0 / (1.0 + rs);
for i in (timeperiod + 1)..n_samples {
let delta = col[i] - col[i - 1];
let (gain, loss) = if delta > 0.0 {
(delta, 0.0)
} else {
(0.0, -delta)
};
avg_gain = (avg_gain * (period_f - 1.0) + gain) / period_f;
avg_loss = (avg_loss * (period_f - 1.0) + loss) / period_f;
let rs = if avg_loss == 0.0 {
f64::MAX
} else {
avg_gain / avg_loss
};
col_result[i] = 100.0 - 100.0 / (1.0 + rs);
}
col_result
}))
run_unary_batch_par(py, data, parallel, |col| {
ferro_ta_core::momentum::rsi(col, timeperiod)
})
}
// ---------------------------------------------------------------------------
@@ -451,40 +222,27 @@ pub fn batch_atr<'py>(
validate_same_shape((n_samples, n_series), arr_l.dim(), "low")?;
validate_same_shape((n_samples, n_series), arr_c.dim(), "close")?;
let high_by_series = transpose_to_series_major(arr_h);
let low_by_series = transpose_to_series_major(arr_l);
let close_by_series = transpose_to_series_major(arr_c);
let h_cols = numpy2d_to_columns(&arr_h);
let l_cols = numpy2d_to_columns(&arr_l);
let c_cols = numpy2d_to_columns(&arr_c);
let col_results: Vec<Vec<f64>> = py.allow_threads(|| {
let process_col = |series_idx: usize| -> Vec<f64> {
let high_row = high_by_series.row(series_idx);
let low_row = low_by_series.row(series_idx);
let close_row = close_by_series.row(series_idx);
let high_col = high_row
.as_slice()
.expect("series-major rows are contiguous");
let low_col = low_row
.as_slice()
.expect("series-major rows are contiguous");
let close_col = close_row
.as_slice()
.expect("series-major rows are contiguous");
ferro_ta_core::volatility::atr(high_col, low_col, close_col, timeperiod)
let process = |i: usize| {
ferro_ta_core::volatility::atr(&h_cols[i], &l_cols[i], &c_cols[i], timeperiod)
};
if parallel {
(0..n_series).into_par_iter().map(process_col).collect()
(0..n_series).into_par_iter().map(process).collect()
} else {
(0..n_series).map(process_col).collect()
(0..n_series).map(process).collect()
}
});
Ok(finish_single_output(py, n_samples, n_series, col_results))
Ok(columns_to_numpy2d(py, n_samples, col_results))
}
// ---------------------------------------------------------------------------
// batch_stoch
// ---------------------------------------------------------------------------
/// Stoch batch result type (slowk, slowd arrays).
type StochBatchResult<'py> = (Bound<'py, PyArray2<f64>>, Bound<'py, PyArray2<f64>>);
#[pyfunction]
@@ -507,40 +265,30 @@ pub fn batch_stoch<'py>(
validate_same_shape((n_samples, n_series), arr_l.dim(), "low")?;
validate_same_shape((n_samples, n_series), arr_c.dim(), "close")?;
let high_by_series = transpose_to_series_major(arr_h);
let low_by_series = transpose_to_series_major(arr_l);
let close_by_series = transpose_to_series_major(arr_c);
let h_cols = numpy2d_to_columns(&arr_h);
let l_cols = numpy2d_to_columns(&arr_l);
let c_cols = numpy2d_to_columns(&arr_c);
let col_results: Vec<(Vec<f64>, Vec<f64>)> = py.allow_threads(|| {
let process_col = |series_idx: usize| -> (Vec<f64>, Vec<f64>) {
let high_row = high_by_series.row(series_idx);
let low_row = low_by_series.row(series_idx);
let close_row = close_by_series.row(series_idx);
let high_col = high_row
.as_slice()
.expect("series-major rows are contiguous");
let low_col = low_row
.as_slice()
.expect("series-major rows are contiguous");
let close_col = close_row
.as_slice()
.expect("series-major rows are contiguous");
let process = |i: usize| {
ferro_ta_core::momentum::stoch(
high_col,
low_col,
close_col,
&h_cols[i],
&l_cols[i],
&c_cols[i],
fastk_period,
slowk_period,
slowd_period,
)
};
if parallel {
(0..n_series).into_par_iter().map(process_col).collect()
(0..n_series).into_par_iter().map(process).collect()
} else {
(0..n_series).map(process_col).collect()
(0..n_series).map(process).collect()
}
});
Ok(finish_pair_output(py, n_samples, n_series, col_results))
let (all_k, all_d): (Vec<Vec<f64>>, Vec<Vec<f64>>) = col_results.into_iter().unzip();
Ok(column_pair_to_numpy2d(py, n_samples, all_k, all_d))
}
// ---------------------------------------------------------------------------
@@ -567,39 +315,29 @@ pub fn batch_adx<'py>(
validate_same_shape((n_samples, n_series), arr_l.dim(), "low")?;
validate_same_shape((n_samples, n_series), arr_c.dim(), "close")?;
let high_by_series = transpose_to_series_major(arr_h);
let low_by_series = transpose_to_series_major(arr_l);
let close_by_series = transpose_to_series_major(arr_c);
let h_cols = numpy2d_to_columns(&arr_h);
let l_cols = numpy2d_to_columns(&arr_l);
let c_cols = numpy2d_to_columns(&arr_c);
let col_results: Vec<Vec<f64>> = py.allow_threads(|| {
let process_col = |series_idx: usize| -> Vec<f64> {
let high_row = high_by_series.row(series_idx);
let low_row = low_by_series.row(series_idx);
let close_row = close_by_series.row(series_idx);
let high_col = high_row
.as_slice()
.expect("series-major rows are contiguous");
let low_col = low_row
.as_slice()
.expect("series-major rows are contiguous");
let close_col = close_row
.as_slice()
.expect("series-major rows are contiguous");
ferro_ta_core::momentum::adx(high_col, low_col, close_col, timeperiod)
let process = |i: usize| {
ferro_ta_core::momentum::adx(&h_cols[i], &l_cols[i], &c_cols[i], timeperiod)
};
if parallel {
(0..n_series).into_par_iter().map(process_col).collect()
(0..n_series).into_par_iter().map(process).collect()
} else {
(0..n_series).map(process_col).collect()
(0..n_series).map(process).collect()
}
});
Ok(finish_single_output(py, n_samples, n_series, col_results))
Ok(columns_to_numpy2d(py, n_samples, col_results))
}
// ---------------------------------------------------------------------------
// grouped 1-D execution
// ---------------------------------------------------------------------------
type IndicatorArrayList = Vec<Py<PyArray1<f64>>>;
#[pyfunction]
#[pyo3(signature = (close, names, timeperiods, parallel = true))]
pub fn run_close_indicators<'py>(
@@ -609,22 +347,35 @@ pub fn run_close_indicators<'py>(
timeperiods: Vec<usize>,
parallel: bool,
) -> PyResult<IndicatorArrayList> {
validate_indicator_requests(&names, &timeperiods)?;
let close_values = close.as_slice()?;
let results: Vec<PyResult<Vec<f64>>> = py.allow_threads(|| {
let run = |idx: usize| compute_close_indicator(&names[idx], close_values, timeperiods[idx]);
if parallel {
(0..names.len()).into_par_iter().map(run).collect()
} else {
(0..names.len()).map(run).collect()
}
});
results
.into_iter()
.map(|result| result.map(|values| values.into_pyarray(py).unbind()))
.collect()
if parallel {
// Parallel path: call core per-indicator in parallel via Rayon
let results: Vec<Result<Vec<f64>, String>> = py.allow_threads(|| {
(0..names.len())
.into_par_iter()
.map(|idx| {
ferro_ta_core::batch::run_close_indicators(
close_values,
&[names[idx].clone()],
&[timeperiods[idx]],
)
.map(|mut v| v.remove(0))
})
.collect()
});
results
.into_iter()
.map(|r| r.map(|v| v.into_pyarray(py).unbind()).map_err(map_core_err))
.collect()
} else {
let results = ferro_ta_core::batch::run_close_indicators(close_values, &names, &timeperiods)
.map_err(map_core_err)?;
Ok(results
.into_iter()
.map(|v| v.into_pyarray(py).unbind())
.collect())
}
}
#[pyfunction]
@@ -638,7 +389,6 @@ pub fn run_hlc_indicators<'py>(
timeperiods: Vec<usize>,
parallel: bool,
) -> PyResult<IndicatorArrayList> {
validate_indicator_requests(&names, &timeperiods)?;
let high_values = high.as_slice()?;
let low_values = low.as_slice()?;
let close_values = close.as_slice()?;
@@ -649,27 +399,40 @@ pub fn run_hlc_indicators<'py>(
));
}
let results: Vec<PyResult<Vec<f64>>> = py.allow_threads(|| {
let run = |idx: usize| {
compute_hlc_indicator(
&names[idx],
high_values,
low_values,
close_values,
timeperiods[idx],
)
};
if parallel {
(0..names.len()).into_par_iter().map(run).collect()
} else {
(0..names.len()).map(run).collect()
}
});
results
.into_iter()
.map(|result| result.map(|values| values.into_pyarray(py).unbind()))
.collect()
if parallel {
let results: Vec<Result<Vec<f64>, String>> = py.allow_threads(|| {
(0..names.len())
.into_par_iter()
.map(|idx| {
ferro_ta_core::batch::run_hlc_indicators(
high_values,
low_values,
close_values,
&[names[idx].clone()],
&[timeperiods[idx]],
)
.map(|mut v| v.remove(0))
})
.collect()
});
results
.into_iter()
.map(|r| r.map(|v| v.into_pyarray(py).unbind()).map_err(map_core_err))
.collect()
} else {
let results = ferro_ta_core::batch::run_hlc_indicators(
high_values,
low_values,
close_values,
&names,
&timeperiods,
)
.map_err(map_core_err)?;
Ok(results
.into_iter()
.map(|v| v.into_pyarray(py).unbind())
.collect())
}
}
// ---------------------------------------------------------------------------
+19 -136
View File
@@ -1,45 +1,10 @@
//! Chunked / out-of-core execution helpers.
//!
//! These functions support running indicators on data that is too large for
//! memory by processing it in chunks. The caller splits a large series into
//! overlapping chunks (overlap = indicator warm-up period), runs an indicator
//! on each chunk, and then stitches the results by trimming the overlap from
//! the front of each chunk's output.
//!
//! Functions
//! ---------
//! - `trim_overlap` — remove the first *overlap* elements from
//! an array (to strip the warm-up from a chunk's indicator output).
//! - `stitch_chunks` — concatenate trimmed chunk results into one
//! array.
//! - `make_chunk_ranges` — compute start/end indices for a series
//! given chunk size and overlap, for use by the Python caller.
//! - `chunk_apply_close_indicator`— run chunked close-only indicators fully in
//! Rust (SMA/EMA/RSI).
//! - `forward_fill_nan` — forward-fill NaN values in a 1-D array.
//! Chunked / out-of-core execution helpers (thin PyO3 wrapper over ferro_ta_core::chunked).
use numpy::{IntoPyArray, PyArray1, PyReadonlyArray1};
use pyo3::exceptions::PyValueError;
use pyo3::prelude::*;
// ---------------------------------------------------------------------------
// trim_overlap
// ---------------------------------------------------------------------------
/// Remove the first *overlap* elements from an array.
///
/// After running an indicator on a chunk that includes a warm-up prefix, the
/// first *overlap* output values are unreliable (NaN or influenced by
/// artificial padding). This function discards them.
///
/// Parameters
/// ----------
/// chunk_out : 1-D float64 array — indicator output for a chunk
/// overlap : int — number of leading elements to discard
///
/// Returns
/// -------
/// 1-D float64 array — the trailing ``len(chunk_out) - overlap`` elements
#[pyfunction]
pub fn trim_overlap<'py>(
py: Python<'py>,
@@ -47,67 +12,32 @@ pub fn trim_overlap<'py>(
overlap: usize,
) -> PyResult<Bound<'py, PyArray1<f64>>> {
let s = chunk_out.as_slice()?;
let n = s.len();
if overlap > n {
if overlap > s.len() {
return Err(PyValueError::new_err(format!(
"overlap ({overlap}) must be <= chunk length ({n})"
"overlap ({overlap}) must be <= chunk length ({})",
s.len()
)));
}
let out = s[overlap..].to_vec();
Ok(out.into_pyarray(py))
let result = ferro_ta_core::chunked::trim_overlap(s, overlap);
Ok(result.into_pyarray(py))
}
// ---------------------------------------------------------------------------
// stitch_chunks
// ---------------------------------------------------------------------------
/// Concatenate a list of trimmed chunk results into a single output array.
///
/// Parameters
/// ----------
/// chunks : list of 1-D float64 arrays — trimmed outputs from each chunk
///
/// Returns
/// -------
/// 1-D float64 array — concatenated result
#[pyfunction]
pub fn stitch_chunks<'py>(
py: Python<'py>,
chunks: Vec<PyReadonlyArray1<'py, f64>>,
) -> PyResult<Bound<'py, PyArray1<f64>>> {
let mut out: Vec<f64> = Vec::new();
for chunk in &chunks {
out.extend_from_slice(chunk.as_slice()?);
}
Ok(out.into_pyarray(py))
let vecs: Vec<Vec<f64>> = chunks
.iter()
.map(|c| c.as_slice().map(|s| s.to_vec()))
.collect::<Result<_, _>>()?;
let refs: Vec<&[f64]> = vecs.iter().map(|v| v.as_slice()).collect();
let result = ferro_ta_core::chunked::stitch_chunks(&refs);
Ok(result.into_pyarray(py))
}
// ---------------------------------------------------------------------------
// make_chunk_ranges
// ---------------------------------------------------------------------------
/// Compute the (start, end) index pairs for chunked processing.
///
/// Each range ``[start, end)`` specifies a slice of the input series that the
/// caller should pass to the indicator function. The first *overlap* elements
/// of each range (except the very first range) are the warm-up prefix from the
/// previous chunk.
///
/// Parameters
/// ----------
/// n : int — total length of the series
/// chunk_size : int — desired number of *output* bars per chunk (>= 1)
/// overlap : int — number of warm-up bars prepended to each chunk (>= 0)
///
/// Returns
/// -------
/// list of (start: int, end: int) pairs as a flattened 1-D int64 array of
/// length 2 × n_chunks. Caller unpacks with ``ranges.reshape(-1, 2)``.
///
/// Example
/// -------
/// For n=10, chunk_size=4, overlap=2 the ranges would cover:
/// [0, 4), [2, 8), [6, 10) (start of chunk 2 = end of prev chunk - overlap)
/// Compute (start, end) index pairs for chunked processing.
#[pyfunction]
pub fn make_chunk_ranges<'py>(
py: Python<'py>,
@@ -118,26 +48,12 @@ pub fn make_chunk_ranges<'py>(
if chunk_size == 0 {
return Err(PyValueError::new_err("chunk_size must be >= 1"));
}
let mut ranges: Vec<i64> = Vec::new();
if n == 0 {
return Ok(ranges.into_pyarray(py));
}
let mut start: usize = 0;
loop {
let end = (start + chunk_size + overlap).min(n);
ranges.push(start as i64);
ranges.push(end as i64);
if end >= n {
break;
}
// Next chunk starts at end - overlap (so the next chunk has its overlap prefix)
start = end.saturating_sub(overlap);
}
Ok(ranges.into_pyarray(py))
let result = ferro_ta_core::chunked::make_chunk_ranges(n, chunk_size, overlap);
Ok(result.into_pyarray(py))
}
// ---------------------------------------------------------------------------
// chunk_apply_close_indicator
// chunk_apply_close_indicator — stays in PyO3 (dispatches to ferro_ta_core indicators)
// ---------------------------------------------------------------------------
fn compute_close_indicator(
@@ -156,18 +72,6 @@ fn compute_close_indicator(
}
/// Run chunked execution for close-only indicators in Rust.
///
/// Parameters
/// ----------
/// series : 1-D float64 array
/// indicator : one of {"SMA", "EMA", "RSI"}
/// timeperiod : indicator period (>= 1)
/// chunk_size : output bars per chunk (>= 1)
/// overlap : warm-up bars prepended to each chunk
///
/// Returns
/// -------
/// 1-D float64 array with the same length as `series`.
#[pyfunction]
#[pyo3(signature = (series, indicator, timeperiod, chunk_size = 10_000, overlap = 100))]
pub fn chunk_apply_close_indicator<'py>(
@@ -227,38 +131,17 @@ pub fn chunk_apply_close_indicator<'py>(
Ok(stitched.into_pyarray(py))
}
// ---------------------------------------------------------------------------
// forward_fill_nan
// ---------------------------------------------------------------------------
/// Forward-fill NaN values in a 1-D array.
///
/// Leading NaN values are preserved until the first non-NaN value appears.
#[pyfunction]
pub fn forward_fill_nan<'py>(
py: Python<'py>,
values: PyReadonlyArray1<'py, f64>,
) -> PyResult<Bound<'py, PyArray1<f64>>> {
let input = values.as_slice()?;
let mut out = Vec::with_capacity(input.len());
let mut last = f64::NAN;
for &value in input {
if value.is_nan() {
out.push(last);
} else {
last = value;
out.push(value);
}
}
Ok(out.into_pyarray(py))
let result = ferro_ta_core::chunked::forward_fill_nan(input);
Ok(result.into_pyarray(py))
}
// ---------------------------------------------------------------------------
// Register
// ---------------------------------------------------------------------------
pub fn register(m: &Bound<'_, PyModule>) -> PyResult<()> {
m.add_function(wrap_pyfunction!(trim_overlap, m)?)?;
m.add_function(wrap_pyfunction!(stitch_chunks, m)?)?;
+9 -99
View File
@@ -1,40 +1,10 @@
//! Crypto and 24/7 market helpers.
//!
//! Functions designed for continuous (24/7) markets such as crypto:
//!
//! - `funding_cumulative_pnl` — cumulative PnL from periodic funding rate
//! payments, given a constant position size.
//! - `continuous_bar_labels` — assign a sequential integer label to each bar
//! based on a fixed period size (e.g. daily UTC buckets).
//! - `mark_session_boundaries` — return indices where the session rolls over
//! (useful for calendar-free resampling on continuous data).
//! Crypto and 24/7 market helpers (thin PyO3 wrapper over ferro_ta_core::crypto).
use numpy::{IntoPyArray, PyArray1, PyReadonlyArray1};
use pyo3::exceptions::PyValueError;
use pyo3::prelude::*;
// ---------------------------------------------------------------------------
// funding_cumulative_pnl
// ---------------------------------------------------------------------------
/// Compute the cumulative PnL from funding rate payments.
///
/// Crypto perpetual contracts charge a periodic funding rate to the holder.
/// If you hold ``position_size`` contracts and the funding rate is ``rate[i]``,
/// the PnL at period *i* is ``-position_size * rate[i]`` (longs pay when rate
/// is positive). This function returns the **cumulative** funding PnL.
///
/// Parameters
/// ----------
/// position_size : 1-D float64 array — signed position size per funding period
/// (positive = long, negative = short). Must have the same length as
/// ``funding_rate``.
/// funding_rate : 1-D float64 array — periodic funding rate (decimal, e.g.
/// 0.0001 = 0.01%).
///
/// Returns
/// -------
/// 1-D float64 array — cumulative funding PnL (same length as inputs).
#[pyfunction]
pub fn funding_cumulative_pnl<'py>(
py: Python<'py>,
@@ -43,41 +13,16 @@ pub fn funding_cumulative_pnl<'py>(
) -> PyResult<Bound<'py, PyArray1<f64>>> {
let pos = position_size.as_slice()?;
let rate = funding_rate.as_slice()?;
let n = pos.len();
if n != rate.len() {
if pos.len() != rate.len() {
return Err(PyValueError::new_err(
"position_size and funding_rate must have the same length",
));
}
let mut out = vec![0.0_f64; n];
let mut cumulative = 0.0_f64;
for i in 0..n {
// Longs pay when rate > 0; shorts receive
cumulative += -pos[i] * rate[i];
out[i] = cumulative;
}
Ok(out.into_pyarray(py))
let result = ferro_ta_core::crypto::funding_cumulative_pnl(pos, rate);
Ok(result.into_pyarray(py))
}
// ---------------------------------------------------------------------------
// continuous_bar_labels
// ---------------------------------------------------------------------------
/// Assign a sequential integer label per bar based on a fixed-size period.
///
/// For 24/7 data (no session gaps), this groups bars into equal-sized buckets:
/// bars 0…(period_bars-1) get label 0, bars period_bars…(2*period_bars-1) get
/// label 1, etc. Useful for resampling continuous data (e.g. group every 24
/// one-hour bars into a "day").
///
/// Parameters
/// ----------
/// n_bars : int — total number of bars
/// period_bars: int — number of bars per period (must be >= 1)
///
/// Returns
/// -------
/// 1-D int64 array of length *n_bars* — period labels (0-based).
#[pyfunction]
pub fn continuous_bar_labels<'py>(
py: Python<'py>,
@@ -87,56 +32,21 @@ pub fn continuous_bar_labels<'py>(
if period_bars == 0 {
return Err(PyValueError::new_err("period_bars must be >= 1"));
}
let out: Vec<i64> = (0..n_bars).map(|i| (i / period_bars) as i64).collect();
Ok(out.into_pyarray(py))
let result = ferro_ta_core::crypto::continuous_bar_labels(n_bars, period_bars);
Ok(result.into_pyarray(py))
}
// ---------------------------------------------------------------------------
// mark_session_boundaries
// ---------------------------------------------------------------------------
/// Return bar indices where a new session begins.
///
/// Given UTC timestamps in nanoseconds (int64), marks boundaries at the start
/// of each new UTC day (midnight). Returns the bar indices where the UTC day
/// changes. Bar 0 is always included as the first boundary.
///
/// Parameters
/// ----------
/// timestamps_ns : 1-D int64 array — UTC timestamps in nanoseconds
/// (e.g. from ``pandas.DatetimeIndex.astype('int64')``).
///
/// Returns
/// -------
/// 1-D int64 array — indices of bars at the start of each new UTC day.
/// Return bar indices where a new UTC day begins.
#[pyfunction]
pub fn mark_session_boundaries<'py>(
py: Python<'py>,
timestamps_ns: PyReadonlyArray1<'py, i64>,
) -> PyResult<Bound<'py, PyArray1<i64>>> {
let ts = timestamps_ns.as_slice()?;
let n = ts.len();
if n == 0 {
return Ok(Vec::<i64>::new().into_pyarray(py));
}
// Nanoseconds per day
const NS_PER_DAY: i64 = 86_400_000_000_000;
let mut out = vec![0i64]; // bar 0 is always a boundary
let mut prev_day = ts[0].div_euclid(NS_PER_DAY);
for (i, &t) in ts.iter().enumerate().skip(1) {
let day = t.div_euclid(NS_PER_DAY);
if day != prev_day {
out.push(i as i64);
prev_day = day;
}
}
Ok(out.into_pyarray(py))
let result = ferro_ta_core::crypto::mark_session_boundaries(ts);
Ok(result.into_pyarray(py))
}
// ---------------------------------------------------------------------------
// Register
// ---------------------------------------------------------------------------
pub fn register(m: &Bound<'_, PyModule>) -> PyResult<()> {
m.add_function(wrap_pyfunction!(funding_cumulative_pnl, m)?)?;
m.add_function(wrap_pyfunction!(continuous_bar_labels, m)?)?;
-187
View File
@@ -1,187 +0,0 @@
// Shared Hilbert Transform Core
// Based on John Ehlers' Discrete Hilbert Transform as implemented in TA-Lib.
// Reference: "Cybernetic Analysis for Stocks and Futures" by J.F. Ehlers
//
// All HT functions share a 63-bar lookback period.
use std::f64::consts::PI;
pub(super) const HT_LOOKBACK: usize = 63;
/// Shared output from the core Hilbert Transform computation.
pub(super) struct HtCore {
pub(super) trendline: Vec<f64>,
pub(super) dc_period: Vec<f64>,
pub(super) dc_phase: Vec<f64>,
pub(super) inphase: Vec<f64>,
pub(super) quadrature: Vec<f64>,
pub(super) trend_mode: Vec<i32>,
}
/// Run the full Hilbert Transform pipeline on a slice of close prices.
pub(super) fn compute_ht_core(prices: &[f64]) -> HtCore {
let n = prices.len();
let mut trendline = vec![f64::NAN; n];
let mut dc_period = vec![f64::NAN; n];
let mut dc_phase = vec![f64::NAN; n];
let mut inphase = vec![f64::NAN; n];
let mut quadrature = vec![f64::NAN; n];
let mut trend_mode = vec![0i32; n];
if n <= HT_LOOKBACK {
return HtCore {
trendline,
dc_period,
dc_phase,
inphase,
quadrature,
trend_mode,
};
}
// Step 1: Smooth the price series (4-bar weighted average)
let mut smooth = vec![0.0f64; n];
for i in 0..n {
smooth[i] = if i >= 3 {
(4.0 * prices[i] + 3.0 * prices[i - 1] + 2.0 * prices[i - 2] + prices[i - 3]) / 10.0
} else {
prices[i]
};
}
// Step 2: Full Hilbert Transform pipeline
let mut detrender = vec![0.0f64; n];
let mut q1 = vec![0.0f64; n];
let mut i1 = vec![0.0f64; n];
let mut ji = vec![0.0f64; n];
let mut jq = vec![0.0f64; n];
let mut i2 = vec![0.0f64; n];
let mut q2 = vec![0.0f64; n];
let mut re = vec![0.0f64; n];
let mut im = vec![0.0f64; n];
let mut period = vec![0.0f64; n];
let mut smooth_period = vec![0.0f64; n];
let mut phase = vec![0.0f64; n];
for i in 6..n {
let prev_period = period[i - 1];
// Alpha coefficient for HT filters depends on the current period estimate
let alpha = 0.075 * prev_period + 0.54;
// Discrete Hilbert Transform of smooth price (detrender)
detrender[i] = (0.0962 * smooth[i] + 0.5769 * smooth[i - 2]
- 0.5769 * smooth[i - 4]
- 0.0962 * smooth[i - 6])
* alpha;
// Q1: HT of detrender
if i >= 12 {
q1[i] = (0.0962 * detrender[i] + 0.5769 * detrender[i - 2]
- 0.5769 * detrender[i - 4]
- 0.0962 * detrender[i - 6])
* alpha;
}
// I1: delayed detrender
if i >= 9 {
i1[i] = detrender[i - 3];
}
// jI: HT of I1
if i >= 15 {
ji[i] = (0.0962 * i1[i] + 0.5769 * i1[i - 2] - 0.5769 * i1[i - 4] - 0.0962 * i1[i - 6])
* alpha;
}
// jQ: HT of Q1
if i >= 18 {
jq[i] = (0.0962 * q1[i] + 0.5769 * q1[i - 2] - 0.5769 * q1[i - 4] - 0.0962 * q1[i - 6])
* alpha;
}
// Phase components
let i2_raw = i1[i] - jq[i];
let q2_raw = q1[i] + ji[i];
// EMA smoothing of I2 and Q2
let i2_prev = i2[i - 1];
let q2_prev = q2[i - 1];
i2[i] = 0.2 * i2_raw + 0.8 * i2_prev;
q2[i] = 0.2 * q2_raw + 0.8 * q2_prev;
// Cross-product for period estimation
let re_raw = i2[i] * i2_prev + q2[i] * q2_prev;
let im_raw = i2[i] * q2_prev - q2[i] * i2_prev;
// EMA smoothing of Re and Im
re[i] = 0.2 * re_raw + 0.8 * re[i - 1];
im[i] = 0.2 * im_raw + 0.8 * im[i - 1];
// Compute period from cross-product of consecutive phasors.
// Uses atan(Im/Re) per Ehlers' convention; guard against negative Re
// which would flip the sign of the period estimate.
let mut p = if re[i] != 0.0 && im[i] != 0.0 && re[i] > 0.0 {
2.0 * PI / (im[i] / re[i]).atan()
} else {
prev_period
};
// Clamp period relative to previous
if prev_period > 0.0 {
if p > 1.5 * prev_period {
p = 1.5 * prev_period;
}
if p < 0.67 * prev_period {
p = 0.67 * prev_period;
}
}
// Hard clamp to [6, 50] bars
p = p.clamp(6.0, 50.0);
// EMA smooth the period
period[i] = 0.2 * p + 0.8 * prev_period;
// Smooth the smoothed period once more
smooth_period[i] = 0.33 * period[i] + 0.67 * smooth_period[i - 1];
// Phase from I1 and Q1
phase[i] = if i1[i] != 0.0 {
q1[i].atan2(i1[i]) * 180.0 / PI
} else if q1[i] > 0.0 {
90.0
} else if q1[i] < 0.0 {
-90.0
} else {
0.0
};
// Write outputs once past lookback
if i >= HT_LOOKBACK {
dc_period[i] = smooth_period[i];
dc_phase[i] = phase[i];
inphase[i] = i1[i];
quadrature[i] = q1[i];
// Trend mode: cycle when SmoothPeriod >= 20, trend when < 20
trend_mode[i] = if smooth_period[i] < 20.0 { 1 } else { 0 };
}
}
// Trendline: average over the current dominant cycle period
for i in HT_LOOKBACK..n {
let sp = smooth_period[i];
let dc = (sp.round() as usize).max(1).min(i + 1);
let sum: f64 = (0..dc).map(|j| smooth[i - j]).sum();
trendline[i] = sum / dc as f64;
}
HtCore {
trendline,
dc_period,
dc_phase,
inphase,
quadrature,
trend_mode,
}
}
+2 -4
View File
@@ -1,4 +1,3 @@
use super::common::compute_ht_core;
use numpy::{IntoPyArray, PyArray1, PyReadonlyArray1};
use pyo3::prelude::*;
@@ -8,7 +7,6 @@ pub fn ht_dcperiod<'py>(
py: Python<'py>,
close: PyReadonlyArray1<'py, f64>,
) -> PyResult<Bound<'py, PyArray1<f64>>> {
let prices = close.as_slice()?;
let core = compute_ht_core(prices);
Ok(core.dc_period.into_pyarray(py))
let result = ferro_ta_core::cycle::ht_dcperiod(close.as_slice()?);
Ok(result.into_pyarray(py))
}
+2 -4
View File
@@ -1,4 +1,3 @@
use super::common::compute_ht_core;
use numpy::{IntoPyArray, PyArray1, PyReadonlyArray1};
use pyo3::prelude::*;
@@ -8,7 +7,6 @@ pub fn ht_dcphase<'py>(
py: Python<'py>,
close: PyReadonlyArray1<'py, f64>,
) -> PyResult<Bound<'py, PyArray1<f64>>> {
let prices = close.as_slice()?;
let core = compute_ht_core(prices);
Ok(core.dc_phase.into_pyarray(py))
let result = ferro_ta_core::cycle::ht_dcphase(close.as_slice()?);
Ok(result.into_pyarray(py))
}
+2 -7
View File
@@ -1,4 +1,3 @@
use super::common::compute_ht_core;
use numpy::{IntoPyArray, PyArray1, PyReadonlyArray1};
use pyo3::prelude::*;
@@ -9,10 +8,6 @@ pub fn ht_phasor<'py>(
py: Python<'py>,
close: PyReadonlyArray1<'py, f64>,
) -> PyResult<(Bound<'py, PyArray1<f64>>, Bound<'py, PyArray1<f64>>)> {
let prices = close.as_slice()?;
let core = compute_ht_core(prices);
Ok((
core.inphase.into_pyarray(py),
core.quadrature.into_pyarray(py),
))
let (inphase, quadrature) = ferro_ta_core::cycle::ht_phasor(close.as_slice()?);
Ok((inphase.into_pyarray(py), quadrature.into_pyarray(py)))
}
+1 -17
View File
@@ -1,7 +1,5 @@
use super::common::{compute_ht_core, HT_LOOKBACK};
use numpy::{IntoPyArray, PyArray1, PyReadonlyArray1};
use pyo3::prelude::*;
use std::f64::consts::PI;
/// Hilbert Transform SineWave. Returns (sine, leadsine) where leadsine leads sine by 45°.
#[pyfunction]
@@ -10,20 +8,6 @@ pub fn ht_sine<'py>(
py: Python<'py>,
close: PyReadonlyArray1<'py, f64>,
) -> PyResult<(Bound<'py, PyArray1<f64>>, Bound<'py, PyArray1<f64>>)> {
let prices = close.as_slice()?;
let n = prices.len();
let core = compute_ht_core(prices);
let mut sine = vec![f64::NAN; n];
let mut lead_sine = vec![f64::NAN; n];
for i in HT_LOOKBACK..n {
if !core.dc_phase[i].is_nan() {
let phase_rad = core.dc_phase[i] * PI / 180.0;
sine[i] = phase_rad.sin();
lead_sine[i] = (phase_rad + PI / 4.0).sin(); // 45-degree lead
}
}
let (sine, lead_sine) = ferro_ta_core::cycle::ht_sine(close.as_slice()?);
Ok((sine.into_pyarray(py), lead_sine.into_pyarray(py)))
}
+2 -4
View File
@@ -1,4 +1,3 @@
use super::common::compute_ht_core;
use numpy::{IntoPyArray, PyArray1, PyReadonlyArray1};
use pyo3::prelude::*;
@@ -8,7 +7,6 @@ pub fn ht_trendline<'py>(
py: Python<'py>,
close: PyReadonlyArray1<'py, f64>,
) -> PyResult<Bound<'py, PyArray1<f64>>> {
let prices = close.as_slice()?;
let core = compute_ht_core(prices);
Ok(core.trendline.into_pyarray(py))
let result = ferro_ta_core::cycle::ht_trendline(close.as_slice()?);
Ok(result.into_pyarray(py))
}
+2 -4
View File
@@ -1,4 +1,3 @@
use super::common::compute_ht_core;
use numpy::{IntoPyArray, PyArray1, PyReadonlyArray1};
use pyo3::prelude::*;
@@ -8,7 +7,6 @@ pub fn ht_trendmode<'py>(
py: Python<'py>,
close: PyReadonlyArray1<'py, f64>,
) -> PyResult<Bound<'py, PyArray1<i32>>> {
let prices = close.as_slice()?;
let core = compute_ht_core(prices);
Ok(core.trend_mode.into_pyarray(py))
let result = ferro_ta_core::cycle::ht_trendmode(close.as_slice()?);
Ok(result.into_pyarray(py))
}
-2
View File
@@ -3,8 +3,6 @@
//!
//! All functions use a 63-bar lookback period (first 63 values are NaN).
mod common;
mod ht_dcperiod;
mod ht_dcphase;
mod ht_phasor;
+39 -546
View File
@@ -1,76 +1,20 @@
//! Extended Indicators — Rust implementations of indicators not in TA-Lib.
//! Extended Indicators — thin PyO3 wrappers delegating to `ferro_ta_core::extended`.
//!
//! All compute-heavy work (sequential loops, rolling windows) is done in Rust.
//! Python wrappers in `python/ferro_ta/extended.py` are thin call-throughs that
//! handle input conversion and pandas/polars wrapping.
//! All compute-heavy work lives in the core crate. These functions convert
//! numpy arrays to slices, call the core, and convert the results back.
#![allow(clippy::type_complexity)]
#![allow(clippy::too_many_arguments)]
use std::collections::VecDeque;
use crate::validation;
use numpy::{IntoPyArray, PyArray1, PyReadonlyArray1};
use pyo3::exceptions::PyValueError;
use pyo3::prelude::*;
// ---------------------------------------------------------------------------
// Internal helpers
// ---------------------------------------------------------------------------
/// Compute ATR array using Wilder smoothing (same as src/volatility/atr.rs).
fn compute_atr(high: &[f64], low: &[f64], close: &[f64], timeperiod: usize) -> Vec<f64> {
let n = high.len();
let mut result = vec![f64::NAN; n];
if n <= timeperiod {
return result;
}
// Seed: SMA of first `timeperiod` true range values
let mut seed_sum = high[0] - low[0]; // first TR has no prev_close
for i in 1..timeperiod {
let hl = high[i] - low[i];
let hc = (high[i] - close[i - 1]).abs();
let lc = (low[i] - close[i - 1]).abs();
seed_sum += hl.max(hc).max(lc);
}
let mut atr = seed_sum / timeperiod as f64;
result[timeperiod - 1] = atr;
let pf = (timeperiod - 1) as f64;
for i in timeperiod..n {
let hl = high[i] - low[i];
let hc = (high[i] - close[i - 1]).abs();
let lc = (low[i] - close[i - 1]).abs();
let tr = hl.max(hc).max(lc);
atr = (atr * pf + tr) / timeperiod as f64;
result[i] = atr;
}
result
}
/// Compute EMA using ferro_ta_core (SMA-seeded, matches Python EMA).
fn compute_ema_ta(prices: &[f64], timeperiod: usize) -> Vec<f64> {
ferro_ta_core::overlap::ema(prices, timeperiod)
}
/// Compute WMA array using ferro_ta_core O(n) implementation.
fn compute_wma(prices: &[f64], timeperiod: usize) -> Vec<f64> {
ferro_ta_core::overlap::wma(prices, timeperiod)
}
// ---------------------------------------------------------------------------
// VWAP
// ---------------------------------------------------------------------------
/// Volume Weighted Average Price (cumulative or rolling).
///
/// Parameters
/// ----------
/// high, low, close, volume : 1-D float64 arrays (equal length)
/// timeperiod : 0 = cumulative from bar 0; >= 1 = rolling window
///
/// Returns
/// -------
/// 1-D float64 array of VWAP values.
#[pyfunction]
#[pyo3(signature = (high, low, close, volume, timeperiod = 0))]
pub fn vwap<'py>(
@@ -85,59 +29,13 @@ pub fn vwap<'py>(
let lo = low.as_slice()?;
let c = close.as_slice()?;
let v = volume.as_slice()?;
let n = h.len();
validation::validate_equal_length(&[
(n, "high"),
(h.len(), "high"),
(lo.len(), "low"),
(c.len(), "close"),
(v.len(), "volume"),
])?;
let mut result = vec![f64::NAN; n];
let mut cum_tpv = 0.0f64;
let mut cum_vol = 0.0f64;
if timeperiod == 0 {
for i in 0..n {
let tp = (h[i] + lo[i] + c[i]) / 3.0;
cum_tpv += tp * v[i];
cum_vol += v[i];
result[i] = if cum_vol != 0.0 {
cum_tpv / cum_vol
} else {
f64::NAN
};
}
} else {
// Pre-compute cumulative sums for O(n) rolling window
let mut cum_tpv_arr = vec![0.0f64; n];
let mut cum_vol_arr = vec![0.0f64; n];
for i in 0..n {
let tp = (h[i] + lo[i] + c[i]) / 3.0;
let tpv = tp * v[i];
cum_tpv_arr[i] = tpv + if i > 0 { cum_tpv_arr[i - 1] } else { 0.0 };
cum_vol_arr[i] = v[i] + if i > 0 { cum_vol_arr[i - 1] } else { 0.0 };
}
for i in (timeperiod - 1)..n {
let prev_tpv = if i >= timeperiod {
cum_tpv_arr[i - timeperiod]
} else {
0.0
};
let prev_vol = if i >= timeperiod {
cum_vol_arr[i - timeperiod]
} else {
0.0
};
let w_tpv = cum_tpv_arr[i] - prev_tpv;
let w_vol = cum_vol_arr[i] - prev_vol;
result[i] = if w_vol != 0.0 {
w_tpv / w_vol
} else {
f64::NAN
};
}
}
let result = ferro_ta_core::extended::vwap(h, lo, c, v, timeperiod);
Ok(result.into_pyarray(py))
}
@@ -145,9 +43,6 @@ pub fn vwap<'py>(
// VWMA
// ---------------------------------------------------------------------------
/// Volume Weighted Moving Average.
///
/// VWMA = sum(close * volume, n) / sum(volume, n)
#[pyfunction]
#[pyo3(signature = (close, volume, timeperiod = 20))]
pub fn vwma<'py>(
@@ -159,32 +54,8 @@ pub fn vwma<'py>(
validation::validate_timeperiod(timeperiod, "timeperiod", 1)?;
let c = close.as_slice()?;
let v = volume.as_slice()?;
let n = c.len();
validation::validate_equal_length(&[(n, "close"), (v.len(), "volume")])?;
let mut cum_cv = vec![0.0f64; n];
let mut cum_v = vec![0.0f64; n];
for i in 0..n {
cum_cv[i] = c[i] * v[i] + if i > 0 { cum_cv[i - 1] } else { 0.0 };
cum_v[i] = v[i] + if i > 0 { cum_v[i - 1] } else { 0.0 };
}
let mut result = vec![f64::NAN; n];
for i in (timeperiod - 1)..n {
let prev_cv = if i >= timeperiod {
cum_cv[i - timeperiod]
} else {
0.0
};
let prev_v = if i >= timeperiod {
cum_v[i - timeperiod]
} else {
0.0
};
let w_cv = cum_cv[i] - prev_cv;
let w_v = cum_v[i] - prev_v;
result[i] = if w_v != 0.0 { w_cv / w_v } else { f64::NAN };
}
validation::validate_equal_length(&[(c.len(), "close"), (v.len(), "volume")])?;
let result = ferro_ta_core::extended::vwma(c, v, timeperiod);
Ok(result.into_pyarray(py))
}
@@ -192,10 +63,6 @@ pub fn vwma<'py>(
// SUPERTREND
// ---------------------------------------------------------------------------
/// ATR-based Supertrend indicator.
///
/// Returns (supertrend_line, direction) where direction is an int8 array:
/// 1 = uptrend, -1 = downtrend, 0 = warmup.
#[pyfunction]
#[pyo3(signature = (high, low, close, timeperiod = 7, multiplier = 3.0))]
pub fn supertrend<'py>(
@@ -210,95 +77,15 @@ pub fn supertrend<'py>(
let h = high.as_slice()?;
let lo = low.as_slice()?;
let c = close.as_slice()?;
let n = h.len();
validation::validate_equal_length(&[(n, "high"), (lo.len(), "low"), (c.len(), "close")])?;
let atr = compute_atr(h, lo, c, timeperiod);
let mut supertrend_out = vec![f64::NAN; n];
let mut direction = vec![0i8; n];
let mut upper_band = vec![f64::NAN; n];
let mut lower_band = vec![f64::NAN; n];
let first_valid = timeperiod - 1;
if first_valid >= n || atr[first_valid].is_nan() {
return Ok((supertrend_out.into_pyarray(py), direction.into_pyarray(py)));
}
// Compute basic bands
let mut upper_basic = vec![f64::NAN; n];
let mut lower_basic = vec![f64::NAN; n];
for i in 0..n {
if !atr[i].is_nan() {
let hl2 = (h[i] + lo[i]) / 2.0;
upper_basic[i] = hl2 + multiplier * atr[i];
lower_basic[i] = hl2 - multiplier * atr[i];
}
}
// Initialize band state at first valid ATR bar
upper_band[first_valid] = upper_basic[first_valid];
lower_band[first_valid] = lower_basic[first_valid];
// Keep supertrend_out NaN and direction 0 for indices 0..timeperiod (warmup)
for i in (first_valid + 1)..n {
if atr[i].is_nan() {
continue;
}
// Adjust lower band
lower_band[i] = if lower_basic[i] > lower_band[i - 1] || c[i - 1] < lower_band[i - 1] {
lower_basic[i]
} else {
lower_band[i - 1]
};
// Adjust upper band
upper_band[i] = if upper_basic[i] < upper_band[i - 1] || c[i - 1] > upper_band[i - 1] {
upper_basic[i]
} else {
upper_band[i - 1]
};
// Direction and output only from index timeperiod (warmup = 0, NaN)
if i >= timeperiod {
let prev_dir = direction[i - 1];
direction[i] = if prev_dir == 0 {
// First output bar: bootstrap with -1 (downtrend) or 1 from price vs band
if c[i] > upper_band[i] {
1
} else {
-1
}
} else if prev_dir == -1 {
if c[i] > upper_band[i] {
1
} else {
-1
}
} else if c[i] < lower_band[i] {
-1
} else {
1
};
supertrend_out[i] = if direction[i] == 1 {
lower_band[i]
} else {
upper_band[i]
};
}
}
Ok((supertrend_out.into_pyarray(py), direction.into_pyarray(py)))
validation::validate_equal_length(&[(h.len(), "high"), (lo.len(), "low"), (c.len(), "close")])?;
let (st, dir) = ferro_ta_core::extended::supertrend(h, lo, c, timeperiod, multiplier);
Ok((st.into_pyarray(py), dir.into_pyarray(py)))
}
// ---------------------------------------------------------------------------
// DONCHIAN
// ---------------------------------------------------------------------------
/// Donchian Channels — rolling highest high / lowest low.
///
/// Returns (upper, middle, lower) arrays.
#[pyfunction]
#[pyo3(signature = (high, low, timeperiod = 20))]
pub fn donchian<'py>(
@@ -314,51 +101,8 @@ pub fn donchian<'py>(
validation::validate_timeperiod(timeperiod, "timeperiod", 1)?;
let h = high.as_slice()?;
let lo = low.as_slice()?;
let n = h.len();
validation::validate_equal_length(&[(n, "high"), (lo.len(), "low")])?;
let mut upper = vec![f64::NAN; n];
let mut lower = vec![f64::NAN; n];
let mut middle = vec![f64::NAN; n];
// Use monotonic deque for O(n) sliding max / min
let mut max_dq: VecDeque<usize> = VecDeque::new();
let mut min_dq: VecDeque<usize> = VecDeque::new();
for i in 0..n {
// Remove out-of-window indices
while max_dq
.front()
.map(|&j| j + timeperiod <= i)
.unwrap_or(false)
{
max_dq.pop_front();
}
while min_dq
.front()
.map(|&j| j + timeperiod <= i)
.unwrap_or(false)
{
min_dq.pop_front();
}
// Maintain decreasing deque for max
while max_dq.back().map(|&j| h[j] <= h[i]).unwrap_or(false) {
max_dq.pop_back();
}
max_dq.push_back(i);
// Maintain increasing deque for min
while min_dq.back().map(|&j| lo[j] >= lo[i]).unwrap_or(false) {
min_dq.pop_back();
}
min_dq.push_back(i);
if i + 1 >= timeperiod {
upper[i] = h[*max_dq.front().unwrap()];
lower[i] = lo[*min_dq.front().unwrap()];
middle[i] = (upper[i] + lower[i]) / 2.0;
}
}
validation::validate_equal_length(&[(h.len(), "high"), (lo.len(), "low")])?;
let (upper, middle, lower) = ferro_ta_core::extended::donchian(h, lo, timeperiod);
Ok((
upper.into_pyarray(py),
middle.into_pyarray(py),
@@ -370,10 +114,6 @@ pub fn donchian<'py>(
// CHOPPINESS_INDEX
// ---------------------------------------------------------------------------
/// Choppiness Index — measures market choppiness vs trending.
///
/// Values near 100 → choppy; near 0 → trending.
/// Leading `timeperiod` values are NaN.
#[pyfunction]
#[pyo3(signature = (high, low, close, timeperiod = 14))]
pub fn choppiness_index<'py>(
@@ -387,73 +127,8 @@ pub fn choppiness_index<'py>(
let h = high.as_slice()?;
let lo = low.as_slice()?;
let c = close.as_slice()?;
let n = h.len();
validation::validate_equal_length(&[(n, "high"), (lo.len(), "low"), (c.len(), "close")])?;
// ATR(1) = True Range per bar
let mut tr = vec![0.0f64; n];
tr[0] = h[0] - lo[0];
for i in 1..n {
let hl = h[i] - lo[i];
let hc = (h[i] - c[i - 1]).abs();
let lc = (lo[i] - c[i - 1]).abs();
tr[i] = hl.max(hc).max(lc);
}
// Cumulative TR for rolling sum
let mut cum_tr = vec![0.0f64; n];
cum_tr[0] = tr[0];
for i in 1..n {
cum_tr[i] = cum_tr[i - 1] + tr[i];
}
let log_n = (timeperiod as f64).log10();
let mut result = vec![f64::NAN; n];
// Rolling max and min using monotonic deques
let mut max_dq: VecDeque<usize> = VecDeque::new();
let mut min_dq: VecDeque<usize> = VecDeque::new();
for i in 0..n {
while max_dq
.front()
.map(|&j| j + timeperiod <= i)
.unwrap_or(false)
{
max_dq.pop_front();
}
while min_dq
.front()
.map(|&j| j + timeperiod <= i)
.unwrap_or(false)
{
min_dq.pop_front();
}
while max_dq.back().map(|&j| h[j] <= h[i]).unwrap_or(false) {
max_dq.pop_back();
}
max_dq.push_back(i);
while min_dq.back().map(|&j| lo[j] >= lo[i]).unwrap_or(false) {
min_dq.pop_back();
}
min_dq.push_back(i);
if i + 1 > timeperiod {
let prev_cum = if i >= timeperiod {
cum_tr[i - timeperiod]
} else {
0.0
};
let sum_tr = cum_tr[i] - prev_cum;
let hh = h[*max_dq.front().unwrap()];
let ll = lo[*min_dq.front().unwrap()];
let hl_range = hh - ll;
if hl_range > 0.0 && log_n > 0.0 {
result[i] = 100.0 * (sum_tr / hl_range).log10() / log_n;
}
}
}
validation::validate_equal_length(&[(h.len(), "high"), (lo.len(), "low"), (c.len(), "close")])?;
let result = ferro_ta_core::extended::choppiness_index(h, lo, c, timeperiod);
Ok(result.into_pyarray(py))
}
@@ -461,9 +136,6 @@ pub fn choppiness_index<'py>(
// KELTNER_CHANNELS
// ---------------------------------------------------------------------------
/// Keltner Channels — EMA ± (multiplier × ATR).
///
/// Returns (upper, middle, lower) arrays.
#[pyfunction]
#[pyo3(signature = (high, low, close, timeperiod = 20, atr_period = 10, multiplier = 2.0))]
pub fn keltner_channels<'py>(
@@ -484,22 +156,9 @@ pub fn keltner_channels<'py>(
let h = high.as_slice()?;
let lo = low.as_slice()?;
let c = close.as_slice()?;
let n = h.len();
validation::validate_equal_length(&[(n, "high"), (lo.len(), "low"), (c.len(), "close")])?;
let middle = compute_ema_ta(c, timeperiod);
let atr = compute_atr(h, lo, c, atr_period);
let mut upper = vec![f64::NAN; n];
let mut lower = vec![f64::NAN; n];
for i in 0..n {
if !middle[i].is_nan() && !atr[i].is_nan() {
let band = multiplier * atr[i];
upper[i] = middle[i] + band;
lower[i] = middle[i] - band;
}
}
validation::validate_equal_length(&[(h.len(), "high"), (lo.len(), "low"), (c.len(), "close")])?;
let (upper, middle, lower) =
ferro_ta_core::extended::keltner_channels(h, lo, c, timeperiod, atr_period, multiplier);
Ok((
upper.into_pyarray(py),
middle.into_pyarray(py),
@@ -511,9 +170,6 @@ pub fn keltner_channels<'py>(
// HULL_MA
// ---------------------------------------------------------------------------
/// Hull Moving Average (HMA).
///
/// Formula: HMA(n) = WMA(2 * WMA(n/2) - WMA(n), sqrt(n))
#[pyfunction]
#[pyo3(signature = (close, timeperiod = 16))]
pub fn hull_ma<'py>(
@@ -523,43 +179,14 @@ pub fn hull_ma<'py>(
) -> PyResult<Bound<'py, PyArray1<f64>>> {
validation::validate_timeperiod(timeperiod, "timeperiod", 1)?;
let c = close.as_slice()?;
let n = c.len();
let half = (timeperiod / 2).max(1);
let sqrt_p = ((timeperiod as f64).sqrt().round() as usize).max(1);
let wma_full = compute_wma(c, timeperiod);
let wma_half = compute_wma(c, half);
// raw = 2 * wma_half - wma_full
let mut raw = vec![f64::NAN; n];
for i in 0..n {
if !wma_full[i].is_nan() && !wma_half[i].is_nan() {
raw[i] = 2.0 * wma_half[i] - wma_full[i];
}
}
// Find first valid index in raw
let first_valid = raw.iter().position(|x| !x.is_nan()).unwrap_or(n);
let mut hull = vec![f64::NAN; n];
if first_valid < n {
let raw_valid = &raw[first_valid..];
let hma_slice = compute_wma(raw_valid, sqrt_p);
for (k, &v) in hma_slice.iter().enumerate() {
hull[first_valid + k] = v;
}
}
Ok(hull.into_pyarray(py))
let result = ferro_ta_core::extended::hull_ma(c, timeperiod);
Ok(result.into_pyarray(py))
}
// ---------------------------------------------------------------------------
// CHANDELIER_EXIT
// ---------------------------------------------------------------------------
/// Chandelier Exit — ATR-based trailing stop levels.
///
/// Returns (long_exit, short_exit) arrays.
#[pyfunction]
#[pyo3(signature = (high, low, close, timeperiod = 22, multiplier = 3.0))]
pub fn chandelier_exit<'py>(
@@ -574,56 +201,9 @@ pub fn chandelier_exit<'py>(
let h = high.as_slice()?;
let lo = low.as_slice()?;
let c = close.as_slice()?;
let n = h.len();
validation::validate_equal_length(&[(n, "high"), (lo.len(), "low"), (c.len(), "close")])?;
let atr = compute_atr(h, lo, c, timeperiod);
// Rolling max/min using monotonic deques
let mut max_dq: VecDeque<usize> = VecDeque::new();
let mut min_dq: VecDeque<usize> = VecDeque::new();
let mut highest_high = vec![f64::NAN; n];
let mut lowest_low = vec![f64::NAN; n];
for i in 0..n {
while max_dq
.front()
.map(|&j| j + timeperiod <= i)
.unwrap_or(false)
{
max_dq.pop_front();
}
while min_dq
.front()
.map(|&j| j + timeperiod <= i)
.unwrap_or(false)
{
min_dq.pop_front();
}
while max_dq.back().map(|&j| h[j] <= h[i]).unwrap_or(false) {
max_dq.pop_back();
}
max_dq.push_back(i);
while min_dq.back().map(|&j| lo[j] >= lo[i]).unwrap_or(false) {
min_dq.pop_back();
}
min_dq.push_back(i);
if i + 1 >= timeperiod {
highest_high[i] = h[*max_dq.front().unwrap()];
lowest_low[i] = lo[*min_dq.front().unwrap()];
}
}
let mut long_exit = vec![f64::NAN; n];
let mut short_exit = vec![f64::NAN; n];
for i in 0..n {
if !highest_high[i].is_nan() && !atr[i].is_nan() {
long_exit[i] = highest_high[i] - multiplier * atr[i];
short_exit[i] = lowest_low[i] + multiplier * atr[i];
}
}
validation::validate_equal_length(&[(h.len(), "high"), (lo.len(), "low"), (c.len(), "close")])?;
let (long_exit, short_exit) =
ferro_ta_core::extended::chandelier_exit(h, lo, c, timeperiod, multiplier);
Ok((long_exit.into_pyarray(py), short_exit.into_pyarray(py)))
}
@@ -631,9 +211,6 @@ pub fn chandelier_exit<'py>(
// ICHIMOKU
// ---------------------------------------------------------------------------
/// Ichimoku Cloud (Ichimoku Kinko Hyo).
///
/// Returns (tenkan, kijun, senkou_a, senkou_b, chikou) arrays.
#[pyfunction]
#[pyo3(signature = (high, low, close, tenkan_period = 9, kijun_period = 26, senkou_b_period = 52, displacement = 26))]
pub fn ichimoku<'py>(
@@ -658,62 +235,16 @@ pub fn ichimoku<'py>(
let h = high.as_slice()?;
let lo = low.as_slice()?;
let c = close.as_slice()?;
let n = h.len();
validation::validate_equal_length(&[(n, "high"), (lo.len(), "low"), (c.len(), "close")])?;
// Helper: rolling (H+L)/2 using monotonic deques
let midpoint_rolling = |period: usize| -> Vec<f64> {
let mut result = vec![f64::NAN; n];
let mut max_dq: VecDeque<usize> = VecDeque::new();
let mut min_dq: VecDeque<usize> = VecDeque::new();
for i in 0..n {
while max_dq.front().map(|&j| j + period <= i).unwrap_or(false) {
max_dq.pop_front();
}
while min_dq.front().map(|&j| j + period <= i).unwrap_or(false) {
min_dq.pop_front();
}
while max_dq.back().map(|&j| h[j] <= h[i]).unwrap_or(false) {
max_dq.pop_back();
}
max_dq.push_back(i);
while min_dq.back().map(|&j| lo[j] >= lo[i]).unwrap_or(false) {
min_dq.pop_back();
}
min_dq.push_back(i);
if i + 1 >= period {
result[i] = (h[*max_dq.front().unwrap()] + lo[*min_dq.front().unwrap()]) / 2.0;
}
}
result
};
let tenkan = midpoint_rolling(tenkan_period);
let kijun = midpoint_rolling(kijun_period);
let raw_b = midpoint_rolling(senkou_b_period);
// Senkou A: (tenkan + kijun) / 2 shifted back `displacement` bars
let mut senkou_a = vec![f64::NAN; n];
if n > displacement {
for i in displacement..n {
if !tenkan[i].is_nan() && !kijun[i].is_nan() {
senkou_a[i - displacement] = (tenkan[i] + kijun[i]) / 2.0;
}
}
}
// Senkou B: raw_b shifted back `displacement` bars
let mut senkou_b = vec![f64::NAN; n];
if n > displacement {
senkou_b[..n - displacement].copy_from_slice(&raw_b[displacement..]);
}
// Chikou: close shifted forward `displacement` bars
let mut chikou = vec![f64::NAN; n];
if n > displacement {
chikou[displacement..].copy_from_slice(&c[..n - displacement]);
}
validation::validate_equal_length(&[(h.len(), "high"), (lo.len(), "low"), (c.len(), "close")])?;
let (tenkan, kijun, senkou_a, senkou_b, chikou) = ferro_ta_core::extended::ichimoku(
h,
lo,
c,
tenkan_period,
kijun_period,
senkou_b_period,
displacement,
);
Ok((
tenkan.into_pyarray(py),
kijun.into_pyarray(py),
@@ -727,10 +258,6 @@ pub fn ichimoku<'py>(
// PIVOT_POINTS
// ---------------------------------------------------------------------------
/// Pivot Points — support / resistance levels computed from previous bar.
///
/// method: "classic" | "fibonacci" | "camarilla"
/// Returns (pivot, r1, s1, r2, s2) arrays.
#[pyfunction]
#[pyo3(signature = (high, low, close, method = "classic"))]
pub fn pivot_points<'py>(
@@ -749,51 +276,17 @@ pub fn pivot_points<'py>(
let h = high.as_slice()?;
let lo = low.as_slice()?;
let c = close.as_slice()?;
let n = h.len();
validation::validate_equal_length(&[(n, "high"), (lo.len(), "low"), (c.len(), "close")])?;
let mut pivot = vec![f64::NAN; n];
let mut r1 = vec![f64::NAN; n];
let mut s1 = vec![f64::NAN; n];
let mut r2 = vec![f64::NAN; n];
let mut s2 = vec![f64::NAN; n];
validation::validate_equal_length(&[(h.len(), "high"), (lo.len(), "low"), (c.len(), "close")])?;
let method_lower = method.to_lowercase();
for i in 1..n {
let ph = h[i - 1];
let pl = lo[i - 1];
let pc = c[i - 1];
let hl = ph - pl;
let p = (ph + pl + pc) / 3.0;
pivot[i] = p;
match method_lower.as_str() {
"classic" => {
r1[i] = 2.0 * p - pl;
s1[i] = 2.0 * p - ph;
r2[i] = p + hl;
s2[i] = p - hl;
}
"fibonacci" => {
r1[i] = p + 0.382 * hl;
s1[i] = p - 0.382 * hl;
r2[i] = p + 0.618 * hl;
s2[i] = p - 0.618 * hl;
}
"camarilla" => {
r1[i] = pc + 1.1 * hl / 12.0;
s1[i] = pc - 1.1 * hl / 12.0;
r2[i] = pc + 1.1 * hl / 6.0;
s2[i] = pc - 1.1 * hl / 6.0;
}
_ => {
return Err(PyValueError::new_err(format!(
"Unknown pivot method '{}'. Use 'classic', 'fibonacci', or 'camarilla'.",
method
)));
}
}
if !matches!(method_lower.as_str(), "classic" | "fibonacci" | "camarilla") {
return Err(PyValueError::new_err(format!(
"Unknown pivot method '{}'. Use 'classic', 'fibonacci', or 'camarilla'.",
method
)));
}
let (pivot, r1, s1, r2, s2) = ferro_ta_core::extended::pivot_points(h, lo, c, method);
Ok((
pivot.into_pyarray(py),
r1.into_pyarray(py),
+8 -129
View File
@@ -1,26 +1,10 @@
//! Rust rolling math operators — O(n) sliding window using monotonic deques.
//!
//! Functions exposed to Python:
//! rolling_sum — Rolling sum over `timeperiod` bars
//! rolling_max — Rolling maximum (O(n) via monotonic deque)
//! rolling_min — Rolling minimum (O(n) via monotonic deque)
//! rolling_maxindex — Index of rolling maximum
//! rolling_minindex — Index of rolling minimum
use std::collections::VecDeque;
//! Rolling math operators (thin PyO3 wrapper over ferro_ta_core::math_ops).
use crate::validation;
use numpy::{IntoPyArray, PyArray1, PyReadonlyArray1};
use pyo3::prelude::*;
// ---------------------------------------------------------------------------
// rolling_sum
// ---------------------------------------------------------------------------
/// Rolling sum over `timeperiod` bars.
///
/// Uses a prefix-sum array for O(n) computation.
/// Leading `timeperiod - 1` values are NaN.
#[pyfunction]
#[pyo3(signature = (real, timeperiod = 30))]
pub fn rolling_sum<'py>(
@@ -30,29 +14,11 @@ pub fn rolling_sum<'py>(
) -> PyResult<Bound<'py, PyArray1<f64>>> {
validation::validate_timeperiod(timeperiod, "timeperiod", 1)?;
let prices = real.as_slice()?;
let n = prices.len();
let mut result = vec![f64::NAN; n];
if n < timeperiod {
return Ok(result.into_pyarray(py));
}
// Prefix sum
let mut cs = vec![0.0f64; n + 1];
for i in 0..n {
cs[i + 1] = cs[i] + prices[i];
}
for i in (timeperiod - 1)..n {
result[i] = cs[i + 1] - cs[i + 1 - timeperiod];
}
let result = ferro_ta_core::math_ops::rolling_sum(prices, timeperiod);
Ok(result.into_pyarray(py))
}
// ---------------------------------------------------------------------------
// rolling_max
// ---------------------------------------------------------------------------
/// Rolling maximum over `timeperiod` bars (O(n) monotonic deque).
///
/// Leading `timeperiod - 1` values are NaN.
#[pyfunction]
#[pyo3(signature = (real, timeperiod = 30))]
pub fn rolling_max<'py>(
@@ -62,34 +28,11 @@ pub fn rolling_max<'py>(
) -> PyResult<Bound<'py, PyArray1<f64>>> {
validation::validate_timeperiod(timeperiod, "timeperiod", 1)?;
let prices = real.as_slice()?;
let n = prices.len();
let mut result = vec![f64::NAN; n];
let mut dq: VecDeque<usize> = VecDeque::new();
for i in 0..n {
// Remove indices out of the window
while dq.front().map(|&j| j + timeperiod <= i).unwrap_or(false) {
dq.pop_front();
}
// Maintain decreasing deque
while dq.back().map(|&j| prices[j] <= prices[i]).unwrap_or(false) {
dq.pop_back();
}
dq.push_back(i);
if i + 1 >= timeperiod {
result[i] = prices[*dq.front().unwrap()];
}
}
let result = ferro_ta_core::math_ops::rolling_max(prices, timeperiod);
Ok(result.into_pyarray(py))
}
// ---------------------------------------------------------------------------
// rolling_min
// ---------------------------------------------------------------------------
/// Rolling minimum over `timeperiod` bars (O(n) monotonic deque).
///
/// Leading `timeperiod - 1` values are NaN.
#[pyfunction]
#[pyo3(signature = (real, timeperiod = 30))]
pub fn rolling_min<'py>(
@@ -99,34 +42,11 @@ pub fn rolling_min<'py>(
) -> PyResult<Bound<'py, PyArray1<f64>>> {
validation::validate_timeperiod(timeperiod, "timeperiod", 1)?;
let prices = real.as_slice()?;
let n = prices.len();
let mut result = vec![f64::NAN; n];
let mut dq: VecDeque<usize> = VecDeque::new();
for i in 0..n {
while dq.front().map(|&j| j + timeperiod <= i).unwrap_or(false) {
dq.pop_front();
}
// Maintain increasing deque
while dq.back().map(|&j| prices[j] >= prices[i]).unwrap_or(false) {
dq.pop_back();
}
dq.push_back(i);
if i + 1 >= timeperiod {
result[i] = prices[*dq.front().unwrap()];
}
}
let result = ferro_ta_core::math_ops::rolling_min(prices, timeperiod);
Ok(result.into_pyarray(py))
}
// ---------------------------------------------------------------------------
// rolling_maxindex
// ---------------------------------------------------------------------------
/// Index of rolling maximum over `timeperiod` bars (O(n) monotonic deque).
///
/// Returns the 0-based index into the input array. During the warmup window
/// the value is `-1` (not valid — mask with warmup period if needed).
/// Index of rolling maximum over `timeperiod` bars.
#[pyfunction]
#[pyo3(signature = (real, timeperiod = 30))]
pub fn rolling_maxindex<'py>(
@@ -136,33 +56,11 @@ pub fn rolling_maxindex<'py>(
) -> PyResult<Bound<'py, PyArray1<i64>>> {
validation::validate_timeperiod(timeperiod, "timeperiod", 1)?;
let prices = real.as_slice()?;
let n = prices.len();
let mut result = vec![-1i64; n];
let mut dq: VecDeque<usize> = VecDeque::new();
for i in 0..n {
while dq.front().map(|&j| j + timeperiod <= i).unwrap_or(false) {
dq.pop_front();
}
while dq.back().map(|&j| prices[j] <= prices[i]).unwrap_or(false) {
dq.pop_back();
}
dq.push_back(i);
if i + 1 >= timeperiod {
result[i] = *dq.front().unwrap() as i64;
}
}
let result = ferro_ta_core::math_ops::rolling_maxindex(prices, timeperiod);
Ok(result.into_pyarray(py))
}
// ---------------------------------------------------------------------------
// rolling_minindex
// ---------------------------------------------------------------------------
/// Index of rolling minimum over `timeperiod` bars (O(n) monotonic deque).
///
/// Returns the 0-based index into the input array. During the warmup window
/// the value is `-1` (not valid — mask with warmup period if needed).
/// Index of rolling minimum over `timeperiod` bars.
#[pyfunction]
#[pyo3(signature = (real, timeperiod = 30))]
pub fn rolling_minindex<'py>(
@@ -172,29 +70,10 @@ pub fn rolling_minindex<'py>(
) -> PyResult<Bound<'py, PyArray1<i64>>> {
validation::validate_timeperiod(timeperiod, "timeperiod", 1)?;
let prices = real.as_slice()?;
let n = prices.len();
let mut result = vec![-1i64; n];
let mut dq: VecDeque<usize> = VecDeque::new();
for i in 0..n {
while dq.front().map(|&j| j + timeperiod <= i).unwrap_or(false) {
dq.pop_front();
}
while dq.back().map(|&j| prices[j] >= prices[i]).unwrap_or(false) {
dq.pop_back();
}
dq.push_back(i);
if i + 1 >= timeperiod {
result[i] = *dq.front().unwrap() as i64;
}
}
let result = ferro_ta_core::math_ops::rolling_minindex(prices, timeperiod);
Ok(result.into_pyarray(py))
}
// ---------------------------------------------------------------------------
// register
// ---------------------------------------------------------------------------
pub fn register(m: &Bound<'_, PyModule>) -> PyResult<()> {
m.add_function(pyo3::wrap_pyfunction!(rolling_sum, m)?)?;
m.add_function(pyo3::wrap_pyfunction!(rolling_max, m)?)?;
+45
View File
@@ -5,6 +5,16 @@ use crate::validation;
use numpy::{IntoPyArray, PyArray1, PyReadonlyArray1};
use pyo3::prelude::*;
/// Six-tuple of bound PyArray1 vectors (PLUS_DM, MINUS_DM, +DI, -DI, DX, ADX).
type AdxAllResult<'py> = (
Bound<'py, PyArray1<f64>>,
Bound<'py, PyArray1<f64>>,
Bound<'py, PyArray1<f64>>,
Bound<'py, PyArray1<f64>>,
Bound<'py, PyArray1<f64>>,
Bound<'py, PyArray1<f64>>,
);
/// Plus Directional Movement (Wilder smoothing).
#[pyfunction]
#[pyo3(signature = (high, low, timeperiod = 14))]
@@ -158,3 +168,38 @@ pub fn adxr<'py>(
let result = ferro_ta_core::momentum::adxr(highs, lows, closes, timeperiod);
Ok(result.into_pyarray(py))
}
/// Compute all six ADX-family outputs in a single TR/PDM/MDM pass.
///
/// Returns (plus_dm, minus_dm, plus_di, minus_di, dx, adx) — six arrays of
/// the same length as the inputs. Use this when you need more than one ADX
/// family output to avoid redundant computation.
#[pyfunction]
#[pyo3(signature = (high, low, close, timeperiod = 14))]
pub fn adx_all<'py>(
py: Python<'py>,
high: PyReadonlyArray1<'py, f64>,
low: PyReadonlyArray1<'py, f64>,
close: PyReadonlyArray1<'py, f64>,
timeperiod: usize,
) -> PyResult<AdxAllResult<'py>> {
validation::validate_timeperiod(timeperiod, "timeperiod", 1)?;
let highs = high.as_slice()?;
let lows = low.as_slice()?;
let closes = close.as_slice()?;
validation::validate_equal_length(&[
(highs.len(), "high"),
(lows.len(), "low"),
(closes.len(), "close"),
])?;
let (pdm, mdm, pdi, mdi, dx, adx) =
ferro_ta_core::momentum::adx_all(highs, lows, closes, timeperiod);
Ok((
pdm.into_pyarray(py),
mdm.into_pyarray(py),
pdi.into_pyarray(py),
mdi.into_pyarray(py),
dx.into_pyarray(py),
adx.into_pyarray(py),
))
}
+1
View File
@@ -47,6 +47,7 @@ pub fn register(m: &Bound<'_, PyModule>) -> PyResult<()> {
m.add_function(pyo3::wrap_pyfunction!(self::adx::dx, m)?)?;
m.add_function(pyo3::wrap_pyfunction!(self::adx::adx, m)?)?;
m.add_function(pyo3::wrap_pyfunction!(self::adx::adxr, m)?)?;
m.add_function(pyo3::wrap_pyfunction!(self::adx::adx_all, m)?)?;
m.add_function(pyo3::wrap_pyfunction!(self::trix::trix, m)?)?;
m.add_function(pyo3::wrap_pyfunction!(self::ultosc::ultosc, m)?)?;
Ok(())
+5 -30
View File
@@ -1,8 +1,6 @@
use numpy::{IntoPyArray, PyArray1, PyReadonlyArray1};
use pyo3::prelude::*;
use super::common::*;
#[pyfunction]
pub fn cdl2crows<'py>(
py: Python<'py>,
@@ -11,33 +9,10 @@ pub fn cdl2crows<'py>(
low: PyReadonlyArray1<'py, f64>,
close: PyReadonlyArray1<'py, f64>,
) -> PyResult<Bound<'py, PyArray1<i32>>> {
let opens = open.as_slice()?;
let highs = high.as_slice()?;
let lows = low.as_slice()?;
let closes = close.as_slice()?;
let n = opens.len();
super::common::validate_ohlc_length(n, highs.len(), lows.len(), closes.len())?;
let mut result = vec![0i32; n];
for i in 2..n {
let (o1, _h1, _l1, c1) = (opens[i - 2], highs[i - 2], lows[i - 2], closes[i - 2]);
let (o2, _h2, _l2, c2) = (opens[i - 1], highs[i - 1], lows[i - 1], closes[i - 1]);
let (o3, _h3, _l3, c3) = (opens[i], highs[i], lows[i], closes[i]);
// Two Crows:
// 1. First candle is a long white (bullish) candle
// 2. Second candle gaps up (opens above first close) and closes lower but still above first close
// 3. Third candle opens within second body and closes within first body
if is_bullish(o1, c1)
&& is_bearish(o2, c2)
&& o2 > c1 // gap up
&& c2 > c1 // second still closes above first close
&& is_bearish(o3, c3)
&& o3 < o2 && o3 > c2 // opens within second body
&& c3 > o1 && c3 < c1
// closes within first body
{
result[i] = -100;
}
}
let o = open.as_slice()?;
let h = high.as_slice()?;
let l = low.as_slice()?;
let c = close.as_slice()?;
let result = ferro_ta_core::pattern::cdl2crows(o, h, l, c);
Ok(result.into_pyarray(py))
}
+5 -54
View File
@@ -1,8 +1,6 @@
use numpy::{IntoPyArray, PyArray1, PyReadonlyArray1};
use pyo3::prelude::*;
use super::common::*;
#[pyfunction]
pub fn cdl3blackcrows<'py>(
py: Python<'py>,
@@ -11,57 +9,10 @@ pub fn cdl3blackcrows<'py>(
low: PyReadonlyArray1<'py, f64>,
close: PyReadonlyArray1<'py, f64>,
) -> PyResult<Bound<'py, PyArray1<i32>>> {
let opens = open.as_slice()?;
let highs = high.as_slice()?;
let lows = low.as_slice()?;
let closes = close.as_slice()?;
let n = opens.len();
super::common::validate_ohlc_length(n, highs.len(), lows.len(), closes.len())?;
let mut result = vec![0i32; n];
for i in 2..n {
let (o1, h1, l1, c1) = (opens[i - 2], highs[i - 2], lows[i - 2], closes[i - 2]);
let (o2, h2, l2, c2) = (opens[i - 1], highs[i - 1], lows[i - 1], closes[i - 1]);
let (o3, h3, l3, c3) = (opens[i], highs[i], lows[i], closes[i]);
let body1 = body_size(o1, c1);
let body2 = body_size(o2, c2);
let body3 = body_size(o3, c3);
let range1 = candle_range(h1, l1);
let range2 = candle_range(h2, l2);
let range3 = candle_range(h3, l3);
// All three candles must be bearish with large bodies
let long_body1 = range1 > 0.0 && body1 >= range1 * 0.6;
let long_body2 = range2 > 0.0 && body2 >= range2 * 0.5;
let long_body3 = range3 > 0.0 && body3 >= range3 * 0.5;
// Each opens within the previous candle's body and closes lower
let open2_in_body1 = o2 < o1 && o2 > c1;
let open3_in_body2 = o3 < o2 && o3 > c2;
// Small upper shadows (closes near the low)
let small_upper1 = upper_shadow(o1, h1, c1) <= body1 * 0.3;
let small_upper2 = upper_shadow(o2, h2, c2) <= body2 * 0.3;
let small_upper3 = upper_shadow(o3, h3, c3) <= body3 * 0.3;
if is_bearish(o1, c1)
&& is_bearish(o2, c2)
&& is_bearish(o3, c3)
&& long_body1
&& long_body2
&& long_body3
&& open2_in_body1
&& open3_in_body2
&& small_upper1
&& small_upper2
&& small_upper3
&& c2 < c1
&& c3 < c2
&& l3 < l2
&& l2 < l1
{
result[i] = -100;
}
}
let o = open.as_slice()?;
let h = high.as_slice()?;
let l = low.as_slice()?;
let c = close.as_slice()?;
let result = ferro_ta_core::pattern::cdl3blackcrows(o, h, l, c);
Ok(result.into_pyarray(py))
}
+5 -36
View File
@@ -1,8 +1,6 @@
use numpy::{IntoPyArray, PyArray1, PyReadonlyArray1};
use pyo3::prelude::*;
use super::common::*;
#[pyfunction]
pub fn cdl3inside<'py>(
py: Python<'py>,
@@ -11,39 +9,10 @@ pub fn cdl3inside<'py>(
low: PyReadonlyArray1<'py, f64>,
close: PyReadonlyArray1<'py, f64>,
) -> PyResult<Bound<'py, PyArray1<i32>>> {
let opens = open.as_slice()?;
let highs = high.as_slice()?;
let lows = low.as_slice()?;
let closes = close.as_slice()?;
let n = opens.len();
super::common::validate_ohlc_length(n, highs.len(), lows.len(), closes.len())?;
let mut result = vec![0i32; n];
for i in 2..n {
let (o1, h1, l1, c1) = (opens[i - 2], highs[i - 2], lows[i - 2], closes[i - 2]);
let (o2, _h2, _l2, c2) = (opens[i - 1], highs[i - 1], lows[i - 1], closes[i - 1]);
let (_o3, _h3, _l3, c3) = (opens[i], highs[i], lows[i], closes[i]);
let body1 = body_size(o1, c1);
let body2 = body_size(o2, c2);
let range1 = candle_range(h1, l1);
let large_body1 = range1 > 0.0 && body1 >= range1 * 0.5;
// Candle 2 body is inside candle 1 body (harami condition)
let body2_high = o2.max(c2);
let body2_low = o2.min(c2);
let body1_high = o1.max(c1);
let body1_low = o1.min(c1);
let inside = body2_high <= body1_high && body2_low >= body1_low && body2 < body1 * 0.5;
// Three Inside Up: C1 bearish, C2 bullish harami, C3 closes above C2 close
if is_bearish(o1, c1) && large_body1 && inside && is_bullish(o2, c2) && c3 > c2 {
result[i] = 100;
}
// Three Inside Down: C1 bullish, C2 bearish harami, C3 closes below C2 close
else if is_bullish(o1, c1) && large_body1 && inside && is_bearish(o2, c2) && c3 < c2 {
result[i] = -100;
}
}
let o = open.as_slice()?;
let h = high.as_slice()?;
let l = low.as_slice()?;
let c = close.as_slice()?;
let result = ferro_ta_core::pattern::cdl3inside(o, h, l, c);
Ok(result.into_pyarray(py))
}
+5 -36
View File
@@ -1,8 +1,6 @@
use numpy::{IntoPyArray, PyArray1, PyReadonlyArray1};
use pyo3::prelude::*;
use super::common::*;
#[pyfunction]
pub fn cdl3linestrike<'py>(
py: Python<'py>,
@@ -11,39 +9,10 @@ pub fn cdl3linestrike<'py>(
low: PyReadonlyArray1<'py, f64>,
close: PyReadonlyArray1<'py, f64>,
) -> PyResult<Bound<'py, PyArray1<i32>>> {
let opens = open.as_slice()?;
let highs = high.as_slice()?;
let lows = low.as_slice()?;
let closes = close.as_slice()?;
let n = opens.len();
super::common::validate_ohlc_length(n, highs.len(), lows.len(), closes.len())?;
let mut result = vec![0i32; n];
for i in 3..n {
let (o0, c0) = (opens[i - 3], closes[i - 3]);
let (o1, c1) = (opens[i - 2], closes[i - 2]);
let (o2, c2) = (opens[i - 1], closes[i - 1]);
let (o3, c3) = (opens[i], closes[i]);
if is_bearish(o0, c0)
&& is_bearish(o1, c1)
&& is_bearish(o2, c2)
&& c1 < c0
&& c2 < c1
&& is_bullish(o3, c3)
&& o3 < c2
&& c3 > o0
{
result[i] = 100;
} else if is_bullish(o0, c0)
&& is_bullish(o1, c1)
&& is_bullish(o2, c2)
&& c1 > c0
&& c2 > c1
&& is_bearish(o3, c3)
&& o3 > c2
&& c3 < o0
{
result[i] = -100;
}
}
let o = open.as_slice()?;
let h = high.as_slice()?;
let l = low.as_slice()?;
let c = close.as_slice()?;
let result = ferro_ta_core::pattern::cdl3linestrike(o, h, l, c);
Ok(result.into_pyarray(py))
}
+5 -31
View File
@@ -1,8 +1,6 @@
use numpy::{IntoPyArray, PyArray1, PyReadonlyArray1};
use pyo3::prelude::*;
use super::common::*;
#[pyfunction]
pub fn cdl3outside<'py>(
py: Python<'py>,
@@ -11,34 +9,10 @@ pub fn cdl3outside<'py>(
low: PyReadonlyArray1<'py, f64>,
close: PyReadonlyArray1<'py, f64>,
) -> PyResult<Bound<'py, PyArray1<i32>>> {
let opens = open.as_slice()?;
let highs = high.as_slice()?;
let lows = low.as_slice()?;
let closes = close.as_slice()?;
let n = opens.len();
super::common::validate_ohlc_length(n, highs.len(), lows.len(), closes.len())?;
let mut result = vec![0i32; n];
for i in 2..n {
let (o1, _h1, _l1, c1) = (opens[i - 2], highs[i - 2], lows[i - 2], closes[i - 2]);
let (o2, _h2, _l2, c2) = (opens[i - 1], highs[i - 1], lows[i - 1], closes[i - 1]);
let (_o3, _h3, _l3, c3) = (opens[i], highs[i], lows[i], closes[i]);
let body1_high = o1.max(c1);
let body1_low = o1.min(c1);
let body2_high = o2.max(c2);
let body2_low = o2.min(c2);
// Engulfing: candle 2 body completely covers candle 1 body
let engulfs = body2_high > body1_high && body2_low < body1_low;
// Three Outside Up: C1 bearish, C2 bullish engulfing, C3 closes above C2
if is_bearish(o1, c1) && is_bullish(o2, c2) && engulfs && c3 > c2 {
result[i] = 100;
}
// Three Outside Down: C1 bullish, C2 bearish engulfing, C3 closes below C2
else if is_bullish(o1, c1) && is_bearish(o2, c2) && engulfs && c3 < c2 {
result[i] = -100;
}
}
let o = open.as_slice()?;
let h = high.as_slice()?;
let l = low.as_slice()?;
let c = close.as_slice()?;
let result = ferro_ta_core::pattern::cdl3outside(o, h, l, c);
Ok(result.into_pyarray(py))
}
+5 -26
View File
@@ -1,8 +1,6 @@
use numpy::{IntoPyArray, PyArray1, PyReadonlyArray1};
use pyo3::prelude::*;
use super::common::*;
#[pyfunction]
pub fn cdl3starsinsouth<'py>(
py: Python<'py>,
@@ -11,29 +9,10 @@ pub fn cdl3starsinsouth<'py>(
low: PyReadonlyArray1<'py, f64>,
close: PyReadonlyArray1<'py, f64>,
) -> PyResult<Bound<'py, PyArray1<i32>>> {
let opens = open.as_slice()?;
let highs = high.as_slice()?;
let lows = low.as_slice()?;
let closes = close.as_slice()?;
let n = opens.len();
super::common::validate_ohlc_length(n, highs.len(), lows.len(), closes.len())?;
let mut result = vec![0i32; n];
for i in 2..n {
let (o0, h0, l0, c0) = (opens[i - 2], highs[i - 2], lows[i - 2], closes[i - 2]);
let (o1, h1, l1, c1) = (opens[i - 1], highs[i - 1], lows[i - 1], closes[i - 1]);
let (o2, h2, l2, c2) = (opens[i], highs[i], lows[i], closes[i]);
if is_bearish(o0, c0)
&& is_bearish(o1, c1)
&& is_bearish(o2, c2)
&& h1 <= h0
&& l1 >= l0
&& h2 <= h1
&& l2 >= l1
&& body_size(o2, c2) <= body_size(o1, c1) * 0.6
&& upper_shadow(o2, h2, c2) <= body_size(o2, c2) * 0.2
{
result[i] = 100;
}
}
let o = open.as_slice()?;
let h = high.as_slice()?;
let l = low.as_slice()?;
let c = close.as_slice()?;
let result = ferro_ta_core::pattern::cdl3starsinsouth(o, h, l, c);
Ok(result.into_pyarray(py))
}
+5 -51
View File
@@ -1,8 +1,6 @@
use numpy::{IntoPyArray, PyArray1, PyReadonlyArray1};
use pyo3::prelude::*;
use super::common::*;
#[pyfunction]
pub fn cdl3whitesoldiers<'py>(
py: Python<'py>,
@@ -11,54 +9,10 @@ pub fn cdl3whitesoldiers<'py>(
low: PyReadonlyArray1<'py, f64>,
close: PyReadonlyArray1<'py, f64>,
) -> PyResult<Bound<'py, PyArray1<i32>>> {
let opens = open.as_slice()?;
let highs = high.as_slice()?;
let lows = low.as_slice()?;
let closes = close.as_slice()?;
let n = opens.len();
super::common::validate_ohlc_length(n, highs.len(), lows.len(), closes.len())?;
let mut result = vec![0i32; n];
for i in 2..n {
let (o1, h1, l1, c1) = (opens[i - 2], highs[i - 2], lows[i - 2], closes[i - 2]);
let (o2, h2, l2, c2) = (opens[i - 1], highs[i - 1], lows[i - 1], closes[i - 1]);
let (o3, h3, l3, c3) = (opens[i], highs[i], lows[i], closes[i]);
let body1 = body_size(o1, c1);
let body2 = body_size(o2, c2);
let body3 = body_size(o3, c3);
let range1 = candle_range(h1, l1);
let range2 = candle_range(h2, l2);
let range3 = candle_range(h3, l3);
let long_body1 = range1 > 0.0 && body1 >= range1 * 0.6;
let long_body2 = range2 > 0.0 && body2 >= range2 * 0.5;
let long_body3 = range3 > 0.0 && body3 >= range3 * 0.5;
let open2_in_body1 = o2 > o1 && o2 < c1;
let open3_in_body2 = o3 > o2 && o3 < c2;
let small_lower1 = lower_shadow(o1, l1, c1) <= body1 * 0.3;
let small_lower2 = lower_shadow(o2, l2, c2) <= body2 * 0.3;
let small_lower3 = lower_shadow(o3, l3, c3) <= body3 * 0.3;
if is_bullish(o1, c1)
&& is_bullish(o2, c2)
&& is_bullish(o3, c3)
&& long_body1
&& long_body2
&& long_body3
&& open2_in_body1
&& open3_in_body2
&& small_lower1
&& small_lower2
&& small_lower3
&& c2 > c1
&& c3 > c2
&& h3 > h2
&& h2 > h1
{
result[i] = 100;
}
}
let o = open.as_slice()?;
let h = high.as_slice()?;
let l = low.as_slice()?;
let c = close.as_slice()?;
let result = ferro_ta_core::pattern::cdl3whitesoldiers(o, h, l, c);
Ok(result.into_pyarray(py))
}
+5 -44
View File
@@ -1,8 +1,6 @@
use numpy::{IntoPyArray, PyArray1, PyReadonlyArray1};
use pyo3::prelude::*;
use super::common::*;
#[pyfunction]
pub fn cdlabandonedbaby<'py>(
py: Python<'py>,
@@ -11,47 +9,10 @@ pub fn cdlabandonedbaby<'py>(
low: PyReadonlyArray1<'py, f64>,
close: PyReadonlyArray1<'py, f64>,
) -> PyResult<Bound<'py, PyArray1<i32>>> {
let opens = open.as_slice()?;
let highs = high.as_slice()?;
let lows = low.as_slice()?;
let closes = close.as_slice()?;
let n = opens.len();
super::common::validate_ohlc_length(n, highs.len(), lows.len(), closes.len())?;
let mut result = vec![0i32; n];
for i in 2..n {
let (o0, h0, l0, c0) = (opens[i - 2], highs[i - 2], lows[i - 2], closes[i - 2]);
let (o1, h1, l1, c1) = (opens[i - 1], highs[i - 1], lows[i - 1], closes[i - 1]);
let (o2, h2, l2, c2) = (opens[i], highs[i], lows[i], closes[i]);
let range0 = candle_range(h0, l0);
let range1 = candle_range(h1, l1);
let range2 = candle_range(h2, l2);
let body0 = body_size(o0, c0);
let body1 = body_size(o1, c1);
let body2 = body_size(o2, c2);
let is_doji1 = range1 > 0.0 && body1 / range1 <= 0.1;
if is_bearish(o0, c0)
&& range0 > 0.0
&& body0 >= range0 * 0.5
&& is_doji1
&& h1 < l0
&& is_bullish(o2, c2)
&& range2 > 0.0
&& body2 >= range2 * 0.5
&& l2 > h1
{
result[i] = 100;
} else if is_bullish(o0, c0)
&& range0 > 0.0
&& body0 >= range0 * 0.5
&& is_doji1
&& l1 > h0
&& is_bearish(o2, c2)
&& range2 > 0.0
&& body2 >= range2 * 0.5
&& h2 < l1
{
result[i] = -100;
}
}
let o = open.as_slice()?;
let h = high.as_slice()?;
let l = low.as_slice()?;
let c = close.as_slice()?;
let result = ferro_ta_core::pattern::cdlabandonedbaby(o, h, l, c);
Ok(result.into_pyarray(py))
}
+5 -33
View File
@@ -1,8 +1,6 @@
use numpy::{IntoPyArray, PyArray1, PyReadonlyArray1};
use pyo3::prelude::*;
use super::common::*;
#[pyfunction]
pub fn cdladvanceblock<'py>(
py: Python<'py>,
@@ -11,36 +9,10 @@ pub fn cdladvanceblock<'py>(
low: PyReadonlyArray1<'py, f64>,
close: PyReadonlyArray1<'py, f64>,
) -> PyResult<Bound<'py, PyArray1<i32>>> {
let opens = open.as_slice()?;
let highs = high.as_slice()?;
let lows = low.as_slice()?;
let closes = close.as_slice()?;
let n = opens.len();
super::common::validate_ohlc_length(n, highs.len(), lows.len(), closes.len())?;
let mut result = vec![0i32; n];
for i in 2..n {
let (o0, h0, _l0, c0) = (opens[i - 2], highs[i - 2], lows[i - 2], closes[i - 2]);
let (o1, h1, _l1, c1) = (opens[i - 1], highs[i - 1], lows[i - 1], closes[i - 1]);
let (o2, h2, _l2, c2) = (opens[i], highs[i], lows[i], closes[i]);
let body0 = body_size(o0, c0);
let body1 = body_size(o1, c1);
let body2 = body_size(o2, c2);
let us0 = upper_shadow(o0, h0, c0);
let us1 = upper_shadow(o1, h1, c1);
let us2 = upper_shadow(o2, h2, c2);
if is_bullish(o0, c0)
&& is_bullish(o1, c1)
&& is_bullish(o2, c2)
&& c1 > c0
&& c2 > c1
&& o1 >= o0
&& o1 <= c0
&& o2 >= o1
&& o2 <= c1
&& (body1 < body0 || body2 < body1 || us2 > us1 || us1 > us0)
{
result[i] = -100;
}
}
let o = open.as_slice()?;
let h = high.as_slice()?;
let l = low.as_slice()?;
let c = close.as_slice()?;
let result = ferro_ta_core::pattern::cdladvanceblock(o, h, l, c);
Ok(result.into_pyarray(py))
}
+5 -23
View File
@@ -1,8 +1,6 @@
use numpy::{IntoPyArray, PyArray1, PyReadonlyArray1};
use pyo3::prelude::*;
use super::common::*;
#[pyfunction]
pub fn cdlbelthold<'py>(
py: Python<'py>,
@@ -11,26 +9,10 @@ pub fn cdlbelthold<'py>(
low: PyReadonlyArray1<'py, f64>,
close: PyReadonlyArray1<'py, f64>,
) -> PyResult<Bound<'py, PyArray1<i32>>> {
let opens = open.as_slice()?;
let highs = high.as_slice()?;
let lows = low.as_slice()?;
let closes = close.as_slice()?;
let n = opens.len();
super::common::validate_ohlc_length(n, highs.len(), lows.len(), closes.len())?;
let mut result = vec![0i32; n];
for i in 0..n {
let (o, h, l, c) = (opens[i], highs[i], lows[i], closes[i]);
let range = candle_range(h, l);
let body = body_size(o, c);
if range == 0.0 {
continue;
}
let long_body = body >= range * 0.6;
if is_bullish(o, c) && long_body && (o - l).abs() <= range * 0.01 {
result[i] = 100;
} else if is_bearish(o, c) && long_body && (h - o).abs() <= range * 0.01 {
result[i] = -100;
}
}
let o = open.as_slice()?;
let h = high.as_slice()?;
let l = low.as_slice()?;
let c = close.as_slice()?;
let result = ferro_ta_core::pattern::cdlbelthold(o, h, l, c);
Ok(result.into_pyarray(py))
}
+5 -43
View File
@@ -1,8 +1,6 @@
use numpy::{IntoPyArray, PyArray1, PyReadonlyArray1};
use pyo3::prelude::*;
use super::common::*;
#[pyfunction]
pub fn cdlbreakaway<'py>(
py: Python<'py>,
@@ -11,46 +9,10 @@ pub fn cdlbreakaway<'py>(
low: PyReadonlyArray1<'py, f64>,
close: PyReadonlyArray1<'py, f64>,
) -> PyResult<Bound<'py, PyArray1<i32>>> {
let opens = open.as_slice()?;
let highs = high.as_slice()?;
let lows = low.as_slice()?;
let closes = close.as_slice()?;
let n = opens.len();
super::common::validate_ohlc_length(n, highs.len(), lows.len(), closes.len())?;
let mut result = vec![0i32; n];
for i in 4..n {
let (o0, h0, l0, c0) = (opens[i - 4], highs[i - 4], lows[i - 4], closes[i - 4]);
let c1 = closes[i - 3];
let c2 = closes[i - 2];
let c3 = closes[i - 1];
let _l3 = lows[i - 1];
let _h3 = highs[i - 1];
let (o4, c4) = (opens[i], closes[i]);
let range0 = candle_range(h0, l0);
let body0 = body_size(o0, c0);
if is_bearish(o0, c0)
&& range0 > 0.0
&& body0 >= range0 * 0.4
&& c1 < l0
&& c2 < c1
&& c3 < c2
&& is_bullish(o4, c4)
&& c4 > c1
&& c4 < c0
{
result[i] = 100;
} else if is_bullish(o0, c0)
&& range0 > 0.0
&& body0 >= range0 * 0.4
&& c1 > h0
&& c2 > c1
&& c3 > c2
&& is_bearish(o4, c4)
&& c4 < c1
&& c4 > c0
{
result[i] = -100;
}
}
let o = open.as_slice()?;
let h = high.as_slice()?;
let l = low.as_slice()?;
let c = close.as_slice()?;
let result = ferro_ta_core::pattern::cdlbreakaway(o, h, l, c);
Ok(result.into_pyarray(py))
}
+5 -25
View File
@@ -1,8 +1,6 @@
use numpy::{IntoPyArray, PyArray1, PyReadonlyArray1};
use pyo3::prelude::*;
use super::common::*;
#[pyfunction]
pub fn cdlclosingmarubozu<'py>(
py: Python<'py>,
@@ -11,28 +9,10 @@ pub fn cdlclosingmarubozu<'py>(
low: PyReadonlyArray1<'py, f64>,
close: PyReadonlyArray1<'py, f64>,
) -> PyResult<Bound<'py, PyArray1<i32>>> {
let opens = open.as_slice()?;
let highs = high.as_slice()?;
let lows = low.as_slice()?;
let closes = close.as_slice()?;
let n = opens.len();
super::common::validate_ohlc_length(n, highs.len(), lows.len(), closes.len())?;
let mut result = vec![0i32; n];
for i in 0..n {
let (o, h, l, c) = (opens[i], highs[i], lows[i], closes[i]);
let range = candle_range(h, l);
if range == 0.0 {
continue;
}
let body = body_size(o, c);
if body < range * 0.4 {
continue;
}
if is_bullish(o, c) && (h - c).abs() <= range * 0.01 {
result[i] = 100;
} else if is_bearish(o, c) && (c - l).abs() <= range * 0.01 {
result[i] = -100;
}
}
let o = open.as_slice()?;
let h = high.as_slice()?;
let l = low.as_slice()?;
let c = close.as_slice()?;
let result = ferro_ta_core::pattern::cdlclosingmarubozu(o, h, l, c);
Ok(result.into_pyarray(py))
}
+5 -38
View File
@@ -1,8 +1,6 @@
use numpy::{IntoPyArray, PyArray1, PyReadonlyArray1};
use pyo3::prelude::*;
use super::common::*;
#[pyfunction]
pub fn cdlconcealbabyswall<'py>(
py: Python<'py>,
@@ -11,41 +9,10 @@ pub fn cdlconcealbabyswall<'py>(
low: PyReadonlyArray1<'py, f64>,
close: PyReadonlyArray1<'py, f64>,
) -> PyResult<Bound<'py, PyArray1<i32>>> {
let opens = open.as_slice()?;
let highs = high.as_slice()?;
let lows = low.as_slice()?;
let closes = close.as_slice()?;
let n = opens.len();
super::common::validate_ohlc_length(n, highs.len(), lows.len(), closes.len())?;
let mut result = vec![0i32; n];
for i in 3..n {
let (o0, h0, l0, c0) = (opens[i - 3], highs[i - 3], lows[i - 3], closes[i - 3]);
let (o1, h1, l1, c1) = (opens[i - 2], highs[i - 2], lows[i - 2], closes[i - 2]);
let (o2, h2, l2, c2) = (opens[i - 1], highs[i - 1], lows[i - 1], closes[i - 1]);
let (o3, h3, l3, c3) = (opens[i], highs[i], lows[i], closes[i]);
let range0 = candle_range(h0, l0);
let range1 = candle_range(h1, l1);
let maru0 = range0 > 0.0
&& upper_shadow(o0, h0, c0) <= range0 * 0.02
&& lower_shadow(o0, l0, c0) <= range0 * 0.02;
let maru1 = range1 > 0.0
&& upper_shadow(o1, h1, c1) <= range1 * 0.02
&& lower_shadow(o1, l1, c1) <= range1 * 0.02;
let gap_down = o2 < c1;
let shadow_into = h2 >= c1;
let engulfs = o3 >= o2 && c3 <= c2 && h3 >= h2 && l3 <= l2;
if is_bearish(o0, c0)
&& is_bearish(o1, c1)
&& maru0
&& maru1
&& is_bearish(o2, c2)
&& gap_down
&& shadow_into
&& is_bearish(o3, c3)
&& engulfs
{
result[i] = 100;
}
}
let o = open.as_slice()?;
let h = high.as_slice()?;
let l = low.as_slice()?;
let c = close.as_slice()?;
let result = ferro_ta_core::pattern::cdlconcealbabyswall(o, h, l, c);
Ok(result.into_pyarray(py))
}

Some files were not shown because too many files have changed in this diff Show More