You cannot select more than 25 topics
Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
201 lines
7.7 KiB
Python
201 lines
7.7 KiB
Python
from __future__ import annotations
|
|
|
|
import importlib.util
|
|
import sys
|
|
from pathlib import Path
|
|
|
|
sys.path.insert(0, str(Path(__file__).resolve().parents[2]))
|
|
|
|
_SCRIPT_PATH = Path(__file__).resolve().parents[2] / "scripts" / "enrich_tier2_features.py"
|
|
_SPEC = importlib.util.spec_from_file_location("enrich_tier2_features", _SCRIPT_PATH)
|
|
assert _SPEC and _SPEC.loader
|
|
_MODULE = importlib.util.module_from_spec(_SPEC)
|
|
_SPEC.loader.exec_module(_MODULE)
|
|
|
|
|
|
def test_extract_short_ratio_history_supports_new_oracle_shape() -> None:
|
|
payload = {
|
|
"symbol": "AAPL",
|
|
"history": [
|
|
{"date": "2026-03-05", "short_ratio": 0.31},
|
|
{"date": "2026-03-06", "short_ratio": 0.29},
|
|
],
|
|
}
|
|
|
|
history = _MODULE._extract_short_ratio_history(payload)
|
|
|
|
assert history == payload["history"]
|
|
|
|
|
|
def test_fetch_short_ratio_uses_history_short_ratio(monkeypatch) -> None:
|
|
class _Response:
|
|
status_code = 200
|
|
|
|
def json(self) -> dict:
|
|
return {
|
|
"symbol": "AAPL",
|
|
"history": [
|
|
{"date": "2026-03-07", "short_ratio": 0.50},
|
|
{"date": "2026-03-06", "short_ratio": 0.40},
|
|
{"date": "2026-03-05", "short_ratio": 0.30},
|
|
{"date": "2026-03-04", "short_ratio": 0.20},
|
|
{"date": "2026-03-03", "short_ratio": 0.10},
|
|
{"date": "2026-03-02", "short_ratio": 0.00},
|
|
],
|
|
}
|
|
|
|
_MODULE._short_ratio_db_cache.clear()
|
|
_MODULE._short_ratio_cache.clear()
|
|
monkeypatch.setattr(_MODULE, "_fetch_short_ratio_history_from_db", lambda ticker: [])
|
|
monkeypatch.setattr(_MODULE, "_fetch_short_ratio_history_from_finra_cdn", lambda ticker, event_date: [])
|
|
monkeypatch.setattr(_MODULE._http, "get", lambda *args, **kwargs: _Response())
|
|
|
|
value = _MODULE.fetch_short_ratio("AAPL", "2026-03-08")
|
|
|
|
assert value == (0.50 + 0.40 + 0.30 + 0.20 + 0.10) / 5
|
|
|
|
|
|
def test_fetch_short_ratio_prefers_db_history(monkeypatch) -> None:
|
|
_MODULE._short_ratio_db_cache.clear()
|
|
_MODULE._short_ratio_cache.clear()
|
|
monkeypatch.setattr(
|
|
_MODULE,
|
|
"_fetch_short_ratio_history_from_db",
|
|
lambda ticker: [
|
|
{"date": "2026-03-07", "short_ratio": 0.50},
|
|
{"date": "2026-03-06", "short_ratio": 0.40},
|
|
{"date": "2026-03-05", "short_ratio": 0.30},
|
|
{"date": "2026-03-04", "short_ratio": 0.20},
|
|
{"date": "2026-03-03", "short_ratio": 0.10},
|
|
],
|
|
)
|
|
monkeypatch.setattr(
|
|
_MODULE._http,
|
|
"get",
|
|
lambda *args, **kwargs: (_ for _ in ()).throw(AssertionError("oracle should not be called")),
|
|
)
|
|
|
|
value = _MODULE.fetch_short_ratio("AAPL", "2026-03-08")
|
|
|
|
assert value == (0.50 + 0.40 + 0.30 + 0.20 + 0.10) / 5
|
|
|
|
|
|
def test_required_short_ratio_days_tracks_event_date() -> None:
|
|
days = _MODULE._required_short_ratio_days("2022-03-02")
|
|
assert 1400 <= days <= 2000
|
|
|
|
|
|
def test_fetch_short_ratio_oracle_uses_dynamic_days(monkeypatch) -> None:
|
|
calls: list[int] = []
|
|
|
|
class _Response:
|
|
status_code = 200
|
|
|
|
def json(self) -> dict:
|
|
return {
|
|
"symbol": "AAPL",
|
|
"history": [
|
|
{"date": "2026-03-07", "short_ratio": 0.50},
|
|
{"date": "2026-03-06", "short_ratio": 0.40},
|
|
{"date": "2026-03-05", "short_ratio": 0.30},
|
|
{"date": "2026-03-04", "short_ratio": 0.20},
|
|
{"date": "2026-03-03", "short_ratio": 0.10},
|
|
],
|
|
}
|
|
|
|
def _fake_get(*args, **kwargs):
|
|
calls.append(int(kwargs["params"]["days"]))
|
|
return _Response()
|
|
|
|
_MODULE._short_ratio_db_cache.clear()
|
|
_MODULE._short_ratio_cache.clear()
|
|
monkeypatch.setattr(_MODULE, "_fetch_short_ratio_history_from_db", lambda ticker: [])
|
|
monkeypatch.setattr(_MODULE, "_fetch_short_ratio_history_from_finra_cdn", lambda ticker, event_date: [])
|
|
monkeypatch.setattr(_MODULE._http, "get", _fake_get)
|
|
|
|
value = _MODULE.fetch_short_ratio("AAPL", "2026-03-08")
|
|
|
|
assert value == (0.50 + 0.40 + 0.30 + 0.20 + 0.10) / 5
|
|
assert calls and calls[0] >= 60
|
|
|
|
|
|
def test_fetch_short_ratio_falls_back_to_finra_daily_files(monkeypatch) -> None:
|
|
class _Response:
|
|
status_code = 200
|
|
|
|
def __init__(self, text: str) -> None:
|
|
self.text = text
|
|
|
|
payloads = {
|
|
"20260307": "Date|Symbol|ShortVolume|ShortExemptVolume|TotalVolume|Market\n20260307|AAPL|50|0|100|Q\n",
|
|
"20260306": "Date|Symbol|ShortVolume|ShortExemptVolume|TotalVolume|Market\n20260306|AAPL|40|0|100|Q\n",
|
|
"20260305": "Date|Symbol|ShortVolume|ShortExemptVolume|TotalVolume|Market\n20260305|AAPL|30|0|100|Q\n",
|
|
"20260304": "Date|Symbol|ShortVolume|ShortExemptVolume|TotalVolume|Market\n20260304|AAPL|20|0|100|Q\n",
|
|
"20260303": "Date|Symbol|ShortVolume|ShortExemptVolume|TotalVolume|Market\n20260303|AAPL|10|0|100|Q\n",
|
|
}
|
|
|
|
def _fake_get(url: str, **_kwargs):
|
|
date_key = url.rsplit("CNMSshvol", 1)[1].split(".txt", 1)[0]
|
|
return _Response(payloads.get(date_key, "Date|Symbol|ShortVolume|ShortExemptVolume|TotalVolume|Market\n"))
|
|
|
|
_MODULE._short_ratio_db_cache.clear()
|
|
_MODULE._short_ratio_cache.clear()
|
|
_MODULE._finra_daily_cache.clear()
|
|
monkeypatch.setattr(_MODULE, "_fetch_short_ratio_history_from_db", lambda ticker: [])
|
|
monkeypatch.setattr(_MODULE, "_fetch_short_ratio_history_from_oracle", lambda ticker, days: [])
|
|
monkeypatch.setattr(_MODULE._http, "get", _fake_get)
|
|
|
|
value = _MODULE.fetch_short_ratio("AAPL", "2026-03-08")
|
|
|
|
assert value == (0.50 + 0.40 + 0.30 + 0.20 + 0.10) / 5
|
|
|
|
|
|
def test_fetch_short_ratio_oracle_timeout_returns_none(monkeypatch) -> None:
|
|
_MODULE._short_ratio_db_cache.clear()
|
|
_MODULE._short_ratio_cache.clear()
|
|
_MODULE._finra_daily_cache.clear()
|
|
_MODULE._oracle_short_ratio_unavailable = False
|
|
monkeypatch.setattr(_MODULE, "_fetch_short_ratio_history_from_db", lambda ticker: [])
|
|
monkeypatch.setattr(_MODULE, "_fetch_short_ratio_history_from_finra_cdn", lambda ticker, event_date: [])
|
|
|
|
def _raise_timeout(*_args, **_kwargs):
|
|
raise _MODULE.requests.ReadTimeout("oracle unavailable")
|
|
|
|
monkeypatch.setattr(_MODULE._http, "get", _raise_timeout)
|
|
|
|
value = _MODULE.fetch_short_ratio("AAPL", "2026-03-08")
|
|
|
|
assert value is None
|
|
assert _MODULE._oracle_short_ratio_unavailable is True
|
|
|
|
|
|
def test_normalize_sector_name_maps_snapshot_aliases() -> None:
|
|
assert _MODULE.normalize_sector_name("Consumer Cyclical") == "Consumer Discretionary"
|
|
assert _MODULE.normalize_sector_name("Financial Services") == "Financials"
|
|
assert _MODULE.normalize_sector_name("Consumer Defensive") == "Consumer Staples"
|
|
assert _MODULE.normalize_sector_name("Basic Materials") == "Materials"
|
|
assert _MODULE.normalize_sector_name("Healthcare") == "Healthcare"
|
|
|
|
|
|
def test_fetch_sector_momentum_uses_normalized_sector_alias(monkeypatch) -> None:
|
|
def _fake_fetch(ticker: str, start: str, end: str) -> list[dict]:
|
|
dates = [f"2026-01-{day:02d}" for day in range(1, 22)]
|
|
if ticker == "XLY":
|
|
return [
|
|
{"date": date, "close": 100.0 + idx}
|
|
for idx, date in enumerate(dates[:-1])
|
|
] + [{"date": dates[-1], "close": 144.0}]
|
|
if ticker == "SPY":
|
|
return [
|
|
{"date": date, "close": 100.0 + idx}
|
|
for idx, date in enumerate(dates[:-1])
|
|
] + [{"date": dates[-1], "close": 132.0}]
|
|
raise AssertionError(f"unexpected ticker {ticker}")
|
|
|
|
monkeypatch.setattr(_MODULE, "_fetch_bars_raw", _fake_fetch)
|
|
_MODULE._sector_etf_cache.clear()
|
|
|
|
value = _MODULE.fetch_sector_momentum("Consumer Cyclical", "2026-01-21")
|
|
|
|
assert value == 0.12
|