218 lines
8.1 KiB
Python
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)
|