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.

270 lines
9.8 KiB
Python

This file contains ambiguous Unicode characters!

This file contains ambiguous Unicode characters that may be confused with others in your current locale. If your use case is intentional and legitimate, you can safely ignore this warning. Use the Escape button to highlight these characters.

"""
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, lightweight column query).
# Select only the 7 columns the endpoints need instead of loading full ORM
# objects — avoids SQLAlchemy instrumentation overhead (~10x lighter per row).
async with AsyncSessionLocal() as db:
result = await db.execute(
select(
AlpacaPriceData.ticker,
AlpacaPriceData.date,
AlpacaPriceData.open,
AlpacaPriceData.high,
AlpacaPriceData.low,
AlpacaPriceData.close,
AlpacaPriceData.volume,
)
.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.all() # lightweight Row namedtuples, not ORM objects
await asyncio.sleep(0) # yield after bulk row creation
data: Dict[str, List] = {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