Initial commit: Trading Portal - FastAPI + React + PostgreSQL

This commit is contained in:
2026-07-03 13:08:09 +00:00
commit 34a1e91541
198 changed files with 35110 additions and 0 deletions
View File
+104
View File
@@ -0,0 +1,104 @@
from __future__ import annotations
from collections.abc import AsyncGenerator
from uuid import UUID
from fastapi import Depends, Header, Request
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.exceptions import InvalidCredentialsException, InvalidTokenException
from app.core.security import decode_token
from app.database import get_db
from app.models import User
async def get_db_session() -> AsyncGenerator[AsyncSession, None]:
"""Provide an async SQLAlchemy database session via FastAPI dependency."""
async for session in get_db():
yield session
async def get_current_user(
request: Request,
db: AsyncSession = Depends(get_db_session),
) -> User:
"""
Extract the Bearer token from the Authorization header, decode it,
and fetch the corresponding user from the database.
Raises ``InvalidCredentialsException`` if the token is missing,
invalid, or the user is not found.
"""
authorization: str | None = request.headers.get("Authorization")
if not authorization or not authorization.startswith("Bearer "):
raise InvalidCredentialsException(
detail="Missing or malformed Authorization header"
)
token = authorization.removeprefix("Bearer ").strip()
try:
payload = decode_token(token)
except Exception:
raise InvalidTokenException(detail="Invalid or expired token")
sub: str | None = payload.get("sub")
if sub is None:
raise InvalidTokenException(detail="Token payload missing subject")
user_id: UUID
try:
user_id = UUID(sub)
except ValueError:
raise InvalidTokenException(detail="Invalid token subject format")
result = await db.execute(select(User).where(User.id == user_id))
user = result.scalar_one_or_none()
if user is None:
raise InvalidCredentialsException(detail="User not found")
return user
async def get_current_active_user(
current_user: User = Depends(get_current_user),
) -> User:
"""Return the current user if active, otherwise raise an error."""
if not current_user.is_active:
raise InvalidCredentialsException(detail="Inactive user account")
return current_user
async def get_current_trader_user(
current_user: User = Depends(get_current_active_user),
) -> User:
"""Return the current user if they are a trader or admin, otherwise raise an error."""
if current_user.role not in ("admin", "trader"):
from app.core.exceptions import ForbiddenException
raise ForbiddenException(detail="Trader or admin privileges required")
return current_user
async def get_current_viewer_user(
current_user: User = Depends(get_current_active_user),
) -> User:
"""Return the current user if active (any role can view)."""
return current_user
async def get_current_admin_user(
current_user: User = Depends(get_current_active_user),
) -> User:
"""Return the current user if they are an admin, otherwise raise an error.
P1-22: Checks both is_admin AND role for frontend/backend consistency.
"""
if not current_user.is_admin and current_user.role != "admin":
from app.core.exceptions import ForbiddenException
raise ForbiddenException(detail="Admin privileges required")
return current_user
+74
View File
@@ -0,0 +1,74 @@
from __future__ import annotations
class AppException(Exception):
"""Base exception for all application-level errors."""
status_code: int = 500
detail: str = "Internal server error"
code: str = "internal_error"
def __init__(
self,
status_code: int | None = None,
detail: str | None = None,
code: str | None = None,
) -> None:
if status_code is not None:
self.status_code = status_code
if detail is not None:
self.detail = detail
if code is not None:
self.code = code
super().__init__(self.detail)
class NotFoundException(AppException):
status_code: int = 404
code: str = "not_found"
detail: str = "Resource not found"
class AuthException(AppException):
status_code: int = 401
code: str = "auth_error"
detail: str = "Authentication error"
class InvalidCredentialsException(AuthException):
code: str = "invalid_credentials"
detail: str = "Invalid username or password"
class TokenExpiredException(AuthException):
code: str = "token_expired"
detail: str = "Token has expired"
class InvalidTokenException(AuthException):
code: str = "invalid_token"
detail: str = "Invalid token"
class ForbiddenException(AppException):
status_code: int = 403
code: str = "forbidden"
detail: str = "Forbidden"
class ValidationException(AppException):
status_code: int = 422
code: str = "validation_error"
detail: str = "Validation error"
class RateLimitException(AppException):
status_code: int = 429
code: str = "rate_limit"
detail: str = "Rate limit exceeded"
class ConflictException(AppException):
status_code: int = 409
code: str = "conflict"
detail: str = "Resource already exists"
+73
View File
@@ -0,0 +1,73 @@
from __future__ import annotations
import time
import structlog
from starlette.requests import Request
from starlette.responses import Response
from starlette.types import ASGIApp, Receive, Scope, Send
logger = structlog.get_logger(__name__)
class RequestLoggingMiddleware:
"""ASGI middleware that logs every request with method, path, status code,
and duration using structlog."""
def __init__(self, app: ASGIApp) -> None:
self.app = app
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
if scope["type"] != "http":
await self.app(scope, receive, send)
return
start = time.perf_counter()
request = Request(scope)
# Wrap send to capture the response status code
status_code: int | None = None
async def send_wrapper(message: dict) -> None:
nonlocal status_code
if message["type"] == "http.response.start":
status_code = message["status"]
await send(message)
try:
await self.app(scope, receive, send_wrapper)
except Exception:
duration = time.perf_counter() - start
logger.error(
"request_error",
method=request.method,
path=request.url.path,
status_code=500,
duration_ms=round(duration * 1000, 2),
)
raise
else:
duration = time.perf_counter() - start
if (status_code or 0) >= 400:
logger.warning(
"request_complete",
method=request.method,
path=request.url.path,
status_code=status_code,
duration_ms=round(duration * 1000, 2),
)
else:
logger.info(
"request_complete",
method=request.method,
path=request.url.path,
status_code=status_code,
duration_ms=round(duration * 1000, 2),
)
def register_middleware(app: ASGIApp) -> None:
"""Convenience helper — add the middleware to a FastAPI app."""
from app.core.middleware import RequestLoggingMiddleware # noqa: F811
app.add_middleware(RequestLoggingMiddleware) # type: ignore[arg-type]
+351
View File
@@ -0,0 +1,351 @@
"""
Security module for the trading portal backend.
Provides password hashing, JWT token management (RS256),
AES-256-CBC encryption for API key storage, and token utilities.
"""
from __future__ import annotations
import hashlib
import uuid
from datetime import datetime, timedelta, timezone
from typing import Optional, Tuple
from fastapi import HTTPException
from jose import JWTError, jwt
from jose.exceptions import ExpiredSignatureError
from passlib.context import CryptContext
from cryptography.hazmat.primitives.ciphers import Cipher, algorithms, modes
from cryptography.hazmat.backends import default_backend
from app.config import settings
# ---------------------------------------------------------------------------
# Password hashing
# ---------------------------------------------------------------------------
_pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto")
def hash_password(password: str) -> str:
"""Hash a plaintext password using bcrypt (synchronous — use in thread pool)."""
return _pwd_context.hash(password)
def verify_password(plain: str, hashed: str) -> bool:
"""Verify a plaintext password against a bcrypt hash (synchronous — use in thread pool)."""
return _pwd_context.verify(plain, hashed)
async def hash_password_async(password: str) -> str:
"""Hash a plaintext password using bcrypt (async, runs in thread pool)."""
import asyncio
loop = asyncio.get_running_loop()
return await loop.run_in_executor(None, _pwd_context.hash, password)
async def verify_password_async(plain: str, hashed: str) -> bool:
"""Verify a plaintext password against a bcrypt hash (async, runs in thread pool)."""
import asyncio
loop = asyncio.get_running_loop()
return await loop.run_in_executor(None, _pwd_context.verify, plain, hashed)
# ---------------------------------------------------------------------------
# JWT RS256 token management — with multi-key rotation support
# ---------------------------------------------------------------------------
#
# PRINCIPLE:
# - Tokens are SIGNED with the current private key and carry a "kid"
# (Key ID = SHA-256 fingerprint of the public key) in the JWT header.
# - Tokens are VERIFIED against ALL public keys in the keys directory.
# ANY valid key can decode the token → old tokens survive key rotation.
# - When rotating: add new key pair, old tokens remain valid until they
# expire naturally. Remove old public keys only after all their tokens
# have expired.
#
# Directory layout:
# /run/secrets/jwt_private.pem ← CURRENT signing key
# /run/secrets/jwt_public_keys/*.pem ← ALL valid public keys
import hashlib
import os
from functools import lru_cache
ALGORITHM = "RS256"
def _load_private_key() -> str:
"""Read the current RSA private key PEM for signing new tokens."""
try:
with open(settings.JWT_PRIVATE_KEY_PATH, "r") as f:
return f.read()
except FileNotFoundError:
raise HTTPException(
status_code=500,
detail=f"JWT private key not found at {settings.JWT_PRIVATE_KEY_PATH}",
)
except OSError as exc:
raise HTTPException(
status_code=500,
detail=f"Failed to read JWT private key: {exc}",
)
def _compute_kid(public_key_pem: str) -> str:
"""Return a Key ID (SHA-256 fingerprint) for a public key PEM."""
return hashlib.sha256(public_key_pem.strip().encode()).hexdigest()[:16]
@lru_cache(maxsize=1)
def _cached_public_key() -> str:
"""Return the CURRENT public key (used for kid computation). Cached."""
keys = _load_all_public_keys()
if not keys:
raise HTTPException(status_code=500, detail="No valid JWT public keys found")
return keys[-1][1] # newest key
def _load_all_public_keys() -> list[tuple[str, str]]:
"""
Load ALL valid public keys from the keys directory.
Returns:
List of (kid, pem_content) tuples, sorted by filename for determinism.
The store allows multiple keys to coexist — old keys still validate
tokens signed before the last rotation.
"""
keys_dir = settings.JWT_PUBLIC_KEYS_DIR
keys: list[tuple[str, str]] = []
try:
for filename in sorted(os.listdir(keys_dir)):
if not filename.endswith(".pem"):
continue
filepath = os.path.join(keys_dir, filename)
try:
with open(filepath, "r") as f:
pem = f.read().strip()
if pem:
kid = _compute_kid(pem)
keys.append((kid, pem))
except OSError:
continue # skip unreadable files
except FileNotFoundError:
pass # no directory yet — handled gracefully
except NotADirectoryError:
raise HTTPException(
status_code=500,
detail=f"JWT_PUBLIC_KEYS_DIR ({keys_dir}) is not a directory",
)
if not keys:
raise HTTPException(
status_code=500,
detail=f"No valid JWT public keys found in {keys_dir}",
)
return keys
def create_access_token(
data: dict,
expires_delta: Optional[timedelta] = None,
) -> str:
"""Create a short-lived JWT access token (RS256).
Includes ``kid`` header so the verifier knows which key to try first.
"""
to_encode = data.copy()
now = datetime.now(timezone.utc)
if expires_delta is not None:
expire = now + expires_delta
else:
expire = now + timedelta(minutes=settings.JWT_ACCESS_TOKEN_EXPIRE_MINUTES)
to_encode.update({"iat": now, "exp": expire, "sub": str(data["sub"])})
private_key = _load_private_key()
public_key = _cached_public_key()
kid = _compute_kid(public_key)
headers = {"kid": kid}
return jwt.encode(to_encode, private_key, algorithm=ALGORITHM, headers=headers)
def create_refresh_token(data: dict) -> str:
"""Create a long-lived JWT refresh token (RS256).
Includes ``kid``, ``type: refresh``, and a unique ``jti``.
"""
to_encode = data.copy()
now = datetime.now(timezone.utc)
expire = now + timedelta(days=settings.JWT_REFRESH_TOKEN_EXPIRE_DAYS)
to_encode.update(
{
"iat": now,
"exp": expire,
"sub": str(data["sub"]),
"type": "refresh",
"jti": generate_jti(),
}
)
private_key = _load_private_key()
public_key = _cached_public_key()
kid = _compute_kid(public_key)
headers = {"kid": kid}
return jwt.encode(to_encode, private_key, algorithm=ALGORITHM, headers=headers)
def decode_token(token: str) -> dict:
"""Decode and verify a JWT token using ALL known public keys.
Tries every valid public key in the directory. If ANY key verifies the
token, it is valid — this is how key rotation works without breaking
existing sessions.
Optimisation: the token's ``kid`` header is used to try the matching key
first before falling back to a full linear scan.
"""
# 1. Extract kid from token header (without verifying signature yet)
try:
unverified_header = jwt.get_unverified_header(token)
token_kid = unverified_header.get("kid")
except JWTError:
token_kid = None
# 2. Load all valid public keys
all_keys = _load_all_public_keys() # [(kid, pem), ...]
# 3. If we have a kid, try the matching key first
if token_kid:
for kid, pem in all_keys:
if kid == token_kid:
try:
return jwt.decode(token, pem, algorithms=[ALGORITHM])
except ExpiredSignatureError:
raise HTTPException(
status_code=401,
detail="Token has expired",
headers={"WWW-Authenticate": "Bearer"},
)
except JWTError:
pass # key mismatch — fall through to full scan
# 4. Fallback: try all keys (handles tokens without kid, or kid mismatch)
for kid, pem in all_keys:
try:
return jwt.decode(token, pem, algorithms=[ALGORITHM])
except ExpiredSignatureError:
raise HTTPException(
status_code=401,
detail="Token has expired",
headers={"WWW-Authenticate": "Bearer"},
)
except JWTError:
continue # try next key
# 5. No key worked
raise HTTPException(
status_code=401,
detail="Invalid or expired token",
headers={"WWW-Authenticate": "Bearer"},
)
# ---------------------------------------------------------------------------
# AES-256-CBC encryption (for API key storage)
# ---------------------------------------------------------------------------
_BACKEND = default_backend()
def generate_encryption_key() -> str:
"""Generate a random 32-byte (256-bit) hex-encoded encryption key.
Print the key to stdout so it can be copied into the ``.env`` file.
"""
key = uuid.uuid4().hex + uuid.uuid4().hex # 64 hex chars = 32 bytes
print(f"Encryption key (save in .env as ENCRYPTION_KEY={key}): {key}")
return key
def _resolve_key(key_hex: Optional[str] = None) -> bytes:
"""Return the AES key as bytes, falling back to settings."""
raw = key_hex if key_hex is not None else settings.ENCRYPTION_KEY
if not raw:
raise HTTPException(
status_code=500,
detail="Encryption key not configured. Set ENCRYPTION_KEY in .env",
)
return bytes.fromhex(raw)
def encrypt_api_key(
api_key: str,
key_hex: Optional[str] = None,
) -> Tuple[str, str]:
"""Encrypt an API key with AES-256-CBC.
Returns ``(ciphertext_hex, iv_hex)``.
"""
key = _resolve_key(key_hex)
iv = uuid.uuid4().bytes # 16 random bytes
cipher = Cipher(algorithms.AES(key), modes.CBC(iv), backend=_BACKEND)
encryptor = cipher.encryptor()
# Pad plaintext to AES block size (16 bytes) using PKCS7
plaintext_bytes = api_key.encode("utf-8")
pad_len = 16 - (len(plaintext_bytes) % 16)
padded = plaintext_bytes + bytes([pad_len] * pad_len)
ciphertext = encryptor.update(padded) + encryptor.finalize()
return ciphertext.hex(), iv.hex()
def decrypt_api_key(
ciphertext_hex: str,
iv_hex: str,
key_hex: Optional[str] = None,
) -> str:
"""Decrypt an AES-256-CBC encrypted API key.
Returns the original plaintext string.
"""
key = _resolve_key(key_hex)
ciphertext = bytes.fromhex(ciphertext_hex)
iv = bytes.fromhex(iv_hex)
cipher = Cipher(algorithms.AES(key), modes.CBC(iv), backend=_BACKEND)
decryptor = cipher.decryptor()
padded = decryptor.update(ciphertext) + decryptor.finalize()
# Remove PKCS7 padding
pad_len = padded[-1]
if pad_len < 1 or pad_len > 16:
raise HTTPException(
status_code=500,
detail="Decryption failed: invalid padding",
)
plaintext_bytes = padded[:-pad_len]
return plaintext_bytes.decode("utf-8")
# ---------------------------------------------------------------------------
# Token utilities
# ---------------------------------------------------------------------------
def generate_token_hash(token: str) -> str:
"""Return the SHA-256 hex digest of a token string."""
return hashlib.sha256(token.encode("utf-8")).hexdigest()
def generate_jti() -> str:
"""Return a UUID4 hex string for use as a JWT token ID."""
return uuid.uuid4().hex