"""Plugin DB migration runner with tenant_id column validation.""" from __future__ import annotations import os from pathlib import Path from typing import Any from sqlalchemy import text, inspect from sqlalchemy.ext.asyncio import AsyncEngine, AsyncSession from app.models.plugin import PluginMigration class MigrationValidationError(Exception): """Raised when a plugin migration creates a table without tenant_id column.""" class MigrationRunner: """Runs plugin SQL migrations and validates that all created tables have tenant_id.""" def __init__(self, engine: AsyncEngine, migrations_dir: str | None = None) -> None: self._engine = engine self._migrations_dir = Path(migrations_dir or os.path.join(os.path.dirname(__file__), "migrations")) def _resolve_migration_path(self, filename: str, plugin_name: str | None = None) -> Path: """Resolve a migration filename to an absolute path. Search order: 1. Plugin-specific migrations dir: app/plugins/builtins//migrations/ 2. Global migrations_dir (default: app/plugins/migrations/) 3. Shared builtins migrations dir: app/plugins/builtins/migrations/ """ # 1. Plugin-specific migrations directory if plugin_name: plugin_migrations_dir = Path(__file__).parent / "builtins" / plugin_name / "migrations" candidate = plugin_migrations_dir / filename if candidate.exists(): return candidate # 2. Global migrations_dir candidate = self._migrations_dir / filename if candidate.exists(): return candidate # 3. Shared builtins migrations directory builtins_dir = Path(__file__).parent / "builtins" / "migrations" candidate2 = builtins_dir / filename if candidate2.exists(): return candidate2 raise FileNotFoundError(f"Migration file not found: {filename}") async def run_migration( self, db: AsyncSession, plugin_name: str, migration_filename: str, tenant_id: Any | None = None, ) -> PluginMigration: """Run a single SQL migration file for a plugin. Reads the SQL file, executes it, validates created tables have tenant_id, and records the migration in plugin_migrations table. Args: db: Async database session. plugin_name: Name of the plugin being migrated. migration_filename: Filename of the migration SQL file. tenant_id: Optional tenant UUID for tenant-scoped migrations. Returns: PluginMigration record. Raises: FileNotFoundError: If migration file doesn't exist. MigrationValidationError: If a created table lacks tenant_id column. """ migration_path = self._resolve_migration_path(migration_filename, plugin_name) sql_content = migration_path.read_text(encoding="utf-8") # Capture existing table names before migration (static analysis of SQL) existing_tables = await self._get_table_names_via_session(db) # Execute the migration SQL # Split on semicolons but handle potential function definitions statements = self._split_sql(sql_content) for stmt in statements: stmt = stmt.strip() if stmt: await db.execute(text(stmt)) await db.flush() # Capture table names after migration (using same session connection) new_tables = await self._get_table_names_via_session(db) created_tables = new_tables - existing_tables # Validate: all newly created tables must have tenant_id column tables_without_tenant = [] for table_name in created_tables: if not await self._table_has_column_via_session(db, table_name, "tenant_id"): tables_without_tenant.append(table_name) if tables_without_tenant: # Rollback the migration await db.rollback() raise MigrationValidationError( f"Plugin '{plugin_name}' migration '{migration_filename}' created tables without tenant_id column: " f"{', '.join(tables_without_tenant)}. All plugin tables must have tenant_id for multi-tenant isolation." ) # Record the migration in plugin_migrations table migration_record = PluginMigration( plugin_name=plugin_name, migration_file=migration_filename, status="applied", ) db.add(migration_record) await db.flush() return migration_record async def run_all_migrations( self, db: AsyncSession, plugin_name: str, migration_files: list[str], tenant_id: Any | None = None, ) -> list[PluginMigration]: """Run all migration files for a plugin in order. Skips already-applied migrations (idempotent). """ # Check which migrations are already applied from sqlalchemy import select result = await db.execute( select(PluginMigration).where( PluginMigration.plugin_name == plugin_name, PluginMigration.status == "applied", ) ) applied_files = {row.migration_file for row in result.scalars().all()} records: list[PluginMigration] = [] for filename in migration_files: if filename in applied_files: continue # Already applied — idempotent skip record = await self.run_migration(db, plugin_name, filename, tenant_id) records.append(record) return records async def drop_plugin_tables(self, db: AsyncSession, plugin_name: str) -> list[str]: """Drop all tables that were created by a plugin's migrations. Args: db: Async database session. plugin_name: Name of the plugin to remove tables for. Returns: List of dropped table names. """ from sqlalchemy import select # Get all migration records for this plugin result = await db.execute( select(PluginMigration).where(PluginMigration.plugin_name == plugin_name) ) migration_records = result.scalars().all() dropped_tables: list[str] = [] for record in migration_records: migration_path = self._resolve_migration_path(record.migration_file, plugin_name) sql_content = migration_path.read_text(encoding="utf-8") # Parse table names from CREATE TABLE statements table_names = self._extract_table_names(sql_content) for table_name in table_names: try: await db.execute(text(f'DROP TABLE IF EXISTS "{table_name}" CASCADE')) dropped_tables.append(table_name) except Exception: pass # Table may already be gone await db.flush() # Remove migration records for record in migration_records: await db.delete(record) await db.flush() return dropped_tables async def _get_table_names_via_session(self, db: AsyncSession) -> set[str]: """Get current table names using the session's own connection. This ensures we see uncommitted DDL changes within the same transaction. """ result = await db.execute( text("SELECT tablename FROM pg_tables WHERE schemaname = 'public'") ) return {row[0] for row in result.fetchall()} async def _table_has_column_via_session( self, db: AsyncSession, table_name: str, column_name: str ) -> bool: """Check if a table has a specific column using the session's connection.""" result = await db.execute( text( "SELECT column_name FROM information_schema.columns " "WHERE table_schema = 'public' AND table_name = :table_name AND column_name = :column_name" ), {"table_name": table_name, "column_name": column_name}, ) return result.fetchone() is not None async def _get_table_names(self) -> set[str]: """Get current table names from the database (separate connection).""" def _get_names(sync_conn): insp = inspect(sync_conn) return set(insp.get_table_names()) async with self._engine.connect() as conn: return await conn.run_sync(_get_names) async def _table_has_column(self, table_name: str, column_name: str) -> bool: """Check if a table has a specific column (separate connection).""" def _has_col(sync_conn): insp = inspect(sync_conn) if table_name not in insp.get_table_names(): return False columns = [col["name"] for col in insp.get_columns(table_name)] return column_name in columns async with self._engine.connect() as conn: return await conn.run_sync(_has_col) @staticmethod def _split_sql(sql: str) -> list[str]: """Split SQL into individual statements, respecting basic semicolon boundaries.""" statements: list[str] = [] current: list[str] = [] in_dollar_quote = False dollar_tag = "" for line in sql.splitlines(): stripped = line.strip() # Handle dollar-quoted blocks (PostgreSQL function bodies) if "$$" in stripped and not in_dollar_quote: parts = stripped.split("$$", 1) if len(parts) > 1: dollar_tag = parts[0].strip() if dollar_tag: in_dollar_quote = True elif stripped.count("$$") >= 2: # Single-line $$ block current.append(line) if stripped.rstrip().endswith(";"): statements.append("\n".join(current)) current = [] continue else: in_dollar_quote = True current.append(line) continue if in_dollar_quote: current.append(line) if "$$" in stripped: in_dollar_quote = False # After closing $$, check if the line ends with ; to flush statement after_dollar = stripped.split("$$", 1)[-1] if "$$" in stripped else "" if after_dollar.rstrip().endswith(";"): statements.append("\n".join(current)) current = [] continue current.append(line) if stripped.endswith(";"): statements.append("\n".join(current)) current = [] if current: remaining = "\n".join(current).strip() if remaining: statements.append(remaining) return statements @staticmethod def _extract_table_names(sql: str) -> list[str]: """Extract table names from CREATE TABLE statements in SQL.""" import re # Match CREATE TABLE [IF NOT EXISTS] "table_name" or CREATE TABLE table_name pattern = r'CREATE\s+TABLE\s+(?:IF\s+NOT\s+EXISTS\s+)?["\']?(\w+)["\']?' return re.findall(pattern, sql, re.IGNORECASE)