@ -9,9 +9,11 @@ from __future__ import annotations
import argparse
import asyncio
import concurrent . futures
import datetime as dt
import json
import sys
from collections import defaultdict
from pathlib import Path
import asyncpg
@ -231,7 +233,10 @@ def fetch_bars(ticker: str, event_date: str) -> list[PriceBar]:
from datetime import datetime , timedelta
end_dt = datetime . strptime ( event_date , " % Y- % m- %d " )
start_dt = end_dt - timedelta ( days = 120 )
bars_raw = _fetch_bars_raw ( ticker , start_dt . strftime ( " % Y- % m- %d " ) , event_date )
start_str = start_dt . strftime ( " % Y- % m- %d " )
# Use full-history cache to avoid per-row HTTP calls
all_bars_raw = _fetch_full_bars_raw ( ticker )
bars_raw = [ b for b in all_bars_raw if start_str < = b . get ( " date " , " " ) < = event_date ]
bars = [ ]
for b in bars_raw :
try :
@ -243,6 +248,44 @@ def fetch_bars(ticker: str, event_date: str) -> list[PriceBar]:
return bars
def _prefetch_ticker_bars ( tickers : list , max_workers : int = 16 ) - > None :
""" Pre-warm full-history bar cache for all unique tickers in parallel. """
unique = sorted ( { str ( t ) for t in tickers if t } )
print ( f " Prefetching price bars for { len ( unique ) } tickers... " )
with concurrent . futures . ThreadPoolExecutor ( max_workers = max_workers ) as pool :
for loaded , _ in enumerate ( pool . map ( _fetch_full_bars_raw , unique ) , start = 1 ) :
if loaded % 100 == 0 or loaded == len ( unique ) :
print ( f " bars prefetch { loaded } / { len ( unique ) } " )
async def _prefetch_short_ratio_db_batch ( tickers : list [ str ] ) - > None :
""" Bulk-fetch short ratio history for all tickers in a single DB query. """
dsn = get_settings ( ) . postgres_dsn . replace ( " +asyncpg " , " " )
conn = await asyncpg . connect ( dsn = dsn )
try :
rows = await conn . fetch (
"""
SELECT ticker_raw ,
trade_date : : text AS date ,
short_volume : : double precision / NULLIF ( total_volume : : double precision , 0.0 ) AS short_ratio
FROM short_sale_daily
WHERE ticker_raw = ANY ( $ 1 )
AND total_volume IS NOT NULL
AND total_volume > 0
ORDER BY trade_date DESC
""" ,
tickers ,
)
finally :
await conn . close ( )
grouped : dict [ str , list [ dict ] ] = defaultdict ( list )
for row in rows :
if row [ " short_ratio " ] is not None :
grouped [ row [ " ticker_raw " ] ] . append ( { " date " : row [ " date " ] , " short_ratio " : float ( row [ " short_ratio " ] ) } )
for ticker in tickers :
_short_ratio_db_cache [ ticker ] = grouped . get ( ticker , [ ] )
def fetch_short_ratio ( ticker : str , event_date : str ) - > float | None :
""" Fetch average short ratio over 5 days before event. """
points = _fetch_short_ratio_history_from_db ( ticker )
@ -377,6 +420,15 @@ def enrich_split(input_path: Path, output_path: Path):
results = { f : [ None ] * n for f in FEATURES }
success = 0
# Pre-fetch price bars and short ratio history for all unique tickers in batch
_prefetch_ticker_bars ( tickers )
unique_tickers = sorted ( { str ( t ) for t in tickers if t } )
try :
asyncio . run ( _prefetch_short_ratio_db_batch ( unique_tickers ) )
print ( f " Short ratio DB prefetch done for { len ( unique_tickers ) } tickers " )
except Exception as e :
print ( f " Short ratio DB prefetch failed (will fall back per-row): { e } " )
for i in range ( n ) :
ticker = tickers [ i ]
event_date = str ( event_dates [ i ] )