273 lines
13 KiB
Python
273 lines
13 KiB
Python
"""Provider abstraction and defensive local CSV cache for daily market bars."""
|
|
|
|
from __future__ import annotations
|
|
|
|
from abc import ABC, abstractmethod
|
|
import csv
|
|
from datetime import datetime
|
|
import io
|
|
from pathlib import Path
|
|
import re
|
|
import time
|
|
from datetime import timedelta
|
|
import zipfile
|
|
import zlib
|
|
|
|
import pandas as pd
|
|
|
|
|
|
class SymbolDataError(RuntimeError):
|
|
pass
|
|
|
|
|
|
class _IncompleteHistoryError(SymbolDataError):
|
|
pass
|
|
|
|
|
|
def yahoo_symbol(symbol: str) -> str:
|
|
symbol = symbol.strip().upper()
|
|
if symbol in {"^SET", "^SET.BK"}:
|
|
return "^SET.BK"
|
|
if symbol.startswith("^") or "." in symbol:
|
|
return symbol
|
|
return f"{symbol}.BK"
|
|
|
|
|
|
def safe_filename(symbol: str) -> str:
|
|
return re.sub(r"[^A-Za-z0-9._-]", "_", symbol.upper())
|
|
|
|
|
|
def normalize_bars(frame: pd.DataFrame) -> pd.DataFrame:
|
|
if frame is None or frame.empty:
|
|
raise SymbolDataError("provider returned no bars")
|
|
out = frame.copy()
|
|
if isinstance(out.columns, pd.MultiIndex):
|
|
out.columns = [column[0] for column in out.columns]
|
|
out.columns = [str(column).strip().lower().replace(" ", "_") for column in out.columns]
|
|
required = {"open", "high", "low", "close", "volume"}
|
|
if not required.issubset(out.columns):
|
|
missing = required - set(out.columns)
|
|
raise SymbolDataError(f"bar data missing columns: {', '.join(sorted(missing))}")
|
|
if "adj_close" not in out:
|
|
out["adj_close"] = out["close"]
|
|
if "dividends" not in out:
|
|
out["dividends"] = 0.0
|
|
if "stock_splits" not in out:
|
|
out["stock_splits"] = 0.0
|
|
out.index = pd.to_datetime(out.index).tz_localize(None).normalize()
|
|
out.index.name = "date"
|
|
for column in ["open", "high", "low", "close", "adj_close", "volume", "dividends", "stock_splits"]:
|
|
out[column] = pd.to_numeric(out[column], errors="coerce")
|
|
out = out[~out.index.duplicated(keep="last")].sort_index()
|
|
valid = out[["open", "high", "low", "close"]].notna().all(axis=1)
|
|
valid &= (out[["open", "high", "low", "close"]] > 0).all(axis=1)
|
|
out = out.loc[valid]
|
|
if out.empty:
|
|
raise SymbolDataError("provider returned no valid OHLC bars")
|
|
out["value_turnover"] = out["close"] * out["volume"]
|
|
# Yahoo Adj Close is its split/dividend-adjusted series. Keep a corresponding
|
|
# adjusted OHLC proxy alongside the unadjusted-by-dividend fields; the simulator
|
|
# deliberately uses the latter and credits cash dividends separately.
|
|
factor = out["adj_close"] / out["close"].where(out["close"] != 0)
|
|
for column in ("open", "high", "low"):
|
|
out[f"adj_{column}"] = out[column] * factor
|
|
return out
|
|
|
|
|
|
def read_siamchart_archives(archives: list[str | Path], symbols: list[str]) -> tuple[dict[str, pd.DataFrame], list[str]]:
|
|
"""Read selected securities from SiamChart's daily EOD ZIP archives.
|
|
|
|
SiamChart stores one CSV per trading date with columns such as ``<TICKER>``
|
|
and ``<DTYYYYMMDD>``. This importer keeps only explicitly requested symbols,
|
|
maps the index ticker ``SET`` to ``^SET``, and skips unreadable daily files
|
|
while returning a warning for each skipped entry.
|
|
"""
|
|
tickers: dict[str, str] = {}
|
|
for requested in symbols:
|
|
symbol = requested.strip().upper()
|
|
if symbol in {"SET", "^SET", "^SET.BK"}:
|
|
symbol, ticker = "^SET", "SET"
|
|
else:
|
|
ticker = symbol
|
|
if not symbol:
|
|
continue
|
|
tickers[ticker] = symbol
|
|
if not tickers:
|
|
raise SymbolDataError("select at least one SiamChart ticker to import")
|
|
|
|
records: dict[str, dict[str, dict[str, float | str]]] = {
|
|
symbol: {} for symbol in tickers.values()
|
|
}
|
|
issues: list[str] = []
|
|
dated_files = 0
|
|
filename_pattern = re.compile(r"set-history_EOD_\d{4}-\d{2}-\d{2}\.csv$", re.IGNORECASE)
|
|
required_columns = {"<TICKER>", "<DTYYYYMMDD>", "<OPEN>", "<HIGH>", "<LOW>", "<CLOSE>", "<VOL>"}
|
|
|
|
for archive_name in archives:
|
|
archive_path = Path(archive_name)
|
|
if not archive_path.is_file():
|
|
raise SymbolDataError(f"SiamChart archive not found: {archive_path}")
|
|
try:
|
|
archive = zipfile.ZipFile(archive_path)
|
|
except (OSError, zipfile.BadZipFile) as exc:
|
|
raise SymbolDataError(f"cannot open SiamChart archive {archive_path}: {exc}") from exc
|
|
with archive:
|
|
for entry in archive.infolist():
|
|
if not filename_pattern.search(Path(entry.filename).name):
|
|
continue
|
|
dated_files += 1
|
|
daily_records: dict[tuple[str, str], dict[str, float | str]] = {}
|
|
try:
|
|
with archive.open(entry) as binary, io.TextIOWrapper(binary, encoding="utf-8-sig", newline="") as stream:
|
|
reader = csv.DictReader(stream)
|
|
columns = set(reader.fieldnames or [])
|
|
if not required_columns.issubset(columns):
|
|
raise ValueError("unexpected CSV columns")
|
|
for row in reader:
|
|
ticker = (row.get("<TICKER>") or "").strip().upper()
|
|
symbol = tickers.get(ticker)
|
|
if symbol is None:
|
|
continue
|
|
raw_date = (row.get("<DTYYYYMMDD>") or "").strip()
|
|
day = datetime.strptime(raw_date, "%Y%m%d").date().isoformat()
|
|
values = {
|
|
"open": float(row["<OPEN>"]),
|
|
"high": float(row["<HIGH>"]),
|
|
"low": float(row["<LOW>"]),
|
|
"close": float(row["<CLOSE>"]),
|
|
"volume": float(row["<VOL>"]),
|
|
}
|
|
if not all(pd.notna(value) for value in values.values()):
|
|
raise ValueError(f"invalid OHLCV for {ticker} on {raw_date}")
|
|
daily_records[(symbol, day)] = {"date": day, **values}
|
|
except (OSError, EOFError, UnicodeDecodeError, csv.Error, ValueError, zipfile.BadZipFile, zlib.error) as exc:
|
|
issues.append(f"{archive_path.name}:{entry.filename}: {exc}")
|
|
continue
|
|
for (symbol, day), record in daily_records.items():
|
|
records[symbol][day] = record
|
|
|
|
if dated_files == 0:
|
|
raise SymbolDataError("no daily set-history_EOD_YYYY-MM-DD.csv files found in the supplied archives")
|
|
imported: dict[str, pd.DataFrame] = {}
|
|
for symbol, by_date in records.items():
|
|
if not by_date:
|
|
raise SymbolDataError(f"no {symbol} rows found in the supplied SiamChart archives")
|
|
frame = pd.DataFrame(sorted(by_date.values(), key=lambda row: row["date"]))
|
|
imported[symbol] = normalize_bars(frame.set_index("date"))
|
|
return imported, issues
|
|
|
|
|
|
class MarketDataProvider(ABC):
|
|
@abstractmethod
|
|
def get_history(self, symbol: str, *, start: str | None = None, end: str | None = None,
|
|
period: str | None = None, refresh: bool = False) -> pd.DataFrame:
|
|
"""Return normalized daily bars, or raise SymbolDataError."""
|
|
|
|
|
|
class YFinanceProvider(MarketDataProvider):
|
|
"""Yahoo daily data. auto_adjust=False and corporate-action columns are explicit."""
|
|
|
|
def __init__(self, cache_dir: str | Path = "data/cache", *, repair: bool = True,
|
|
retries: int = 2, retry_delay: float = 1.0) -> None:
|
|
self.cache_dir, self.repair = Path(cache_dir), repair
|
|
self.retries, self.retry_delay = retries, retry_delay
|
|
|
|
def _cache_path(self, symbol: str) -> Path:
|
|
return self.cache_dir / f"{safe_filename(yahoo_symbol(symbol))}.csv"
|
|
|
|
@staticmethod
|
|
def _period_start(period: str | None) -> pd.Timestamp | None:
|
|
if not period:
|
|
return None
|
|
now = pd.Timestamp.now().normalize()
|
|
if period == "ytd":
|
|
return pd.Timestamp(year=now.year, month=1, day=1)
|
|
if period == "max":
|
|
return pd.Timestamp("1980-01-01")
|
|
match = re.fullmatch(r"(\d+)(d|wk|mo|y)", period)
|
|
if not match:
|
|
return None
|
|
amount, unit = int(match.group(1)), match.group(2)
|
|
delta = {"d": timedelta(days=amount), "wk": timedelta(weeks=amount),
|
|
"mo": timedelta(days=30 * amount), "y": timedelta(days=365 * amount)}[unit]
|
|
return now - delta
|
|
|
|
@classmethod
|
|
def _cache_covers(cls, frame: pd.DataFrame, *, start: str | None, end: str | None,
|
|
period: str | None) -> bool:
|
|
lower = pd.Timestamp(start) if start else cls._period_start(period or "5y")
|
|
upper = pd.Timestamp(end) if end else pd.Timestamp.now().normalize()
|
|
# Market holidays/weekends can leave a few calendar days between the last bar and end.
|
|
if lower is not None and frame.index.min() > lower + pd.Timedelta(days=10):
|
|
return False
|
|
if frame.index.max() < upper - pd.Timedelta(days=10):
|
|
return False
|
|
return True
|
|
|
|
def get_history(self, symbol: str, *, start: str | None = None, end: str | None = None,
|
|
period: str | None = None, refresh: bool = False) -> pd.DataFrame:
|
|
yf_symbol, cache_path = yahoo_symbol(symbol), self._cache_path(symbol)
|
|
cached: pd.DataFrame | None = None
|
|
if cache_path.exists() and not refresh:
|
|
cached = normalize_bars(pd.read_csv(cache_path, index_col="date", parse_dates=True))
|
|
if self._cache_covers(cached, start=start, end=end, period=period):
|
|
result = cached
|
|
if start:
|
|
result = result.loc[result.index >= pd.Timestamp(start)]
|
|
if end:
|
|
result = result.loc[result.index < pd.Timestamp(end)]
|
|
if not result.empty:
|
|
return result
|
|
try:
|
|
import yfinance as yf
|
|
except ImportError as exc:
|
|
raise SymbolDataError("yfinance is not installed; run python -m pip install -r requirements.txt") from exc
|
|
last_error: Exception | None = None
|
|
for attempt in range(self.retries + 1):
|
|
try:
|
|
downloaded = yf.download(
|
|
yf_symbol, start=start, end=end,
|
|
period=period or (None if start else "5y"), interval="1d",
|
|
auto_adjust=False, actions=True, repair=self.repair,
|
|
progress=False, threads=False, multi_level_index=False,
|
|
)
|
|
normalized = normalize_bars(downloaded)
|
|
if cached is not None:
|
|
normalized = normalize_bars(pd.concat([cached, normalized]).sort_index())
|
|
if yf_symbol == "^SET.BK" and period != "max" and not self._cache_covers(
|
|
normalized, start=start, end=end, period=period):
|
|
raise _IncompleteHistoryError(
|
|
f"{yf_symbol}: Yahoo returned {len(normalized)} bar(s) but did not cover "
|
|
"the requested history; refusing an incomplete SET benchmark"
|
|
)
|
|
if yf_symbol == "^SET.BK" and period == "max" and len(normalized) < 1500:
|
|
raise _IncompleteHistoryError(
|
|
f"{yf_symbol}: Yahoo returned only {len(normalized)} bars for period=max; "
|
|
"refusing an incomplete SET benchmark"
|
|
)
|
|
self.cache_dir.mkdir(parents=True, exist_ok=True)
|
|
normalized.to_csv(cache_path, index_label="date")
|
|
if start:
|
|
normalized = normalized.loc[normalized.index >= pd.Timestamp(start)]
|
|
if end:
|
|
normalized = normalized.loc[normalized.index < pd.Timestamp(end)]
|
|
return normalized
|
|
except _IncompleteHistoryError:
|
|
raise
|
|
except Exception as exc:
|
|
last_error = exc
|
|
if attempt < self.retries:
|
|
time.sleep(self.retry_delay * (attempt + 1))
|
|
raise SymbolDataError(f"{yf_symbol}: Yahoo download failed: {last_error}") from last_error
|
|
|
|
def download_many(self, symbols: list[str], *, start: str | None = None, end: str | None = None,
|
|
period: str | None = None, refresh: bool = False) -> tuple[dict[str, pd.DataFrame], dict[str, str]]:
|
|
data: dict[str, pd.DataFrame] = {}
|
|
errors: dict[str, str] = {}
|
|
for symbol in symbols:
|
|
try:
|
|
data[symbol] = self.get_history(symbol, start=start, end=end, period=period, refresh=refresh)
|
|
except SymbolDataError as exc:
|
|
errors[symbol] = str(exc)
|
|
return data, errors
|