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:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user