"""FastAPI main application for Invest Copilot BFF.""" import logging from contextlib import asynccontextmanager from fastapi import FastAPI, Request from fastapi.middleware.cors import CORSMiddleware from fastapi.responses import JSONResponse from config import settings from database import init_db, close_db from routers import search, stocks, prices, sec_filings, insider_trades, peers, sentiment from routers import watchlists, strategies, screeners, sectors, alerts, stream, dashboard, data_sync, auth from services.rate_limiter import check_rate_limit, DEFAULT_RULES logging.basicConfig(level=logging.INFO) logger = logging.getLogger(__name__) @asynccontextmanager async def lifespan(app: FastAPI): """Startup and shutdown events.""" logger.info("Starting Invest Copilot backend...") await init_db() logger.info("Database initialized") # Start background Celery tasks on startup from tasks.ingest_prices import ingest_prices_task from tasks.sector_scan import sector_scan_task from config import settings import asyncio import httpx from database import execute_query, execute_command from services.market_data import MarketDataService from tasks.ingest_sec import ingest_sec_filings_task async def _seed_if_empty(): """Seed database with sample data if empty.""" count = await execute_query("SELECT count(*) AS cnt FROM stock_profiles") if count and int(count[0]["cnt"]) == 0: logger.info("Database appears empty — running seed...") try: from seed_data import seed await seed() logger.info("Seed complete") except Exception as e: logger.error(f"Seed failed: {e}") async def _ingest_initial_prices(): """Ingest price data for seed stocks using yfinance (no API key needed).""" tickers = ["AAPL", "MSFT", "GOOGL", "AMZN", "NVDA", "META", "TSLA", "JPM", "V", "JNJ"] service = MarketDataService() for ticker in tickers: try: prices = await service.get_time_series(ticker) for price in prices: await execute_command( """ INSERT INTO prices (ticker, date, open, high, low, close, volume, adjusted_close) VALUES ($1, $2, $3, $4, $5, $6, $7, $8) ON CONFLICT (ticker, date) DO UPDATE SET open = EXCLUDED.open, high = EXCLUDED.high, low = EXCLUDED.low, close = EXCLUDED.close, volume = EXCLUDED.volume, adjusted_close = EXCLUDED.adjusted_close """, ( ticker, price.get("date", ""), price.get("open", 0), price.get("high", 0), price.get("low", 0), price.get("close", 0), price.get("volume", 0), price.get("close", 0), ), ) logger.info(f"Ingested {len(prices)} prices for {ticker}") except Exception as e: logger.error(f"Price ingestion failed for {ticker}: {e}") async def _run_sector_scan(): """Run initial sector scan.""" try: from services.rotation_service import RotationService service = RotationService() await service.scan_sectors() logger.info("Initial sector scan complete") except Exception as e: logger.error(f"Sector scan failed: {e}") # Run seed and data ingestion in background asyncio.create_task(_seed_if_empty()) asyncio.create_task(_ingest_initial_prices()) asyncio.create_task(_run_sector_scan()) yield await close_db() logger.info("Backend shutting down") app = FastAPI( title="Invest Copilot API", version="0.1.0", description="AI-native investment research and portfolio copilot", lifespan=lifespan, ) app.add_middleware( CORSMiddleware, allow_origins=["http://localhost:3000", "http://localhost:8000"], allow_credentials=True, allow_methods=["*"], allow_headers=["*"], ) def _get_client_ip(request: Request) -> str: """Extract client IP from request, handling proxies.""" # Check X-Forwarded-For header first (behind proxy) forwarded_for = request.headers.get("x-forwarded-for") if forwarded_for: return forwarded_for.split(",")[0].strip() # Check X-Real-IP (common in some proxies) real_ip = request.headers.get("x-real-ip") if real_ip: return real_ip # Fall back to client host return request.client.host if request.client else "unknown" def _get_rate_limit_rule(request: Request) -> tuple: """Determine rate limit rule and key for a request.""" path = request.url.path method = request.method # Auth endpoints get strict per-IP limiting if "/api/v1/auth/" in path: return "auth", _get_client_ip(request) # Heavy operations (backtest, screener run) if any(endpoint in path for endpoint in ["/backtest", "/run"]): return "heavy", _get_client_ip(request) # Authenticated endpoints - check for Bearer token auth_header = request.headers.get("authorization", "") if auth_header.startswith("Bearer "): # For now, rate limit by IP for authenticated users too # In production, decode JWT to get user ID return "authenticated", _get_client_ip(request) # Public endpoints return "public", _get_client_ip(request) async def rate_limit_middleware(request: Request, call_next): """Middleware to enforce rate limiting on API requests.""" # Skip rate limiting for health checks if request.url.path == "/health": response = await call_next(request) return response rule_name, key = _get_rate_limit_rule(request) allowed, headers = check_rate_limit(rule_name, key) if not allowed: # Return 429 Too Many Requests response = JSONResponse( status_code=429, content={ "detail": "Rate limit exceeded. Please try again later.", "retry_after": int(headers.get("Retry-After", 60)), }, headers=headers, ) return response # Process the request response = await call_next(request) # Add rate limit headers to successful responses for header_name, header_value in headers.items(): response.headers[header_name] = header_value return response # Register rate limiting middleware app.middleware("http")(rate_limit_middleware) @app.get("/health") async def health(): return {"status": "ok", "version": "0.1.0"} # Register routers app.include_router(search.router, prefix="/api/v1", tags=["Search"]) app.include_router(stocks.router, prefix="/api/v1", tags=["Stocks"]) app.include_router(prices.router, prefix="/api/v1", tags=["Prices"]) app.include_router(sec_filings.router, prefix="/api/v1", tags=["SEC Filings"]) app.include_router(insider_trades.router, prefix="/api/v1", tags=["Insider Trades"]) app.include_router(peers.router, prefix="/api/v1", tags=["Peers"]) app.include_router(sentiment.router, prefix="/api/v1", tags=["Sentiment"]) app.include_router(watchlists.router, prefix="/api/v1", tags=["Watchlists"]) app.include_router(strategies.router, prefix="/api/v1", tags=["Strategies"]) app.include_router(screeners.router, prefix="/api/v1", tags=["Screeners"]) app.include_router(sectors.router, prefix="/api/v1", tags=["Sectors"]) app.include_router(alerts.router, prefix="/api/v1", tags=["Alerts"]) app.include_router(stream.router, prefix="/api/v1", tags=["Stream"]) app.include_router(dashboard.router, prefix="/api/v1", tags=["Dashboard"]) app.include_router(data_sync.router, prefix="/api/v1", tags=["Data Sync"]) app.include_router(auth.router, prefix="/api/v1", tags=["Authentication"]) if __name__ == "__main__": import uvicorn uvicorn.run(app, host="0.0.0.0", port=8000)