Implement row-level security (RLS) context management for database sessions. Refactor invoice processing and reverting tasks to utilize scoped database sessions with RLS context. Update middleware to extract and set company ID from requests. Enhance task dispatching to propagate RLS context via Celery headers. Update architecture documentation to reflect RLS implementation details.

This commit is contained in:
2026-04-24 18:01:46 -05:00
parent 7acf77994f
commit 73a6d92863
13 changed files with 1052 additions and 306 deletions

View File

@@ -1,19 +1,37 @@
import logging
import time
from typing import Callable
from fastapi import Request, Response
from typing import Callable, Optional
from fastapi import Request
from fastapi.responses import JSONResponse
from starlette.middleware.base import BaseHTTPMiddleware
from .config import settings
from .database import CoreSessionLocal
from .database import scoped_core_db
from .security import get_tenant_from_token, verify_token
logger = logging.getLogger(__name__)
def _extract_company_id(request: Request) -> Optional[int]:
"""Obtiene ``company_id`` activa desde header ``X-Company-Id`` o cookie.
El frontend guarda la compañía activa en la cookie ``active_company_id``
(ver ``frontend/src/lib/stores/company.svelte.ts``). El header es la
ruta explícita para clientes no-browser.
"""
header_value = request.headers.get("X-Company-Id")
raw = header_value or request.cookies.get("active_company_id")
if not raw:
return None
try:
return int(raw)
except (TypeError, ValueError):
return None
class TenantMiddleware(BaseHTTPMiddleware):
async def dispatch(self, request: Request, call_next: Callable):
# Rutas públicas que no requieren tenant
# Permitir acceso sin autenticación a rutas de documentación y salud
doc_prefixes = ["/api/redoc", "/api/openapi.json"]
public_prefixes = [
"/api/v1/auth",
@@ -27,14 +45,12 @@ class TenantMiddleware(BaseHTTPMiddleware):
path = request.url.path
# 3. Bypass para rutas públicas y docs
if any(path == prefix or path.startswith(prefix + "/") for prefix in doc_prefixes):
return await call_next(request)
if any(path == prefix or (prefix != "/" and path.startswith(prefix)) for prefix in public_prefixes):
return await call_next(request)
# 4. Validación estricta de Token (solo para lo que no es público ni OPTIONS)
auth_header = request.headers.get("Authorization")
if not auth_header or not auth_header.startswith("Bearer "):
return JSONResponse(
@@ -50,9 +66,10 @@ class TenantMiddleware(BaseHTTPMiddleware):
try:
user_info = verify_token(token)
tenant_id = get_tenant_from_token(user_info)
request.state.tenant_id = tenant_id
request.state.user_info = user_info
request.state.company_id = _extract_company_id(request)
except Exception as e:
logger.error(f"❌ Tenant validation error: {str(e)}")
return JSONResponse(
@@ -64,20 +81,16 @@ class TenantMiddleware(BaseHTTPMiddleware):
}
)
# 5. Continuar con la petición real
return await call_next(request)
class LicenseValidationMiddleware(BaseHTTPMiddleware):
"""
Middleware para validar la licencia del tenant antes de procesar requests
"""
"""Middleware para validar la licencia del tenant antes de procesar requests."""
async def dispatch(self, request: Request, call_next: Callable):
if not settings.LICENSE_CHECK_ENABLED:
return await call_next(request)
# Rutas que no requieren validación de licencia
exempt_paths = [
"/api/docs",
"/api/redoc",
@@ -92,7 +105,6 @@ class LicenseValidationMiddleware(BaseHTTPMiddleware):
"/api/v1/core/users/avatar",
]
# Verificar si la ruta está exenta (comparación exacta o prefijo)
is_exempt = False
for path in exempt_paths:
if request.url.path == path or (
@@ -104,34 +116,32 @@ class LicenseValidationMiddleware(BaseHTTPMiddleware):
if is_exempt:
return await call_next(request)
# Obtener tenant_id del request state (debe ser seteado por TenantMiddleware)
tenant_id = getattr(request.state, "tenant_id", None)
if not tenant_id:
return await call_next(request) # Dejamos que TenantMiddleware maneje esto
return await call_next(request)
# Validar licencia
db = CoreSessionLocal()
# core.licenses / core.license_usage están bajo RLS por tenant_id:
# se abre la sesión con contexto explícito para que LicenseService
# vea las filas del tenant actual.
try:
# Importar aquí para evitar imports circulares
from api.v1.modules.core.licenses.service import LicenseService
with scoped_core_db(tenant_id=tenant_id) as db:
from api.v1.modules.core.licenses.service import LicenseService
license_service = LicenseService(db)
license_info = license_service.validate_license(tenant_id)
license_service = LicenseService(db)
license_info = license_service.validate_license(tenant_id)
if not license_info["is_valid"]:
return JSONResponse(
status_code=402,
content={
"error": "HTTP_ERROR",
"message": f"License validation failed: {license_info['reason']}",
"status_code": 402,
}
)
# Agregar info de licencia al request state
request.state.license_info = license_info
if not license_info["is_valid"]:
return JSONResponse(
status_code=402,
content={
"error": "HTTP_ERROR",
"message": f"License validation failed: {license_info['reason']}",
"status_code": 402,
}
)
request.state.license_info = license_info
except Exception as e:
logger.error(f"License validation error: {str(e)}")
return JSONResponse(
@@ -142,17 +152,12 @@ class LicenseValidationMiddleware(BaseHTTPMiddleware):
"status_code": 500,
}
)
finally:
db.close()
response = await call_next(request)
return response
return await call_next(request)
class RequestLoggingMiddleware(BaseHTTPMiddleware):
"""
Middleware para logging de requests
"""
"""Middleware para logging de requests."""
async def dispatch(self, request: Request, call_next: Callable):
start_time = time.time()
@@ -170,12 +175,10 @@ class RequestLoggingMiddleware(BaseHTTPMiddleware):
):
return await call_next(request)
# Log request
logger.info(f"Request: {request.method} {request.url.path}")
response = await call_next(request)
# Log response
process_time = time.time() - start_time
logger.info(
f"Response: {request.method} {request.url.path} "
@@ -183,7 +186,6 @@ class RequestLoggingMiddleware(BaseHTTPMiddleware):
f"Duration: {process_time:.3f}s"
)
# Agregar header con tiempo de procesamiento
response.headers["X-Process-Time"] = str(process_time)
return response