""" 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 asynccontextmanager, contextmanager from contextvars import ContextVar from typing import AsyncGenerator, Dict, Generator, Optional 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 from .config import settings logger = logging.getLogger(__name__) Base = declarative_base() core_engine = create_engine( settings.core_database_url, pool_pre_ping=True, pool_size=10, max_overflow=20, echo=False, ) CoreSessionLocal = sessionmaker( autocommit=False, autoflush=False, bind=core_engine) async_core_engine = create_async_engine( settings.async_core_database_url, pool_pre_ping=True, pool_size=10, max_overflow=20, echo=settings.DEBUG, ) AsyncCoreSessionLocal = async_sessionmaker( async_core_engine, class_=AsyncSession, expire_on_commit=False ) _tenant_engines: Dict[str, any] = {} 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. """ 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 reset_rls_context_tokens(token_t, token_c) -> None: """Restaura ContextVars de RLS de forma segura entre hilos/tareas asyncio. ``ContextVar.reset`` exige que el token se cree y restaure en el mismo contexto lógico; en rutas FastAPI async + dependencias síncronas con ``yield`` (thread pool) el ``finally`` puede ejecutarse en otro contexto y lanzar ``ValueError`` (mensaje: "was created in a different Context"). En ese caso degradamos a ``set(None)``, igual que ``task_postrun`` en ``core/celery_app.py``. """ try: rls_tenant_var.reset(token_t) rls_company_var.reset(token_c) except (ValueError, RuntimeError): rls_tenant_var.set(None) rls_company_var.set(None) 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. No se escriben las ContextVars de RLS aquí: las dependencias síncronas con ``yield`` se ejecutan vía ``contextmanager_in_threadpool`` (hilo worker) y mezclar ``ContextVar.set`` / ``reset`` entre ese hilo y el bucle asyncio provoca ``ValueError: ... was created in a different Context``. El aislamiento RLS se aplica con ``session.info`` (véase ``after_begin`` y audit listeners). """ 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(request: Request = None) -> AsyncGenerator[AsyncSession, None]: """Dependency async para obtener sesión con contexto RLS.""" tenant_id, company_id = _extract_rls_context(request) prev_tenant = rls_tenant_var.get() prev_company = rls_company_var.get() rls_tenant_var.set(tenant_id) rls_company_var.set(company_id) try: 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() finally: rls_tenant_var.set(prev_tenant) rls_company_var.set(prev_company) @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: await session.close() def get_tenant_engine(tenant_id: int, db_config: dict): """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( db_url, pool_pre_ping=True, pool_size=5, max_overflow=10 ) return _tenant_engines[tenant_id] @contextmanager 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. 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: db = CoreSessionLocal() db.info[RLS_TENANT_KEY] = tenant_id else: engine = get_tenant_engine(tenant_id, db_config) SessionLocal = sessionmaker( autocommit=False, autoflush=False, bind=engine) db = SessionLocal() try: yield db finally: db.close() def init_db(): """Inicializa las tablas de la base de datos core.""" try: Base.metadata.create_all(bind=core_engine, checkfirst=True) except ProgrammingError as e: if "already exists" in str(e): logger.warning( f"Algunas tablas ya existen en la base de datos: {e}") else: raise async def init_async_db(): """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)