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.

242 lines
8.2 KiB
Python

"""
Alpaca bars → AlpacaPriceData conversion and DB storage
"""
import logging
from datetime import datetime, timezone
from typing import Dict, List, Optional
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy import select, and_, func
from sqlalchemy.dialects.postgresql import insert as pg_insert
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,
db: AsyncSession,
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")
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
# Build rows for batch insert
rows = []
for bar in bars:
bar_dt = _parse_bar_timestamp(bar["t"])
rows.append({
"ticker": ticker,
"date": bar_dt,
"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
# Batch insert — skip duplicates via ON CONFLICT DO NOTHING
stmt = pg_insert(AlpacaPriceData).values(rows)
stmt = stmt.on_conflict_do_nothing(constraint='uq_alpaca_price_data')
result = await db.execute(stmt)
await db.commit()
inserted = result.rowcount
if inserted:
logger.info(f"Alpaca: inserted {inserted} bars for {ticker}")
return inserted
async def get_or_fetch_multi_bars(
self,
db: AsyncSession,
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.
1. Check which tickers are missing data (max_date < end_dt) in DB.
2. Fetch only missing tickers from Alpaca and upsert.
3. Read all data from DB and return as Dict[ticker → rows].
"""
upper_tickers = [t.upper() for t in tickers]
if not force_refresh:
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
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",
})
if rows:
# Chunk to stay under asyncpg's 32767-param limit (~11 params/row → 2900/chunk)
CHUNK = 2900
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 db.commit()
logger.info(
f"Alpaca multi-bars: stored {len(rows)} rows for {len(need_fetch)} tickers"
)
# Read back from DB — filtered by interval to avoid mixing 1d/5m/etc.
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()
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