""" Alpaca bars → AlpacaPriceData conversion and DB storage """ import asyncio import logging from datetime import datetime, timezone from typing import Dict, List, Optional from sqlalchemy import select, and_, func from sqlalchemy.dialects.postgresql import insert as pg_insert from app.core.database import AsyncSessionLocal from app.models.alpaca_price import AlpacaPriceData from app.services.alpaca_client import AlpacaClient, normalize_ticker logger = logging.getLogger(__name__) class AlpacaPriceService: """Alpaca bars → AlpacaPriceData conversion and DB storage""" def __init__(self, client: Optional[AlpacaClient] = None): self.client = client or AlpacaClient() def is_available(self) -> bool: return self.client.is_configured() async def fetch_and_store_bars( self, ticker: str, start_date: datetime, end_date: datetime, interval: str = "1d", ) -> int: """ Fetch bars from Alpaca and upsert into AlpacaPriceData table. Returns: Number of newly inserted records. """ ticker = ticker.upper() start_str = start_date.strftime("%Y-%m-%d") end_str = end_date.strftime("%Y-%m-%d") # Phase 1: fetch from Alpaca (no DB session held) bars = await self.client.get_bars( symbol=ticker, timeframe=interval, start=start_str, end=end_str, ) if not bars: logger.warning(f"Alpaca returned 0 bars for {ticker}") return 0 rows = [] for bar in bars: bar_dt = _parse_bar_timestamp(bar["t"]) rows.append({ "ticker": ticker, "date": bar_dt, "interval": interval, "open": float(bar.get("o", 0)), "high": float(bar.get("h", 0)), "low": float(bar.get("l", 0)), "close": float(bar.get("c", 0)), "volume": float(bar.get("v", 0)), "vwap": float(bar["vw"]) if bar.get("vw") else None, "trade_count": int(bar["n"]) if bar.get("n") else None, "data_source": "ALPACA", }) if not rows: return 0 # Phase 2: batch insert with short-lived session # 14 params/row → CHUNK=2300 (2300×14=32,200 < 32,767 asyncpg limit) CHUNK = 2300 inserted = 0 async with AsyncSessionLocal() as db: for i in range(0, len(rows), CHUNK): stmt = pg_insert(AlpacaPriceData).values(rows[i : i + CHUNK]) stmt = stmt.on_conflict_do_nothing(constraint='uq_alpaca_price_data') result = await db.execute(stmt) inserted += result.rowcount await asyncio.sleep(0) # yield event loop between chunks await db.commit() if inserted: logger.info(f"Alpaca: inserted {inserted} bars for {ticker}") return inserted async def get_or_fetch_multi_bars( self, tickers: List[str], start_dt: datetime, end_dt: datetime, interval: str = "1d", force_refresh: bool = False, feed: Optional[str] = None, ) -> Dict[str, List[AlpacaPriceData]]: """ DB-first multi-ticker daily bars. Session-per-phase: DB connections are held only during short DB operations, never during external Alpaca HTTP calls. 1. Short session: check which tickers are missing data in DB. 2. Fetch only missing tickers from Alpaca (no session held). 3. Short session: upsert fetched rows. 4. Short session: read back and return as Dict[ticker → rows]. """ upper_tickers = [t.upper() for t in tickers] # Phase 1: coverage check (short session) if not force_refresh: async with AsyncSessionLocal() as db: result = await db.execute( select( AlpacaPriceData.ticker, func.max(AlpacaPriceData.date).label("max_date"), ) .where( and_( AlpacaPriceData.ticker.in_(upper_tickers), AlpacaPriceData.interval == interval, AlpacaPriceData.date >= start_dt, AlpacaPriceData.date <= end_dt, ) ) .group_by(AlpacaPriceData.ticker) ) coverage = {row.ticker: row.max_date for row in result} need_fetch = [ t for t in upper_tickers if coverage.get(t) is None or coverage[t].date() < end_dt.date() ] else: need_fetch = upper_tickers # Phase 2: Alpaca HTTP (no DB session held) if need_fetch: start_str = start_dt.strftime("%Y-%m-%d") end_str = end_dt.strftime("%Y-%m-%d") # reverse-map: normalized Alpaca symbol → original input ticker reverse_map = {normalize_ticker(t): t for t in need_fetch} raw = await self.client.get_multi_bars( symbols=need_fetch, timeframe=interval, start=start_str, end=end_str, feed=feed, ) rows = [] for alpaca_sym, bar_list in raw.items(): original = reverse_map.get(alpaca_sym, alpaca_sym) for bar in bar_list: bar_dt = _parse_bar_timestamp(bar["t"]) rows.append({ "ticker": original, "date": bar_dt, "interval": interval, "open": float(bar.get("o", 0)), "high": float(bar.get("h", 0)), "low": float(bar.get("l", 0)), "close": float(bar.get("c", 0)), "volume": float(bar.get("v", 0)), "vwap": float(bar["vw"]) if bar.get("vw") else None, "trade_count": int(bar["n"]) if bar.get("n") else None, "data_source": "ALPACA", }) # Phase 3: upsert (short session) if rows: # 14 params/row → CHUNK=2300 (2300×14=32,200 < 32,767 asyncpg limit) CHUNK = 2300 async with AsyncSessionLocal() as db: for i in range(0, len(rows), CHUNK): stmt = pg_insert(AlpacaPriceData).values(rows[i : i + CHUNK]) stmt = stmt.on_conflict_do_nothing(constraint="uq_alpaca_price_data") await db.execute(stmt) await asyncio.sleep(0) # yield event loop between chunks await db.commit() logger.info( f"Alpaca multi-bars: stored {len(rows)} rows for {len(need_fetch)} tickers" ) # Phase 4: read back (short session) — filtered by interval to avoid mixing 1d/5m/etc. async with AsyncSessionLocal() as db: result = await db.execute( select(AlpacaPriceData) .where( and_( AlpacaPriceData.ticker.in_(upper_tickers), AlpacaPriceData.interval == interval, AlpacaPriceData.date >= start_dt, AlpacaPriceData.date <= end_dt, ) ) .order_by(AlpacaPriceData.ticker, AlpacaPriceData.date) ) db_rows = result.scalars().all() await asyncio.sleep(0) # yield after bulk ORM object creation data: Dict[str, List[AlpacaPriceData]] = {t: [] for t in upper_tickers} for row in db_rows: data[row.ticker].append(row) return data async def fetch_bars_raw( self, ticker: str, interval: str = "1d", start_date: Optional[datetime] = None, end_date: Optional[datetime] = None, feed: Optional[str] = None, ) -> List[Dict]: """ Return raw bar dicts without touching the DB (useful for intraday / non-persistent use). """ start_str = start_date.strftime("%Y-%m-%d") if start_date else None end_str = end_date.strftime("%Y-%m-%d") if end_date else None bars = await self.client.get_bars( symbol=ticker.upper(), timeframe=interval, start=start_str, end=end_str, feed=feed, ) return [ { "timestamp": bar["t"], "open": bar.get("o"), "high": bar.get("h"), "low": bar.get("l"), "close": bar.get("c"), "volume": bar.get("v"), "vwap": bar.get("vw"), "trade_count": bar.get("n"), } for bar in bars ] def _parse_bar_timestamp(ts_str: str) -> datetime: """Parse Alpaca bar timestamp (RFC-3339) into a timezone-aware datetime.""" # Alpaca returns e.g. "2024-01-02T05:00:00Z" dt = datetime.fromisoformat(ts_str.replace("Z", "+00:00")) if dt.tzinfo is None: dt = dt.replace(tzinfo=timezone.utc) return dt