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,10 +116,8 @@ 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()
@ -153,7 +151,8 @@ 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)
async with AsyncSessionLocal() as db:
missing_periods = await price_service._check_missing_periods( missing_periods = await price_service._check_missing_periods(
db, request.ticker.upper(), start_date, end_date, request.interval db, request.ticker.upper(), start_date, end_date, request.interval
) )
@ -169,9 +168,8 @@ async def get_price_data(
# 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
@ -441,7 +437,7 @@ async def get_price_data_simple(
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,7 +519,6 @@ 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"""
@ -556,9 +551,6 @@ async def get_bulk_price_data(
# 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,
@ -573,7 +565,6 @@ async def get_bulk_price_data(
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,9 +74,11 @@ 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
async with AsyncSessionLocal() as db:
for i in range(0, len(rows), CHUNK): for i in range(0, len(rows), CHUNK):
stmt = pg_insert(AlpacaPriceData).values(rows[i : i + CHUNK]) stmt = pg_insert(AlpacaPriceData).values(rows[i : i + CHUNK])
stmt = stmt.on_conflict_do_nothing(constraint='uq_alpaca_price_data') stmt = stmt.on_conflict_do_nothing(constraint='uq_alpaca_price_data')
@ -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,13 +103,19 @@ 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:
async with AsyncSessionLocal() as db:
result = await db.execute( result = await db.execute(
select( select(
AlpacaPriceData.ticker, AlpacaPriceData.ticker,
@ -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,12 +174,11 @@ 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
async with AsyncSessionLocal() as db:
for i in range(0, len(rows), CHUNK): for i in range(0, len(rows), CHUNK):
stmt = pg_insert(AlpacaPriceData).values(rows[i : i + CHUNK]) stmt = pg_insert(AlpacaPriceData).values(rows[i : i + CHUNK])
stmt = stmt.on_conflict_do_nothing(constraint="uq_alpaca_price_data") stmt = stmt.on_conflict_do_nothing(constraint="uq_alpaca_price_data")
@ -182,7 +188,8 @@ class AlpacaPriceService:
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.
async with AsyncSessionLocal() as db:
result = await db.execute( result = await db.execute(
select(AlpacaPriceData) select(AlpacaPriceData)
.where( .where(

@ -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,22 +59,18 @@ 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)
async with AsyncSessionLocal() as db:
missing_periods = await self._check_missing_periods( missing_periods = await self._check_missing_periods(
db, ticker, start_date, end_date, interval db, ticker, start_date, end_date, interval
) )
@ -83,12 +79,17 @@ class PriceDataService:
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)
db, 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()
# Retrieve data from database # Phase 4: read from DB (short session)
async with AsyncSessionLocal() as db:
price_data = await self._get_price_data_from_db( price_data = await self._get_price_data_from_db(
db, ticker, start_date, end_date, interval db, ticker, start_date, end_date, interval
) )
@ -182,31 +183,26 @@ 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:
Returns a DataFrame or None if no data was returned.
"""
logger.info(f"Fetching price data for {ticker} from {start_date} to {end_date}") logger.info(f"Fetching price data for {ticker} from {start_date} to {end_date}")
# Create yfinance-plus ticker object
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 for daily data when using date strings. # yfinance's `end` parameter is exclusive — add +1 day to include end_date.
# Add +1 day to include the intended end_date day in the results.
from datetime import timedelta
end_inclusive = end_date + timedelta(days=1) end_inclusive = end_date + timedelta(days=1)
end_str = end_inclusive.strftime('%Y-%m-%d') end_str = end_inclusive.strftime('%Y-%m-%d')
# Run yfinance-plus in executor to avoid blocking
loop = asyncio.get_event_loop() loop = asyncio.get_event_loop()
hist_data = await _run_with_timeout( hist_data = await _run_with_timeout(
loop.run_in_executor( loop.run_in_executor(
@ -217,28 +213,19 @@ class PriceDataService:
interval=interval, interval=interval,
auto_adjust=True, auto_adjust=True,
prepost=False, prepost=False,
period=None # Explicitly set period to None when using start/end dates period=None
) )
), ),
timeout_seconds=30, timeout_seconds=30,
description=f"history {ticker} {start_str}:{end_str}" description=f"history {ticker} {start_str}:{end_str}"
) )
if hist_data.empty: if hist_data is None or hist_data.empty:
logger.warning(f"No price data returned for {ticker}") logger.warning(f"No price data returned for {ticker}")
return return None
# 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}") logger.info(f"Fetched {len(hist_data)} price records for {ticker}")
return hist_data
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,23 +645,30 @@ 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:
async with AsyncSessionLocal() as db:
missing_tickers = await self._batch_check_missing_periods( missing_tickers = await self._batch_check_missing_periods(
db, chunk_tickers, start_date, end_date, interval 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()
# Step 3: Batch retrieve all data from database for this chunk # 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( ticker_data_map = await self._batch_get_price_data_from_db(
db, chunk_tickers, start_date, end_date, interval db, chunk_tickers, start_date, end_date, interval
) )
@ -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")
# Convert dates to strings Returns a list of (chunk_tickers, bulk_data) tuples for the caller to
store with short-lived sessions.
"""
results = []
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')
# 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
# Smaller chunks for yfinance API calls to avoid overwhelming the service
total_tickers = len(tickers) total_tickers = len(tickers)
if total_tickers <= 10: if total_tickers <= 10:
chunk_size = total_tickers # Single chunk for very small batches chunk_size = total_tickers
elif total_tickers <= 50: elif total_tickers <= 50:
chunk_size = 15 # Small chunks for moderate batches chunk_size = 15
else: else:
chunk_size = 20 # Standard chunks for large batches chunk_size = 20
for i in range(0, len(tickers), chunk_size):
chunk_tickers = tickers[i:i + chunk_size]
logger.info(f"Processing chunk {i//chunk_size + 1}: {len(chunk_tickers)} tickers") logger.info(f"Starting bulk fetch for {total_tickers} tickers")
# Use yfinance-plus bulk download for i in range(0, total_tickers, chunk_size):
chunk_tickers = tickers[i:i + chunk_size]
_chunk_str = ' '.join(chunk_tickers) _chunk_str = ' '.join(chunk_tickers)
logger.info(f"Fetching chunk {i//chunk_size + 1}: {len(chunk_tickers)} 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))
except Exception as e:
logger.error(f"Error fetching chunk {i//chunk_size + 1}: {e}")
# Process and store data for each ticker in the chunk
await self._process_bulk_data(db, chunk_tickers, bulk_data, interval)
# Small delay to be nice to the API
await asyncio.sleep(0.1) await asyncio.sleep(0.1)
await db.commit() logger.info(f"Completed bulk fetch for {total_tickers} tickers ({len(results)} chunks succeeded)")
logger.info(f"Successfully completed bulk fetch for {len(tickers)} tickers") return results
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,
@ -1085,7 +1048,7 @@ class PriceDataService:
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