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:
@@ -0,0 +1,302 @@
|
||||
"""enable_rls_tenant_company
|
||||
|
||||
Habilita Row-Level Security en las tablas multi-tenant conforme al skill
|
||||
`aduanasoft-dev-standards` (sección 10). Las políticas dependen de dos
|
||||
GUCs que la aplicación establece por transacción con `SET LOCAL`:
|
||||
|
||||
- ``app.tenant_id`` (ID del tenant actual, obligatorio para aislamiento)
|
||||
- ``app.company_id`` (ID de la compañía activa; opcional — si no está fijado
|
||||
la política permite todas las compañías del tenant, útil para vistas de
|
||||
selector de compañía / bootstrap de sesión)
|
||||
|
||||
Las funciones SQL viven en el esquema ``app`` y retornan ``NULL`` cuando la
|
||||
GUC correspondiente está vacía, lo que hace que las comparaciones
|
||||
``col = app.current_xxx_id()`` devuelvan 0 filas sin contexto (fail-closed
|
||||
para ``tenant_id``).
|
||||
|
||||
Las tablas ``core.tenants`` y ``core.user_tenants`` NO quedan bajo RLS: son
|
||||
necesarias para el bootstrap de la sesión (obtener tenant del JWT y listar
|
||||
los tenants del usuario en el selector).
|
||||
|
||||
La migración instala ``FORCE ROW LEVEL SECURITY`` para que las políticas
|
||||
apliquen también al owner — los superusuarios (p. ej. ``postgres`` en dev)
|
||||
siguen haciendo bypass por diseño de PostgreSQL; en producción la API debe
|
||||
conectarse con un rol sin BYPASSRLS.
|
||||
|
||||
Revision ID: d1a2b3c4e5f6
|
||||
Revises: c8d9e0f1a2b3
|
||||
Create Date: 2026-04-24 17:00:00.000000
|
||||
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision: str = "d1a2b3c4e5f6"
|
||||
down_revision: Union[str, Sequence[str], None] = "c8d9e0f1a2b3"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
TABLES_TENANT_ONLY: list[tuple[str, str]] = [
|
||||
("a76", "company"),
|
||||
("core", "license_usage"),
|
||||
("core", "licenses"),
|
||||
]
|
||||
|
||||
TABLES_TENANT_AND_COMPANY: list[tuple[str, str]] = [
|
||||
("a24", "balance_movement"),
|
||||
("a24", "discharge_detail"),
|
||||
("a24", "discharge_header"),
|
||||
("a24", "discharge_scrap"),
|
||||
("a24", "fa_classes"),
|
||||
("a24", "fa_item_lines"),
|
||||
("a24", "fa_partes"),
|
||||
("a24", "inv_aphis_characteristic"),
|
||||
("a24", "inv_aphis_containers"),
|
||||
("a24", "inv_aphis_entities"),
|
||||
("a24", "inv_aphis_general"),
|
||||
("a24", "inv_aphis_lpcos"),
|
||||
("a24", "inv_aphis_routing"),
|
||||
("a24", "inv_aphis_stype_pitems"),
|
||||
("a24", "inv_bom"),
|
||||
("a24", "inv_classes"),
|
||||
("a24", "inv_parte_paises"),
|
||||
("a24", "inv_partes"),
|
||||
("a76", "app_settings"),
|
||||
("a76", "audit_logs"),
|
||||
("a76", "canadian_tariff_fractions"),
|
||||
("a76", "classes"),
|
||||
("a76", "classification_concepts"),
|
||||
("a76", "clients_and_providers"),
|
||||
("a76", "clients_and_providers_address"),
|
||||
("a76", "clients_and_providers_programs"),
|
||||
("a76", "concept_manifestations"),
|
||||
("a76", "concepts"),
|
||||
("a76", "country_rule_oct"),
|
||||
("a76", "ctm_receipts"),
|
||||
("a76", "customs_broker_concepts"),
|
||||
("a76", "customs_brokers"),
|
||||
("a76", "customs_brokers_personnel"),
|
||||
("a76", "customs_brokers_vu"),
|
||||
("a76", "depreciation_catalog"),
|
||||
("a76", "document_types_digitization"),
|
||||
("a76", "doda"),
|
||||
("a76", "doda_american_pedimentos"),
|
||||
("a76", "doda_container_seals"),
|
||||
("a76", "doda_containers"),
|
||||
("a76", "doda_pedimentos"),
|
||||
("a76", "driver"),
|
||||
("a76", "electronic_notices"),
|
||||
("a76", "equivalencies"),
|
||||
("a76", "equivalency_items"),
|
||||
("a76", "error_catalogs"),
|
||||
("a76", "error_classifications"),
|
||||
("a76", "exchange_rate"),
|
||||
("a76", "fa_location_ext"),
|
||||
("a76", "fda_affirmation_codes"),
|
||||
("a76", "fda_catalog"),
|
||||
("a76", "fda_constituent_elements"),
|
||||
("a76", "fda_lot_production"),
|
||||
("a76", "fda_specifications"),
|
||||
("a76", "fraction_rule_octave"),
|
||||
("a76", "historical_tariff_fractions"),
|
||||
("a76", "identifier_details"),
|
||||
("a76", "identifiers"),
|
||||
("a76", "inpc"),
|
||||
("a76", "invoice_collections"),
|
||||
("a76", "invoice_compliance_mx"),
|
||||
("a76", "invoice_financials"),
|
||||
("a76", "invoice_header"),
|
||||
("a76", "invoice_logistics"),
|
||||
("a76", "invoice_sales_details"),
|
||||
("a76", "invoice_settings"),
|
||||
("a76", "item_line_series"),
|
||||
("a76", "item_lines"),
|
||||
("a76", "item_presets"),
|
||||
("a76", "legends"),
|
||||
("a76", "location"),
|
||||
("a76", "manifest_anexos"),
|
||||
("a76", "manifest_drivers"),
|
||||
("a76", "manifests"),
|
||||
("a76", "multi_currency_types"),
|
||||
("a76", "octave_balance"),
|
||||
("a76", "packages"),
|
||||
("a76", "packing_lists"),
|
||||
("a76", "parts"),
|
||||
("a76", "pedimento_config_additional"),
|
||||
("a76", "pedimento_config_calculations"),
|
||||
("a76", "pedimento_config_parameters"),
|
||||
("a76", "pedimento_config_surcharges"),
|
||||
("a76", "pedimento_config_update_rectification"),
|
||||
("a76", "pedimento_config_updates"),
|
||||
("a76", "pedimento_containers"),
|
||||
("a76", "pedimento_contributions"),
|
||||
("a76", "pedimento_customs_offices"),
|
||||
("a76", "pedimento_dates"),
|
||||
("a76", "pedimento_decrementables"),
|
||||
("a76", "pedimento_guides"),
|
||||
("a76", "pedimento_incrementables"),
|
||||
("a76", "pedimento_indexes"),
|
||||
("a76", "pedimento_packages"),
|
||||
("a76", "pedimento_payments"),
|
||||
("a76", "pedimento_rectification_destination"),
|
||||
("a76", "pedimento_rectification_origin"),
|
||||
("a76", "pedimento_seals"),
|
||||
("a76", "pedimento_transport_carriers"),
|
||||
("a76", "pedimento_transport_means"),
|
||||
("a76", "pedimento_validation"),
|
||||
("a76", "pedimentos"),
|
||||
("a76", "permission_rule_oct"),
|
||||
("a76", "permission_rule_octave"),
|
||||
("a76", "ports"),
|
||||
("a76", "prevalidators"),
|
||||
("a76", "previous_fractions"),
|
||||
("a76", "seal"),
|
||||
("a76", "sectors"),
|
||||
("a76", "signatures"),
|
||||
("a76", "subassembly_entries"),
|
||||
("a76", "trailer"),
|
||||
("a76", "transporter"),
|
||||
("a76", "unit_conversions"),
|
||||
("a76", "units_of_measure"),
|
||||
("a76", "units_of_measure_general"),
|
||||
("a76", "us_tariff_fractions"),
|
||||
("a76", "value_manifestations"),
|
||||
("a76", "vehicle"),
|
||||
("core", "company_roles"),
|
||||
("core", "role_permissions"),
|
||||
("core", "user_company_permissions"),
|
||||
("core", "user_company_roles"),
|
||||
("public", "warning_fractions"),
|
||||
]
|
||||
|
||||
TABLES_COMPANY_ONLY: list[tuple[str, str]] = [
|
||||
("a24", "inv_aphis_catalog"),
|
||||
("a76", "company_address"),
|
||||
("a76", "company_certification"),
|
||||
("a76", "company_cfdi"),
|
||||
("a76", "company_digital_certificate"),
|
||||
("a76", "company_electronic_agent"),
|
||||
("a76", "company_prevalidator"),
|
||||
]
|
||||
|
||||
|
||||
POLICY_TENANT_ONLY = "tenant_isolation"
|
||||
POLICY_TENANT_COMPANY = "tenant_company_isolation"
|
||||
POLICY_COMPANY_ONLY = "company_isolation"
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Habilita RLS con políticas de aislamiento por tenant_id / company_id."""
|
||||
op.execute("CREATE SCHEMA IF NOT EXISTS app")
|
||||
|
||||
op.execute(
|
||||
"""
|
||||
CREATE OR REPLACE FUNCTION app.current_tenant_id() RETURNS INTEGER
|
||||
LANGUAGE sql STABLE AS $$
|
||||
SELECT NULLIF(current_setting('app.tenant_id', true), '')::INTEGER
|
||||
$$
|
||||
"""
|
||||
)
|
||||
op.execute(
|
||||
"""
|
||||
CREATE OR REPLACE FUNCTION app.current_company_id() RETURNS INTEGER
|
||||
LANGUAGE sql STABLE AS $$
|
||||
SELECT NULLIF(current_setting('app.company_id', true), '')::INTEGER
|
||||
$$
|
||||
"""
|
||||
)
|
||||
|
||||
for schema, table in TABLES_TENANT_ONLY:
|
||||
op.execute(f'ALTER TABLE "{schema}"."{table}" ENABLE ROW LEVEL SECURITY')
|
||||
op.execute(f'ALTER TABLE "{schema}"."{table}" FORCE ROW LEVEL SECURITY')
|
||||
op.execute(
|
||||
f"""
|
||||
CREATE POLICY {POLICY_TENANT_ONLY} ON "{schema}"."{table}"
|
||||
USING (tenant_id = app.current_tenant_id())
|
||||
WITH CHECK (tenant_id = app.current_tenant_id())
|
||||
"""
|
||||
)
|
||||
|
||||
for schema, table in TABLES_TENANT_AND_COMPANY:
|
||||
op.execute(f'ALTER TABLE "{schema}"."{table}" ENABLE ROW LEVEL SECURITY')
|
||||
op.execute(f'ALTER TABLE "{schema}"."{table}" FORCE ROW LEVEL SECURITY')
|
||||
op.execute(
|
||||
f"""
|
||||
CREATE POLICY {POLICY_TENANT_COMPANY} ON "{schema}"."{table}"
|
||||
USING (
|
||||
tenant_id = app.current_tenant_id()
|
||||
AND (
|
||||
app.current_company_id() IS NULL
|
||||
OR company_id = app.current_company_id()
|
||||
)
|
||||
)
|
||||
WITH CHECK (
|
||||
tenant_id = app.current_tenant_id()
|
||||
AND (
|
||||
app.current_company_id() IS NULL
|
||||
OR company_id = app.current_company_id()
|
||||
)
|
||||
)
|
||||
"""
|
||||
)
|
||||
|
||||
for schema, table in TABLES_COMPANY_ONLY:
|
||||
op.execute(f'ALTER TABLE "{schema}"."{table}" ENABLE ROW LEVEL SECURITY')
|
||||
op.execute(f'ALTER TABLE "{schema}"."{table}" FORCE ROW LEVEL SECURITY')
|
||||
op.execute(
|
||||
f"""
|
||||
CREATE POLICY {POLICY_COMPANY_ONLY} ON "{schema}"."{table}"
|
||||
USING (
|
||||
EXISTS (
|
||||
SELECT 1 FROM a76.company c
|
||||
WHERE c.id = "{schema}"."{table}".company_id
|
||||
AND c.tenant_id = app.current_tenant_id()
|
||||
)
|
||||
AND (
|
||||
app.current_company_id() IS NULL
|
||||
OR company_id = app.current_company_id()
|
||||
)
|
||||
)
|
||||
WITH CHECK (
|
||||
EXISTS (
|
||||
SELECT 1 FROM a76.company c
|
||||
WHERE c.id = "{schema}"."{table}".company_id
|
||||
AND c.tenant_id = app.current_tenant_id()
|
||||
)
|
||||
AND (
|
||||
app.current_company_id() IS NULL
|
||||
OR company_id = app.current_company_id()
|
||||
)
|
||||
)
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Revierte: drop policies, deshabilita RLS y elimina helpers."""
|
||||
for schema, table in TABLES_COMPANY_ONLY:
|
||||
op.execute(
|
||||
f'DROP POLICY IF EXISTS {POLICY_COMPANY_ONLY} ON "{schema}"."{table}"'
|
||||
)
|
||||
op.execute(f'ALTER TABLE "{schema}"."{table}" NO FORCE ROW LEVEL SECURITY')
|
||||
op.execute(f'ALTER TABLE "{schema}"."{table}" DISABLE ROW LEVEL SECURITY')
|
||||
|
||||
for schema, table in TABLES_TENANT_AND_COMPANY:
|
||||
op.execute(
|
||||
f'DROP POLICY IF EXISTS {POLICY_TENANT_COMPANY} ON "{schema}"."{table}"'
|
||||
)
|
||||
op.execute(f'ALTER TABLE "{schema}"."{table}" NO FORCE ROW LEVEL SECURITY')
|
||||
op.execute(f'ALTER TABLE "{schema}"."{table}" DISABLE ROW LEVEL SECURITY')
|
||||
|
||||
for schema, table in TABLES_TENANT_ONLY:
|
||||
op.execute(
|
||||
f'DROP POLICY IF EXISTS {POLICY_TENANT_ONLY} ON "{schema}"."{table}"'
|
||||
)
|
||||
op.execute(f'ALTER TABLE "{schema}"."{table}" NO FORCE ROW LEVEL SECURITY')
|
||||
op.execute(f'ALTER TABLE "{schema}"."{table}" DISABLE ROW LEVEL SECURITY')
|
||||
|
||||
op.execute("DROP FUNCTION IF EXISTS app.current_company_id()")
|
||||
op.execute("DROP FUNCTION IF EXISTS app.current_tenant_id()")
|
||||
op.execute("DROP SCHEMA IF EXISTS app")
|
||||
@@ -331,7 +331,11 @@ def digitalizar_expediente_archivo(
|
||||
"request_data": body.model_dump(),
|
||||
"company_id": company_id,
|
||||
"tenant_id": tenant_id,
|
||||
}
|
||||
},
|
||||
headers={
|
||||
"rls_tenant_id": str(int(tenant_id)),
|
||||
"rls_company_id": str(int(company_id)),
|
||||
},
|
||||
)
|
||||
|
||||
return DigitalizarResponse(
|
||||
@@ -367,7 +371,11 @@ def registrar_digitalizacion(
|
||||
"request_data": {"rfc_consulta": body.rfc_consulta},
|
||||
"company_id": company_id,
|
||||
"tenant_id": tenant_id,
|
||||
}
|
||||
},
|
||||
headers={
|
||||
"rls_tenant_id": str(int(tenant_id)),
|
||||
"rls_company_id": str(int(company_id)),
|
||||
},
|
||||
)
|
||||
launched.append({"id": record_id, "task_id": task.id})
|
||||
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
from celery import Task
|
||||
|
||||
from core.celery_app import celery_app
|
||||
from core.database import CoreSessionLocal
|
||||
from core.database import scoped_core_db
|
||||
from core.exceptions import ValidationException
|
||||
|
||||
from api.v1.modules.a76.invoices.models import InvoiceHeader, InvoiceStatus
|
||||
@@ -18,9 +18,8 @@ def process_export_invoice_task(self: Task, invoice_id: int, tenant_id: str, com
|
||||
Procesa una factura de exportación ejecutando todas las validaciones y
|
||||
actualizaciones del proceso principal de exportación con reporte de progreso.
|
||||
"""
|
||||
db = CoreSessionLocal()
|
||||
with scoped_core_db(tenant_id=int(tenant_id), company_id=int(company_id)) as db:
|
||||
try:
|
||||
# ── Paso 1: Cargar factura ────────────────────────────────────────────
|
||||
_progress(self, 5, "Cargando factura...")
|
||||
invoice: InvoiceHeader | None = db.get(InvoiceHeader, invoice_id)
|
||||
|
||||
@@ -56,5 +55,3 @@ def process_export_invoice_task(self: Task, invoice_id: int, tenant_id: str, com
|
||||
except Exception as exc:
|
||||
db.rollback()
|
||||
raise exc
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
from celery import Task
|
||||
|
||||
from core.celery_app import celery_app
|
||||
from core.database import CoreSessionLocal
|
||||
from core.database import scoped_core_db
|
||||
from core.exceptions import ErrorCollector, ValidationException
|
||||
|
||||
from api.v1.modules.a76.invoices.models import InvoiceHeader
|
||||
@@ -26,9 +26,8 @@ def revert_invoice_task(
|
||||
validaciones y reversiones del proceso principal (revert/main_process) con
|
||||
reporte de progreso.
|
||||
"""
|
||||
db = CoreSessionLocal()
|
||||
with scoped_core_db(tenant_id=int(tenant_id), company_id=int(company_id)) as db:
|
||||
try:
|
||||
# ── Paso 1: Cargar factura ────────────────────────────────────────────
|
||||
_progress(self, 5, "Cargando factura...")
|
||||
invoice: InvoiceHeader | None = db.get(InvoiceHeader, invoice_id)
|
||||
if invoice is None:
|
||||
@@ -49,7 +48,6 @@ def revert_invoice_task(
|
||||
|
||||
errors = ErrorCollector()
|
||||
|
||||
# ── Paso 2: Pre-validaciones ──────────────────────────────────────────
|
||||
_progress(self, 10, "Validando estatus de la factura...")
|
||||
lines = pre_validators(db, invoice, tenant_id, company_id, errors)
|
||||
if not lines:
|
||||
@@ -61,7 +59,6 @@ def revert_invoice_task(
|
||||
)
|
||||
errors.raise_if_errors()
|
||||
|
||||
# ── Paso 3: Validar cantidades y ejecutar reversión ───────────────────
|
||||
_progress(self, 40, "Verificando saldos de partidas...")
|
||||
sql_errors = revert_process(
|
||||
db=db,
|
||||
@@ -73,7 +70,6 @@ def revert_invoice_task(
|
||||
cancelled_by=cancelled_by,
|
||||
)
|
||||
|
||||
# ── Paso 4: Confirmar transacción ─────────────────────────────────────
|
||||
_progress(self, 95, "Anulando saldos de inventario y confirmando...")
|
||||
db.flush()
|
||||
db.commit()
|
||||
@@ -94,5 +90,3 @@ def revert_invoice_task(
|
||||
except Exception as exc:
|
||||
db.rollback()
|
||||
raise exc
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
@@ -3,7 +3,7 @@ import logging
|
||||
from celery import Task
|
||||
|
||||
from core.celery_app import celery_app
|
||||
from core.database import CoreSessionLocal
|
||||
from core.database import scoped_core_db
|
||||
from core.exceptions import ValidationException
|
||||
|
||||
from api.v1.modules.a76.invoices.models import InvoiceHeader, InvoiceStatus
|
||||
@@ -22,10 +22,8 @@ def process_invoice_task(self: Task, invoice_id: int, tenant_id: str, company_id
|
||||
Procesa una factura de importación ejecutando todas las validaciones y
|
||||
actualizaciones del proceso principal (main_process) con reporte de progreso.
|
||||
"""
|
||||
db = CoreSessionLocal()
|
||||
with scoped_core_db(tenant_id=int(tenant_id), company_id=int(company_id)) as db:
|
||||
try:
|
||||
# ── Paso 1: Cargar factura ────────────────────────────────────────────
|
||||
# ── Paso 1: Cargar factura ────────────────────────────────────────────
|
||||
_progress(self, 5, "Cargando factura...")
|
||||
invoice: InvoiceHeader | None = db.get(InvoiceHeader, invoice_id)
|
||||
|
||||
@@ -44,18 +42,15 @@ def process_invoice_task(self: Task, invoice_id: int, tenant_id: str, company_id
|
||||
"errors": [{"field": "status", "message": "Factura ya procesada."}],
|
||||
}
|
||||
|
||||
# ── Paso 2: Ejecutar Proceso Principal ───────────────────────────────
|
||||
# Unificamos lógica: El task solo llama al main_process centralizado.
|
||||
_progress(self, 20, "Iniciando procesamiento de factura...")
|
||||
result = main_process(
|
||||
db=db,
|
||||
invoice=invoice,
|
||||
tenant_id=tenant_id,
|
||||
company_id=company_id,
|
||||
username=username
|
||||
username=username,
|
||||
)
|
||||
|
||||
# ── Paso 3: Confirmar transacción ─────────────────────────────────────
|
||||
_progress(self, 95, "Confirmando cambios...")
|
||||
db.commit()
|
||||
|
||||
@@ -73,5 +68,3 @@ def process_invoice_task(self: Task, invoice_id: int, tenant_id: str, company_id
|
||||
db.rollback()
|
||||
logger.error(f"Error en process_invoice_task: {str(exc)}", exc_info=True)
|
||||
raise exc
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
from celery import Task
|
||||
|
||||
from core.celery_app import celery_app
|
||||
from core.database import CoreSessionLocal
|
||||
from core.database import scoped_core_db
|
||||
from core.exceptions import ErrorCollector, ValidationException
|
||||
|
||||
from api.v1.modules.a76.invoices.models import InvoiceHeader
|
||||
@@ -26,9 +26,8 @@ def revert_invoice_task(
|
||||
validaciones y reversiones del proceso principal (revert/main_process) con
|
||||
reporte de progreso.
|
||||
"""
|
||||
db = CoreSessionLocal()
|
||||
with scoped_core_db(tenant_id=int(tenant_id), company_id=int(company_id)) as db:
|
||||
try:
|
||||
# ── Paso 1: Cargar factura ────────────────────────────────────────────
|
||||
_progress(self, 5, "Cargando factura...")
|
||||
invoice: InvoiceHeader | None = db.get(InvoiceHeader, invoice_id)
|
||||
if invoice is None:
|
||||
@@ -49,7 +48,6 @@ def revert_invoice_task(
|
||||
|
||||
errors = ErrorCollector()
|
||||
|
||||
# ── Paso 2: Pre-validaciones ──────────────────────────────────────────
|
||||
_progress(self, 10, "Validando estatus de la factura...")
|
||||
lines = pre_validators(db, invoice, tenant_id, company_id, errors)
|
||||
if not lines:
|
||||
@@ -61,7 +59,6 @@ def revert_invoice_task(
|
||||
)
|
||||
errors.raise_if_errors()
|
||||
|
||||
# ── Paso 3: Validar cantidades y ejecutar reversión ───────────────────
|
||||
_progress(self, 40, "Verificando saldos de partidas...")
|
||||
sql_errors = revert_process(
|
||||
db=db,
|
||||
@@ -72,7 +69,6 @@ def revert_invoice_task(
|
||||
errors=errors,
|
||||
)
|
||||
|
||||
# ── Paso 4: Confirmar transacción ─────────────────────────────────────
|
||||
_progress(self, 95, "Anulando saldos de inventario y confirmando...")
|
||||
db.flush()
|
||||
db.commit()
|
||||
@@ -93,5 +89,3 @@ def revert_invoice_task(
|
||||
except Exception as exc:
|
||||
db.rollback()
|
||||
raise exc
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
@@ -3,6 +3,8 @@ from typing import Any
|
||||
from celery import Task
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from core.database import rls_company_var, rls_tenant_var
|
||||
|
||||
from .service import TaskTrackerService
|
||||
|
||||
|
||||
@@ -21,7 +23,25 @@ def track_and_dispatch(
|
||||
task_id: str | None = None,
|
||||
meta_payload: dict[str, Any] | None = None,
|
||||
):
|
||||
celery_task = task.apply_async(args=args or [], kwargs=kwargs or {}, task_id=task_id)
|
||||
# Propaga contexto RLS vía Celery headers (leídos en task_prerun) y
|
||||
# ContextVars (para modo eager, donde before_task_publish no dispara).
|
||||
headers = {"rls_tenant_id": str(int(tenant_id))}
|
||||
if company_id is not None:
|
||||
headers["rls_company_id"] = str(int(company_id))
|
||||
|
||||
token_t = rls_tenant_var.set(int(tenant_id))
|
||||
token_c = rls_company_var.set(int(company_id) if company_id is not None else None)
|
||||
try:
|
||||
celery_task = task.apply_async(
|
||||
args=args or [],
|
||||
kwargs=kwargs or {},
|
||||
task_id=task_id,
|
||||
headers=headers,
|
||||
)
|
||||
finally:
|
||||
rls_tenant_var.reset(token_t)
|
||||
rls_company_var.reset(token_c)
|
||||
|
||||
tracker = TaskTrackerService(db)
|
||||
tracker.register_dispatch(
|
||||
task_id=celery_task.id,
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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(
|
||||
@@ -53,6 +69,7 @@ class TenantMiddleware(BaseHTTPMiddleware):
|
||||
|
||||
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,16 +116,16 @@ 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
|
||||
with scoped_core_db(tenant_id=tenant_id) as db:
|
||||
from api.v1.modules.core.licenses.service import LicenseService
|
||||
|
||||
license_service = LicenseService(db)
|
||||
@@ -129,9 +141,7 @@ class LicenseValidationMiddleware(BaseHTTPMiddleware):
|
||||
}
|
||||
)
|
||||
|
||||
# Agregar info de licencia al request state
|
||||
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
|
||||
|
||||
202
backend/tests/integration/test_rls_tenant_company.py
Normal file
202
backend/tests/integration/test_rls_tenant_company.py
Normal file
@@ -0,0 +1,202 @@
|
||||
"""Pruebas de aislamiento Row-Level Security por ``tenant_id`` y ``company_id``.
|
||||
|
||||
Las políticas se crean en la migración
|
||||
``d1a2b3c4e5f6_enable_rls_tenant_company`` y dependen de las GUCs
|
||||
``app.tenant_id`` / ``app.company_id`` fijadas por la aplicación.
|
||||
|
||||
Como ``postgres`` hace BYPASSRLS por defecto, los tests cambian de rol
|
||||
dentro de la transacción a un usuario sin ese atributo (``anexo76_rls_test``)
|
||||
antes de contar filas. Si la conexión de pruebas no tiene privilegios para
|
||||
crear el rol, el set completo se salta con ``pytest.skip``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from typing import Iterator
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import text
|
||||
from sqlalchemy.exc import DBAPIError, ProgrammingError
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from tests.conftest import TestingSessionLocal, engine
|
||||
from tests.fixtures.builders import ensure_tenant_company
|
||||
|
||||
|
||||
TEST_ROLE = "anexo76_rls_test"
|
||||
RLS_SCHEMAS = ("core", "a76", "a24", "public")
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def rls_role() -> str:
|
||||
"""Crea (idempotente) un rol no-superusuario con los privilegios mínimos
|
||||
necesarios para ejercitar las políticas RLS en los tests."""
|
||||
try:
|
||||
with engine.begin() as conn:
|
||||
conn.execute(
|
||||
text(
|
||||
f"""
|
||||
DO $$
|
||||
BEGIN
|
||||
IF NOT EXISTS (SELECT 1 FROM pg_roles WHERE rolname = '{TEST_ROLE}') THEN
|
||||
CREATE ROLE {TEST_ROLE} NOLOGIN NOBYPASSRLS;
|
||||
END IF;
|
||||
END $$;
|
||||
"""
|
||||
)
|
||||
)
|
||||
for schema in RLS_SCHEMAS:
|
||||
conn.execute(text(f'GRANT USAGE ON SCHEMA "{schema}" TO {TEST_ROLE}'))
|
||||
conn.execute(
|
||||
text(
|
||||
f'GRANT SELECT, INSERT, UPDATE, DELETE ON ALL TABLES IN SCHEMA "{schema}" TO {TEST_ROLE}'
|
||||
)
|
||||
)
|
||||
conn.execute(
|
||||
text(
|
||||
f'GRANT USAGE, SELECT ON ALL SEQUENCES IN SCHEMA "{schema}" TO {TEST_ROLE}'
|
||||
)
|
||||
)
|
||||
except ProgrammingError as exc:
|
||||
pytest.skip(f"No hay privilegio para preparar rol de pruebas RLS: {exc}")
|
||||
return TEST_ROLE
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def rls_session() -> Iterator[Session]:
|
||||
"""Sesión con transacción externa + rollback (no contamina BD real)."""
|
||||
connection = engine.connect()
|
||||
transaction = connection.begin()
|
||||
session = TestingSessionLocal(bind=connection)
|
||||
try:
|
||||
yield session
|
||||
finally:
|
||||
session.close()
|
||||
transaction.rollback()
|
||||
connection.close()
|
||||
|
||||
|
||||
def _allocate_id() -> int:
|
||||
return 1_700_000_000 + (uuid.uuid4().int % 90_000_000)
|
||||
|
||||
|
||||
def _bootstrap_two_tenants(session: Session) -> tuple[tuple[int, int], tuple[int, int]]:
|
||||
tid_a, tid_b = _allocate_id(), _allocate_id()
|
||||
cid_a, cid_b = _allocate_id(), _allocate_id()
|
||||
ensure_tenant_company(session, tenant_id=tid_a, company_id=cid_a)
|
||||
ensure_tenant_company(session, tenant_id=tid_b, company_id=cid_b)
|
||||
return (tid_a, cid_a), (tid_b, cid_b)
|
||||
|
||||
|
||||
def _seed_clients(session: Session, tenant_id: int, company_id: int, *, count: int) -> None:
|
||||
"""Inserta filas en ``a76.clients_and_providers`` con SQL crudo para
|
||||
aislarnos del ORM y medir filtrado puro a nivel de BD."""
|
||||
for i in range(count):
|
||||
session.execute(
|
||||
text(
|
||||
"""
|
||||
INSERT INTO a76.clients_and_providers
|
||||
(tenant_id, company_id, name, client_or_provider, is_active)
|
||||
VALUES
|
||||
(:tid, :cid, :name, 'BOTH', true)
|
||||
"""
|
||||
),
|
||||
{"tid": tenant_id, "cid": company_id, "name": f"Seed {tenant_id}/{company_id}/{i}"},
|
||||
)
|
||||
session.flush()
|
||||
|
||||
|
||||
def _count_under_role(
|
||||
session: Session,
|
||||
tenant_id: int | None,
|
||||
company_id: int | None,
|
||||
*,
|
||||
role: str,
|
||||
table: str = "a76.clients_and_providers",
|
||||
) -> int:
|
||||
"""Cuenta filas tras cambiar a un rol sin BYPASSRLS y fijar el contexto."""
|
||||
session.execute(text(f"SET LOCAL ROLE {role}"))
|
||||
session.execute(
|
||||
text("SELECT set_config('app.tenant_id', :t, true), set_config('app.company_id', :c, true)"),
|
||||
{
|
||||
"t": "" if tenant_id is None else str(tenant_id),
|
||||
"c": "" if company_id is None else str(company_id),
|
||||
},
|
||||
)
|
||||
try:
|
||||
return int(session.execute(text(f"SELECT count(*) FROM {table}")).scalar() or 0)
|
||||
finally:
|
||||
session.execute(text("RESET ROLE"))
|
||||
|
||||
|
||||
def test_tenant_isolation_hides_rows_from_other_tenant(
|
||||
rls_session: Session, rls_role: str
|
||||
) -> None:
|
||||
(tid_a, cid_a), (tid_b, cid_b) = _bootstrap_two_tenants(rls_session)
|
||||
_seed_clients(rls_session, tid_a, cid_a, count=3)
|
||||
_seed_clients(rls_session, tid_b, cid_b, count=2)
|
||||
|
||||
visible_a = _count_under_role(rls_session, tid_a, None, role=rls_role)
|
||||
visible_b = _count_under_role(rls_session, tid_b, None, role=rls_role)
|
||||
|
||||
assert visible_a == 3, "Tenant A debe ver solo sus 3 filas"
|
||||
assert visible_b == 2, "Tenant B debe ver solo sus 2 filas"
|
||||
|
||||
|
||||
def test_company_scope_narrows_within_tenant(
|
||||
rls_session: Session, rls_role: str
|
||||
) -> None:
|
||||
(tid_a, cid_a), _ = _bootstrap_two_tenants(rls_session)
|
||||
|
||||
second_company_id = _allocate_id()
|
||||
ensure_tenant_company(rls_session, tenant_id=tid_a, company_id=second_company_id)
|
||||
|
||||
_seed_clients(rls_session, tid_a, cid_a, count=4)
|
||||
_seed_clients(rls_session, tid_a, second_company_id, count=7)
|
||||
|
||||
without_company = _count_under_role(rls_session, tid_a, None, role=rls_role)
|
||||
with_company_a = _count_under_role(rls_session, tid_a, cid_a, role=rls_role)
|
||||
with_second = _count_under_role(rls_session, tid_a, second_company_id, role=rls_role)
|
||||
|
||||
assert without_company == 11, "Sin company context el tenant ve ambas compañías"
|
||||
assert with_company_a == 4
|
||||
assert with_second == 7
|
||||
|
||||
|
||||
def test_missing_tenant_context_returns_zero_rows(
|
||||
rls_session: Session, rls_role: str
|
||||
) -> None:
|
||||
(tid_a, cid_a), _ = _bootstrap_two_tenants(rls_session)
|
||||
_seed_clients(rls_session, tid_a, cid_a, count=5)
|
||||
|
||||
visible = _count_under_role(rls_session, None, None, role=rls_role)
|
||||
assert visible == 0, "Sin app.tenant_id la política debe devolver 0 filas"
|
||||
|
||||
|
||||
def test_insert_violation_respects_tenant_policy(
|
||||
rls_session: Session, rls_role: str
|
||||
) -> None:
|
||||
"""La cláusula ``WITH CHECK`` debe rechazar inserts fuera de contexto."""
|
||||
(tid_a, cid_a), (tid_b, _) = _bootstrap_two_tenants(rls_session)
|
||||
|
||||
rls_session.execute(text(f"SET LOCAL ROLE {rls_role}"))
|
||||
rls_session.execute(
|
||||
text("SELECT set_config('app.tenant_id', :t, true), set_config('app.company_id', :c, true)"),
|
||||
{"t": str(tid_a), "c": str(cid_a)},
|
||||
)
|
||||
try:
|
||||
with pytest.raises(DBAPIError):
|
||||
rls_session.execute(
|
||||
text(
|
||||
"""
|
||||
INSERT INTO a76.clients_and_providers
|
||||
(tenant_id, company_id, name, client_or_provider, is_active)
|
||||
VALUES
|
||||
(:tid, :cid, 'Cross-tenant attempt', 'BOTH', true)
|
||||
"""
|
||||
),
|
||||
{"tid": tid_b, "cid": cid_a},
|
||||
)
|
||||
finally:
|
||||
rls_session.rollback()
|
||||
@@ -395,6 +395,75 @@ elif tenant.type == "DEDICATED":
|
||||
- **Row-level security** en BD compartida
|
||||
- **BD dedicada** para mayor aislamiento (enterprise)
|
||||
|
||||
### Row-Level Security (RLS) en BD compartida
|
||||
|
||||
> Convención alineada al skill `aduanasoft-dev-standards` (sección 10).
|
||||
> La capa API sigue siendo responsable del control fino (roles/permisos
|
||||
> con Keycloak + `PermissionService`); RLS añade **defensa en profundidad**
|
||||
> a nivel de BD para que un bug en un `WHERE` no permita salirse del tenant.
|
||||
|
||||
#### Variables de sesión (`SET LOCAL`)
|
||||
|
||||
| GUC | Origen | Comportamiento RLS |
|
||||
|-----|--------|--------------------|
|
||||
| `app.tenant_id` | JWT (`TenantMiddleware`) → `request.state.tenant_id` | Obligatoria. Si está vacía, `app.current_tenant_id()` retorna `NULL` y las políticas devuelven `0` filas (fail-closed). |
|
||||
| `app.company_id` | Header `X-Company-Id` o cookie `active_company_id` | Opcional. Si está vacía, el tenant ve **todas sus compañías** (útil para selectores de compañía y bootstrap). |
|
||||
|
||||
Ambas se fijan con `SET LOCAL` al inicio de cada transacción —
|
||||
**nunca** con `SET` global, para no contaminar conexiones del pool.
|
||||
|
||||
Helpers SQL definidos por la migración `d1a2b3c4e5f6_enable_rls_tenant_company`:
|
||||
|
||||
```sql
|
||||
CREATE FUNCTION app.current_tenant_id() RETURNS INTEGER LANGUAGE sql STABLE AS
|
||||
$$ SELECT NULLIF(current_setting('app.tenant_id', true), '')::INTEGER $$;
|
||||
CREATE FUNCTION app.current_company_id() RETURNS INTEGER LANGUAGE sql STABLE AS
|
||||
$$ SELECT NULLIF(current_setting('app.company_id', true), '')::INTEGER $$;
|
||||
```
|
||||
|
||||
#### Tipos de política
|
||||
|
||||
1. **Solo `tenant_id`** (p.ej. `a76.company`, `core.licenses`):
|
||||
`tenant_id = app.current_tenant_id()`.
|
||||
2. **`tenant_id` + `company_id`** (`TenantScopedMixin`, mayoría de tablas
|
||||
`a24/`a76/`core`): además exige `company_id = app.current_company_id()`
|
||||
cuando esa GUC está fijada.
|
||||
3. **Solo `company_id`** (algunas tablas `a76.company_*`): valida el
|
||||
`tenant_id` indirectamente vía `EXISTS` contra `a76.company`.
|
||||
|
||||
Todas las tablas usan `FORCE ROW LEVEL SECURITY` para que la política
|
||||
aplique también al owner. Las únicas tablas core **excluidas** son
|
||||
`core.tenants` y `core.user_tenants` — necesarias para el bootstrap del
|
||||
selector de tenant antes de tener contexto fijado.
|
||||
|
||||
#### Propagación del contexto
|
||||
|
||||
| Camino | Cómo se fija el contexto |
|
||||
|--------|--------------------------|
|
||||
| HTTP request | `TenantMiddleware` rellena `request.state.tenant_id`/`company_id`; `get_core_db` / `get_async_core_db` leen esos valores y los guardan en `Session.info`. Un listener `after_begin` ejecuta `SET LOCAL` por transacción. |
|
||||
| `LicenseValidationMiddleware` | Usa `scoped_core_db(tenant_id=...)` para que la consulta de licencia entre con contexto RLS válido. |
|
||||
| Tareas Celery | `track_and_dispatch` inyecta `rls_tenant_id` / `rls_company_id` en los headers del task; los signals `task_prerun`/`task_postrun` los copian a `ContextVar`s del worker, que el listener `after_begin` consume como fallback. Tareas críticas (imports/exports de invoices, expediente) abren la sesión con `scoped_core_db(tenant_id=..., company_id=...)`. |
|
||||
| Tests | Las suites de pytest pueden usar `scoped_core_db(...)` o emular el flujo con `set_config('app.tenant_id', ...)` antes del query. Hay un set de tests en `backend/tests/integration/test_rls_tenant_company.py` que valida aislamiento A vs B usando un rol sin `BYPASSRLS`. |
|
||||
|
||||
#### Reparto de responsabilidades
|
||||
|
||||
| Capa | Decide |
|
||||
|------|--------|
|
||||
| **API (FastAPI + Keycloak + `PermissionService`)** | Roles, permisos por compañía, accesos a recursos concretos (`validate_access_to_resource`), reglas de negocio. |
|
||||
| **RLS (PostgreSQL)** | Límite estructural duro: `tenant_id` y `company_id`. **No** modela roles/permisos para evitar duplicar lógica fina con la API. |
|
||||
|
||||
#### Operación / DevOps
|
||||
|
||||
- En **producción** la API debe conectar con un rol **sin** `BYPASSRLS`
|
||||
(`postgres` superusuario lo bypassea por diseño). El `docker-compose.yml`
|
||||
de desarrollo usa `postgres` deliberadamente para no romper migraciones;
|
||||
los tests crean un rol `anexo76_rls_test` para ejercitar las políticas.
|
||||
- Los jobs/ETL/migraciones que necesiten ver todos los tenants deben usar
|
||||
un rol técnico explícito con `BYPASSRLS` o fijar `app.tenant_id` por
|
||||
iteración — nunca asumir que la sesión global "ve todo".
|
||||
- La migración `d1a2b3c4e5f6_enable_rls_tenant_company` tiene `downgrade()`
|
||||
completo (drop policies + `DISABLE ROW LEVEL SECURITY`) para revertir.
|
||||
|
||||
### Validación de Licencias
|
||||
- Middleware verifica en cada request:
|
||||
- ✓ Licencia activa
|
||||
|
||||
Reference in New Issue
Block a user