Initial commit: Trading Portal - FastAPI + React + PostgreSQL
This commit is contained in:
Executable
Executable
+104
@@ -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
|
||||
Executable
+74
@@ -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"
|
||||
Executable
+73
@@ -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]
|
||||
Executable
+351
@@ -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
|
||||
Reference in New Issue
Block a user