"""Phase 2 acceptance: adapter contract, checkpoint resume, calibration maths. The real-model smoke test is marked `slow` and skipped unless `chronos` is importable, so the default suite stays fast and offline. """ from __future__ import annotations import json import sys from pathlib import Path import numpy as np import pandas as pd import pytest sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) from src import config from src.adapters import ( AdapterError, Forecast, ModelNotAllowed, PlaceholderAdapter, build_windows, get_adapter, validate_model_id, ) from src.metrics import calibration_coverage, calibration_error, directional_accuracy from src.store import SignalStore, validate_signal_frame from scripts.seed_store import Checkpoint, SeedTarget, plan_v1, seed_target def series(n=600, seed=0, start="2023-01-01"): rng = np.random.default_rng(seed) return pd.Series( 100 * np.exp(np.cumsum(rng.normal(0.0005, 0.02, n))), index=pd.date_range(start, periods=n, freq="D", tz="UTC"), ) def price_frame(n=600, seed=0): close = series(n, seed) open_ = close.shift(1).fillna(close.iloc[0] * 0.999) return pd.DataFrame({ "open": open_, "high": pd.concat([open_, close], axis=1).max(axis=1) * 1.004, "low": pd.concat([open_, close], axis=1).min(axis=1) * 0.996, "close": close, "volume": 1000.0, "source": "synthetic", }) # -------------------------------------------------------------------------- # Allow-list: no arbitrary code execution # -------------------------------------------------------------------------- def test_unknown_adapter_family_is_refused(): with pytest.raises(ModelNotAllowed, match="not allowed"): get_adapter("evil-custom", "someone/backdoor") def test_allowed_families_are_exactly_the_configured_set(): assert set(config.ALLOWED_ADAPTER_FAMILIES) == {"chronos", "timesfm", "baseline"} @pytest.mark.parametrize("bad", [ "", "no-slash", "a/b/c", "../../etc/passwd", "owner/../name", "owner/name;rm -rf /", "owner/na me", "owner/$(whoami)", ]) def test_malformed_model_ids_are_rejected(bad): with pytest.raises(AdapterError): validate_model_id(bad) @pytest.mark.parametrize("good", [ "amazon/chronos-bolt-small", "google/timesfm-2.0-500m-pytorch", "org/model_v1.2", ]) def test_well_formed_model_ids_pass(good): assert validate_model_id(good) == good # -------------------------------------------------------------------------- # Windowing is causal # -------------------------------------------------------------------------- def test_window_ends_at_its_own_timestamp(): s = series(n=100) stamps, wins = build_windows(s, context_len=30) assert wins.shape == (len(stamps), 30) # The window stored at `t` must end with the value observed at `t`. for i, ts in enumerate(list(stamps)[:5]): assert wins[i, -1] == pytest.approx(float(s.loc[ts])) def test_windows_never_include_future_values(): s = series(n=200) stamps, wins = build_windows(s, context_len=50) pos = {ts: i for i, ts in enumerate(s.index)} for i, ts in enumerate(list(stamps)[:10]): expected = s.to_numpy()[pos[ts] - 49: pos[ts] + 1] assert np.allclose(wins[i], expected) def test_too_short_series_yields_no_windows(): stamps, wins = build_windows(series(n=20), context_len=50) assert len(stamps) == 0 and wins.shape[0] == 0 # -------------------------------------------------------------------------- # Placeholder adapter # -------------------------------------------------------------------------- def test_placeholder_is_labelled_and_deterministic(): a = PlaceholderAdapter("synthetic/placeholder").load() assert a.inference_version() == config.PLACEHOLDER_VERSION _, wins = build_windows(series(n=300), 100) first, second = a.predict(wins[:20]), a.predict(wins[:20]) assert np.allclose(first.q50, second.q50) assert np.allclose(first.q10, second.q10) def test_placeholder_output_is_schema_valid(): a = PlaceholderAdapter("synthetic/placeholder").load() stamps, wins = build_windows(series(n=300), 100) df = a.predict(wins[:50]).as_frame(stamps[:50], a.inference_version()) out = validate_signal_frame(df) assert len(out) == 50 assert (out["inference_version"] == config.PLACEHOLDER_VERSION).all() assert ((out["q10"] <= out["q50"]) & (out["q50"] <= out["q90"])).all() def test_forecast_frame_sorts_crossed_quantiles(): f = Forecast(q10=np.array([5.0]), q50=np.array([1.0]), q90=np.array([3.0]), context_len=10) df = f.as_frame(pd.DatetimeIndex(["2024-01-01"], tz="UTC"), "v1") assert df["q10"].iloc[0] <= df["q50"].iloc[0] <= df["q90"].iloc[0] def test_forecast_rejects_mismatched_timestamp_count(): f = Forecast(q10=np.zeros(3), q50=np.zeros(3), q90=np.zeros(3), context_len=10) with pytest.raises(AdapterError, match="!="): f.as_frame(pd.DatetimeIndex(["2024-01-01"], tz="UTC"), "v1") def test_inference_version_pins_model_and_revision(): a = get_adapter("chronos", "amazon/chronos-bolt-small", revision="abc123") a._resolved_revision = "abc123" b = get_adapter("chronos", "amazon/chronos-bolt-small", revision="def456") b._resolved_revision = "def456" assert a.inference_version() != b.inference_version() assert a.inference_version() == a.inference_version() # -------------------------------------------------------------------------- # Checkpoint resume # -------------------------------------------------------------------------- def test_checkpoint_round_trips(tmp_path): p = tmp_path / "ckpt.json" c = Checkpoint.load(p) c.mark("m|BTC-USD|1d", pd.Timestamp("2024-06-01", tz="UTC")) again = Checkpoint.load(p) assert again.last_ts("m|BTC-USD|1d") == pd.Timestamp("2024-06-01", tz="UTC") def test_corrupt_checkpoint_starts_fresh_instead_of_crashing(tmp_path): p = tmp_path / "ckpt.json" p.write_text("{not json") assert Checkpoint.load(p).done == {} def test_checkpoint_records_failures_separately(tmp_path): c = Checkpoint.load(tmp_path / "c.json") c.mark_failed("m|ETH-USD|1h", "boom") assert Checkpoint.load(tmp_path / "c.json").failed["m|ETH-USD|1h"] == "boom" def test_seed_resumes_from_checkpoint_and_skips_finished_work(tmp_path, monkeypatch): """An interrupted seed must not redo inference it already paid for.""" store = SignalStore(repo_id=None, local_root=tmp_path / "store", offline=True) px = price_frame(n=400) store.write_prices("BTC-USD", "1d", px.reset_index(names="ts")) monkeypatch.setitem(config.SEED_MODELS, "test-model", config.ModelSpec( slug="test-model", model_id="test/model", family="placeholder", display="Test", context_len=100, )) target = SeedTarget("test-model", "BTC-USD", "1d", 3.0, placeholder=True) ckpt = Checkpoint.load(tmp_path / "ckpt.json") msg = seed_target(store, target, ckpt, batch_size=64, force_placeholder=True) assert msg.startswith("OK") first_rows = len(store.get_signals("test-model", "BTC-USD", "1d")) assert first_rows > 0 assert ckpt.last_ts(target.key) is not None # Second call: the checkpoint says complete, so nothing more is computed. msg2 = seed_target(store, target, ckpt, batch_size=64, force_placeholder=True) assert "SKIP" in msg2 assert len(store.get_signals("test-model", "BTC-USD", "1d")) == first_rows def test_seed_skips_targets_already_in_the_manifest(tmp_path, monkeypatch): store = SignalStore(repo_id=None, local_root=tmp_path / "store", offline=True) store.write_prices("BTC-USD", "1d", price_frame(n=400).reset_index(names="ts")) monkeypatch.setitem(config.SEED_MODELS, "test-model", config.ModelSpec( slug="test-model", model_id="test/model", family="placeholder", display="Test", context_len=100, )) target = SeedTarget("test-model", "BTC-USD", "1d", 3.0, placeholder=True) seed_target(store, target, Checkpoint.load(tmp_path / "a.json"), batch_size=64, force_placeholder=True) # Fresh checkpoint, but the manifest already covers the range. msg = seed_target(store, target, Checkpoint.load(tmp_path / "b.json"), batch_size=64, force_placeholder=True) assert "already covered" in msg def test_seed_skips_when_there_is_no_price_coverage(tmp_path, monkeypatch): store = SignalStore(repo_id=None, local_root=tmp_path / "store", offline=True) monkeypatch.setitem(config.SEED_MODELS, "test-model", config.ModelSpec( slug="test-model", model_id="test/model", family="placeholder", display="Test", context_len=100, )) msg = seed_target(store, SeedTarget("test-model", "ETH-USD", "1d", 1.0), Checkpoint.load(tmp_path / "c.json"), force_placeholder=True) assert "no price coverage" in msg def test_v1_plan_covers_the_specified_universe(): targets = plan_v1() assets = {t.asset for t in targets} assert {"BTC-USD", "ETH-USD", "SOL-USD"} <= assets assert {"SPY", "QQQ", "NVDA"} <= assets assert {t.timeframe for t in targets} == {"1d", "1h", "15m"} # The v1 seed is entirely real: batched inference was cheap enough that no # slice needs a synthetic placeholder. assert not any(t.placeholder for t in targets) assert len({t.model_slug for t in targets}) >= 2 # -------------------------------------------------------------------------- # Calibration maths against a synthetic series of KNOWN coverage # -------------------------------------------------------------------------- def test_calibration_recovers_a_known_80_percent_coverage(): """Construct a series where exactly 80% of actuals sit inside the band.""" n = 1000 idx = pd.DatetimeIndex(pd.date_range("2024-01-01", periods=n, freq="D", tz="UTC")) lower = pd.Series(90.0, index=idx) upper = pd.Series(110.0, index=idx) actual = pd.Series(100.0, index=idx) # inside actual.iloc[:200] = 500.0 # 20% outside, by construction cov = calibration_coverage(actual, lower, upper) assert cov == pytest.approx(0.80, abs=1e-12) assert calibration_error(cov, 0.80) == pytest.approx(0.0, abs=1e-12) @pytest.mark.parametrize("frac", [0.0, 0.25, 0.5, 0.9, 1.0]) def test_calibration_recovers_any_known_coverage(frac): n = 200 idx = pd.date_range("2024-01-01", periods=n, freq="D", tz="UTC") lower, upper = pd.Series(0.0, index=idx), pd.Series(1.0, index=idx) actual = pd.Series(5.0, index=idx) inside = int(round(n * frac)) actual.iloc[:inside] = 0.5 assert calibration_coverage(actual, lower, upper) == pytest.approx(frac, abs=1e-12) def test_overconfident_band_reads_as_undercoverage(): n = 500 idx = pd.date_range("2024-01-01", periods=n, freq="D", tz="UTC") rng = np.random.default_rng(3) actual = pd.Series(rng.normal(0, 1, n), index=idx) # A band far narrower than the true spread must score well under 0.80. lower, upper = pd.Series(-0.05, index=idx), pd.Series(0.05, index=idx) cov = calibration_coverage(actual, lower, upper) assert cov < 0.20 assert calibration_error(cov, 0.80) < -0.5 def test_calibration_of_a_correctly_specified_normal_band(): """A true 10th/90th percentile band on normal data covers ~80%.""" n = 20_000 idx = pd.date_range("2000-01-01", periods=n, freq="D", tz="UTC") rng = np.random.default_rng(11) actual = pd.Series(rng.normal(0, 1, n), index=idx) lower = pd.Series(-1.2815515655446004, index=idx) upper = pd.Series(1.2815515655446004, index=idx) assert calibration_coverage(actual, lower, upper) == pytest.approx(0.80, abs=0.02) def test_directional_accuracy_is_perfect_for_a_perfect_forecast(): idx = pd.date_range("2024-01-01", periods=100, freq="D", tz="UTC") ref = pd.Series(np.linspace(100, 200, 100), index=idx) actual_next = ref.shift(-1).ffill() assert directional_accuracy(actual_next, actual_next, ref) == pytest.approx(1.0) def test_directional_accuracy_is_zero_for_a_perfectly_wrong_forecast(): idx = pd.date_range("2024-01-01", periods=100, freq="D", tz="UTC") ref = pd.Series(np.linspace(100, 200, 100), index=idx) actual_next = ref.shift(-1).ffill() inverted = ref - (actual_next - ref) assert directional_accuracy(actual_next, inverted, ref) == pytest.approx(0.0) # -------------------------------------------------------------------------- # Real-model smoke test # -------------------------------------------------------------------------- def _chronos_available() -> bool: try: import chronos # noqa: F401 return True except Exception: return False @pytest.mark.slow @pytest.mark.skipif(not _chronos_available(), reason="chronos-forecasting not installed") def test_chronos_100_step_run_is_schema_valid(): """Phase 2 acceptance: 100 steps on the real model, output schema-valid.""" a = get_adapter("chronos", "amazon/chronos-bolt-small", context_len=256).load() assert a.resolved_revision not in ("", "unpinned") stamps, wins = build_windows(series(n=600), 256) stamps, wins = stamps[:100], wins[:100] df = a.predict(wins).as_frame(stamps, a.inference_version()) out = validate_signal_frame(df) assert len(out) == 100 assert ((out["q10"] <= out["q50"]) & (out["q50"] <= out["q90"])).all() assert (out["context_len"] == 256).all() assert out["inference_version"].nunique() == 1 assert out["inference_version"].iloc[0] != config.PLACEHOLDER_VERSION assert np.isfinite(out[["q10", "q50", "q90"]].to_numpy()).all() # -------------------------------------------------------------------------- # Chronos-2 # -------------------------------------------------------------------------- # # `predict_quantiles` returns a different shape per pipeline, and the # difference is not cosmetic: # # Bolt / T5 one stacked tensor, (batch, horizon, quantiles) # Chronos-2 a LIST of per-item tensors, (n_variates, horizon, quantiles) # # Calling .float() on the list raises AttributeError, which is what # amazon/chronos-2 did before `_to_array` existed. def test_bolt_style_tensor_is_normalised(): import numpy as np from src.adapters import ChronosAdapter stacked = np.zeros((4, 1, 3), dtype="float32") stacked[:, 0, 1] = 5.0 # q50 for every row out = ChronosAdapter._to_array(stacked) assert out.shape == (4, 3) assert np.allclose(out[:, 1], 5.0) def test_chronos2_list_of_tensors_is_normalised(): import numpy as np from src.adapters import ChronosAdapter # Four items, each (n_variates=1, horizon=1, quantiles=3). listed = [np.array([[[1.0, 2.0, 3.0]]], dtype="float32") for _ in range(4)] out = ChronosAdapter._to_array(listed) assert out.shape == (4, 3), "Chronos-2's list shape was not handled" assert np.allclose(out[0], [1.0, 2.0, 3.0]) def test_an_unexpected_result_type_raises_a_named_error(): from src.adapters import AdapterError, ChronosAdapter with pytest.raises(AdapterError): ChronosAdapter._to_array({"not": "a tensor"}) @pytest.mark.parametrize("model_id,is_c2,chunk_attr", [ ("amazon/chronos-2", True, "CHRONOS2_CHUNK"), ("amazon/chronos-bolt-base", False, "BOLT_CHUNK"), ("amazon/chronos-t5-small", False, "T5_CHUNK"), ]) def test_chunk_size_per_variant(model_id, is_c2, chunk_attr): from src.adapters import ChronosAdapter chunk = getattr(ChronosAdapter, chunk_attr) adapter = ChronosAdapter.__new__(ChronosAdapter) adapter.model_id = model_id assert adapter._is_chronos2 is is_c2 assert adapter.chunk_size == chunk def test_chronos2_does_not_ask_for_samples(): """Only the original T5 Chronos samples paths; asking Chronos-2 for `num_samples` is an unexpected keyword.""" from src.adapters import ChronosAdapter adapter = ChronosAdapter.__new__(ChronosAdapter) adapter.model_id = "amazon/chronos-2" assert adapter._is_bolt is False and adapter._is_chronos2 is True