diff --git a/app/api/v1/endpoints/alpaca.py b/app/api/v1/endpoints/alpaca.py index e47f0bc..4b56722 100644 --- a/app/api/v1/endpoints/alpaca.py +++ b/app/api/v1/endpoints/alpaca.py @@ -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) diff --git a/app/api/v1/endpoints/price.py b/app/api/v1/endpoints/price.py index 60652d1..2285e80 100644 --- a/app/api/v1/endpoints/price.py +++ b/app/api/v1/endpoints/price.py @@ -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,14 +116,12 @@ 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() - + # Resolve time parameters to get start and end dates from app.utils.date_utils import 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" 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: data_source = "yfinance-fresh" @@ -165,13 +164,12 @@ async def get_price_data( data_source = "yfinance-partial" else: data_source = "database-cache" - + # 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 @@ -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." } ) - + if not period and not (start_date and end_date): raise HTTPException( 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." } ) - + # Create request based on provided parameters if period: request = PriceDataRequest( @@ -440,8 +436,8 @@ async def get_price_data_simple( interval=interval, force_refresh=force_refresh ) - - return await get_price_data(request, response, db) + + return await get_price_data(request, response) @router.post( "/data/bulk", @@ -523,13 +519,12 @@ 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""" - + # Use the updated service that handles period resolution price_service = PriceDataService() - + # Resolve time parameters to get start and end dates from app.utils.date_utils import resolve_time_parameters start_date, end_date = resolve_time_parameters( @@ -538,13 +533,13 @@ async def get_bulk_price_data( quarters=request.quarters, period=request.period ) - + # Validate date range - ensure both dates are timezone-aware if start_date and start_date.tzinfo is None: start_date = start_date.replace(tzinfo=timezone.utc) if end_date and end_date.tzinfo is None: end_date = end_date.replace(tzinfo=timezone.utc) - + if start_date and end_date and start_date >= end_date: raise HTTPException( status_code=400, @@ -553,12 +548,9 @@ async def get_bulk_price_data( "message": "Start date must be before end date" } ) - + # 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, @@ -567,13 +559,12 @@ async def get_bulk_price_data( "message": "Cannot request data for future dates" } ) - + # Use optimized bulk processing method (300s endpoint-level timeout) import asyncio as _asyncio 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, diff --git a/app/services/alpaca_price_service.py b/app/services/alpaca_price_service.py index 16a6ee9..7a4d24f 100644 --- a/app/services/alpaca_price_service.py +++ b/app/services/alpaca_price_service.py @@ -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: diff --git a/app/services/financial_service.py b/app/services/financial_service.py index f7e4857..2ab250d 100644 --- a/app/services/financial_service.py +++ b/app/services/financial_service.py @@ -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}") diff --git a/app/services/price_data_service.py b/app/services/price_data_service.py index 3e842ad..ee3b020 100644 --- a/app/services/price_data_service.py +++ b/app/services/price_data_service.py @@ -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 - - 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 - + Get price data from database or fetch from Yahoo Finance if needed. + + 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 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}" - ) - - if hist_data.empty: - logger.warning(f"No price data returned for {ticker}") - 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 + """Fetch price data from Yahoo Finance (no DB operations). + + Returns a DataFrame or None if no data was returned. + """ + logger.info(f"Fetching price data for {ticker} from {start_date} to {end_date}") + + yf_ticker = yf.Ticker(ticker) + + 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') + + 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}" + ) + + if hist_data is None or hist_data.empty: + logger.warning(f"No price data returned for {ticker}") + return None + + 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 - ) - - # Step 2: If we have missing data, use bulk yfinance fetch + async with AsyncSessionLocal() as db: + missing_tickers = await self._batch_check_missing_periods( + db, chunk_tickers, start_date, end_date, interval + ) + + # 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") - - # Convert dates to strings - 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() - - # 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] - - logger.info(f"Processing chunk {i//chunk_size + 1}: {len(chunk_tickers)} tickers") - - # Use yfinance-plus bulk download - _chunk_str = ' '.join(chunk_tickers) + ) -> List[Tuple[List[str], object]]: + """Fetch bulk price data from yfinance (no DB operations). + + 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') + + loop = asyncio.get_event_loop() + + 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"Fetching chunk {i//chunk_size + 1}: {len(chunk_tickers)} 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" ) - - # 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") - - except Exception as e: - logger.error(f"Error in bulk fetch: {str(e)}") - await db.rollback() - raise + results.append((chunk_tickers, bulk_data)) + except Exception as e: + logger.error(f"Error fetching chunk {i//chunk_size + 1}: {e}") + + await asyncio.sleep(0.1) + + 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, @@ -1074,18 +1037,18 @@ class PriceDataService: ) -> Tuple[List, int, int]: """Fallback to individual processing if bulk processing fails""" from app.schemas.financial import BulkPriceDataItem, PriceDataResponse, PriceDataPoint - + logger.warning("Falling back to individual ticker processing") - + results = [] successful_count = 0 failed_count = 0 - + for ticker in tickers: 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 diff --git a/app/services/real_sec_financial_service.py b/app/services/real_sec_financial_service.py index 5589113..94baed2 100644 --- a/app/services/real_sec_financial_service.py +++ b/app/services/real_sec_financial_service.py @@ -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(