@ -9,9 +9,11 @@ from __future__ import annotations
import argparse
import argparse
import asyncio
import asyncio
import concurrent . futures
import datetime as dt
import datetime as dt
import json
import json
import sys
import sys
from collections import defaultdict
from pathlib import Path
from pathlib import Path
import asyncpg
import asyncpg
@ -231,7 +233,10 @@ def fetch_bars(ticker: str, event_date: str) -> list[PriceBar]:
from datetime import datetime , timedelta
from datetime import datetime , timedelta
end_dt = datetime . strptime ( event_date , " % Y- % m- %d " )
end_dt = datetime . strptime ( event_date , " % Y- % m- %d " )
start_dt = end_dt - timedelta ( days = 120 )
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 = [ ]
bars = [ ]
for b in bars_raw :
for b in bars_raw :
try :
try :
@ -243,6 +248,44 @@ def fetch_bars(ticker: str, event_date: str) -> list[PriceBar]:
return bars
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 :
def fetch_short_ratio ( ticker : str , event_date : str ) - > float | None :
""" Fetch average short ratio over 5 days before event. """
""" Fetch average short ratio over 5 days before event. """
points = _fetch_short_ratio_history_from_db ( ticker )
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 }
results = { f : [ None ] * n for f in FEATURES }
success = 0
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 ) :
for i in range ( n ) :
ticker = tickers [ i ]
ticker = tickers [ i ]
event_date = str ( event_dates [ i ] )
event_date = str ( event_dates [ i ] )