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.

249 lines
8.6 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 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,
"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
# Batch insert with chunking (14 params/row → CHUNK=2300, 2300×14=32,200 < 32,767)
CHUNK = 2300
inserted = 0
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 db.commit()
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 (14 params/row → 2300/chunk)
# 14 params: id/ticker/date/open/high/low/close/volume/vwap/trade_count/
# interval/data_source/created_at/updated_at
# 2300 × 14 = 32,200 < 32,767
CHUNK = 2300
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