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

@@ -10,6 +10,9 @@ from .database import (
get_tenant_db,
init_async_db,
init_db,
scoped_async_core_db,
scoped_core_db,
set_rls_context,
)
from .security import (
get_current_active_user,
@@ -25,6 +28,9 @@ __all__ = [
"get_core_db",
"get_async_core_db",
"get_tenant_db",
"scoped_core_db",
"scoped_async_core_db",
"set_rls_context",
"init_db",
"init_async_db",
"verify_token",

View File

@@ -1,5 +1,8 @@
import os
from celery import Celery
from celery.signals import task_postrun, task_prerun
from core.database import rls_company_var, rls_tenant_var
valkey_url = os.getenv("VALKEY_URL", "redis://valkey:6379/0")
@@ -13,6 +16,59 @@ celery_app = Celery(
)
celery_app.set_default()
_RLS_TOKENS_ATTR = "_rls_context_tokens"
def _coerce_int(value) -> int | None:
if value is None or value == "":
return None
try:
return int(value)
except (TypeError, ValueError):
return None
@task_prerun.connect
def _set_rls_context_from_task(task_id=None, task=None, args=None, kwargs=None, **_):
"""Fija las ContextVars de RLS para la ejecución de la tarea.
Las rutas propagan ``tenant_id`` / ``company_id`` vía Celery headers en
:func:`track_and_dispatch`. Aquí los materializamos en ContextVars para
que cualquier sesión que se abra durante la tarea (incluidos los helpers
``scoped_core_db`` y llamadas directas a ``CoreSessionLocal()``) aplique
``SET LOCAL`` automáticamente.
"""
headers = {}
request = getattr(task, "request", None) if task is not None else None
if request is not None:
headers = getattr(request, "headers", None) or {}
tenant_id = _coerce_int(headers.get("rls_tenant_id"))
company_id = _coerce_int(headers.get("rls_company_id"))
token_t = rls_tenant_var.set(tenant_id)
token_c = rls_company_var.set(company_id)
setattr(task, _RLS_TOKENS_ATTR, (token_t, token_c))
@task_postrun.connect
def _reset_rls_context_from_task(task_id=None, task=None, **_):
"""Restaura las ContextVars al terminar la tarea (evita fuga entre tareas
cuando un worker reutiliza el mismo hilo)."""
tokens = getattr(task, _RLS_TOKENS_ATTR, None) if task is not None else None
if tokens is None:
return
token_t, token_c = tokens
try:
rls_tenant_var.reset(token_t)
rls_company_var.reset(token_c)
except ValueError:
rls_tenant_var.set(None)
rls_company_var.set(None)
finally:
delattr(task, _RLS_TOKENS_ATTR)
# ----------------------------------------------------------------------------
# Import models in correct order for SQLAlchemy relationship resolution
# MUST happen AFTER celery_app exists to avoid circular imports during

View File

@@ -2,13 +2,21 @@
Configuración de base de datos con soporte multi-tenant
- Base de datos compartida (core_db) para tenants pequeños/medianos
- Bases de datos dedicadas para clientes enterprise
Row-Level Security (RLS):
Para respetar el aislamiento por tenant/company definido en PostgreSQL,
cada sesión fija las GUCs ``app.tenant_id`` y ``app.company_id`` vía
``SET LOCAL`` al inicio de cada transacción. El listener
``after_begin`` aplica el contexto guardado en ``Session.info``.
"""
import logging
from contextlib import contextmanager
from contextlib import asynccontextmanager, contextmanager
from contextvars import ContextVar
from typing import AsyncGenerator, Dict, Generator, Optional
from sqlalchemy import create_engine
from fastapi import Request
from sqlalchemy import create_engine, event, text
from sqlalchemy.exc import ProgrammingError
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
from sqlalchemy.orm import Session, declarative_base, sessionmaker
@@ -17,10 +25,8 @@ from .config import settings
logger = logging.getLogger(__name__)
# Base declarativa para modelos ORM
Base = declarative_base()
# Engine y SessionLocal para base de datos core (sincrónico)
core_engine = create_engine(
settings.core_database_url,
pool_pre_ping=True,
@@ -32,7 +38,6 @@ core_engine = create_engine(
CoreSessionLocal = sessionmaker(
autocommit=False, autoflush=False, bind=core_engine)
# Engine asíncrono para operaciones async
async_core_engine = create_async_engine(
settings.async_core_database_url,
pool_pre_ping=True,
@@ -45,27 +50,144 @@ AsyncCoreSessionLocal = async_sessionmaker(
async_core_engine, class_=AsyncSession, expire_on_commit=False
)
# Cache de engines para tenants con BD dedicada
_tenant_engines: Dict[str, any] = {}
def get_core_db() -> Generator[Session, None, None]:
RLS_TENANT_KEY = "rls_tenant_id"
RLS_COMPANY_KEY = "rls_company_id"
rls_tenant_var: ContextVar[Optional[int]] = ContextVar("rls_tenant_id", default=None)
rls_company_var: ContextVar[Optional[int]] = ContextVar("rls_company_id", default=None)
def _apply_rls_context(connection, tenant_id: Optional[int], company_id: Optional[int]) -> None:
"""Ejecuta ``SET LOCAL`` en la transacción activa para fijar el contexto RLS."""
tenant_value = "" if tenant_id is None else str(int(tenant_id))
company_value = "" if company_id is None else str(int(company_id))
connection.execute(
text(
"SELECT set_config('app.tenant_id', :t, true), "
"set_config('app.company_id', :c, true)"
),
{"t": tenant_value, "c": company_value},
)
def _resolve_context(session) -> tuple[Optional[int], Optional[int]]:
"""Selecciona tenant_id/company_id desde ``session.info`` y, si faltan,
desde las ContextVars (usadas por tareas Celery vía task_prerun)."""
tenant_id = session.info.get(RLS_TENANT_KEY)
company_id = session.info.get(RLS_COMPANY_KEY)
if tenant_id is None:
tenant_id = rls_tenant_var.get()
if company_id is None:
company_id = rls_company_var.get()
return tenant_id, company_id
@event.listens_for(Session, "after_begin")
def _after_begin(session: Session, transaction, connection) -> None: # type: ignore[no-untyped-def]
"""Aplica ``SET LOCAL`` en cada transacción nueva.
Cubre sesiones síncronas y asíncronas porque ``AsyncSession`` envuelve
internamente una ``Session`` que hereda de esta clase.
"""
Dependency para obtener sesión de base de datos core (compartida)
Uso en FastAPI: db: Session = Depends(get_core_db)
tenant_id, company_id = _resolve_context(session)
if tenant_id is None and company_id is None:
return
_apply_rls_context(connection, tenant_id, company_id)
def set_rls_context(
session: Session,
tenant_id: Optional[int] = None,
company_id: Optional[int] = None,
) -> None:
"""Guarda el contexto RLS en la sesión y, si hay transacción abierta, lo aplica.
Útil para endpoints que validan acceso a una compañía específica después de
crear la sesión (por ejemplo rutas que reciben ``company_id`` en el path).
"""
session.info[RLS_TENANT_KEY] = tenant_id
session.info[RLS_COMPANY_KEY] = company_id
if session.in_transaction():
_apply_rls_context(session.connection(), tenant_id, company_id)
def _extract_rls_context(request: Optional[Request]) -> tuple[Optional[int], Optional[int]]:
"""Recupera ``tenant_id`` / ``company_id`` del estado del request (o de cookies)."""
if request is None:
return None, None
tenant_id = getattr(request.state, "tenant_id", None)
company_id = getattr(request.state, "company_id", None)
if company_id is None:
cookie_value = request.cookies.get("active_company_id")
if cookie_value:
try:
company_id = int(cookie_value)
except (TypeError, ValueError):
company_id = None
return tenant_id, company_id
def get_core_db(request: Request = None) -> Generator[Session, None, None]:
"""Dependency para obtener sesión síncrona con contexto RLS.
FastAPI inyecta ``Request`` automáticamente; los llamadores existentes que
escriben ``db: Session = Depends(get_core_db)`` siguen funcionando sin
cambios porque ``Request`` se resuelve en la capa de dependencia.
"""
tenant_id, company_id = _extract_rls_context(request)
db = CoreSessionLocal()
db.info[RLS_TENANT_KEY] = tenant_id
db.info[RLS_COMPANY_KEY] = company_id
try:
yield db
finally:
db.close()
async def get_async_core_db() -> AsyncGenerator[AsyncSession, None]:
"""
Dependency para obtener sesión asíncrona de base de datos core
"""
async def get_async_core_db(request: Request = None) -> AsyncGenerator[AsyncSession, None]:
"""Dependency async para obtener sesión con contexto RLS."""
tenant_id, company_id = _extract_rls_context(request)
async with AsyncCoreSessionLocal() as session:
session.info[RLS_TENANT_KEY] = tenant_id
session.info[RLS_COMPANY_KEY] = company_id
try:
yield session
finally:
await session.close()
@contextmanager
def scoped_core_db(
tenant_id: Optional[int] = None,
company_id: Optional[int] = None,
) -> Generator[Session, None, None]:
"""Abre una sesión síncrona con contexto RLS explícito.
Pensado para tareas Celery, comandos de mantenimiento o cualquier camino
fuera del ciclo de request HTTP. El contexto se aplica con ``SET LOCAL``
en cada transacción.
"""
db = CoreSessionLocal()
db.info[RLS_TENANT_KEY] = tenant_id
db.info[RLS_COMPANY_KEY] = company_id
try:
yield db
finally:
db.close()
@asynccontextmanager
async def scoped_async_core_db(
tenant_id: Optional[int] = None,
company_id: Optional[int] = None,
) -> AsyncGenerator[AsyncSession, None]:
"""Variante async de :func:`scoped_core_db`."""
async with AsyncCoreSessionLocal() as session:
session.info[RLS_TENANT_KEY] = tenant_id
session.info[RLS_COMPANY_KEY] = company_id
try:
yield session
finally:
@@ -73,16 +195,7 @@ async def get_async_core_db() -> AsyncGenerator[AsyncSession, None]:
def get_tenant_engine(tenant_id: int, db_config: dict):
"""
Obtiene o crea un engine para un tenant con BD dedicada
Args:
tenant_id: ID del tenant
db_config: Configuración de BD {host, port, name, user, password}
Returns:
Engine de SQLAlchemy para el tenant
"""
"""Obtiene o crea un engine para un tenant con BD dedicada."""
if tenant_id not in _tenant_engines:
db_url = f"postgresql://{db_config['user']}:{db_config['password']}@{db_config['host']}:{db_config['port']}/{db_config['name']}"
_tenant_engines[tenant_id] = create_engine(
@@ -95,21 +208,16 @@ def get_tenant_engine(tenant_id: int, db_config: dict):
def get_tenant_db(
tenant_id: int, db_config: Optional[dict] = None
) -> Generator[Session, None, None]:
"""
Context manager para obtener sesión de BD de un tenant específico
"""Context manager para obtener sesión de BD de un tenant específico.
Si db_config es None, usa la BD core (compartida)
Si db_config está presente, usa la BD dedicada del tenant
Uso:
with get_tenant_db(tenant_id, config) as db:
# operaciones con db
Si ``db_config`` es ``None`` usa la BD core compartida y aplica RLS con
el ``tenant_id`` recibido. Si el tenant tiene BD dedicada, el aislamiento
es físico y no se fija contexto RLS (no hay columna ``tenant_id``).
"""
if db_config is None:
# Tenant en BD compartida
db = CoreSessionLocal()
db.info[RLS_TENANT_KEY] = tenant_id
else:
# Tenant con BD dedicada
engine = get_tenant_engine(tenant_id, db_config)
SessionLocal = sessionmaker(
autocommit=False, autoflush=False, bind=engine)
@@ -122,13 +230,10 @@ def get_tenant_db(
def init_db():
"""
Inicializa las tablas de la base de datos core
"""
"""Inicializa las tablas de la base de datos core."""
try:
Base.metadata.create_all(bind=core_engine, checkfirst=True)
except ProgrammingError as e:
# Si la tabla ya existe, es seguro continuar
if "already exists" in str(e):
logger.warning(
f"Algunas tablas ya existen en la base de datos: {e}")
@@ -137,8 +242,6 @@ def init_db():
async def init_async_db():
"""
Inicializa las tablas de la base de datos core (async)
"""
"""Inicializa las tablas de la base de datos core (async)."""
async with async_core_engine.begin() as conn:
await conn.run_sync(Base.metadata.create_all)

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