refactor: session-per-phase — DB 커넥션을 외부 API 호출 중 해제

서비스 메서드가 db: AsyncSession을 인자로 받아 외부 API 호출(yfinance 30-60s,
Alpaca HTTP) 중에도 DB 커넥션을 잡고 있던 구조를 제거.

변경 패턴 (Session-per-phase):
  Before: Endpoint(db) → Service(db) → DB check → API call(30s 세션 유지) → DB store
  After:  Endpoint()   → Service()   → DB check(세션1) → API call(세션 없음) → DB store(세션2)

변경 파일:
- alpaca_price_service.py: fetch_and_store_bars, get_or_fetch_multi_bars에서 db 제거
- price_data_service.py: get_or_update_price_data, get_multiple_tickers_data_optimized에서
  db 제거; _fetch_price_data(새), _bulk_fetch_price_data(새, lambda 클로저 버그 수정);
  _fetch_and_store_price_data, _bulk_fetch_and_store_price_data, get_multiple_tickers_data 제거
- alpaca.py, price.py: Depends(get_db) 제거 (GET /latest 제외)
- financial_service.py, real_sec_financial_service.py: 호출 인자 정리

결과: "idle in transaction" 커넥션 0개, 풀 고갈 원인 근본 해결

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
main
I Luk Kim 4 months ago
parent 7efe6704b3
commit c203f3920b

@ -6,8 +6,7 @@ from datetime import date, datetime, timezone, timedelta
from typing import Optional from typing import Optional
from zoneinfo import ZoneInfo from zoneinfo import ZoneInfo
from fastapi import APIRouter, Depends, HTTPException, Query from fastapi import APIRouter, HTTPException, Query
from sqlalchemy.ext.asyncio import AsyncSession
_ET = ZoneInfo("America/New_York") _ET = ZoneInfo("America/New_York")
_MARKET_CLOSE_HOUR = 16 # 4:00 PM ET _MARKET_CLOSE_HOUR = 16 # 4:00 PM ET
@ -19,7 +18,6 @@ def _market_closed_for(d: date) -> bool:
market_close_et = datetime(d.year, d.month, d.day, _MARKET_CLOSE_HOUR, 0, tzinfo=_ET) market_close_et = datetime(d.year, d.month, d.day, _MARKET_CLOSE_HOUR, 0, tzinfo=_ET)
return now_et >= market_close_et return now_et >= market_close_et
from app.core.database import get_db
from app.schemas.financial import ( from app.schemas.financial import (
AlpacaMultiBarsResponse, AlpacaMultiBarsResponse,
AlpacaMultiSnapshotResponse, AlpacaMultiSnapshotResponse,
@ -93,7 +91,6 @@ async def get_alpaca_intraday_multi(
start_date: Optional[date] = Query(None, description="Start date (YYYY-MM-DD). Default: yesterday"), start_date: Optional[date] = Query(None, description="Start date (YYYY-MM-DD). Default: yesterday"),
end_date: Optional[date] = Query(None, description="End date (YYYY-MM-DD). Must be before today. Default: yesterday"), end_date: Optional[date] = Query(None, description="End date (YYYY-MM-DD). Must be before today. Default: yesterday"),
force_refresh: bool = Query(False, description="Re-fetch from Alpaca even if DB has data"), force_refresh: bool = Query(False, description="Re-fetch from Alpaca even if DB has data"),
db: AsyncSession = Depends(get_db),
): ):
"""Multi-ticker historical intraday bars via Alpaca SIP (up to today after market close).""" """Multi-ticker historical intraday bars via Alpaca SIP (up to today after market close)."""
symbols = [s.strip().upper() for s in tickers.split(",") if s.strip()] symbols = [s.strip().upper() for s in tickers.split(",") if s.strip()]
@ -125,7 +122,7 @@ async def get_alpaca_intraday_multi(
try: try:
data = await svc.get_or_fetch_multi_bars( data = await svc.get_or_fetch_multi_bars(
db, symbols, start_dt, end_dt, interval, force_refresh, feed="sip" symbols, start_dt, end_dt, interval, force_refresh, feed="sip"
) )
except Exception as e: except Exception as e:
err = str(e) err = str(e)
@ -182,7 +179,6 @@ async def get_alpaca_intraday_multi(
async def get_alpaca_intraday_today( async def get_alpaca_intraday_today(
tickers: str = Query(..., description="Comma-separated tickers, e.g. AAPL,MSFT,BF-B"), tickers: str = Query(..., description="Comma-separated tickers, e.g. AAPL,MSFT,BF-B"),
interval: str = Query("5m", description="Interval: 1m, 5m, 15m, 30m, 1h"), interval: str = Query("5m", description="Interval: 1m, 5m, 15m, 30m, 1h"),
db: AsyncSession = Depends(get_db),
): ):
"""Today's real-time intraday bars via Alpaca IEX (always re-fetches latest).""" """Today's real-time intraday bars via Alpaca IEX (always re-fetches latest)."""
symbols = [s.strip().upper() for s in tickers.split(",") if s.strip()] symbols = [s.strip().upper() for s in tickers.split(",") if s.strip()]
@ -199,7 +195,7 @@ async def get_alpaca_intraday_today(
try: try:
data = await svc.get_or_fetch_multi_bars( data = await svc.get_or_fetch_multi_bars(
db, symbols, start_dt, end_dt, interval, force_refresh=True, feed="iex" symbols, start_dt, end_dt, interval, force_refresh=True, feed="iex"
) )
except Exception as e: except Exception as e:
err = str(e) err = str(e)

@ -7,7 +7,7 @@ from typing import List, Optional
from fastapi import APIRouter, Depends, HTTPException, Query, Response from fastapi import APIRouter, Depends, HTTPException, Query, Response
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from app.core.database import get_db from app.core.database import get_db, AsyncSessionLocal
from app.schemas.financial import ( from app.schemas.financial import (
PriceDataRequest, PriceDataRequest,
PriceDataResponse, PriceDataResponse,
@ -116,14 +116,12 @@ router = APIRouter()
async def get_price_data( async def get_price_data(
request: PriceDataRequest, request: PriceDataRequest,
response: Response, response: Response,
db: AsyncSession = Depends(get_db)
): ):
"""Get price data for a ticker using period, quarters, or date range""" """Get price data for a ticker using period, quarters, or date range"""
try: try:
# Use the updated service that handles period resolution # Use the updated service that handles period resolution
price_service = PriceDataService() price_service = PriceDataService()
# Resolve time parameters to get start and end dates # Resolve time parameters to get start and end dates
from app.utils.date_utils import resolve_time_parameters from app.utils.date_utils import resolve_time_parameters
start_date, end_date = resolve_time_parameters( start_date, end_date = resolve_time_parameters(
@ -153,11 +151,12 @@ async def get_price_data(
response.headers["X-Data-Source"] = "redis-cache" response.headers["X-Data-Source"] = "redis-cache"
return cached_body return cached_body
# Check if we have existing data to determine source # Check if we have existing data to determine X-Data-Source header (short session)
missing_periods = await price_service._check_missing_periods( async with AsyncSessionLocal() as db:
db, request.ticker.upper(), start_date, end_date, request.interval missing_periods = await price_service._check_missing_periods(
) db, request.ticker.upper(), start_date, end_date, request.interval
)
# Determine data source # Determine data source
if request.force_refresh: if request.force_refresh:
data_source = "yfinance-fresh" data_source = "yfinance-fresh"
@ -165,13 +164,12 @@ async def get_price_data(
data_source = "yfinance-partial" data_source = "yfinance-partial"
else: else:
data_source = "database-cache" data_source = "database-cache"
# Add data source header # Add data source header
response.headers["X-Data-Source"] = data_source response.headers["X-Data-Source"] = data_source
# Get price data using resolved dates # Get price data using resolved dates (service manages its own sessions)
price_data = await price_service.get_or_update_price_data( price_data = await price_service.get_or_update_price_data(
db,
request.ticker, request.ticker,
start_date, start_date,
end_date, end_date,
@ -314,7 +312,6 @@ async def get_multi_ticker_daily_bars(
end_date: date = Query(..., description="End date (YYYY-MM-DD)"), end_date: date = Query(..., description="End date (YYYY-MM-DD)"),
interval: str = Query("1d", description="Bar interval: 1d, 1w, 1mo"), interval: str = Query("1d", description="Bar interval: 1d, 1w, 1mo"),
force_refresh: bool = Query(False, description="Re-fetch from Alpaca even if DB has data"), force_refresh: bool = Query(False, description="Re-fetch from Alpaca even if DB has data"),
db: AsyncSession = Depends(get_db),
): ):
"""Multi-ticker daily bars via Alpaca with DB storage (ORB engine interface).""" """Multi-ticker daily bars via Alpaca with DB storage (ORB engine interface)."""
symbols = [s.strip().upper() for s in tickers.split(",") if s.strip()] symbols = [s.strip().upper() for s in tickers.split(",") if s.strip()]
@ -332,7 +329,7 @@ async def get_multi_ticker_daily_bars(
try: try:
data = await svc.get_or_fetch_multi_bars( data = await svc.get_or_fetch_multi_bars(
db, symbols, start_dt, end_dt, interval, force_refresh symbols, start_dt, end_dt, interval, force_refresh
) )
except Exception as e: except Exception as e:
err = str(e) err = str(e)
@ -396,7 +393,6 @@ async def get_price_data_simple(
end: Optional[date] = Query(None, description="Alias for end_date"), end: Optional[date] = Query(None, description="Alias for end_date"),
interval: str = Query("1d", description="Data interval: 1d, 1w, 1m, 5d, 1h, etc."), interval: str = Query("1d", description="Data interval: 1d, 1w, 1m, 5d, 1h, etc."),
force_refresh: bool = Query(False, description="Force refresh from Yahoo Finance"), force_refresh: bool = Query(False, description="Force refresh from Yahoo Finance"),
db: AsyncSession = Depends(get_db)
): ):
"""Simplified GET endpoint for price data""" """Simplified GET endpoint for price data"""
# Resolve aliases: start/end → start_date/end_date # Resolve aliases: start/end → start_date/end_date
@ -414,7 +410,7 @@ async def get_price_data_simple(
"message": "Cannot specify both period and date range. Use either period OR start_date+end_date." "message": "Cannot specify both period and date range. Use either period OR start_date+end_date."
} }
) )
if not period and not (start_date and end_date): if not period and not (start_date and end_date):
raise HTTPException( raise HTTPException(
status_code=400, status_code=400,
@ -423,7 +419,7 @@ async def get_price_data_simple(
"message": "Must specify either period OR both start_date and end_date." "message": "Must specify either period OR both start_date and end_date."
} }
) )
# Create request based on provided parameters # Create request based on provided parameters
if period: if period:
request = PriceDataRequest( request = PriceDataRequest(
@ -440,8 +436,8 @@ async def get_price_data_simple(
interval=interval, interval=interval,
force_refresh=force_refresh force_refresh=force_refresh
) )
return await get_price_data(request, response, db) return await get_price_data(request, response)
@router.post( @router.post(
"/data/bulk", "/data/bulk",
@ -523,13 +519,12 @@ async def get_price_data_simple(
) )
async def get_bulk_price_data( async def get_bulk_price_data(
request: BulkPriceDataRequest, request: BulkPriceDataRequest,
db: AsyncSession = Depends(get_db)
): ):
"""Get price data for multiple tickers""" """Get price data for multiple tickers"""
# Use the updated service that handles period resolution # Use the updated service that handles period resolution
price_service = PriceDataService() price_service = PriceDataService()
# Resolve time parameters to get start and end dates # Resolve time parameters to get start and end dates
from app.utils.date_utils import resolve_time_parameters from app.utils.date_utils import resolve_time_parameters
start_date, end_date = resolve_time_parameters( start_date, end_date = resolve_time_parameters(
@ -538,13 +533,13 @@ async def get_bulk_price_data(
quarters=request.quarters, quarters=request.quarters,
period=request.period period=request.period
) )
# Validate date range - ensure both dates are timezone-aware # Validate date range - ensure both dates are timezone-aware
if start_date and start_date.tzinfo is None: if start_date and start_date.tzinfo is None:
start_date = start_date.replace(tzinfo=timezone.utc) start_date = start_date.replace(tzinfo=timezone.utc)
if end_date and end_date.tzinfo is None: if end_date and end_date.tzinfo is None:
end_date = end_date.replace(tzinfo=timezone.utc) end_date = end_date.replace(tzinfo=timezone.utc)
if start_date and end_date and start_date >= end_date: if start_date and end_date and start_date >= end_date:
raise HTTPException( raise HTTPException(
status_code=400, status_code=400,
@ -553,12 +548,9 @@ async def get_bulk_price_data(
"message": "Start date must be before end date" "message": "Start date must be before end date"
} }
) )
# Future date check # Future date check
current_time = datetime.now(timezone.utc) current_time = datetime.now(timezone.utc)
# Make start_date timezone-aware if it's naive
if start_date and start_date.tzinfo is None:
start_date = start_date.replace(tzinfo=timezone.utc)
if start_date and start_date > current_time: if start_date and start_date > current_time:
raise HTTPException( raise HTTPException(
status_code=400, status_code=400,
@ -567,13 +559,12 @@ async def get_bulk_price_data(
"message": "Cannot request data for future dates" "message": "Cannot request data for future dates"
} }
) )
# Use optimized bulk processing method (300s endpoint-level timeout) # Use optimized bulk processing method (300s endpoint-level timeout)
import asyncio as _asyncio import asyncio as _asyncio
try: try:
results, successful_count, failed_count = await _asyncio.wait_for( results, successful_count, failed_count = await _asyncio.wait_for(
price_service.get_multiple_tickers_data_optimized( price_service.get_multiple_tickers_data_optimized(
db=db,
tickers=request.tickers, tickers=request.tickers,
start_date=start_date, start_date=start_date,
end_date=end_date, end_date=end_date,

@ -6,10 +6,10 @@ import logging
from datetime import datetime, timezone from datetime import datetime, timezone
from typing import Dict, List, Optional from typing import Dict, List, Optional
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy import select, and_, func from sqlalchemy import select, and_, func
from sqlalchemy.dialects.postgresql import insert as pg_insert from sqlalchemy.dialects.postgresql import insert as pg_insert
from app.core.database import AsyncSessionLocal
from app.models.alpaca_price import AlpacaPriceData from app.models.alpaca_price import AlpacaPriceData
from app.services.alpaca_client import AlpacaClient, normalize_ticker from app.services.alpaca_client import AlpacaClient, normalize_ticker
@ -27,7 +27,6 @@ class AlpacaPriceService:
async def fetch_and_store_bars( async def fetch_and_store_bars(
self, self,
db: AsyncSession,
ticker: str, ticker: str,
start_date: datetime, start_date: datetime,
end_date: datetime, end_date: datetime,
@ -43,6 +42,7 @@ class AlpacaPriceService:
start_str = start_date.strftime("%Y-%m-%d") start_str = start_date.strftime("%Y-%m-%d")
end_str = end_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( bars = await self.client.get_bars(
symbol=ticker, symbol=ticker,
timeframe=interval, timeframe=interval,
@ -54,7 +54,6 @@ class AlpacaPriceService:
logger.warning(f"Alpaca returned 0 bars for {ticker}") logger.warning(f"Alpaca returned 0 bars for {ticker}")
return 0 return 0
# Build rows for batch insert
rows = [] rows = []
for bar in bars: for bar in bars:
bar_dt = _parse_bar_timestamp(bar["t"]) bar_dt = _parse_bar_timestamp(bar["t"])
@ -75,15 +74,17 @@ class AlpacaPriceService:
if not rows: if not rows:
return 0 return 0
# Batch insert with chunking (14 params/row → CHUNK=2300, 2300×14=32,200 < 32,767) # Phase 2: batch insert with short-lived session
# 14 params/row → CHUNK=2300 (2300×14=32,200 < 32,767 asyncpg limit)
CHUNK = 2300 CHUNK = 2300
inserted = 0 inserted = 0
for i in range(0, len(rows), CHUNK): async with AsyncSessionLocal() as db:
stmt = pg_insert(AlpacaPriceData).values(rows[i : i + CHUNK]) for i in range(0, len(rows), CHUNK):
stmt = stmt.on_conflict_do_nothing(constraint='uq_alpaca_price_data') stmt = pg_insert(AlpacaPriceData).values(rows[i : i + CHUNK])
result = await db.execute(stmt) stmt = stmt.on_conflict_do_nothing(constraint='uq_alpaca_price_data')
inserted += result.rowcount result = await db.execute(stmt)
await db.commit() inserted += result.rowcount
await db.commit()
if inserted: if inserted:
logger.info(f"Alpaca: inserted {inserted} bars for {ticker}") logger.info(f"Alpaca: inserted {inserted} bars for {ticker}")
@ -92,7 +93,6 @@ class AlpacaPriceService:
async def get_or_fetch_multi_bars( async def get_or_fetch_multi_bars(
self, self,
db: AsyncSession,
tickers: List[str], tickers: List[str],
start_dt: datetime, start_dt: datetime,
end_dt: datetime, end_dt: datetime,
@ -103,29 +103,35 @@ class AlpacaPriceService:
""" """
DB-first multi-ticker daily bars. DB-first multi-ticker daily bars.
1. Check which tickers are missing data (max_date < end_dt) in DB. Session-per-phase: DB connections are held only during short DB operations,
2. Fetch only missing tickers from Alpaca and upsert. never during external Alpaca HTTP calls.
3. Read all data from DB and return as Dict[ticker rows].
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] upper_tickers = [t.upper() for t in tickers]
# Phase 1: coverage check (short session)
if not force_refresh: if not force_refresh:
result = await db.execute( async with AsyncSessionLocal() as db:
select( result = await db.execute(
AlpacaPriceData.ticker, select(
func.max(AlpacaPriceData.date).label("max_date"), AlpacaPriceData.ticker,
) func.max(AlpacaPriceData.date).label("max_date"),
.where( )
and_( .where(
AlpacaPriceData.ticker.in_(upper_tickers), and_(
AlpacaPriceData.interval == interval, AlpacaPriceData.ticker.in_(upper_tickers),
AlpacaPriceData.date >= start_dt, AlpacaPriceData.interval == interval,
AlpacaPriceData.date <= end_dt, AlpacaPriceData.date >= start_dt,
AlpacaPriceData.date <= end_dt,
)
) )
.group_by(AlpacaPriceData.ticker)
) )
.group_by(AlpacaPriceData.ticker) coverage = {row.ticker: row.max_date for row in result}
)
coverage = {row.ticker: row.max_date for row in result}
need_fetch = [ need_fetch = [
t for t in upper_tickers t for t in upper_tickers
if coverage.get(t) is None or coverage[t].date() < end_dt.date() if coverage.get(t) is None or coverage[t].date() < end_dt.date()
@ -133,6 +139,7 @@ class AlpacaPriceService:
else: else:
need_fetch = upper_tickers need_fetch = upper_tickers
# Phase 2: Alpaca HTTP (no DB session held)
if need_fetch: if need_fetch:
start_str = start_dt.strftime("%Y-%m-%d") start_str = start_dt.strftime("%Y-%m-%d")
end_str = end_dt.strftime("%Y-%m-%d") end_str = end_dt.strftime("%Y-%m-%d")
@ -167,35 +174,35 @@ class AlpacaPriceService:
"data_source": "ALPACA", "data_source": "ALPACA",
}) })
# Phase 3: upsert (short session)
if rows: if rows:
# Chunk to stay under asyncpg's 32767-param limit (14 params/row → 2300/chunk) # 14 params/row → CHUNK=2300 (2300×14=32,200 < 32,767 asyncpg limit)
# 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 CHUNK = 2300
for i in range(0, len(rows), CHUNK): async with AsyncSessionLocal() as db:
stmt = pg_insert(AlpacaPriceData).values(rows[i : i + CHUNK]) for i in range(0, len(rows), CHUNK):
stmt = stmt.on_conflict_do_nothing(constraint="uq_alpaca_price_data") stmt = pg_insert(AlpacaPriceData).values(rows[i : i + CHUNK])
await db.execute(stmt) stmt = stmt.on_conflict_do_nothing(constraint="uq_alpaca_price_data")
await db.commit() await db.execute(stmt)
await db.commit()
logger.info( logger.info(
f"Alpaca multi-bars: stored {len(rows)} rows for {len(need_fetch)} tickers" 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. # Phase 4: read back (short session) — filtered by interval to avoid mixing 1d/5m/etc.
result = await db.execute( async with AsyncSessionLocal() as db:
select(AlpacaPriceData) result = await db.execute(
.where( select(AlpacaPriceData)
and_( .where(
AlpacaPriceData.ticker.in_(upper_tickers), and_(
AlpacaPriceData.interval == interval, AlpacaPriceData.ticker.in_(upper_tickers),
AlpacaPriceData.date >= start_dt, AlpacaPriceData.interval == interval,
AlpacaPriceData.date <= end_dt, AlpacaPriceData.date >= start_dt,
AlpacaPriceData.date <= end_dt,
)
) )
.order_by(AlpacaPriceData.ticker, AlpacaPriceData.date)
) )
.order_by(AlpacaPriceData.ticker, AlpacaPriceData.date) db_rows = result.scalars().all()
)
db_rows = result.scalars().all()
data: Dict[str, List[AlpacaPriceData]] = {t: [] for t in upper_tickers} data: Dict[str, List[AlpacaPriceData]] = {t: [] for t in upper_tickers}
for row in db_rows: for row in db_rows:

@ -469,7 +469,7 @@ class FinancialService:
if not price_data: if not price_data:
try: try:
price_data = await self.price_service.get_or_update_price_data( price_data = await self.price_service.get_or_update_price_data(
db, ticker, start_date, end_date, "1d", force_refresh=False ticker, start_date, end_date, "1d", force_refresh=False
) )
except Exception as e: except Exception as e:
logger.warning(f"Could not fetch price data for {ticker}: {e}") logger.warning(f"Could not fetch price data for {ticker}: {e}")

@ -14,6 +14,7 @@ import os
# Add parent directory to path for imports # Add parent directory to path for imports
sys.path.append(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))) sys.path.append(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))))
from app.core.database import AsyncSessionLocal
from app.models.financial import PriceData from app.models.financial import PriceData
from app.schemas.financial import DataSource, ErrorType from app.schemas.financial import DataSource, ErrorType
from app.utils.date_utils import parse_period, quarters_to_date_range, resolve_time_parameters from app.utils.date_utils import parse_period, quarters_to_date_range, resolve_time_parameters
@ -51,7 +52,6 @@ class PriceDataService:
async def get_or_update_price_data( async def get_or_update_price_data(
self, self,
db: AsyncSession,
ticker: str, ticker: str,
start_date: datetime, start_date: datetime,
end_date: datetime, end_date: datetime,
@ -59,40 +59,41 @@ class PriceDataService:
force_refresh: bool = False force_refresh: bool = False
) -> List[PriceData]: ) -> List[PriceData]:
""" """
Get price data from database or fetch from Yahoo Finance if needed Get price data from database or fetch from Yahoo Finance if needed.
Args: Session-per-phase: DB connections are held only during short DB operations,
db: Database session never during yfinance calls (which can take 30s+).
ticker: Stock ticker symbol
start_date: Start date for data retrieval
end_date: End date for data retrieval
interval: Data interval (1d, 1w, 1m, 1h, etc.)
force_refresh: Force refresh data from Yahoo Finance
Returns: Returns:
List of PriceData objects List of PriceData objects
""" """
ticker = ticker.upper() ticker = ticker.upper()
# Check if we need to fetch new data # Phase 1: check missing periods (short session)
missing_periods = await self._check_missing_periods( async with AsyncSessionLocal() as db:
db, ticker, start_date, end_date, interval missing_periods = await self._check_missing_periods(
) db, ticker, start_date, end_date, interval
)
if missing_periods or force_refresh: if missing_periods or force_refresh:
if not self.yf_available: if not self.yf_available:
raise ValueError("Yahoo Finance (yfinance-plus) data source not available") raise ValueError("Yahoo Finance (yfinance-plus) data source not available")
# Fetch data from Yahoo Finance using yfinance-plus # Phase 2: fetch from yfinance (no session held)
await self._fetch_and_store_price_data( hist_data = await self._fetch_price_data(ticker, start_date, end_date, interval)
if hist_data is not None and not hist_data.empty:
# Phase 3: store in DB (short session)
async with AsyncSessionLocal() as db:
await self._store_price_data(db, ticker, hist_data, interval)
await db.commit()
# Phase 4: read from DB (short session)
async with AsyncSessionLocal() as db:
price_data = await self._get_price_data_from_db(
db, ticker, start_date, end_date, interval db, ticker, start_date, end_date, interval
) )
# Retrieve data from database
price_data = await self._get_price_data_from_db(
db, ticker, start_date, end_date, interval
)
return price_data return price_data
async def _check_missing_periods( async def _check_missing_periods(
@ -182,63 +183,49 @@ class PriceDataService:
return expected_dates return expected_dates
async def _fetch_and_store_price_data( async def _fetch_price_data(
self, self,
db: AsyncSession,
ticker: str, ticker: str,
start_date: datetime, start_date: datetime,
end_date: datetime, end_date: datetime,
interval: str interval: str
): ):
"""Fetch price data from Yahoo Finance using yfinance-plus and store in database""" """Fetch price data from Yahoo Finance (no DB operations).
try:
logger.info(f"Fetching price data for {ticker} from {start_date} to {end_date}") Returns a DataFrame or None if no data was returned.
"""
# Create yfinance-plus ticker object logger.info(f"Fetching price data for {ticker} from {start_date} to {end_date}")
yf_ticker = yf.Ticker(ticker)
yf_ticker = yf.Ticker(ticker)
# Fetch historical data
# Convert dates to strings in YYYY-MM-DD format start_str = start_date.strftime('%Y-%m-%d')
start_str = start_date.strftime('%Y-%m-%d') # yfinance's `end` parameter is exclusive — add +1 day to include end_date.
# yfinance's `end` parameter is exclusive for daily data when using date strings. end_inclusive = end_date + timedelta(days=1)
# Add +1 day to include the intended end_date day in the results. end_str = end_inclusive.strftime('%Y-%m-%d')
from datetime import timedelta
end_inclusive = end_date + timedelta(days=1) loop = asyncio.get_event_loop()
end_str = end_inclusive.strftime('%Y-%m-%d') hist_data = await _run_with_timeout(
loop.run_in_executor(
# Run yfinance-plus in executor to avoid blocking None,
loop = asyncio.get_event_loop() lambda: yf_ticker.history(
hist_data = await _run_with_timeout( start=start_str,
loop.run_in_executor( end=end_str,
None, interval=interval,
lambda: yf_ticker.history( auto_adjust=True,
start=start_str, prepost=False,
end=end_str, period=None
interval=interval, )
auto_adjust=True, ),
prepost=False, timeout_seconds=30,
period=None # Explicitly set period to None when using start/end dates description=f"history {ticker} {start_str}:{end_str}"
) )
),
timeout_seconds=30, if hist_data is None or hist_data.empty:
description=f"history {ticker} {start_str}:{end_str}" logger.warning(f"No price data returned for {ticker}")
) return None
if hist_data.empty: logger.info(f"Fetched {len(hist_data)} price records for {ticker}")
logger.warning(f"No price data returned for {ticker}") return hist_data
return
# Store data in database
await self._store_price_data(db, ticker, hist_data, interval)
await db.commit()
logger.info(f"Successfully stored {len(hist_data)} price records for {ticker}")
except Exception as e:
logger.error(f"Error fetching price data for {ticker}: {str(e)}")
await db.rollback()
raise
async def get_quote(self, ticker: str, use_prepost: bool = True) -> Dict: async def get_quote(self, ticker: str, use_prepost: bool = True) -> Dict:
"""Get latest quote using yfinance-plus .info fields with fallback to fast history last row.""" """Get latest quote using yfinance-plus .info fields with fallback to fast history last row."""
@ -597,33 +584,8 @@ class PriceDataService:
logger.error(f"Error fetching ticker info for {ticker}: {str(e)}") logger.error(f"Error fetching ticker info for {ticker}: {str(e)}")
raise raise
async def get_multiple_tickers_data(
self,
db: AsyncSession,
tickers: List[str],
start_date: datetime,
end_date: datetime,
interval: str = "1d",
force_refresh: bool = False
) -> Dict[str, List[PriceData]]:
"""Get price data for multiple tickers (legacy method)"""
results = {}
for ticker in tickers:
try:
data = await self.get_or_update_price_data(
db, ticker, start_date, end_date, interval, force_refresh
)
results[ticker] = data
except Exception as e:
logger.error(f"Error fetching data for {ticker}: {str(e)}")
results[ticker] = []
return results
async def get_multiple_tickers_data_optimized( async def get_multiple_tickers_data_optimized(
self, self,
db: AsyncSession,
tickers: List[str], tickers: List[str],
start_date: datetime, start_date: datetime,
end_date: datetime, end_date: datetime,
@ -683,26 +645,33 @@ class PriceDataService:
logger.info(f"Processing chunk {chunk_num}/{total_chunks}: {len(chunk_tickers)} tickers") logger.info(f"Processing chunk {chunk_num}/{total_chunks}: {len(chunk_tickers)} tickers")
# Step 1: Batch check missing periods for chunk # Phase 1: batch check missing periods (short session)
missing_tickers = [] missing_tickers = []
if force_refresh: if force_refresh:
missing_tickers = chunk_tickers.copy() missing_tickers = chunk_tickers.copy()
else: else:
missing_tickers = await self._batch_check_missing_periods( async with AsyncSessionLocal() as db:
db, chunk_tickers, start_date, end_date, interval missing_tickers = await self._batch_check_missing_periods(
) db, chunk_tickers, start_date, end_date, interval
)
# Step 2: If we have missing data, use bulk yfinance fetch
# Phase 2: bulk yfinance fetch (no session held)
if missing_tickers and self.yf_available: if missing_tickers and self.yf_available:
logger.info(f"Bulk fetching price data for {len(missing_tickers)} tickers in chunk {chunk_num}") logger.info(f"Bulk fetching price data for {len(missing_tickers)} tickers in chunk {chunk_num}")
await self._bulk_fetch_and_store_price_data( chunk_data_list = await self._bulk_fetch_price_data(
db, missing_tickers, start_date, end_date, interval missing_tickers, start_date, end_date, interval
)
# Phase 3: store each sub-chunk with its own short-lived session
for sub_chunk_tickers, bulk_data in chunk_data_list:
async with AsyncSessionLocal() as db:
await self._process_bulk_data(db, sub_chunk_tickers, bulk_data, interval)
await db.commit()
# Phase 4: batch retrieve all data from DB (short session)
async with AsyncSessionLocal() as db:
ticker_data_map = await self._batch_get_price_data_from_db(
db, chunk_tickers, start_date, end_date, interval
) )
# Step 3: Batch retrieve all data from database for this chunk
ticker_data_map = await self._batch_get_price_data_from_db(
db, chunk_tickers, start_date, end_date, interval
)
# Step 4: Process results for this chunk # Step 4: Process results for this chunk
chunk_results = [] chunk_results = []
@ -792,7 +761,7 @@ class PriceDataService:
logger.error(f"Error in optimized bulk processing: {str(e)}") logger.error(f"Error in optimized bulk processing: {str(e)}")
# Fallback to individual processing # Fallback to individual processing
return await self._fallback_individual_processing( return await self._fallback_individual_processing(
db, tickers, start_date, end_date, interval, force_refresh tickers, start_date, end_date, interval, force_refresh
) )
async def _batch_check_missing_periods( async def _batch_check_missing_periods(
@ -861,72 +830,67 @@ class PriceDataService:
logger.info(f"Found {len(missing_tickers)} tickers needing data refresh out of {len(tickers)}") logger.info(f"Found {len(missing_tickers)} tickers needing data refresh out of {len(tickers)}")
return missing_tickers return missing_tickers
async def _bulk_fetch_and_store_price_data( async def _bulk_fetch_price_data(
self, self,
db: AsyncSession,
tickers: List[str], tickers: List[str],
start_date: datetime, start_date: datetime,
end_date: datetime, end_date: datetime,
interval: str interval: str
): ) -> List[Tuple[List[str], object]]:
"""Optimized bulk fetch using yfinance-plus bulk features""" """Fetch bulk price data from yfinance (no DB operations).
try:
logger.info(f"Starting bulk fetch for {len(tickers)} tickers") Returns a list of (chunk_tickers, bulk_data) tuples for the caller to
store with short-lived sessions.
# Convert dates to strings """
start_str = start_date.strftime('%Y-%m-%d') results = []
end_str = end_date.strftime('%Y-%m-%d') start_str = start_date.strftime('%Y-%m-%d')
end_str = end_date.strftime('%Y-%m-%d')
# Use yfinance-plus bulk download feature
loop = asyncio.get_event_loop() loop = asyncio.get_event_loop()
# Use adaptive chunk size for yfinance API calls based on ticker count total_tickers = len(tickers)
# Smaller chunks for yfinance API calls to avoid overwhelming the service if total_tickers <= 10:
total_tickers = len(tickers) chunk_size = total_tickers
if total_tickers <= 10: elif total_tickers <= 50:
chunk_size = total_tickers # Single chunk for very small batches chunk_size = 15
elif total_tickers <= 50: else:
chunk_size = 15 # Small chunks for moderate batches chunk_size = 20
else:
chunk_size = 20 # Standard chunks for large batches logger.info(f"Starting bulk fetch for {total_tickers} tickers")
for i in range(0, len(tickers), chunk_size):
chunk_tickers = tickers[i:i + chunk_size] for i in range(0, total_tickers, chunk_size):
chunk_tickers = tickers[i:i + chunk_size]
logger.info(f"Processing chunk {i//chunk_size + 1}: {len(chunk_tickers)} tickers") _chunk_str = ' '.join(chunk_tickers)
# Use yfinance-plus bulk download logger.info(f"Fetching chunk {i//chunk_size + 1}: {len(chunk_tickers)} tickers")
_chunk_str = ' '.join(chunk_tickers)
try:
bulk_data = await _run_with_timeout( bulk_data = await _run_with_timeout(
loop.run_in_executor( loop.run_in_executor(
None, None,
lambda: yf.download( # lambda default arg binds _chunk_str at definition time (closure bug fix)
tickers=_chunk_str, lambda cs=_chunk_str: yf.download(
tickers=cs,
start=start_str, start=start_str,
end=end_str, end=end_str,
interval=interval, interval=interval,
auto_adjust=True, auto_adjust=True,
prepost=False, prepost=False,
group_by='ticker', group_by='ticker',
threads=True # Enable multi-threading threads=True,
) )
), ),
timeout_seconds=60, timeout_seconds=60,
description=f"bulk_download {len(chunk_tickers)} tickers" description=f"bulk_download {len(chunk_tickers)} tickers"
) )
results.append((chunk_tickers, bulk_data))
# Process and store data for each ticker in the chunk except Exception as e:
await self._process_bulk_data(db, chunk_tickers, bulk_data, interval) logger.error(f"Error fetching chunk {i//chunk_size + 1}: {e}")
# Small delay to be nice to the API await asyncio.sleep(0.1)
await asyncio.sleep(0.1)
logger.info(f"Completed bulk fetch for {total_tickers} tickers ({len(results)} chunks succeeded)")
await db.commit() return results
logger.info(f"Successfully completed bulk fetch for {len(tickers)} tickers")
except Exception as e:
logger.error(f"Error in bulk fetch: {str(e)}")
await db.rollback()
raise
async def _process_bulk_data( async def _process_bulk_data(
self, self,
@ -1065,7 +1029,6 @@ class PriceDataService:
async def _fallback_individual_processing( async def _fallback_individual_processing(
self, self,
db: AsyncSession,
tickers: List[str], tickers: List[str],
start_date: datetime, start_date: datetime,
end_date: datetime, end_date: datetime,
@ -1074,18 +1037,18 @@ class PriceDataService:
) -> Tuple[List, int, int]: ) -> Tuple[List, int, int]:
"""Fallback to individual processing if bulk processing fails""" """Fallback to individual processing if bulk processing fails"""
from app.schemas.financial import BulkPriceDataItem, PriceDataResponse, PriceDataPoint from app.schemas.financial import BulkPriceDataItem, PriceDataResponse, PriceDataPoint
logger.warning("Falling back to individual ticker processing") logger.warning("Falling back to individual ticker processing")
results = [] results = []
successful_count = 0 successful_count = 0
failed_count = 0 failed_count = 0
for ticker in tickers: for ticker in tickers:
try: try:
# Get price data # Get price data
price_data = await self.get_or_update_price_data( price_data = await self.get_or_update_price_data(
db, ticker, start_date, end_date, interval, force_refresh ticker, start_date, end_date, interval, force_refresh
) )
# Convert to response models # Convert to response models

@ -170,7 +170,7 @@ class RealSECFinancialService:
""" """
try: try:
return await self.price_service.get_or_update_price_data( return await self.price_service.get_or_update_price_data(
db, ticker, start_date, end_date, "1d", force_refresh=False ticker, start_date, end_date, "1d", force_refresh=False
) )
except Exception as e: except Exception as e:
logger.warning( logger.warning(

Loading…
Cancel
Save