diff --git a/prestart.sh b/prestart.sh index ce8cd05..1c85fdd 100644 --- a/prestart.sh +++ b/prestart.sh @@ -101,54 +101,8 @@ asyncio.run(set_password()) " fi -echo "[prestart] Adding owner_id to plugin tables if missing..." -python3 -c " -import asyncio, os -from sqlalchemy.ext.asyncio import create_async_engine -from sqlalchemy import text - -async def fix_owner_id(): - db_url = os.environ.get('MIGRATION_DATABASE_URL', os.environ.get('DATABASE_URL', '')) - if not db_url: - print('[prestart] WARNING: No DB URL for owner_id fix') - return - engine = create_async_engine(db_url) - try: - async with engine.begin() as conn: - # Find all tables with tenant_id but without owner_id - result = await conn.execute(text(\"\"\" - SELECT t.table_name - FROM information_schema.tables t - JOIN information_schema.columns c ON t.table_name = c.table_name - WHERE t.table_schema = 'public' - AND c.column_name = 'tenant_id' - AND t.table_type = 'BASE TABLE' - AND NOT EXISTS ( - SELECT 1 FROM information_schema.columns - WHERE table_name = t.table_name AND column_name = 'owner_id' - ) - ORDER BY t.table_name - \"\"\")) - tables = [row[0] for row in result.fetchall()] - for table in tables: - await conn.execute(text( - f'ALTER TABLE {table} ADD COLUMN IF NOT EXISTS owner_id UUID REFERENCES users(id) ON DELETE SET NULL' - )) - await conn.execute(text( - f'CREATE INDEX IF NOT EXISTS ix_{table}_owner ON {table} (owner_id)' - )) - print(f'[prestart] Added owner_id to {table}') - if not tables: - print('[prestart] All tables already have owner_id.') - else: - print(f'[prestart] Added owner_id to {len(tables)} tables.') - except Exception as e: - print(f'[prestart] WARNING: owner_id fix failed: {e}') - finally: - await engine.dispose() - -asyncio.run(fix_owner_id()) -" +echo "[prestart] Syncing plugin table schemas..." +python3 /app/scripts/sync_plugin_schema.py echo "[prestart] Seeding admin user if not exists..." python3 /app/scripts/seed_admin.py || echo "[prestart] WARNING: Admin seed failed (may already exist)" diff --git a/scripts/sync_plugin_schema.py b/scripts/sync_plugin_schema.py new file mode 100644 index 0000000..bf7d2a4 --- /dev/null +++ b/scripts/sync_plugin_schema.py @@ -0,0 +1,288 @@ +#!/usr/bin/env python3 +"""Sync plugin table schemas with ORM model definitions. + +This script compares SQLAlchemy ORM model columns (from plugin models) with +the actual database columns and adds any missing columns or indexes via +ALTER TABLE. It is idempotent and safe to run on every container startup. + +Usage: + python3 scripts/sync_plugin_schema.py + +Environment: + MIGRATION_DATABASE_URL or DATABASE_URL — async SQLAlchemy URL (e.g. postgresql+asyncpg://...) +""" + +from __future__ import annotations + +import importlib +import os +import pkgutil +import sys +import logging +from typing import Any + +from sqlalchemy import inspect, text +from sqlalchemy.dialects.postgresql import dialect as pg_dialect +from sqlalchemy.ext.asyncio import create_async_engine +from sqlalchemy.orm import DeclarativeBase + +logger = logging.getLogger("sync_plugin_schema") +logging.basicConfig(level=logging.INFO, format="[sync] %(levelname)s: %(message)s") + +# ── Plugin model discovery ────────────────────────────────────────────────── + +def _import_all_plugin_models() -> list[type[DeclarativeBase]]: + """Import all models from app.plugins.builtins.*.models modules. + + Returns a list of SQLAlchemy ORM model classes (subclasses of Base). + """ + models: list[type[DeclarativeBase]] = [] + try: + import app.plugins.builtins as builtins_pkg + except ImportError as exc: + logger.error("Cannot import app.plugins.builtins: %s", exc) + return models + + # Import the builtins package to trigger model registrations + for importer, modname, ispkg in pkgutil.iter_modules(builtins_pkg.__path__): + if not ispkg: + continue + module_path = f"app.plugins.builtins.{modname}.models" + try: + importlib.import_module(module_path) + logger.debug("Imported %s", module_path) + except ImportError: + # Some plugins may not have a models.py — skip silently + logger.debug("No models module for plugin '%s'", modname) + except Exception as exc: + logger.warning("Error importing %s: %s", module_path, exc) + + # Now collect all model classes from Base.metadata + from app.core.db import Base + # Filter to only plugin tables — we identify them by checking if the table + # name appears in any plugin models module. We use metadata.tables which + # contains all registered tables. + for table_name, table in Base.metadata.tables.items(): + # We only process tables that are defined in plugin modules. + # Core models are handled by alembic migrations. + # We check the module of the model class. + for mapper_class in Base.registry.mappers: + cls = mapper_class.class_ + if ( + hasattr(cls, "__tablename__") + and cls.__tablename__ == table_name + and "app.plugins.builtins" in cls.__module__ + ): + models.append(cls) + break + + return models + + +# ── Column type compilation ────────────────────────────────────────────────── + +def _compile_column_type(col_type: Any) -> str: + """Compile a SQLAlchemy column type to its PostgreSQL DDL string.""" + try: + return col_type.compile(dialect=pg_dialect()) + except Exception: + # Fallback: try generic compile + try: + return str(col_type) + except Exception: + return "TEXT" + + +# ── ALTER TABLE column DDL builder ────────────────────────────────────────── + +def _build_column_ddl(column: Any) -> str: + """Build the ADD COLUMN DDL fragment for a single SQLAlchemy Column. + + Handles: type, NULL/NOT NULL, DEFAULT, server_default, ForeignKey REFERENCES. + """ + parts: list[str] = [] + + # Column name (quote to be safe) + col_name = f'"{column.name}"' + parts.append(col_name) + + # Column type + type_str = _compile_column_type(column.type) + parts.append(type_str) + + # Foreign key references + for fk in column.foreign_keys: + ref_table = fk.column.table.name + ref_col = fk.column.name + ref_clause = f'REFERENCES "{ref_table}"("{ref_col}")' + # ondelete behavior + if fk.ondelete: + ref_clause += f' ON DELETE {fk.ondelete.upper()}' + parts.append(ref_clause) + + # Nullable + if column.nullable: + parts.append("NULL") + else: + parts.append("NOT NULL") + + # server_default (DB-side default) + if column.server_default is not None: + server_default_text = column.server_default.arg + # server_default can be a text() clause or a func + if hasattr(server_default_text, "compile"): + try: + compiled = server_default_text.compile(dialect=pg_dialect()) + parts.append(f"DEFAULT {compiled}") + except Exception: + parts.append(f"DEFAULT {server_default_text}") + else: + parts.append(f"DEFAULT {server_default_text}") + elif column.default is not None and column.default.is_scalar: + # Python-side default (only add if it's a simple scalar) + val = column.default.arg + if isinstance(val, str): + parts.append(f"DEFAULT '{val}'") + elif isinstance(val, bool): + parts.append(f"DEFAULT {str(val).lower()}") + elif val is None: + pass # Don't add DEFAULT NULL explicitly + else: + parts.append(f"DEFAULT {val}") + + return " ".join(parts) + + +# ── Index DDL builder ──────────────────────────────────────────────────────── + +def _build_index_ddl(index: Any) -> str: + """Build CREATE INDEX IF NOT EXISTS DDL for a SQLAlchemy Index.""" + col_names = [] + for col in index.columns: + col_names.append(f'"{col.name}"') + cols_str = ", ".join(col_names) + return f'CREATE INDEX IF NOT EXISTS "{index.name}" ON "{index.table.name}" ({cols_str})' + + +# ── Main sync logic ────────────────────────────────────────────────────────── + +async def sync_plugin_schemas() -> None: + """Compare ORM models with DB and add missing columns/indexes.""" + db_url = os.environ.get("MIGRATION_DATABASE_URL") or os.environ.get("DATABASE_URL", "") + if not db_url: + logger.error("No MIGRATION_DATABASE_URL or DATABASE_URL set — skipping sync") + return + + # Ensure asyncpg dialect + if db_url.startswith("postgresql://"): + db_url = db_url.replace("postgresql://", "postgresql+asyncpg://", 1) + + logger.info("Importing plugin models...") + model_classes = _import_all_plugin_models() + logger.info("Found %d plugin model classes", len(model_classes)) + + if not model_classes: + logger.warning("No plugin models found — nothing to sync") + return + + engine = create_async_engine(db_url, echo=False) + added_columns = 0 + added_indexes = 0 + + try: + async with engine.begin() as conn: + # Get existing tables in the public schema + result = await conn.execute( + text( + "SELECT table_name FROM information_schema.tables " + "WHERE table_schema = 'public' AND table_type = 'BASE TABLE'" + ) + ) + existing_tables = {row[0] for row in result.fetchall()} + + for model_cls in model_classes: + table_name = model_cls.__tablename__ + if table_name not in existing_tables: + logger.info("Table '%s' does not exist yet — skipping", table_name) + continue + + # Get actual DB columns + result = await conn.execute( + text( + "SELECT column_name FROM information_schema.columns " + "WHERE table_schema = 'public' AND table_name = :tbl" + ), + {"tbl": table_name}, + ) + db_columns = {row[0] for row in result.fetchall()} + + # Get ORM columns via inspection + mapper = inspect(model_cls) + orm_column_names = {col.name for col in mapper.columns} + + # Find missing columns + missing_cols = orm_column_names - db_columns + if missing_cols: + for col in mapper.columns: + if col.name not in missing_cols: + continue + ddl_fragment = _build_column_ddl(col) + alter_sql = f'ALTER TABLE "{table_name}" ADD COLUMN IF NOT EXISTS {ddl_fragment}' + try: + await conn.execute(text(alter_sql)) + logger.info("Added column '%s.%s' (%s)", table_name, col.name, _compile_column_type(col.type)) + added_columns += 1 + except Exception as exc: + logger.error("Failed to add column '%s.%s': %s", table_name, col.name, exc) + else: + logger.debug("Table '%s' — all columns present", table_name) + + # Get existing indexes + result = await conn.execute( + text( + "SELECT indexname FROM pg_indexes " + "WHERE schemaname = 'public' AND tablename = :tbl" + ), + {"tbl": table_name}, + ) + db_indexes = {row[0] for row in result.fetchall()} + + # Check ORM-defined indexes from __table_args__ + table = model_cls.__table__ + for idx in table.indexes: + if idx.name and idx.name not in db_indexes: + index_ddl = _build_index_ddl(idx) + try: + await conn.execute(text(index_ddl)) + logger.info("Added index '%s' on table '%s'", idx.name, table_name) + added_indexes += 1 + except Exception as exc: + logger.error("Failed to add index '%s' on '%s': %s", idx.name, table_name, exc) + + # Also check for column-level indexes (index=True on mapped_column) + for col in mapper.columns: + for idx in col.indexes: + if idx.name and idx.name not in db_indexes: + index_ddl = _build_index_ddl(idx) + try: + await conn.execute(text(index_ddl)) + logger.info("Added index '%s' on table '%s'", idx.name, table_name) + added_indexes += 1 + except Exception as exc: + logger.error("Failed to add index '%s' on '%s': %s", idx.name, table_name, exc) + + logger.info( + "Schema sync complete: %d columns added, %d indexes added", + added_columns, + added_indexes, + ) + except Exception as exc: + logger.error("Schema sync failed: %s", exc) + raise + finally: + await engine.dispose() + + +if __name__ == "__main__": + import asyncio + asyncio.run(sync_plugin_schemas())