Files
2026-09-29 14:08:43 +07:00

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