diff --git a/Dockerfile b/Dockerfile index 851e680..75d2ab5 100644 --- a/Dockerfile +++ b/Dockerfile @@ -24,4 +24,4 @@ RUN mkdir -p /app/data EXPOSE 18000 # Run the application -CMD ["python", "-m", "uvicorn", "app.main:app", "--host", "0.0.0.0", "--port", "18000", "--reload"] \ No newline at end of file +CMD ["python", "-m", "uvicorn", "app.main:app", "--host", "0.0.0.0", "--port", "18000", "--workers", "4"] \ No newline at end of file diff --git a/app/api/v1/endpoints/attention.py b/app/api/v1/endpoints/attention.py index 68cd109..2f1ccac 100644 --- a/app/api/v1/endpoints/attention.py +++ b/app/api/v1/endpoints/attention.py @@ -16,11 +16,12 @@ Admin endpoints: import logging from datetime import date, timedelta -from fastapi import APIRouter, Depends, HTTPException, Query +from fastapi import APIRouter, Depends, HTTPException, Query, Response from sqlalchemy import func, select from sqlalchemy.ext.asyncio import AsyncSession from app.core.database import get_db +from app.utils.cache import with_cache from app.models.attention import AttentionFeaturesDaily, CompanyEntityMap, GdeltArticleRaw from app.services.attention.gdelt_collector import GDELT_EARLIEST_DATE from app.schemas.attention import ( @@ -237,8 +238,10 @@ async def admin_collect_gdelt( **Example**: `GET /attention/entity/AAPL` """, ) +@with_cache(namespace="attention:entity", ttl=3600, key_params=["ticker"]) async def get_entity( ticker: str, + response: Response = None, db: AsyncSession = Depends(get_db), ) -> EntityResolveResponse: ticker = ticker.upper() @@ -297,9 +300,11 @@ async def get_entity( 500: {"description": "Feature materialization or collection error"}, }, ) +@with_cache(namespace="attention:event", ttl=3600, key_params=["ticker", "event_date"]) async def get_event_attention( ticker: str, event_date: date = Query(..., description="Event date in YYYY-MM-DD format"), + response: Response = None, db: AsyncSession = Depends(get_db), ) -> EventAttentionResponse: ticker = ticker.upper() diff --git a/app/core/config.py b/app/core/config.py index 159704d..e6fb778 100644 --- a/app/core/config.py +++ b/app/core/config.py @@ -41,6 +41,8 @@ class Settings(BaseSettings): "postgresql+asyncpg://stockoracle:stockoracle2024@localhost:15433/stock_oracle" ) DATABASE_ECHO: bool = False + DB_POOL_SIZE: int = int(os.getenv("DB_POOL_SIZE", "20")) + DB_MAX_OVERFLOW: int = int(os.getenv("DB_MAX_OVERFLOW", "30")) # Redis REDIS_URL: str = os.getenv("REDIS_URL", "redis://localhost:16380/0") diff --git a/app/core/database.py b/app/core/database.py index 0217c78..863583e 100644 --- a/app/core/database.py +++ b/app/core/database.py @@ -11,8 +11,8 @@ engine = create_async_engine( settings.DATABASE_URL, echo=settings.DATABASE_ECHO, future=True, - pool_size=10, - max_overflow=20, + pool_size=settings.DB_POOL_SIZE, + max_overflow=settings.DB_MAX_OVERFLOW, pool_timeout=30, pool_pre_ping=True, pool_recycle=3600, diff --git a/app/core/http_client.py b/app/core/http_client.py new file mode 100644 index 0000000..b760026 --- /dev/null +++ b/app/core/http_client.py @@ -0,0 +1,33 @@ +""" +Shared aiohttp ClientSession singleton. + +Reuses a single TCP connection pool across all HTTP clients, +avoiding repeated TCP/TLS handshake overhead per request. +""" + +import logging +from typing import Optional + +import aiohttp + +logger = logging.getLogger(__name__) + +_session: Optional[aiohttp.ClientSession] = None + + +async def get_http_session() -> aiohttp.ClientSession: + """Get or create the shared aiohttp ClientSession.""" + global _session + if _session is None or _session.closed: + _session = aiohttp.ClientSession( + timeout=aiohttp.ClientTimeout(total=30), + ) + return _session + + +async def close_http_session() -> None: + """Close the shared session. Call on app shutdown.""" + global _session + if _session and not _session.closed: + await _session.close() + _session = None diff --git a/app/main.py b/app/main.py index bce2c01..645d4b7 100644 --- a/app/main.py +++ b/app/main.py @@ -6,6 +6,7 @@ import os from contextlib import asynccontextmanager from fastapi import FastAPI, HTTPException from fastapi.middleware.cors import CORSMiddleware +from starlette.middleware.gzip import GZipMiddleware from fastapi.responses import RedirectResponse from app.core.config import settings @@ -36,6 +37,11 @@ async def lifespan(app: FastAPI): stop_scheduler() except Exception: pass + try: + from app.core.http_client import close_http_session + await close_http_session() + except Exception: + pass await engine.dispose() # Create FastAPI app @@ -61,6 +67,9 @@ app.add_middleware( allow_headers=["*"], ) +# GZip compression for responses > 1KB +app.add_middleware(GZipMiddleware, minimum_size=1000) + # Include API router app.include_router(api_router, prefix=settings.API_PREFIX) diff --git a/app/middleware/error_logger.py b/app/middleware/error_logger.py index 0637cf6..c578961 100644 --- a/app/middleware/error_logger.py +++ b/app/middleware/error_logger.py @@ -105,8 +105,8 @@ class ErrorLoggingMiddleware(BaseHTTPMiddleware): start_time = time.time() # Store request details for potential error logging - request_info = await self._extract_request_info(request) - + request_info = self._extract_request_info_minimal(request) + logger.info(f"Processing request {request_id}: {request.method} {request.url.path}") try: @@ -123,7 +123,10 @@ class ErrorLoggingMiddleware(BaseHTTPMiddleware): # Check if response indicates an error (4xx or 5xx) if response.status_code >= 400: logger.info(f"Error response detected: {response.status_code} for request {request_id}") - + + # Full extraction for error logging (body + headers) + request_info = await self._extract_request_info(request) + # Try to capture response body for errors # We need to consume the response body and recreate it from starlette.responses import Response @@ -183,7 +186,10 @@ class ErrorLoggingMiddleware(BaseHTTPMiddleware): except Exception as e: # Log unexpected errors response_time_ms = (time.time() - start_time) * 1000 - + + # Full extraction for error logging + request_info = await self._extract_request_info(request) + await self._log_error( request_id=request_id, request_info=request_info, @@ -206,8 +212,21 @@ class ErrorLoggingMiddleware(BaseHTTPMiddleware): headers={"X-Request-ID": request_id} ) + def _extract_request_info_minimal(self, request: Request) -> dict: + """Extract minimal request info (sync, no body read) for success-path logging.""" + return { + "endpoint": str(request.url.path), + "method": request.method, + "path": str(request.url), + "query_params": dict(request.query_params) if request.query_params else None, + "request_body": None, + "headers": None, + "user_agent": request.headers.get("user-agent"), + "client_ip": request.client.host if request.client else None, + } + async def _extract_request_info(self, request: Request) -> dict: - """Extract request information for logging""" + """Extract full request information for error logging""" # Get request body if present body = None diff --git a/app/services/finra_short_volume_service.py b/app/services/finra_short_volume_service.py index 4cf6a9b..4ae7b4c 100644 --- a/app/services/finra_short_volume_service.py +++ b/app/services/finra_short_volume_service.py @@ -10,6 +10,8 @@ import aiohttp from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy import select, and_, func, desc +from app.core.http_client import get_http_session + from app.models.finra_short_volume import FinraShortVolume logger = logging.getLogger(__name__) @@ -34,15 +36,15 @@ class FinraShortVolumeService: url = f"{FINRA_CDN_BASE}/{filename}" try: - async with aiohttp.ClientSession() as session: - async with session.get(url, timeout=aiohttp.ClientTimeout(total=30)) as resp: - if resp.status == 404: - logger.debug(f"FINRA file not found (holiday/weekend): {filename}") - return None - resp.raise_for_status() - text = await resp.text() - logger.info(f"FINRA: downloaded {filename} ({len(text)} bytes)") - return text + session = await get_http_session() + async with session.get(url, timeout=aiohttp.ClientTimeout(total=30)) as resp: + if resp.status == 404: + logger.debug(f"FINRA file not found (holiday/weekend): {filename}") + return None + resp.raise_for_status() + text = await resp.text() + logger.info(f"FINRA: downloaded {filename} ({len(text)} bytes)") + return text except aiohttp.ClientError as e: logger.error(f"FINRA download error for {filename}: {e}") return None diff --git a/app/services/news_social_service.py b/app/services/news_social_service.py index 58edf5f..bfffab7 100644 --- a/app/services/news_social_service.py +++ b/app/services/news_social_service.py @@ -10,6 +10,8 @@ Aggregates news and social media data from multiple sources for sentiment analys import asyncio import aiohttp import logging + +from app.core.http_client import get_http_session from datetime import datetime, timedelta from typing import Dict, List, Optional, Any, Union from dataclasses import dataclass @@ -326,42 +328,42 @@ class NewsSocialService: "language": "en" } - async with aiohttp.ClientSession() as session: - async with session.get(url, params=params) as response: - if response.status == 200: - data = await response.json() - - articles = [] - for item in data.get('articles', []): - try: - # Parse NewsAPI format - make timezone naive for consistency - published_str = item['publishedAt'].replace('Z', '+00:00') - published_at = datetime.fromisoformat(published_str).replace(tzinfo=None) - - article = NewsArticle( - title=item.get('title', ''), - summary=item.get('description', ''), - content=item.get('content', ''), - url=item.get('url', ''), - source="NewsAPI", - published_at=published_at, - author=item.get('author', ''), - image_url=item.get('urlToImage', '') - ) - - articles.append(article) - - except Exception as e: - logger.warning(f"Error parsing NewsAPI article: {e}") - continue - - logger.info(f"Retrieved {len(articles)} articles from NewsAPI") - return articles - - else: - error_data = await response.json() - logger.error(f"NewsAPI error {response.status}: {error_data}") - raise NewsAPIError(f"NewsAPI returned {response.status}: {error_data}") + session = await get_http_session() + async with session.get(url, params=params) as response: + if response.status == 200: + data = await response.json() + + articles = [] + for item in data.get('articles', []): + try: + # Parse NewsAPI format - make timezone naive for consistency + published_str = item['publishedAt'].replace('Z', '+00:00') + published_at = datetime.fromisoformat(published_str).replace(tzinfo=None) + + article = NewsArticle( + title=item.get('title', ''), + summary=item.get('description', ''), + content=item.get('content', ''), + url=item.get('url', ''), + source="NewsAPI", + published_at=published_at, + author=item.get('author', ''), + image_url=item.get('urlToImage', '') + ) + + articles.append(article) + + except Exception as e: + logger.warning(f"Error parsing NewsAPI article: {e}") + continue + + logger.info(f"Retrieved {len(articles)} articles from NewsAPI") + return articles + + else: + error_data = await response.json() + logger.error(f"NewsAPI error {response.status}: {error_data}") + raise NewsAPIError(f"NewsAPI returned {response.status}: {error_data}") except Exception as e: logger.error(f"Error fetching NewsAPI articles for {ticker}: {e}") @@ -405,53 +407,53 @@ class NewsSocialService: "User-Agent": "StockOracle/1.0.0" } - async with aiohttp.ClientSession() as session: - async with session.get(url, params=params, headers=headers) as response: - if response.status == 200: - data = await response.json() - - for item in data.get('data', {}).get('children', []): - try: - post_data = item.get('data', {}) - - # Filter out posts that are too old - created_utc = post_data.get('created_utc', 0) - post_date = datetime.fromtimestamp(created_utc) - - if (datetime.now() - post_date).days > days_back: - continue - - # Skip removed/deleted posts - if post_data.get('removed_by_category') or post_data.get('selftext') == '[removed]': - continue - - post = SocialPost( - title=post_data.get('title', ''), - content=post_data.get('selftext', ''), - url=f"https://reddit.com{post_data.get('permalink', '')}", - platform="Reddit", - author=post_data.get('author', ''), - published_at=post_date, - score=post_data.get('score', 0), - comments_count=post_data.get('num_comments', 0), - upvotes=post_data.get('ups', 0), - downvotes=post_data.get('downs', 0), - subreddit=post_data.get('subreddit', '') - ) - - all_posts.append(post) - - except Exception as e: - logger.warning(f"Error parsing Reddit post: {e}") + session = await get_http_session() + async with session.get(url, params=params, headers=headers) as response: + if response.status == 200: + data = await response.json() + + for item in data.get('data', {}).get('children', []): + try: + post_data = item.get('data', {}) + + # Filter out posts that are too old + created_utc = post_data.get('created_utc', 0) + post_date = datetime.fromtimestamp(created_utc) + + if (datetime.now() - post_date).days > days_back: + continue + + # Skip removed/deleted posts + if post_data.get('removed_by_category') or post_data.get('selftext') == '[removed]': continue - - elif response.status == 401: - logger.error("Reddit API authentication failed") - # Try to refresh token - self._reddit_token = None - await self._ensure_reddit_token() - else: - logger.warning(f"Reddit API error for r/{subreddit}: {response.status}") + + post = SocialPost( + title=post_data.get('title', ''), + content=post_data.get('selftext', ''), + url=f"https://reddit.com{post_data.get('permalink', '')}", + platform="Reddit", + author=post_data.get('author', ''), + published_at=post_date, + score=post_data.get('score', 0), + comments_count=post_data.get('num_comments', 0), + upvotes=post_data.get('ups', 0), + downvotes=post_data.get('downs', 0), + subreddit=post_data.get('subreddit', '') + ) + + all_posts.append(post) + + except Exception as e: + logger.warning(f"Error parsing Reddit post: {e}") + continue + + elif response.status == 401: + logger.error("Reddit API authentication failed") + # Try to refresh token + self._reddit_token = None + await self._ensure_reddit_token() + else: + logger.warning(f"Reddit API error for r/{subreddit}: {response.status}") except Exception as e: logger.warning(f"Error fetching from r/{subreddit}: {e}") @@ -497,21 +499,21 @@ class NewsSocialService: auth = aiohttp.BasicAuth(self.reddit_client_id, self.reddit_client_secret) - async with aiohttp.ClientSession() as session: - async with session.post(auth_url, data=auth_data, auth=auth, headers=headers) as response: - if response.status == 200: - token_data = await response.json() - - self._reddit_token = token_data.get('access_token') - expires_in = token_data.get('expires_in', 3600) - self._reddit_token_expiry = datetime.now() + timedelta(seconds=expires_in - 60) - - logger.info("Successfully obtained Reddit access token") - - else: - error_data = await response.text() - logger.error(f"Reddit auth error {response.status}: {error_data}") - raise RedditAPIError(f"Failed to authenticate with Reddit: {response.status}") + session = await get_http_session() + async with session.post(auth_url, data=auth_data, auth=auth, headers=headers) as response: + if response.status == 200: + token_data = await response.json() + + self._reddit_token = token_data.get('access_token') + expires_in = token_data.get('expires_in', 3600) + self._reddit_token_expiry = datetime.now() + timedelta(seconds=expires_in - 60) + + logger.info("Successfully obtained Reddit access token") + + else: + error_data = await response.text() + logger.error(f"Reddit auth error {response.status}: {error_data}") + raise RedditAPIError(f"Failed to authenticate with Reddit: {response.status}") except Exception as e: logger.error(f"Error obtaining Reddit token: {e}") diff --git a/app/services/overlay/yahoo_rss_adapter.py b/app/services/overlay/yahoo_rss_adapter.py index bdc7cae..215da30 100644 --- a/app/services/overlay/yahoo_rss_adapter.py +++ b/app/services/overlay/yahoo_rss_adapter.py @@ -11,6 +11,8 @@ import aiohttp from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy import select +from app.core.http_client import get_http_session + from app.models.overlay_raw_event import OverlayHeadlineEvent from app.services.overlay.entity_resolver import EntityResolver from app.core.overlay_config import YAHOO_RSS_FEEDS, TOP_50_SYMBOLS @@ -33,18 +35,18 @@ class YahooRSSAdapter: return [] try: - async with aiohttp.ClientSession() as session: - headers = { - "User-Agent": "Mozilla/5.0 StockOracle/1.0", - "Accept": "application/rss+xml, application/xml, text/xml", - } - async with session.get( - url, headers=headers, timeout=aiohttp.ClientTimeout(total=30) - ) as resp: - if resp.status != 200: - logger.warning(f"Yahoo RSS: non-200 from {url}: {resp.status}") - return [] - text = await resp.text() + session = await get_http_session() + headers = { + "User-Agent": "Mozilla/5.0 StockOracle/1.0", + "Accept": "application/rss+xml, application/xml, text/xml", + } + async with session.get( + url, headers=headers, timeout=aiohttp.ClientTimeout(total=30) + ) as resp: + if resp.status != 200: + logger.warning(f"Yahoo RSS: non-200 from {url}: {resp.status}") + return [] + text = await resp.text() feed = feedparser.parse(text) entries = [] diff --git a/app/services/overlay/youtube_adapter.py b/app/services/overlay/youtube_adapter.py index 3d7558c..46ef14d 100644 --- a/app/services/overlay/youtube_adapter.py +++ b/app/services/overlay/youtube_adapter.py @@ -9,6 +9,8 @@ from typing import Dict, List, Optional import aiohttp from sqlalchemy.ext.asyncio import AsyncSession + +from app.core.http_client import get_http_session from sqlalchemy import select from app.core.config import settings @@ -53,15 +55,15 @@ class YouTubeAdapter: "key": self.api_key, } try: - async with aiohttp.ClientSession() as session: - async with session.get( - YOUTUBE_SEARCH_URL, params=params, timeout=aiohttp.ClientTimeout(total=30) - ) as resp: - if resp.status == 403: - logger.warning("YouTube API: quota exceeded or invalid key") - return [] - resp.raise_for_status() - data = await resp.json() + session = await get_http_session() + async with session.get( + YOUTUBE_SEARCH_URL, params=params, timeout=aiohttp.ClientTimeout(total=30) + ) as resp: + if resp.status == 403: + logger.warning("YouTube API: quota exceeded or invalid key") + return [] + resp.raise_for_status() + data = await resp.json() return data.get("items", []) except Exception as e: logger.error(f"YouTube search error for channel {channel_id}: {e}") @@ -77,12 +79,12 @@ class YouTubeAdapter: "key": self.api_key, } try: - async with aiohttp.ClientSession() as session: - async with session.get( - YOUTUBE_VIDEOS_URL, params=params, timeout=aiohttp.ClientTimeout(total=30) - ) as resp: - resp.raise_for_status() - data = await resp.json() + session = await get_http_session() + async with session.get( + YOUTUBE_VIDEOS_URL, params=params, timeout=aiohttp.ClientTimeout(total=30) + ) as resp: + resp.raise_for_status() + data = await resp.json() stats = {} for item in data.get("items", []): vid_id = item["id"] diff --git a/app/services/price_data_service.py b/app/services/price_data_service.py index e1e0b81..b2716b9 100644 --- a/app/services/price_data_service.py +++ b/app/services/price_data_service.py @@ -588,7 +588,12 @@ class PriceDataService: # Convert to response models price_points = [ - PriceDataPoint.model_validate(pd) for pd in price_data + PriceDataPoint.model_construct( + date=pd.date.date() if isinstance(pd.date, datetime) else pd.date, + open=pd.open, high=pd.high, low=pd.low, close=pd.close, + volume=pd.volume, adjusted_close=pd.adjusted_close, + data_source=pd.data_source.value if hasattr(pd.data_source, 'value') else pd.data_source, + ) for pd in price_data ] # Calculate actual date range from returned data @@ -821,39 +826,48 @@ class PriceDataService: ticker_data, interval: str ): - """Store individual ticker data with optimized batch operations""" - # Batch check existing dates to avoid individual DB queries - existing_dates = await self._get_existing_dates_for_ticker(db, ticker) - - new_records = [] - for date, row in ticker_data.iterrows(): - # Convert pandas timestamp to datetime - price_date = date.to_pydatetime() + """Store individual ticker data using batch upsert (INSERT ... ON CONFLICT DO NOTHING).""" + from sqlalchemy.dialects.postgresql import insert as pg_insert + import uuid as _uuid + + def _safe(val): + if val is None: + return None + try: + v = float(val) + return None if pd.isna(v) else v + except (TypeError, ValueError): + return None + + now = datetime.now(timezone.utc) + rows = [] + for date_idx, row in ticker_data.iterrows(): + price_date = date_idx.to_pydatetime() if price_date.tzinfo is None: price_date = price_date.replace(tzinfo=timezone.utc) - - # Skip if already exists - if price_date.date() in existing_dates: - continue - - # Prepare new record - price_record = PriceData( - ticker=ticker, - date=price_date, - open=float(row.get('Open', 0)) if not pd.isna(row.get('Open')) else None, - high=float(row.get('High', 0)) if not pd.isna(row.get('High')) else None, - low=float(row.get('Low', 0)) if not pd.isna(row.get('Low')) else None, - close=float(row.get('Close', 0)) if not pd.isna(row.get('Close')) else 0, - volume=float(row.get('Volume', 0)) if not pd.isna(row.get('Volume')) else None, - adjusted_close=float(row.get('Close', 0)) if not pd.isna(row.get('Close')) else None, - data_source=DataSource.YAHOO_FINANCE - ) - new_records.append(price_record) - - # Batch insert new records - if new_records: - db.add_all(new_records) - logger.info(f"Added {len(new_records)} new price records for {ticker}") + + rows.append({ + 'id': _uuid.uuid4(), + 'ticker': ticker, + 'date': price_date, + 'open': _safe(row.get('Open')), + 'high': _safe(row.get('High')), + 'low': _safe(row.get('Low')), + 'close': _safe(row.get('Close')) or 0.0, + 'volume': _safe(row.get('Volume')), + 'adjusted_close': _safe(row.get('Close')), + 'data_source': DataSource.YAHOO_FINANCE.value, + 'created_at': now, + 'updated_at': now, + }) + + if not rows: + return + + stmt = pg_insert(PriceData).values(rows) + stmt = stmt.on_conflict_do_nothing(constraint='uq_price_data') + await db.execute(stmt) + logger.info(f"Batch upserted {len(rows)} price records for {ticker}") async def _get_existing_dates_for_ticker( self, @@ -927,7 +941,12 @@ class PriceDataService: # Convert to response models price_points = [ - PriceDataPoint.model_validate(pd) for pd in price_data + PriceDataPoint.model_construct( + date=pd.date.date() if isinstance(pd.date, datetime) else pd.date, + open=pd.open, high=pd.high, low=pd.low, close=pd.close, + volume=pd.volume, adjusted_close=pd.adjusted_close, + data_source=pd.data_source.value if hasattr(pd.data_source, 'value') else pd.data_source, + ) for pd in price_data ] # Calculate actual date range from returned data diff --git a/docker-compose.yml b/docker-compose.yml index 8046c2f..d2b4f16 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -23,7 +23,7 @@ services: redis: image: redis:7-alpine container_name: stock_oracle_cache - command: redis-server --appendonly yes + command: redis-server --appendonly yes --maxmemory 512mb --maxmemory-policy allkeys-lru ports: - "16380:6379" # Unique port to avoid conflicts volumes: