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 zoneinfo import ZoneInfo
from fastapi import APIRouter, Depends, HTTPException, Query
from sqlalchemy.ext.asyncio import AsyncSession
from fastapi import APIRouter, HTTPException, Query
_ET = ZoneInfo("America/New_York")
_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)
return now_et >= market_close_et
from app.core.database import get_db
from app.schemas.financial import (
AlpacaMultiBarsResponse,
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"),
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"),
db: AsyncSession = Depends(get_db),
):
"""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()]
@ -125,7 +122,7 @@ async def get_alpaca_intraday_multi(
try:
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:
err = str(e)
@ -182,7 +179,6 @@ async def get_alpaca_intraday_multi(
async def get_alpaca_intraday_today(
tickers: str = Query(..., description="Comma-separated tickers, e.g. AAPL,MSFT,BF-B"),
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)."""
symbols = [s.strip().upper() for s in tickers.split(",") if s.strip()]
@ -199,7 +195,7 @@ async def get_alpaca_intraday_today(
try:
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:
err = str(e)

@ -7,7 +7,7 @@ from typing import List, Optional
from fastapi import APIRouter, Depends, HTTPException, Query, Response
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 (
PriceDataRequest,
PriceDataResponse,
@ -116,10 +116,8 @@ router = APIRouter()
async def get_price_data(
request: PriceDataRequest,
response: Response,
db: AsyncSession = Depends(get_db)
):
"""Get price data for a ticker using period, quarters, or date range"""
try:
# Use the updated service that handles period resolution
price_service = PriceDataService()
@ -153,10 +151,11 @@ async def get_price_data(
response.headers["X-Data-Source"] = "redis-cache"
return cached_body
# Check if we have existing data to determine source
missing_periods = await price_service._check_missing_periods(
db, request.ticker.upper(), start_date, end_date, request.interval
)
# 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(
db, request.ticker.upper(), start_date, end_date, request.interval
)
# Determine data source
if request.force_refresh:
@ -169,9 +168,8 @@ async def get_price_data(
# Add data source header
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(
db,
request.ticker,
start_date,
end_date,
@ -314,7 +312,6 @@ async def get_multi_ticker_daily_bars(
end_date: date = Query(..., description="End date (YYYY-MM-DD)"),
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"),
db: AsyncSession = Depends(get_db),
):
"""Multi-ticker daily bars via Alpaca with DB storage (ORB engine interface)."""
symbols = [s.strip().upper() for s in tickers.split(",") if s.strip()]
@ -332,7 +329,7 @@ async def get_multi_ticker_daily_bars(
try:
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:
err = str(e)
@ -396,7 +393,6 @@ async def get_price_data_simple(
end: Optional[date] = Query(None, description="Alias for end_date"),
interval: str = Query("1d", description="Data interval: 1d, 1w, 1m, 5d, 1h, etc."),
force_refresh: bool = Query(False, description="Force refresh from Yahoo Finance"),
db: AsyncSession = Depends(get_db)
):
"""Simplified GET endpoint for price data"""
# Resolve aliases: start/end → start_date/end_date
@ -441,7 +437,7 @@ async def get_price_data_simple(
force_refresh=force_refresh
)
return await get_price_data(request, response, db)
return await get_price_data(request, response)
@router.post(
"/data/bulk",
@ -523,7 +519,6 @@ async def get_price_data_simple(
)
async def get_bulk_price_data(
request: BulkPriceDataRequest,
db: AsyncSession = Depends(get_db)
):
"""Get price data for multiple tickers"""
@ -556,9 +551,6 @@ async def get_bulk_price_data(
# Future date check
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:
raise HTTPException(
status_code=400,
@ -573,7 +565,6 @@ async def get_bulk_price_data(
try:
results, successful_count, failed_count = await _asyncio.wait_for(
price_service.get_multiple_tickers_data_optimized(
db=db,
tickers=request.tickers,
start_date=start_date,
end_date=end_date,

@ -6,10 +6,10 @@ 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.core.database import AsyncSessionLocal
from app.models.alpaca_price import AlpacaPriceData
from app.services.alpaca_client import AlpacaClient, normalize_ticker
@ -27,7 +27,6 @@ class AlpacaPriceService:
async def fetch_and_store_bars(
self,
db: AsyncSession,
ticker: str,
start_date: datetime,
end_date: datetime,
@ -43,6 +42,7 @@ class AlpacaPriceService:
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,
@ -54,7 +54,6 @@ class AlpacaPriceService:
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"])
@ -75,15 +74,17 @@ class AlpacaPriceService:
if not rows:
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
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()
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 db.commit()
if inserted:
logger.info(f"Alpaca: inserted {inserted} bars for {ticker}")
@ -92,7 +93,6 @@ class AlpacaPriceService:
async def get_or_fetch_multi_bars(
self,
db: AsyncSession,
tickers: List[str],
start_dt: datetime,
end_dt: datetime,
@ -103,29 +103,35 @@ class AlpacaPriceService:
"""
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].
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:
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,
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)
)
.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 = [
t for t in upper_tickers
if coverage.get(t) is None or coverage[t].date() < end_dt.date()
@ -133,6 +139,7 @@ class AlpacaPriceService:
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")
@ -167,35 +174,35 @@ class AlpacaPriceService:
"data_source": "ALPACA",
})
# Phase 3: upsert (short session)
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
# 14 params/row → CHUNK=2300 (2300×14=32,200 < 32,767 asyncpg limit)
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()
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 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,
# Phase 4: read back (short session) — filtered by interval to avoid mixing 1d/5m/etc.
async with AsyncSessionLocal() as db:
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)
)
.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}
for row in db_rows:

@ -469,7 +469,7 @@ class FinancialService:
if not price_data:
try:
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:
logger.warning(f"Could not fetch price data for {ticker}: {e}")

@ -14,6 +14,7 @@ import os
# Add parent directory to path for imports
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.schemas.financial import DataSource, ErrorType
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(
self,
db: AsyncSession,
ticker: str,
start_date: datetime,
end_date: datetime,
@ -59,40 +59,41 @@ class PriceDataService:
force_refresh: bool = False
) -> 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:
db: Database session
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
Session-per-phase: DB connections are held only during short DB operations,
never during yfinance calls (which can take 30s+).
Returns:
List of PriceData objects
"""
ticker = ticker.upper()
# Check if we need to fetch new data
missing_periods = await self._check_missing_periods(
db, ticker, start_date, end_date, interval
)
# Phase 1: check missing periods (short session)
async with AsyncSessionLocal() as db:
missing_periods = await self._check_missing_periods(
db, ticker, start_date, end_date, interval
)
if missing_periods or force_refresh:
if not self.yf_available:
raise ValueError("Yahoo Finance (yfinance-plus) data source not available")
# Fetch data from Yahoo Finance using yfinance-plus
await self._fetch_and_store_price_data(
# Phase 2: fetch from yfinance (no session held)
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
)
# Retrieve data from database
price_data = await self._get_price_data_from_db(
db, ticker, start_date, end_date, interval
)
return price_data
async def _check_missing_periods(
@ -182,63 +183,49 @@ class PriceDataService:
return expected_dates
async def _fetch_and_store_price_data(
async def _fetch_price_data(
self,
db: AsyncSession,
ticker: str,
start_date: datetime,
end_date: datetime,
interval: str
):
"""Fetch price data from Yahoo Finance using yfinance-plus and store in database"""
try:
logger.info(f"Fetching price data for {ticker} from {start_date} to {end_date}")
# Create yfinance-plus ticker object
yf_ticker = yf.Ticker(ticker)
"""Fetch price data from Yahoo Finance (no DB operations).
# Fetch historical data
# Convert dates to strings in YYYY-MM-DD format
start_str = start_date.strftime('%Y-%m-%d')
# yfinance's `end` parameter is exclusive for daily data when using date strings.
# 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_str = end_inclusive.strftime('%Y-%m-%d')
# Run yfinance-plus in executor to avoid blocking
loop = asyncio.get_event_loop()
hist_data = await _run_with_timeout(
loop.run_in_executor(
None,
lambda: yf_ticker.history(
start=start_str,
end=end_str,
interval=interval,
auto_adjust=True,
prepost=False,
period=None # Explicitly set period to None when using start/end dates
)
),
timeout_seconds=30,
description=f"history {ticker} {start_str}:{end_str}"
)
Returns a DataFrame or None if no data was returned.
"""
logger.info(f"Fetching price data for {ticker} from {start_date} to {end_date}")
if hist_data.empty:
logger.warning(f"No price data returned for {ticker}")
return
yf_ticker = yf.Ticker(ticker)
# Store data in database
await self._store_price_data(db, ticker, hist_data, interval)
start_str = start_date.strftime('%Y-%m-%d')
# yfinance's `end` parameter is exclusive — add +1 day to include end_date.
end_inclusive = end_date + timedelta(days=1)
end_str = end_inclusive.strftime('%Y-%m-%d')
await db.commit()
loop = asyncio.get_event_loop()
hist_data = await _run_with_timeout(
loop.run_in_executor(
None,
lambda: yf_ticker.history(
start=start_str,
end=end_str,
interval=interval,
auto_adjust=True,
prepost=False,
period=None
)
),
timeout_seconds=30,
description=f"history {ticker} {start_str}:{end_str}"
)
logger.info(f"Successfully stored {len(hist_data)} price records for {ticker}")
if hist_data is None or hist_data.empty:
logger.warning(f"No price data returned for {ticker}")
return None
except Exception as e:
logger.error(f"Error fetching price data for {ticker}: {str(e)}")
await db.rollback()
raise
logger.info(f"Fetched {len(hist_data)} price records for {ticker}")
return hist_data
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."""
@ -597,33 +584,8 @@ class PriceDataService:
logger.error(f"Error fetching ticker info for {ticker}: {str(e)}")
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(
self,
db: AsyncSession,
tickers: List[str],
start_date: datetime,
end_date: datetime,
@ -683,26 +645,33 @@ class PriceDataService:
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 = []
if force_refresh:
missing_tickers = chunk_tickers.copy()
else:
missing_tickers = await self._batch_check_missing_periods(
db, chunk_tickers, start_date, end_date, interval
)
async with AsyncSessionLocal() as db:
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:
logger.info(f"Bulk fetching price data for {len(missing_tickers)} tickers in chunk {chunk_num}")
await self._bulk_fetch_and_store_price_data(
db, missing_tickers, start_date, end_date, interval
chunk_data_list = await self._bulk_fetch_price_data(
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
chunk_results = []
@ -792,7 +761,7 @@ class PriceDataService:
logger.error(f"Error in optimized bulk processing: {str(e)}")
# Fallback to 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(
@ -861,72 +830,67 @@ class PriceDataService:
logger.info(f"Found {len(missing_tickers)} tickers needing data refresh out of {len(tickers)}")
return missing_tickers
async def _bulk_fetch_and_store_price_data(
async def _bulk_fetch_price_data(
self,
db: AsyncSession,
tickers: List[str],
start_date: datetime,
end_date: datetime,
interval: str
):
"""Optimized bulk fetch using yfinance-plus bulk features"""
try:
logger.info(f"Starting bulk fetch for {len(tickers)} tickers")
) -> List[Tuple[List[str], object]]:
"""Fetch bulk price data from yfinance (no DB operations).
# Convert dates to strings
start_str = start_date.strftime('%Y-%m-%d')
end_str = end_date.strftime('%Y-%m-%d')
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')
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)
if total_tickers <= 10:
chunk_size = total_tickers # Single chunk for very small batches
elif total_tickers <= 50:
chunk_size = 15 # Small chunks for moderate batches
else:
chunk_size = 20 # Standard chunks for large batches
for i in range(0, len(tickers), chunk_size):
chunk_tickers = tickers[i:i + chunk_size]
total_tickers = len(tickers)
if total_tickers <= 10:
chunk_size = total_tickers
elif total_tickers <= 50:
chunk_size = 15
else:
chunk_size = 20
logger.info(f"Starting bulk fetch for {total_tickers} tickers")
for i in range(0, total_tickers, chunk_size):
chunk_tickers = tickers[i:i + chunk_size]
_chunk_str = ' '.join(chunk_tickers)
logger.info(f"Processing chunk {i//chunk_size + 1}: {len(chunk_tickers)} tickers")
logger.info(f"Fetching chunk {i//chunk_size + 1}: {len(chunk_tickers)} tickers")
# Use yfinance-plus bulk download
_chunk_str = ' '.join(chunk_tickers)
try:
bulk_data = await _run_with_timeout(
loop.run_in_executor(
None,
lambda: yf.download(
tickers=_chunk_str,
# lambda default arg binds _chunk_str at definition time (closure bug fix)
lambda cs=_chunk_str: yf.download(
tickers=cs,
start=start_str,
end=end_str,
interval=interval,
auto_adjust=True,
prepost=False,
group_by='ticker',
threads=True # Enable multi-threading
threads=True,
)
),
timeout_seconds=60,
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 db.commit()
logger.info(f"Successfully completed bulk fetch for {len(tickers)} tickers")
await asyncio.sleep(0.1)
except Exception as e:
logger.error(f"Error in bulk fetch: {str(e)}")
await db.rollback()
raise
logger.info(f"Completed bulk fetch for {total_tickers} tickers ({len(results)} chunks succeeded)")
return results
async def _process_bulk_data(
self,
@ -1065,7 +1029,6 @@ class PriceDataService:
async def _fallback_individual_processing(
self,
db: AsyncSession,
tickers: List[str],
start_date: datetime,
end_date: datetime,
@ -1085,7 +1048,7 @@ class PriceDataService:
try:
# Get 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

@ -170,7 +170,7 @@ class RealSECFinancialService:
"""
try:
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:
logger.warning(

Loading…
Cancel
Save