Files
invest-copilot/src/backend/main.py
T

218 lines
8.1 KiB
Python

"""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)