103 lines
5.2 KiB
Python
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)
|