Files
SET100-Trading-system/tests/test_data_universe.py
2026-09-29 14:08:43 +07:00

103 lines
5.2 KiB
Python

import unittest
from pathlib import Path
from unittest.mock import patch
import pandas as pd
from primethai.cli import _accurate_context_bars, _read_bars
from primethai.data import SymbolDataError, YFinanceProvider, normalize_bars, read_siamchart_archives, yahoo_symbol
from primethai.universe import MembershipError, MembershipPeriod, PointInTimeUniverse
class DataUniverseTests(unittest.TestCase):
def test_yahoo_symbol_mapping_and_normalization_preserve_adjustments_and_actions(self):
self.assertEqual(yahoo_symbol("ADVANC"), "ADVANC.BK")
self.assertEqual(yahoo_symbol("^SET"), "^SET.BK")
idx = pd.DatetimeIndex(["2024-01-02", "2024-01-03"], tz="UTC")
raw = pd.DataFrame({"Open": [10, 5], "High": [11, 5.5], "Low": [9, 4.5],
"Close": [10, 5], "Adj Close": [9.8, 5], "Volume": [100, 200],
"Dividends": [0, .5], "Stock Splits": [0, 2]}, index=idx)
out = normalize_bars(raw)
self.assertIsNone(out.index.tz)
self.assertEqual(out.iloc[1]["adj_close"], 5)
self.assertEqual(out.iloc[1]["adj_open"], 5)
self.assertEqual(out.iloc[1]["dividends"], .5)
self.assertEqual(out.iloc[1]["stock_splits"], 2)
def test_missing_or_empty_provider_data_is_reported_safely(self):
with self.assertRaisesRegex(SymbolDataError, "no bars"):
normalize_bars(pd.DataFrame())
with self.assertRaisesRegex(SymbolDataError, "missing columns"):
normalize_bars(pd.DataFrame({"Close": [1]}, index=pd.date_range("2024-01-01", periods=1)))
def test_local_cache_avoids_second_yahoo_request(self):
root = Path(__file__).resolve().parent / "fixtures" / "cache"
provider = YFinanceProvider(root, retries=0)
result = provider.get_history("ADVANC", start="2024-01-02", end="2024-01-04")
self.assertEqual(len(result), 2)
self.assertEqual(result.iloc[-1]["close"], 11)
def test_partial_set_benchmark_history_is_refused(self):
class FakeYFinance:
@staticmethod
def download(*args, **kwargs):
index = pd.DatetimeIndex(["2026-09-25"])
return pd.DataFrame({"Open": [1603.58], "High": [1613.45],
"Low": [1603.46], "Close": [1607.63],
"Volume": [0]}, index=index)
root = Path(__file__).resolve().parent / "fixtures" / "cache"
provider = YFinanceProvider(root, retries=0)
with patch.dict("sys.modules", {"yfinance": FakeYFinance}):
with self.assertRaisesRegex(SymbolDataError, "incomplete SET benchmark"):
provider.get_history("^SET", start="2024-01-01", end="2024-02-01", refresh=True)
def test_manual_set_csv_accepts_exchange_style_headers_and_missing_volume(self):
directory = Path(__file__).resolve().parent / "fixtures" / "set_index_csv"
bars = _read_bars(directory, "^SET")
self.assertEqual(len(bars), 2)
self.assertEqual(bars.iloc[-1]["close"], 1607.63)
self.assertEqual(bars["volume"].sum(), 0)
def test_siamchart_zip_import_selects_symbols_and_reports_bad_daily_files(self):
archive_path = Path(__file__).resolve().parent / "fixtures" / "siamchart_eod_test.zip"
imported, issues = read_siamchart_archives([archive_path], ["^SET", "ADVANC"])
self.assertEqual(imported["^SET"].iloc[0]["close"], 1405)
self.assertEqual(imported["ADVANC"].iloc[0]["close"], 201)
self.assertEqual(imported["^SET"].iloc[0]["dividends"], 0)
self.assertEqual(len(issues), 1)
self.assertIn("set-history_EOD_2024-01-03.csv", issues[0])
def test_point_in_time_membership_and_coverage(self):
path = Path(__file__).resolve().parent / "fixtures" / "members.csv"
universe = PointInTimeUniverse.from_csv(path)
self.assertEqual(universe.members("2024-03-01"), ["A"])
self.assertEqual(universe.members("2024-08-01"), ["B"])
self.assertEqual(universe.sector("A", "2024-03-01"), "Banking")
self.assertFalse(universe.covers(pd.to_datetime(["2023-12-29", "2024-02-01"])))
with self.assertRaisesRegex(MembershipError, "does not cover"):
universe.require_coverage(pd.to_datetime(["2024-01-02", "2025-01-02"]))
def test_accurate_context_clips_full_index_to_pit_end_and_keeps_warmup(self):
dates = pd.bdate_range("2024-01-02", periods=300)
pit_end = dates[279]
membership = PointInTimeUniverse([
MembershipPeriod("A", dates[0], pit_end),
])
set_bars = pd.DataFrame({"close": range(len(dates))}, index=dates)
clipped = _accurate_context_bars(
set_bars, membership, dates[250].date().isoformat(),
(pit_end + pd.Timedelta(days=1)).date().isoformat(),
)
self.assertEqual(clipped.index.min(), dates[0])
self.assertEqual(clipped.index.max(), pit_end)
self.assertEqual(len(clipped), 280)
def test_overlapping_membership_intervals_are_rejected(self):
path = Path(__file__).resolve().parent / "fixtures" / "overlap.csv"
with self.assertRaisesRegex(MembershipError, "overlapping"):
PointInTimeUniverse.from_csv(path)