chart/tests/test_aggregator.py

73 lines
2.6 KiB
Python

from datetime import datetime
from zoneinfo import ZoneInfo
from app.bars.aggregator import Aggregator
from app.bars.models import Bar, Timeframe
ET = ZoneInfo("America/New_York")
def minute(t: int, price: float = 100, volume: int = 1) -> Bar:
return Bar(Timeframe.M1, t, price, price + 1, price - 1, price + 0.5, volume, True, "ES=F", "replay")
def test_aggregates_all_timeframes_and_closes_on_later_bucket():
aggregator = Aggregator()
first = aggregator.update(minute(0, 100, 2))
second = aggregator.update(minute(60, 102, 3))
boundary = aggregator.update(minute(900, 105, 4))
assert {bar.tf for bar in first} == set(Timeframe)
forming_15m = [bar for bar in second if bar.tf is Timeframe.M15][-1]
assert (forming_15m.o, forming_15m.h, forming_15m.l, forming_15m.c, forming_15m.v) == (
100,
103,
99,
102.5,
5,
)
bars_15m = [bar for bar in boundary if bar.tf is Timeframe.M15]
assert [(bar.t, bar.closed) for bar in bars_15m] == [(0, True), (900, False)]
def test_gap_closes_previous_bucket_without_synthesizing_empty_bars():
aggregator = Aggregator([Timeframe.M5])
aggregator.update(minute(0))
output = aggregator.update(minute(1800))
assert [(bar.t, bar.closed) for bar in output] == [(0, True), (1800, False)]
def test_replay_is_deterministic():
tape = [minute(t, 100 + index) for index, t in enumerate(range(0, 3600, 60))]
def run():
aggregator = Aggregator()
return [bar.to_dict() for source in tape for bar in aggregator.update(source)]
assert run() == run()
def test_closed_yahoo_hours_form_the_right_cme_daily_bar_across_1800_et():
def hour(local_time: str, o: float, h: float, low: float, c: float, volume: int) -> Bar:
t = int(datetime.fromisoformat(local_time).replace(tzinfo=ET).timestamp())
return Bar(Timeframe.H1, t, o, h, low, c, volume, True, "ES=F", "yahoo")
aggregator = Aggregator([Timeframe.H1, Timeframe.D1])
tape = [
hour("2026-01-12T17:00:00", 90, 94, 89, 93, 5),
hour("2026-01-12T18:00:00", 100, 104, 98, 103, 10),
hour("2026-01-13T17:00:00", 103, 110, 97, 108, 20),
hour("2026-01-13T18:00:00", 120, 125, 119, 124, 40),
]
emitted = [bar for source in tape for bar in aggregator.update(source)]
session_start = int(datetime(2026, 1, 12, 18, tzinfo=ET).timestamp())
daily = next(
bar for bar in emitted
if bar.tf is Timeframe.D1 and bar.t == session_start and bar.closed
)
assert (daily.o, daily.h, daily.l, daily.c, daily.v) == (100, 110, 97, 108, 30)
assert (daily.symbol, daily.source) == ("ES=F", "yahoo")