diff --git a/backend/alembic/versions/d1a2b3c4e5f6_enable_rls_tenant_company.py b/backend/alembic/versions/d1a2b3c4e5f6_enable_rls_tenant_company.py new file mode 100644 index 00000000..9d334b10 --- /dev/null +++ b/backend/alembic/versions/d1a2b3c4e5f6_enable_rls_tenant_company.py @@ -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") diff --git a/backend/api/v1/modules/a76/expediente_archivos/routes.py b/backend/api/v1/modules/a76/expediente_archivos/routes.py index af9f6c01..c553cfee 100644 --- a/backend/api/v1/modules/a76/expediente_archivos/routes.py +++ b/backend/api/v1/modules/a76/expediente_archivos/routes.py @@ -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}) diff --git a/backend/api/v1/modules/a76/invoices/exports/process/task.py b/backend/api/v1/modules/a76/invoices/exports/process/task.py index 434f9ffd..935b58ff 100644 --- a/backend/api/v1/modules/a76/invoices/exports/process/task.py +++ b/backend/api/v1/modules/a76/invoices/exports/process/task.py @@ -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,43 +18,40 @@ 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() - try: - # ── Paso 1: Cargar factura ──────────────────────────────────────────── - _progress(self, 5, "Cargando factura...") - invoice: InvoiceHeader | None = db.get(InvoiceHeader, invoice_id) - - if invoice is None: + with scoped_core_db(tenant_id=int(tenant_id), company_id=int(company_id)) as db: + try: + _progress(self, 5, "Cargando factura...") + invoice: InvoiceHeader | None = db.get(InvoiceHeader, invoice_id) + + if invoice is None: + return { + "status": "error", + "message": f"Factura con id {invoice_id} no encontrada.", + "errors": [], + } + + _progress(self, 10, "Verificando estatus de seguridad...") + if invoice.status == InvoiceStatus.PROCESSED: + return { + "status": "error", + "message": f"La factura {invoice.invoice_number} ya se encuentra procesada.", + "errors": [{"field": "status", "message": "Factura ya procesada."}], + } + + _progress(self, 15, "Iniciando proceso principal de exportación...") + result = main_process(db, invoice, tenant_id, company_id, username=username) + + db.commit() + _progress(self, 100, "Proceso completado.") + return {**result, "invoice_id": invoice_id} + + except ValidationException as exc: + db.rollback() return { - "status": "error", - "message": f"Factura con id {invoice_id} no encontrada.", - "errors": [], + "status": "validation_error", + "message": exc.message, + "errors": exc.errors, } - - _progress(self, 10, "Verificando estatus de seguridad...") - if invoice.status == InvoiceStatus.PROCESSED: - return { - "status": "error", - "message": f"La factura {invoice.invoice_number} ya se encuentra procesada.", - "errors": [{"field": "status", "message": "Factura ya procesada."}], - } - - _progress(self, 15, "Iniciando proceso principal de exportación...") - result = main_process(db, invoice, tenant_id, company_id, username=username) - - db.commit() - _progress(self, 100, "Proceso completado.") - return {**result, "invoice_id": invoice_id} - - except ValidationException as exc: - db.rollback() - return { - "status": "validation_error", - "message": exc.message, - "errors": exc.errors, - } - except Exception as exc: - db.rollback() - raise exc - finally: - db.close() + except Exception as exc: + db.rollback() + raise exc diff --git a/backend/api/v1/modules/a76/invoices/exports/revert/task.py b/backend/api/v1/modules/a76/invoices/exports/revert/task.py index f6e157e5..e2e8ead4 100644 --- a/backend/api/v1/modules/a76/invoices/exports/revert/task.py +++ b/backend/api/v1/modules/a76/invoices/exports/revert/task.py @@ -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,73 +26,67 @@ def revert_invoice_task( validaciones y reversiones del proceso principal (revert/main_process) con reporte de progreso. """ - db = CoreSessionLocal() - try: - # ── Paso 1: Cargar factura ──────────────────────────────────────────── - _progress(self, 5, "Cargando factura...") - invoice: InvoiceHeader | None = db.get(InvoiceHeader, invoice_id) - if invoice is None: - return { - "status": "error", - "message": f"Factura con id {invoice_id} no encontrada.", - "errors": [], - } + with scoped_core_db(tenant_id=int(tenant_id), company_id=int(company_id)) as db: + try: + _progress(self, 5, "Cargando factura...") + invoice: InvoiceHeader | None = db.get(InvoiceHeader, invoice_id) + if invoice is None: + return { + "status": "error", + "message": f"Factura con id {invoice_id} no encontrada.", + "errors": [], + } - _progress(self, 10, "Verificando estatus de seguridad...") - from api.v1.modules.a76.invoices.models import InvoiceStatus - if invoice.status != InvoiceStatus.PROCESSED: - return { - "status": "error", - "message": f"La factura {invoice.invoice_number} no se puede revertir porque no está procesada.", - "errors": [{"field": "status", "message": "Factura no procesada."}], - } + _progress(self, 10, "Verificando estatus de seguridad...") + from api.v1.modules.a76.invoices.models import InvoiceStatus + if invoice.status != InvoiceStatus.PROCESSED: + return { + "status": "error", + "message": f"La factura {invoice.invoice_number} no se puede revertir porque no está procesada.", + "errors": [{"field": "status", "message": "Factura no procesada."}], + } - errors = ErrorCollector() + 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: - errors.add_error( - field="line_items", - message="La factura no contiene partidas para revertir", - solution=["Verifique que la factura tenga partidas antes de intentar revertirla"], - code="NO_LINE_ITEMS", + _progress(self, 10, "Validando estatus de la factura...") + lines = pre_validators(db, invoice, tenant_id, company_id, errors) + if not lines: + errors.add_error( + field="line_items", + message="La factura no contiene partidas para revertir", + solution=["Verifique que la factura tenga partidas antes de intentar revertirla"], + code="NO_LINE_ITEMS", + ) + errors.raise_if_errors() + + _progress(self, 40, "Verificando saldos de partidas...") + sql_errors = revert_process( + db=db, + invoice=invoice, + lines=lines, + tenant_id=tenant_id, + company_id=company_id, + errors=errors, + cancelled_by=cancelled_by, ) - 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, - invoice=invoice, - lines=lines, - tenant_id=tenant_id, - company_id=company_id, - errors=errors, - cancelled_by=cancelled_by, - ) + _progress(self, 95, "Anulando saldos de inventario y confirmando...") + db.flush() + db.commit() - # ── Paso 4: Confirmar transacción ───────────────────────────────────── - _progress(self, 95, "Anulando saldos de inventario y confirmando...") - db.flush() - db.commit() + return { + "status": "success", + "invoice_id": invoice_id, + "sql_errors": sql_errors, + } - return { - "status": "success", - "invoice_id": invoice_id, - "sql_errors": sql_errors, - } - - except ValidationException as exc: - db.rollback() - return { - "status": "validation_error", - "message": exc.message, - "errors": exc.errors, - } - except Exception as exc: - db.rollback() - raise exc - finally: - db.close() + except ValidationException as exc: + db.rollback() + return { + "status": "validation_error", + "message": exc.message, + "errors": exc.errors, + } + except Exception as exc: + db.rollback() + raise exc diff --git a/backend/api/v1/modules/a76/invoices/imports/process/task.py b/backend/api/v1/modules/a76/invoices/imports/process/task.py index 2c96b180..2e7f56aa 100644 --- a/backend/api/v1/modules/a76/invoices/imports/process/task.py +++ b/backend/api/v1/modules/a76/invoices/imports/process/task.py @@ -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,56 +22,49 @@ 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() - try: - # ── Paso 1: Cargar factura ──────────────────────────────────────────── - # ── Paso 1: Cargar factura ──────────────────────────────────────────── - _progress(self, 5, "Cargando factura...") - invoice: InvoiceHeader | None = db.get(InvoiceHeader, invoice_id) - - if invoice is None: + with scoped_core_db(tenant_id=int(tenant_id), company_id=int(company_id)) as db: + try: + _progress(self, 5, "Cargando factura...") + invoice: InvoiceHeader | None = db.get(InvoiceHeader, invoice_id) + + if invoice is None: + return { + "status": "error", + "message": f"Factura con id {invoice_id} no encontrada.", + "errors": [], + } + + _progress(self, 10, "Verificando estatus de seguridad...") + if invoice.status == InvoiceStatus.PROCESSED: + return { + "status": "error", + "message": f"La factura {invoice.invoice_number} ya se encuentra procesada.", + "errors": [{"field": "status", "message": "Factura ya procesada."}], + } + + _progress(self, 20, "Iniciando procesamiento de factura...") + result = main_process( + db=db, + invoice=invoice, + tenant_id=tenant_id, + company_id=company_id, + username=username, + ) + + _progress(self, 95, "Confirmando cambios...") + db.commit() + + _progress(self, 100, "Proceso completado.") + return result + + except ValidationException as exc: + db.rollback() return { - "status": "error", - "message": f"Factura con id {invoice_id} no encontrada.", - "errors": [], + "status": "validation_error", + "message": exc.message, + "errors": exc.errors, } - - _progress(self, 10, "Verificando estatus de seguridad...") - if invoice.status == InvoiceStatus.PROCESSED: - return { - "status": "error", - "message": f"La factura {invoice.invoice_number} ya se encuentra procesada.", - "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 - ) - - # ── Paso 3: Confirmar transacción ───────────────────────────────────── - _progress(self, 95, "Confirmando cambios...") - db.commit() - - _progress(self, 100, "Proceso completado.") - return result - - except ValidationException as exc: - db.rollback() - return { - "status": "validation_error", - "message": exc.message, - "errors": exc.errors, - } - except Exception as exc: - db.rollback() - logger.error(f"Error en process_invoice_task: {str(exc)}", exc_info=True) - raise exc - finally: - db.close() + except Exception as exc: + db.rollback() + logger.error(f"Error en process_invoice_task: {str(exc)}", exc_info=True) + raise exc diff --git a/backend/api/v1/modules/a76/invoices/imports/revert/task.py b/backend/api/v1/modules/a76/invoices/imports/revert/task.py index b586155e..c6bc4d69 100644 --- a/backend/api/v1/modules/a76/invoices/imports/revert/task.py +++ b/backend/api/v1/modules/a76/invoices/imports/revert/task.py @@ -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,72 +26,66 @@ def revert_invoice_task( validaciones y reversiones del proceso principal (revert/main_process) con reporte de progreso. """ - db = CoreSessionLocal() - try: - # ── Paso 1: Cargar factura ──────────────────────────────────────────── - _progress(self, 5, "Cargando factura...") - invoice: InvoiceHeader | None = db.get(InvoiceHeader, invoice_id) - if invoice is None: - return { - "status": "error", - "message": f"Factura con id {invoice_id} no encontrada.", - "errors": [], - } + with scoped_core_db(tenant_id=int(tenant_id), company_id=int(company_id)) as db: + try: + _progress(self, 5, "Cargando factura...") + invoice: InvoiceHeader | None = db.get(InvoiceHeader, invoice_id) + if invoice is None: + return { + "status": "error", + "message": f"Factura con id {invoice_id} no encontrada.", + "errors": [], + } - _progress(self, 10, "Verificando estatus de seguridad...") - from api.v1.modules.a76.invoices.models import InvoiceStatus - if invoice.status != InvoiceStatus.PROCESSED: - return { - "status": "error", - "message": f"La factura {invoice.invoice_number} no se puede revertir porque no está procesada.", - "errors": [{"field": "status", "message": "Factura no procesada."}], - } + _progress(self, 10, "Verificando estatus de seguridad...") + from api.v1.modules.a76.invoices.models import InvoiceStatus + if invoice.status != InvoiceStatus.PROCESSED: + return { + "status": "error", + "message": f"La factura {invoice.invoice_number} no se puede revertir porque no está procesada.", + "errors": [{"field": "status", "message": "Factura no procesada."}], + } - errors = ErrorCollector() + 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: - errors.add_error( - field="line_items", - message="La factura no contiene partidas para revertir", - solution=["Verifique que la factura tenga partidas antes de intentar revertirla"], - code="NO_LINE_ITEMS", + _progress(self, 10, "Validando estatus de la factura...") + lines = pre_validators(db, invoice, tenant_id, company_id, errors) + if not lines: + errors.add_error( + field="line_items", + message="La factura no contiene partidas para revertir", + solution=["Verifique que la factura tenga partidas antes de intentar revertirla"], + code="NO_LINE_ITEMS", + ) + errors.raise_if_errors() + + _progress(self, 40, "Verificando saldos de partidas...") + sql_errors = revert_process( + db=db, + invoice=invoice, + lines=lines, + tenant_id=tenant_id, + company_id=company_id, + errors=errors, ) - 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, - invoice=invoice, - lines=lines, - tenant_id=tenant_id, - company_id=company_id, - errors=errors, - ) + _progress(self, 95, "Anulando saldos de inventario y confirmando...") + db.flush() + db.commit() - # ── Paso 4: Confirmar transacción ───────────────────────────────────── - _progress(self, 95, "Anulando saldos de inventario y confirmando...") - db.flush() - db.commit() + return { + "status": "success", + "invoice_id": invoice_id, + "sql_errors": sql_errors, + } - return { - "status": "success", - "invoice_id": invoice_id, - "sql_errors": sql_errors, - } - - except ValidationException as exc: - db.rollback() - return { - "status": "validation_error", - "message": exc.message, - "errors": exc.errors, - } - except Exception as exc: - db.rollback() - raise exc - finally: - db.close() + except ValidationException as exc: + db.rollback() + return { + "status": "validation_error", + "message": exc.message, + "errors": exc.errors, + } + except Exception as exc: + db.rollback() + raise exc diff --git a/backend/api/v1/modules/core/tasks_tracking/dispatch.py b/backend/api/v1/modules/core/tasks_tracking/dispatch.py index f37632b6..a0f60d87 100644 --- a/backend/api/v1/modules/core/tasks_tracking/dispatch.py +++ b/backend/api/v1/modules/core/tasks_tracking/dispatch.py @@ -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, diff --git a/backend/core/__init__.py b/backend/core/__init__.py index 233883c3..9739594b 100644 --- a/backend/core/__init__.py +++ b/backend/core/__init__.py @@ -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", diff --git a/backend/core/celery_app.py b/backend/core/celery_app.py index 85552bd6..f9379f5a 100644 --- a/backend/core/celery_app.py +++ b/backend/core/celery_app.py @@ -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 diff --git a/backend/core/database.py b/backend/core/database.py index 2bd4ba4c..4f02f241 100644 --- a/backend/core/database.py +++ b/backend/core/database.py @@ -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) diff --git a/backend/core/middleware.py b/backend/core/middleware.py index 9893b57f..e4bf47a9 100644 --- a/backend/core/middleware.py +++ b/backend/core/middleware.py @@ -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 diff --git a/backend/tests/integration/test_rls_tenant_company.py b/backend/tests/integration/test_rls_tenant_company.py new file mode 100644 index 00000000..cdd9c641 --- /dev/null +++ b/backend/tests/integration/test_rls_tenant_company.py @@ -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() diff --git a/docs/ARCHITECTURE.md b/docs/ARCHITECTURE.md index 6e4332ff..234ce227 100644 --- a/docs/ARCHITECTURE.md +++ b/docs/ARCHITECTURE.md @@ -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