From 957eb3378dae3fdad4cfc8c6cdc8f97d9262a247 Mon Sep 17 00:00:00 2001 From: Adrian Ehrsam Date: Mon, 5 Oct 2026 20:06:34 +0000 Subject: [PATCH 1/3] Support --dialect mssql in pgdb migrate check/apply, bump to 0.12.0 Adds an MSSQL implementation of the migrate primitives (GO-batch splitting, single-transaction apply, OBJECT_ID-based verification and already-applied detection, [schema].[table] tracking). Postgres remains the default. Co-Authored-By: Claude Sonnet 5.5 --- README.md | 32 ++++- pgdevkit/cli.py | 57 ++++++-- pgdevkit/migrate.py | 104 +++++++++++---- pgdevkit/migrate_mssql.py | 221 ++++++++++++++++++++++++++++++ pyproject.toml | 2 +- tests/test_migrate_mssql.py | 222 +++++++++++++++++++++++++++++++ tests/test_migrate_mssql_live.py | 49 +++++++ uv.lock | 2 +- 8 files changed, 648 insertions(+), 41 deletions(-) create mode 100644 pgdevkit/migrate_mssql.py create mode 100644 tests/test_migrate_mssql.py create mode 100644 tests/test_migrate_mssql_live.py diff --git a/README.md b/README.md index de56ad0..afdf560 100644 --- a/README.md +++ b/README.md @@ -303,16 +303,38 @@ REPL. ## `pgdb migrate` Applies numbered, forward-only SQL migration files from a directory to a live -Postgres database, tracking each one in a `schema.table` (default -`public.schema_migrations`) so repeat runs only apply what's pending. Postgres only — -not available for `--dialect mssql`. +Postgres or MSSQL database, tracking each one in a `schema.table` (default +`public.schema_migrations`, `dbo.schema_migrations` for MSSQL) so repeat runs only +apply what's pending. ```bash pgdb migrate check path/to/database/_migration_scripts --url postgresql://user:pass@host:port/db pgdb migrate apply path/to/database/_migration_scripts --url postgresql://user:pass@host:port/db ``` -`--entra-user` works the same as `pgdb compare` (see above). The tracking +### MSSQL + +Pass `--dialect mssql` (needs the `mssql` extra) to both `check` and `apply`, with +`--url` as an ODBC connection string: + +```bash +pgdb migrate apply path/to/database/_migration_scripts --dialect mssql \ + --url "Server=host,1433;Database=db;UID=user;PWD=pass" +``` + +- Migration files are T-SQL: they are split into batches on standalone `GO` lines + (needed for e.g. `CREATE VIEW`, which must be first in its batch), and all batches + of one file run in a single transaction that is rolled back if any batch fails. +- The tracking table needs `filename nvarchar(450) primary key, applied_at + datetimeoffset not null default sysdatetimeoffset(), applied_by nvarchar(128) not + null default suser_sname()`. As on Postgres, a migration creating it can bootstrap it. +- `--ask` auto-detects "already done" for single-statement `GO` batches that create a + table, view or schema, or add a column; anything else is asked about. +- Post-apply verification checks every `CREATE TABLE` target via `OBJECT_ID`. +- `--entra-user` is Postgres-only, and `pgdevkit.migrate.missing_privileges` is not + available for MSSQL. + +For Postgres, `--entra-user` works the same as `pgdb compare` (see above). The tracking table needs `filename text primary key, applied_at timestamptz not null default now(), applied_by text not null default current_user` (a migration file that creates it, in the same directory, is the usual way to bootstrap @@ -352,7 +374,7 @@ for a real one. `pgdevkit.migrate` is also usable directly as a library — `list_migration_files`, `applied_migrations`, `pending_migrations`, and `apply_migration` are the same functions the CLI calls, so a project can script around them without shelling -out. +out. The database-touching ones take `dialect="postgres" | "mssql"`. ## `pgdevkit.db` — helpers for application code diff --git a/pgdevkit/cli.py b/pgdevkit/cli.py index 65c205f..06eef55 100644 --- a/pgdevkit/cli.py +++ b/pgdevkit/cli.py @@ -47,6 +47,9 @@ "--env", help="Environment to apply: skips any ..sql file (e.g. grants.prod.sql); untagged files always apply", ) +_MIGRATE_DIALECT_OPTION = typer.Option( + "postgres", "--dialect", help="postgres (default) or mssql; mssql needs the mssql extra and uses --url as an ODBC connection string" +) _MIGRATE_ENV_OPTION = typer.Option( None, "--env", @@ -376,30 +379,53 @@ def testdb_list_orphaned() -> None: console.print(name) +def _migrate_target( + url: str, entra_user: str | None, dialect: str, tracking_table: str | None, migrations_dir: Path +) -> tuple[str, str, str]: + """Resolve (conninfo, dialect name, tracking table) for a migrate command.""" + try: + resolved = get_backend(dialect).dialect + except ValueError as e: + err_console.print(f"[red]Error:[/red] {e}") + raise typer.Exit(2) + if resolved.name == "mssql" and entra_user: + err_console.print("[red]Error:[/red] --entra-user is only supported for the postgres dialect") + raise typer.Exit(2) + conninfo = url if resolved.name == "mssql" else build_conninfo(url, entra_user) + tracking_table = tracking_table or migrate.default_tracking_table(migrations_dir, dialect=resolved) + return conninfo, resolved.name, tracking_table + + @migrate_app.command("check") def migrate_check( migrations_dir: Path = typer.Argument(..., help="Directory of numbered .sql migration files"), - url: str = typer.Option(..., "--url", help="PostgreSQL DSN (postgresql://user:pass@host:port/db)"), + url: str = typer.Option( + ..., + "--url", + help="PostgreSQL DSN (postgresql://user:pass@host:port/db), or with --dialect mssql " + "a connection string (Server=host,1433;Database=db;UID=user;PWD=pass)", + ), entra_user: str | None = typer.Option(None, "--entra-user", help="Azure Entra user (triggers token auth)"), tracking_table: str | None = typer.Option( None, "--tracking-table", help="schema.table recording applied migrations " - "(default: tool.pgdevkit.migrations_table in pyproject.toml, else public.schema_migrations)", + "(default: tool.pgdevkit.migrations_table in pyproject.toml, else public.schema_migrations " + "-- dbo.schema_migrations with --dialect mssql)", ), area: list[str] = _AREA_OPTION, exclude_area: list[str] = _EXCLUDE_AREA_OPTION, schema: list[str] = _SCHEMA_OPTION, exclude_schema: list[str] = _EXCLUDE_SCHEMA_OPTION, env: str | None = _MIGRATE_ENV_OPTION, + dialect: str = _MIGRATE_DIALECT_OPTION, ) -> None: """List which migration files under migrations_dir are applied vs. pending.""" if not migrations_dir.is_dir(): err_console.print(f"[red]Error:[/red] {migrations_dir} is not a directory") raise typer.Exit(2) - conninfo = build_conninfo(url, entra_user) - tracking_table = tracking_table or migrate.default_tracking_table(migrations_dir) + conninfo, dialect, tracking_table = _migrate_target(url, entra_user, dialect, tracking_table, migrations_dir) local_files = migrate.list_migration_files( migrations_dir, areas=_as_set(area), @@ -409,7 +435,7 @@ def migrate_check( env=env, ) try: - applied = migrate.applied_migrations(conninfo, tracking_table) + applied = migrate.applied_migrations(conninfo, tracking_table, dialect) except migrate.TrackingTableMissing: err_console.print(f"[yellow]⚠[/yellow] {tracking_table} not found — nothing recorded as applied yet") applied = {} @@ -433,13 +459,19 @@ def migrate_check( @migrate_app.command("apply") def migrate_apply( migrations_dir: Path = typer.Argument(..., help="Directory of numbered .sql migration files"), - url: str = typer.Option(..., "--url", help="PostgreSQL DSN (postgresql://user:pass@host:port/db)"), + url: str = typer.Option( + ..., + "--url", + help="PostgreSQL DSN (postgresql://user:pass@host:port/db), or with --dialect mssql " + "a connection string (Server=host,1433;Database=db;UID=user;PWD=pass)", + ), entra_user: str | None = typer.Option(None, "--entra-user", help="Azure Entra user (triggers token auth)"), tracking_table: str | None = typer.Option( None, "--tracking-table", help="schema.table recording applied migrations " - "(default: tool.pgdevkit.migrations_table in pyproject.toml, else public.schema_migrations)", + "(default: tool.pgdevkit.migrations_table in pyproject.toml, else public.schema_migrations " + "-- dbo.schema_migrations with --dialect mssql)", ), file: str | None = typer.Option( None, "--file", help="Apply only this one filename (relative to migrations_dir) instead of all pending" @@ -451,14 +483,14 @@ def migrate_apply( schema: list[str] = _SCHEMA_OPTION, exclude_schema: list[str] = _EXCLUDE_SCHEMA_OPTION, env: str | None = _MIGRATE_ENV_OPTION, + dialect: str = _MIGRATE_DIALECT_OPTION, ) -> None: """Apply pending migration files, in filename order, tracking each in tracking_table.""" if not migrations_dir.is_dir(): err_console.print(f"[red]Error:[/red] {migrations_dir} is not a directory") raise typer.Exit(2) - conninfo = build_conninfo(url, entra_user) - tracking_table = tracking_table or migrate.default_tracking_table(migrations_dir) + conninfo, dialect, tracking_table = _migrate_target(url, entra_user, dialect, tracking_table, migrations_dir) areas, exclude_areas = _as_set(area), _as_set(exclude_area) schemas, exclude_schemas = _as_set(schema), _as_set(exclude_schema) target_desc = url.rsplit("@", 1)[-1] if "@" in url else url @@ -478,6 +510,7 @@ def migrate_apply( schemas=schemas, exclude_schemas=exclude_schemas, env=env, + dialect=dialect, ) except migrate.TrackingTableMissing: err_console.print( @@ -517,7 +550,9 @@ def worker() -> None: for path, already_done in iter(work_q.get, None): if not stop.is_set(): try: - result = migrate.apply_migration(conninfo, path, tracking_table, already_done=already_done) + result = migrate.apply_migration( + conninfo, path, tracking_table, already_done=already_done, dialect=dialect + ) except Exception as e: # noqa: BLE001 failure = e stop.set() @@ -544,7 +579,7 @@ def worker() -> None: break already_done = False if ask: - already_done = migrate.already_fully_applied(conninfo, path) + already_done = migrate.already_fully_applied(conninfo, path, dialect) if already_done: bar_write(f"Auto: {path.name} is already fully present in the database — marking as already done") else: diff --git a/pgdevkit/migrate.py b/pgdevkit/migrate.py index 31df1ef..15fcbc2 100644 --- a/pgdevkit/migrate.py +++ b/pgdevkit/migrate.py @@ -1,7 +1,9 @@ """Apply numbered, forward-only SQL migration files to a live Postgres database, tracked in a `schema.table` (default `public.schema_migrations`) so re-runs only -apply what's pending. Ported from a hand-rolled per-project script — MSSQL is not -supported yet. +apply what's pending. Ported from a hand-rolled per-project script. Every function that +touches a database takes `dialect="postgres"` (default) or `"mssql"`; the MSSQL +implementation lives in `.migrate_mssql` and is only imported when asked for, so the +`mssql` extra stays optional. """ from __future__ import annotations @@ -19,12 +21,12 @@ from psycopg import sql as pg_sql from .areas import filter_by_area +from .dialect import Dialect, resolve_dialect from .envtag import env_allowed from .schemas import filter_by_schema from .sql_text import strip_line_comments _IDENTIFIER = r"[A-Za-z_][A-Za-z0-9_]*" -_DEFAULT_TRACKING_TABLE = "public.schema_migrations" class MigrationVerificationError(RuntimeError): @@ -47,6 +49,10 @@ def _tracking_table_identifier(tracking_table: str) -> pg_sql.Identifier: return pg_sql.Identifier(schema, table) +def _is_mssql(dialect: str | Dialect) -> bool: + return resolve_dialect(dialect).name == "mssql" + + def _find_pyproject(start: Path) -> Path | None: for directory in [start, *start.parents]: candidate = directory / "pyproject.toml" @@ -55,17 +61,18 @@ def _find_pyproject(start: Path) -> Path | None: return None -def default_tracking_table(start: Path | None = None) -> str: +def default_tracking_table(start: Path | None = None, dialect: str | Dialect = "postgres") -> str: """The project's configured tracking table: `[tool.pgdevkit].migrations_table` in the nearest pyproject.toml at or above `start` (default: cwd), or "public.schema_migrations" - if neither is set. Lets a project fix its tracking table once instead of passing + if neither is set ("dbo.schema_migrations" for MSSQL). Lets a project fix its tracking table once instead of passing --tracking-table on every `pgdb migrate` invocation.""" + default = f"{resolve_dialect(dialect).default_schema}.schema_migrations" pyproject = _find_pyproject((start or Path.cwd()).resolve()) if pyproject is None: - return _DEFAULT_TRACKING_TABLE + return default data = tomllib.loads(pyproject.read_text(encoding="utf-8")) section = data.get("tool", {}).get("pgdevkit", {}) - return section.get("migrations_table", _DEFAULT_TRACKING_TABLE) + return section.get("migrations_table", default) def _split_sql(sql: str) -> list[str]: @@ -213,13 +220,22 @@ def _target_exists(con: psycopg.Connection, target: tuple[str, ...]) -> bool: return bool(row and row[0]) -def already_fully_applied(conninfo: str, path: Path) -> bool: +def already_fully_applied(conninfo: str, path: Path, dialect: str | Dialect = "postgres") -> bool: """Whether every statement in this migration is a recognized create-if-missing shape (table/index/sequence/view/schema, or add-column) AND its target already exists in the database — i.e. re-running the migration would do nothing. Used by `--ask` to auto-answer "already done" without prompting, so migrations that are trivially no-ops don't interrupt review. A single unrecognized or not-yet-applied statement means - False — this never guesses.""" + False — this never guesses. For MSSQL the unit is a `GO`-separated batch and the + recognized shapes are table, view, schema and ADD column.""" + if _is_mssql(dialect): + from . import migrate_mssql + + batches = migrate_mssql.split_batches(path.read_text(encoding="utf-8")) + mssql_targets = [migrate_mssql.idempotent_target(b) for b in batches] + if not batches or any(t is None for t in mssql_targets): + return False + return migrate_mssql.targets_exist(conninfo, cast(list[tuple[str, ...]], mssql_targets)) stmts = _split_sql(path.read_text(encoding="utf-8")) if not stmts: return False @@ -250,8 +266,14 @@ def list_migration_files( return [f for f in files if env_allowed(f, env)] -def applied_migrations(conninfo: str, tracking_table: str) -> dict[str, tuple[datetime, str]]: +def applied_migrations( + conninfo: str, tracking_table: str, dialect: str | Dialect = "postgres" +) -> dict[str, tuple[datetime, str]]: """Filename -> (applied_at, applied_by) for every migration recorded in the tracking table.""" + if _is_mssql(dialect): + from . import migrate_mssql + + return migrate_mssql.applied_migrations(conninfo, tracking_table) table = _tracking_table_identifier(tracking_table) with psycopg.connect(conninfo) as con: try: @@ -273,8 +295,9 @@ def pending_migrations( schemas: frozenset[str] | None = None, exclude_schemas: frozenset[str] | None = None, env: str | None = None, + dialect: str | Dialect = "postgres", ) -> list[Path]: - applied = applied_migrations(conninfo, tracking_table) + applied = applied_migrations(conninfo, tracking_table, dialect) files = list_migration_files( migrations_dir, areas=areas, @@ -286,9 +309,15 @@ def pending_migrations( return [p for p in files if p.name not in applied] -def record_applied(conninfo: str, tracking_table: str, filename: str) -> bool: +def record_applied( + conninfo: str, tracking_table: str, filename: str, dialect: str | Dialect = "postgres" +) -> bool: """Best-effort insert into the tracking table. Returns False without raising if the tracking table doesn't exist yet — e.g. this migration is the one that creates it.""" + if _is_mssql(dialect): + from . import migrate_mssql + + return migrate_mssql.record_applied(conninfo, tracking_table, filename) table = _tracking_table_identifier(tracking_table) with psycopg.connect(conninfo) as con: try: @@ -303,17 +332,25 @@ def record_applied(conninfo: str, tracking_table: str, filename: str) -> bool: return False -def created_table_names(sql: str) -> list[str]: +def created_table_names(sql: str, dialect: str | Dialect = "postgres") -> list[str]: """Table names any CREATE TABLE statement in this raw SQL script targets. Public wrapper around the same detection `apply_migration` uses internally, for callers that run a script directly (e.g. via `execute_sql_script`) instead of through a tracked migration file, and still want to know what tables -- if any -- it created.""" + if _is_mssql(dialect): + from . import migrate_mssql + + return migrate_mssql.created_table_names(migrate_mssql.split_batches(sql)) return _created_table_names(_split_sql(sql)) -def verify_created_tables(conninfo: str, stmts: list[str]) -> list[str]: +def verify_created_tables(conninfo: str, stmts: list[str], dialect: str | Dialect = "postgres") -> list[str]: """Table names from this migration's CREATE TABLE statements that do NOT exist in the - database. Empty means everything landed.""" + database. Empty means everything landed. For MSSQL, `stmts` are `GO`-separated batches.""" + if _is_mssql(dialect): + from . import migrate_mssql + + return migrate_mssql.missing_tables(conninfo, migrate_mssql.created_table_names(stmts)) tables = _created_table_names(stmts) if not tables: return [] @@ -326,12 +363,16 @@ def verify_created_tables(conninfo: str, stmts: list[str]) -> list[str]: return missing -def missing_privileges(conninfo: str, role: str, tables: list[str], privilege: str = "select") -> list[str]: +def missing_privileges( + conninfo: str, role: str, tables: list[str], privilege: str = "select", dialect: str | Dialect = "postgres" +) -> list[str]: """Table names from `tables` that `role` cannot currently exercise `privilege` on (checked via Postgres's own `has_table_privilege`). Empty means the role can access all of them. For guarding against the classic "migration creates a table, nobody grants it to the app's runtime role" gap: check the tables a migration just created against the role that will actually query them at runtime.""" + if _is_mssql(dialect): + raise NotImplementedError("missing_privileges is only available for the postgres dialect") if not tables: return [] missing = [] @@ -350,12 +391,18 @@ def _execute_stmts(conninfo: str, stmts: list[str]) -> None: con.commit() -def execute_sql_script(conninfo: str, sql: str) -> None: +def execute_sql_script(conninfo: str, sql: str, dialect: str | Dialect = "postgres") -> None: """Run a raw SQL script as one committed transaction, split into statements the same statement-boundary-safe way apply_migration is (dollar-quoted blocks, string literals, and line comments never get split mid-statement). Unlike apply_migration, this does no tracking-table bookkeeping and isn't forward-only -- for scripts meant to re-run - every time, like an idempotent `GRANT ... ON ALL TABLES IN SCHEMA` privilege sync.""" + every time, like an idempotent `GRANT ... ON ALL TABLES IN SCHEMA` privilege sync. + For MSSQL the script is split on `GO` lines instead and runs as one transaction.""" + if _is_mssql(dialect): + from . import migrate_mssql + + migrate_mssql.execute_batches(conninfo, migrate_mssql.split_batches(sql)) + return _execute_stmts(conninfo, _split_sql(sql)) @@ -372,25 +419,36 @@ def apply_migration( tracking_table: str, *, already_done: bool = False, + dialect: str | Dialect = "postgres", ) -> ApplyResult: sql = path.read_text(encoding="utf-8") filename = path.name - stmts = _split_sql(sql) + mssql = _is_mssql(dialect) + if mssql: + from . import migrate_mssql + + stmts = migrate_mssql.split_batches(sql) + else: + stmts = _split_sql(sql) if not already_done: - _execute_stmts(conninfo, stmts) + if mssql: + migrate_mssql.execute_batches(conninfo, stmts) + else: + _execute_stmts(conninfo, stmts) # Tracking insert is a separate connection/transaction so a missing tracking table # never rolls back the DDL that was just applied. - record_applied(conninfo, tracking_table, filename) + record_applied(conninfo, tracking_table, filename, dialect) if already_done: return ApplyResult(filename, executed=False, verified_tables=[]) - missing = verify_created_tables(conninfo, stmts) + missing = verify_created_tables(conninfo, stmts, dialect) if missing: raise MigrationVerificationError( f"{filename}: table(s) not found after apply — migration may have rolled back: " + ", ".join(missing) ) - return ApplyResult(filename, executed=True, verified_tables=_created_table_names(stmts)) + created = migrate_mssql.created_table_names(stmts) if mssql else _created_table_names(stmts) + return ApplyResult(filename, executed=True, verified_tables=created) diff --git a/pgdevkit/migrate_mssql.py b/pgdevkit/migrate_mssql.py new file mode 100644 index 0000000..2fc2603 --- /dev/null +++ b/pgdevkit/migrate_mssql.py @@ -0,0 +1,221 @@ +"""MSSQL flavour of the `pgdevkit.migrate` primitives: the same forward-only, +tracking-table-backed migration flow, built on `mssql-python` and T-SQL. + +Differences from the Postgres path that live here: + +- A migration is split into batches on standalone `GO` lines (T-SQL has no usable + statement separator for e.g. `CREATE VIEW`, which must be first in its batch) instead + of on semicolons, and all batches run in one transaction. +- Existence checks use `OBJECT_ID` / `SCHEMA_ID` / `COL_LENGTH`. +- The tracking table is `[schema].[table]`; its columns are + `filename nvarchar(450) primary key, applied_at datetimeoffset default + sysdatetimeoffset(), applied_by nvarchar(128) default suser_sname()`. + +Imported lazily by `pgdevkit.migrate` so the `mssql` extra is only needed when +`--dialect mssql` is actually used. +""" + +from __future__ import annotations + +import re +from datetime import datetime +from typing import Any + +import mssql_python +import sqlglot + +from .sql_text import strip_line_comments +from .testdb.query import split_tsql_batches + +_IDENTIFIER = r"[A-Za-z_][A-Za-z0-9_]*" +_NAME = r"(?:\[[^\]]+\]|[\w\"]+)(?:\.(?:\[[^\]]+\]|[\w\"]+))*" +_CREATE_TABLE_RE = re.compile(rf"CREATE\s+TABLE\s+({_NAME})", re.IGNORECASE) +_CREATE_VIEW_RE = re.compile(rf"CREATE\s+VIEW\s+({_NAME})", re.IGNORECASE) +_CREATE_SCHEMA_RE = re.compile(r"CREATE\s+SCHEMA\s+(\[[^\]]+\]|\w+)", re.IGNORECASE) +_ADD_COLUMN_RE = re.compile( + rf"ALTER\s+TABLE\s+({_NAME})\s+ADD\s+(?!(?:CONSTRAINT|PRIMARY|FOREIGN|UNIQUE|CHECK|DEFAULT)\b)" + rf"(\[[^\]]+\]|{_IDENTIFIER})\s", + re.IGNORECASE, +) + + +def tracking_table_parts(tracking_table: str) -> tuple[str, str]: + """Parse 'schema.table' into its two (unquoted) parts.""" + if not re.fullmatch(rf"{_IDENTIFIER}\.{_IDENTIFIER}", tracking_table): + raise ValueError(f"tracking_table must look like schema.table, got {tracking_table!r}") + schema, _, table = tracking_table.partition(".") + return schema, table + + +def _tracking_ident(tracking_table: str) -> str: + from .db.mssql_sql import qualified + + return qualified(*tracking_table_parts(tracking_table)) + + +def _unbracket(name: str) -> str: + return name[1:-1].replace("]]", "]") if name.startswith("[") else name.strip('"') + + +def split_batches(sql: str) -> list[str]: + return split_tsql_batches(sql) + + +def connect(conninfo: str) -> Any: + return mssql_python.connect(conninfo) + + +def _scalar(con: Any, sql: str, params: tuple = ()) -> Any: + cur = con.cursor() + try: + cur.execute(sql, params) + row = cur.fetchone() + return row[0] if row else None + finally: + cur.close() + + +def tracking_table_exists(con: Any, tracking_table: str) -> bool: + schema, table = tracking_table_parts(tracking_table) + return _scalar(con, "select object_id(?, N'U')", (f"[{schema}].[{table}]",)) is not None + + +def applied_migrations(conninfo: str, tracking_table: str) -> dict[str, tuple[datetime, str]]: + from .migrate import TrackingTableMissing + + con = connect(conninfo) + try: + if not tracking_table_exists(con, tracking_table): + raise TrackingTableMissing(tracking_table) + cur = con.cursor() + try: + cur.execute( + f"select filename, applied_at, applied_by from {_tracking_ident(tracking_table)} " + "order by applied_at" + ) + return {r[0]: (r[1], r[2]) for r in cur.fetchall()} + finally: + cur.close() + finally: + con.close() + + +def record_applied(conninfo: str, tracking_table: str, filename: str) -> bool: + """Best-effort insert into the tracking table (idempotent per filename). Returns False + without raising if the tracking table doesn't exist yet.""" + con = connect(conninfo) + try: + if not tracking_table_exists(con, tracking_table): + return False + ident = _tracking_ident(tracking_table) + cur = con.cursor() + try: + cur.execute( + f"if not exists (select 1 from {ident} where filename = ?) " + f"insert into {ident} (filename) values (?)", + (filename, filename), + ) + finally: + cur.close() + con.commit() + return True + finally: + con.close() + + +def execute_batches(conninfo: str, batches: list[str]) -> None: + """Run every batch in one transaction; roll back if any batch fails.""" + con = connect(conninfo) + try: + cur = con.cursor() + try: + for batch in batches: + cur.execute(batch) + finally: + cur.close() + con.commit() + except BaseException: + con.rollback() + raise + finally: + con.close() + + +def created_table_names(batches: list[str]) -> list[str]: + """Table names any CREATE TABLE in these batches targets. Temp tables (#x) are skipped + since they don't outlive the batch.""" + names: list[str] = [] + for batch in batches: + stripped = strip_line_comments(batch) + names.extend( + m.group(1) for m in _CREATE_TABLE_RE.finditer(stripped) if not m.group(1).startswith("#") + ) + return names + + +def _single_statement(batch: str) -> bool: + try: + return len([s for s in sqlglot.parse(batch, dialect="tsql") if s]) == 1 + except Exception: # noqa: BLE001 + return False + + +def _has_top_level_comma(stmt: str) -> bool: + depth = 0 + for c in stmt: + if c == "(": + depth += 1 + elif c == ")": + depth -= 1 + elif c == "," and depth == 0: + return True + return False + + +def idempotent_target(batch: str) -> tuple[str, ...] | None: + """MSSQL counterpart of `migrate._idempotent_target`: ("relation", name) for a table or + view, ("schema", name), or ("column", table, column) for an ADD column. Anything else + (CREATE OR ALTER, indexes, procedures, data changes, a batch holding more than one + statement, a multi-column ADD, ...) is None -- never guessed.""" + stripped = strip_line_comments(batch).strip() + if re.search(r"\bOR\s+ALTER\b", stripped, re.IGNORECASE) or not _single_statement(stripped): + return None + for regex in (_CREATE_TABLE_RE, _CREATE_VIEW_RE): + m = regex.match(stripped) + if m: + return ("relation", m.group(1)) + m = _CREATE_SCHEMA_RE.match(stripped) + if m: + return ("schema", _unbracket(m.group(1))) + m = _ADD_COLUMN_RE.match(stripped) + if m and not _has_top_level_comma(stripped): + return ("column", m.group(1), _unbracket(m.group(2))) + return None + + +def _target_exists(con: Any, target: tuple[str, ...]) -> bool: + kind = target[0] + if kind == "relation": + return _scalar(con, "select object_id(?)", (target[1],)) is not None + if kind == "schema": + return _scalar(con, "select schema_id(?)", (target[1],)) is not None + return _scalar(con, "select col_length(?, ?)", (target[1], target[2])) is not None + + +def targets_exist(conninfo: str, targets: list[tuple[str, ...]]) -> bool: + con = connect(conninfo) + try: + return all(_target_exists(con, t) for t in targets) + finally: + con.close() + + +def missing_tables(conninfo: str, tables: list[str]) -> list[str]: + """Table names (from `created_table_names`) that do NOT exist in the database.""" + if not tables: + return [] + con = connect(conninfo) + try: + return [t for t in tables if _scalar(con, "select object_id(?, N'U')", (t,)) is None] + finally: + con.close() diff --git a/pyproject.toml b/pyproject.toml index 47e56b4..a5177b4 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -11,7 +11,7 @@ packages = ["pgdevkit"] [project] name = "pgdevkit" -version = "0.11.0" +version = "0.12.0" description = "A helper for developing with Postgres" readme = "README.md" requires-python = ">=3.14" diff --git a/tests/test_migrate_mssql.py b/tests/test_migrate_mssql.py new file mode 100644 index 0000000..c1a8047 --- /dev/null +++ b/tests/test_migrate_mssql.py @@ -0,0 +1,222 @@ +from __future__ import annotations + +from pathlib import Path +from typing import Any + +import pytest +from typer.testing import CliRunner + +from pgdevkit import migrate, migrate_mssql +from pgdevkit.cli import app + +runner = CliRunner() + + +class FakeCursor: + def __init__(self, con: FakeConnection) -> None: + self.con = con + self._row: tuple | None = None + self._rows: list[tuple] = [] + + def execute(self, sql: str, params: tuple = ()) -> None: + self.con.log.append((sql, params)) + if self.con.fail_on and self.con.fail_on in sql: + raise RuntimeError("boom") + self._row, self._rows = None, [] + if sql.startswith("select object_id"): + name = params[0] + self._row = (1,) if name in self.con.objects else (None,) + elif sql.startswith("select schema_id"): + self._row = (1,) if params[0] in self.con.schemas else (None,) + elif sql.startswith("select col_length"): + self._row = (4,) if params in self.con.columns else (None,) + elif sql.startswith("select filename"): + self._rows = self.con.applied_rows + + def fetchone(self) -> tuple | None: + return self._row + + def fetchall(self) -> list[tuple]: + return self._rows + + def close(self) -> None: + pass + + +class FakeConnection: + def __init__(self, **kw: Any) -> None: + self.log: list[tuple[str, tuple]] = [] + self.objects: set[str] = kw.get("objects", set()) + self.schemas: set[str] = kw.get("schemas", set()) + self.columns: set[tuple] = kw.get("columns", set()) + self.applied_rows: list[tuple] = kw.get("applied_rows", []) + self.fail_on: str | None = kw.get("fail_on") + self.commits = 0 + self.rollbacks = 0 + + def cursor(self) -> FakeCursor: + return FakeCursor(self) + + def commit(self) -> None: + self.commits += 1 + + def rollback(self) -> None: + self.rollbacks += 1 + + def close(self) -> None: + pass + + def sql(self) -> list[str]: + return [s for s, _ in self.log] + + +@pytest.fixture +def fake(monkeypatch: pytest.MonkeyPatch) -> FakeConnection: + con = FakeConnection() + monkeypatch.setattr(migrate_mssql, "connect", lambda conninfo: con) + return con + + +def test_created_table_names_handles_brackets_and_skips_temp_tables(): + batches = ["CREATE TABLE [app].[widgets] (id int)", "CREATE TABLE #tmp (id int)", "CREATE TABLE dbo.b (id int)"] + assert migrate_mssql.created_table_names(batches) == ["[app].[widgets]", "dbo.b"] + + +def test_created_table_names_ignores_comments(): + assert migrate_mssql.created_table_names(["-- CREATE TABLE x (id int)\nSELECT 1"]) == [] + + +@pytest.mark.parametrize( + ("batch", "expected"), + [ + ("CREATE TABLE app.widgets (id int)", ("relation", "app.widgets")), + ("CREATE VIEW [app].[v] AS SELECT 1 AS a", ("relation", "[app].[v]")), + ("CREATE SCHEMA app", ("schema", "app")), + ("CREATE SCHEMA [app]", ("schema", "app")), + ("ALTER TABLE app.widgets ADD name nvarchar(10) NULL", ("column", "app.widgets", "name")), + ("ALTER TABLE app.widgets ADD [name] decimal(10,2)", ("column", "app.widgets", "name")), + ("ALTER TABLE app.widgets ADD a int, b int", None), + ("ALTER TABLE app.widgets ADD CONSTRAINT pk PRIMARY KEY (id)", None), + ("CREATE OR ALTER VIEW app.v AS SELECT 1 AS a", None), + ("CREATE INDEX ix ON app.widgets (id)", None), + ("CREATE TABLE a (id int); INSERT INTO a VALUES (1)", None), + ("INSERT INTO app.widgets (id) VALUES (1)", None), + ], +) +def test_idempotent_target(batch: str, expected: tuple | None): + assert migrate_mssql.idempotent_target(batch) == expected + + +def test_default_tracking_table_for_mssql_is_dbo(tmp_path: Path): + assert migrate.default_tracking_table(tmp_path, dialect="mssql") == "dbo.schema_migrations" + assert migrate.default_tracking_table(tmp_path) == "public.schema_migrations" + + +def test_apply_migration_splits_on_go_runs_one_transaction_and_records(fake: FakeConnection, tmp_path: Path): + fake.objects = {"app.widgets", "[dbo].[schema_migrations]"} + path = tmp_path / "001_widgets.sql" + path.write_text( + "CREATE TABLE app.widgets (id int);\nGO\nCREATE VIEW app.v AS SELECT id FROM app.widgets;\nGO\n", + encoding="utf-8", + ) + + result = migrate.apply_migration("conn", path, "dbo.schema_migrations", dialect="mssql") + + assert result.executed and result.verified_tables == ["app.widgets"] + statements = fake.sql() + assert statements[0] == "CREATE TABLE app.widgets (id int);" + assert statements[1].startswith("CREATE VIEW app.v") + assert fake.commits >= 2 # migration transaction + tracking insert + insert = next(s for s in statements if "insert into" in s) + assert "[dbo].[schema_migrations]" in insert + assert "not exists" in insert # idempotent record + + +def test_apply_migration_rolls_back_on_failure_and_does_not_record(fake: FakeConnection, tmp_path: Path): + fake.fail_on = "BAD" + path = tmp_path / "001_bad.sql" + path.write_text("SELECT 1\nGO\nBAD STATEMENT\n", encoding="utf-8") + + with pytest.raises(RuntimeError, match="boom"): + migrate.apply_migration("conn", path, "dbo.schema_migrations", dialect="mssql") + + assert fake.rollbacks == 1 and fake.commits == 0 + assert not any("insert into" in s for s in fake.sql()) + + +def test_apply_migration_raises_when_created_table_missing(fake: FakeConnection, tmp_path: Path): + path = tmp_path / "001_widgets.sql" + path.write_text("CREATE TABLE app.widgets (id int)\n", encoding="utf-8") + with pytest.raises(migrate.MigrationVerificationError, match="app.widgets"): + migrate.apply_migration("conn", path, "dbo.schema_migrations", dialect="mssql") + + +def test_apply_migration_tolerates_missing_tracking_table(fake: FakeConnection, tmp_path: Path): + path = tmp_path / "001_tracking.sql" + path.write_text("CREATE TABLE dbo.schema_migrations (filename nvarchar(450))\n", encoding="utf-8") + fake.objects = {"dbo.schema_migrations"} # exists for verification, not as [dbo].[schema_migrations] + result = migrate.apply_migration("conn", path, "dbo.schema_migrations", dialect="mssql") + assert result.executed + assert not any("insert into" in s for s in fake.sql()) + + +def test_applied_migrations_raises_tracking_table_missing(fake: FakeConnection): + with pytest.raises(migrate.TrackingTableMissing): + migrate.applied_migrations("conn", "dbo.schema_migrations", "mssql") + + +def test_applied_migrations_reads_rows(fake: FakeConnection): + fake.objects = {"[dbo].[schema_migrations]"} + fake.applied_rows = [("001_a.sql", "2026-01-01", "sa")] + assert migrate.applied_migrations("conn", "dbo.schema_migrations", "mssql") == { + "001_a.sql": ("2026-01-01", "sa") + } + + +def test_already_fully_applied(fake: FakeConnection, tmp_path: Path): + path = tmp_path / "001.sql" + path.write_text("CREATE SCHEMA app\nGO\nCREATE TABLE app.t (id int)\nGO\n", encoding="utf-8") + assert not migrate.already_fully_applied("conn", path, "mssql") + fake.schemas, fake.objects = {"app"}, {"app.t"} + assert migrate.already_fully_applied("conn", path, "mssql") + + +def test_already_fully_applied_false_for_unrecognized_batch(fake: FakeConnection, tmp_path: Path): + path = tmp_path / "001.sql" + path.write_text("CREATE TABLE app.t (id int)\nGO\nINSERT INTO app.t VALUES (1)\n", encoding="utf-8") + fake.objects = {"app.t"} + assert not migrate.already_fully_applied("conn", path, "mssql") + + +def test_tracking_table_must_be_schema_dot_table(fake: FakeConnection): + with pytest.raises(ValueError): + migrate.applied_migrations("conn", "x]; drop table y;--", "mssql") + + +def test_missing_privileges_is_postgres_only(): + with pytest.raises(NotImplementedError): + migrate.missing_privileges("conn", "role", ["a.b"], dialect="mssql") + + +def test_cli_check_and_apply_with_mssql_dialect(fake: FakeConnection, tmp_path: Path): + (tmp_path / "001_a.sql").write_text("CREATE SCHEMA app\n", encoding="utf-8") + fake.objects = {"[dbo].[schema_migrations]"} + fake.applied_rows = [] + + result = runner.invoke(app, ["migrate", "check", str(tmp_path), "--url", "Server=x", "--dialect", "mssql"]) + assert result.exit_code == 0, result.output + assert "1 pending" in result.output + + result = runner.invoke(app, ["migrate", "apply", str(tmp_path), "--url", "Server=x", "--dialect", "mssql", "-y"]) + assert result.exit_code == 0, result.output + assert "CREATE SCHEMA app" in fake.sql() + assert any("insert into [dbo].[schema_migrations]" in s for s in fake.sql()) + + +def test_cli_rejects_entra_user_and_unknown_dialect(tmp_path: Path): + result = runner.invoke( + app, ["migrate", "check", str(tmp_path), "--url", "Server=x", "--dialect", "mssql", "--entra-user", "a@b.c"] + ) + assert result.exit_code == 2 + result = runner.invoke(app, ["migrate", "check", str(tmp_path), "--url", "x", "--dialect", "oracle"]) + assert result.exit_code == 2 diff --git a/tests/test_migrate_mssql_live.py b/tests/test_migrate_mssql_live.py new file mode 100644 index 0000000..defe608 --- /dev/null +++ b/tests/test_migrate_mssql_live.py @@ -0,0 +1,49 @@ +from __future__ import annotations + +from pathlib import Path + +import pytest + +from pgdevkit import migrate +from pgdevkit.testdb.api import clean_testdb, ensure_testdb, status +from tests.testdb.conftest import _make_project, requires_mssql + +# Selected only by the dedicated `mssql-test` CI job (see test_compare_mssql_live.py). +pytestmark = pytest.mark.mssql + +_TRACKING = """\ +CREATE TABLE dbo.schema_migrations ( + filename nvarchar(450) NOT NULL PRIMARY KEY, + applied_at datetimeoffset NOT NULL DEFAULT sysdatetimeoffset(), + applied_by nvarchar(128) NOT NULL DEFAULT suser_sname() +) +""" + + +@requires_mssql +def test_migrations_apply_and_track_against_live_mssql(tmp_path: Path): + project = _make_project(tmp_path, "mssqlmigrate", "main", engine="mssql") + migrations = tmp_path / "migrations" + migrations.mkdir() + (migrations / "001_tracking.sql").write_text(_TRACKING, encoding="utf-8") + (migrations / "002_widgets.sql").write_text( + "CREATE SCHEMA mig_app\nGO\nCREATE TABLE mig_app.widget (id int NOT NULL PRIMARY KEY)\nGO\n" + "CREATE VIEW mig_app.widget_v AS SELECT id FROM mig_app.widget\nGO\n", + encoding="utf-8", + ) + try: + ensure_testdb(project) + dsn = status(project)["dsn"] + tracking = "dbo.schema_migrations" + + with pytest.raises(migrate.TrackingTableMissing): + migrate.applied_migrations(dsn, tracking, "mssql") + + for path in migrate.list_migration_files(migrations): + migrate.apply_migration(dsn, path, tracking, dialect="mssql") + + assert set(migrate.applied_migrations(dsn, tracking, "mssql")) == {"001_tracking.sql", "002_widgets.sql"} + assert migrate.pending_migrations(migrations, dsn, tracking, dialect="mssql") == [] + assert migrate.already_fully_applied(dsn, migrations / "002_widgets.sql", "mssql") is True + finally: + clean_testdb(project) diff --git a/uv.lock b/uv.lock index 17a3840..9e73ced 100644 --- a/uv.lock +++ b/uv.lock @@ -313,7 +313,7 @@ wheels = [ [[package]] name = "pgdevkit" -version = "0.11.0" +version = "0.12.0" source = { editable = "." } dependencies = [ { name = "docker" }, From 3d28168c6cd052a3231dbfad64189f55bf5d8af6 Mon Sep 17 00:00:00 2001 From: Adrian Ehrsam Date: Mon, 5 Oct 2026 20:18:38 +0000 Subject: [PATCH 2/3] Support --entra-user with --dialect mssql in pgdb migrate Appends Authentication=ActiveDirectoryDefault so mssql-python acquires the Entra token via DefaultAzureCredential. Co-Authored-By: Claude Sonnet 5.5 --- README.md | 8 ++++++-- pgdevkit/cli.py | 23 +++++++++++++++++------ pgdevkit/connection.py | 16 ++++++++++++++++ tests/test_migrate_mssql.py | 27 ++++++++++++++++++++++++++- 4 files changed, 65 insertions(+), 9 deletions(-) diff --git a/README.md b/README.md index afdf560..bd0f68d 100644 --- a/README.md +++ b/README.md @@ -331,8 +331,12 @@ pgdb migrate apply path/to/database/_migration_scripts --dialect mssql \ - `--ask` auto-detects "already done" for single-statement `GO` batches that create a table, view or schema, or add a column; anything else is asked about. - Post-apply verification checks every `CREATE TABLE` target via `OBJECT_ID`. -- `--entra-user` is Postgres-only, and `pgdevkit.migrate.missing_privileges` is not - available for MSSQL. +- `--entra-user` appends `Authentication=ActiveDirectoryDefault` to the connection + string, so mssql-python fetches an Entra ID token through azure-identity's + `DefaultAzureCredential` (install `pgdevkit[mssql,azure]`). The identity is whatever + that credential chain resolves — the flag's value only switches Entra auth on. It + can't be combined with an `Authentication=` already in the connection string. +- `pgdevkit.migrate.missing_privileges` is not available for MSSQL. For Postgres, `--entra-user` works the same as `pgdb compare` (see above). The tracking table needs `filename text primary key, applied_at timestamptz not null diff --git a/pgdevkit/cli.py b/pgdevkit/cli.py index 06eef55..fd8c2d2 100644 --- a/pgdevkit/cli.py +++ b/pgdevkit/cli.py @@ -15,7 +15,7 @@ from . import migrate, stats, testdb from .backends import get_backend -from .connection import build_conninfo +from .connection import build_conninfo, build_mssql_conninfo from .diff import DiffKind, compute_diff from .fetch_missing import SUBFOLDER, find_missing_objects, layer_folder_for, reconstruct_ddl from .parser import parse_directory @@ -388,10 +388,11 @@ def _migrate_target( except ValueError as e: err_console.print(f"[red]Error:[/red] {e}") raise typer.Exit(2) - if resolved.name == "mssql" and entra_user: - err_console.print("[red]Error:[/red] --entra-user is only supported for the postgres dialect") + try: + conninfo = (build_mssql_conninfo if resolved.name == "mssql" else build_conninfo)(url, entra_user) + except ValueError as e: + err_console.print(f"[red]Error:[/red] {e}") raise typer.Exit(2) - conninfo = url if resolved.name == "mssql" else build_conninfo(url, entra_user) tracking_table = tracking_table or migrate.default_tracking_table(migrations_dir, dialect=resolved) return conninfo, resolved.name, tracking_table @@ -405,7 +406,12 @@ def migrate_check( help="PostgreSQL DSN (postgresql://user:pass@host:port/db), or with --dialect mssql " "a connection string (Server=host,1433;Database=db;UID=user;PWD=pass)", ), - entra_user: str | None = typer.Option(None, "--entra-user", help="Azure Entra user (triggers token auth)"), + entra_user: str | None = typer.Option( + None, + "--entra-user", + help="Azure Entra user (triggers token auth); with --dialect mssql this adds " + "Authentication=ActiveDirectoryDefault and the value itself is not used", + ), tracking_table: str | None = typer.Option( None, "--tracking-table", @@ -465,7 +471,12 @@ def migrate_apply( help="PostgreSQL DSN (postgresql://user:pass@host:port/db), or with --dialect mssql " "a connection string (Server=host,1433;Database=db;UID=user;PWD=pass)", ), - entra_user: str | None = typer.Option(None, "--entra-user", help="Azure Entra user (triggers token auth)"), + entra_user: str | None = typer.Option( + None, + "--entra-user", + help="Azure Entra user (triggers token auth); with --dialect mssql this adds " + "Authentication=ActiveDirectoryDefault and the value itself is not used", + ), tracking_table: str | None = typer.Option( None, "--tracking-table", diff --git a/pgdevkit/connection.py b/pgdevkit/connection.py index 1036f06..c6f968f 100644 --- a/pgdevkit/connection.py +++ b/pgdevkit/connection.py @@ -1,5 +1,6 @@ from __future__ import annotations +import re from typing import Literal from urllib.parse import quote, urlparse, urlunparse @@ -67,6 +68,21 @@ def get_azure_postgres_password( return token.token +def build_mssql_conninfo(conn_str: str, entra_user: str | None = None) -> str: + """MSSQL counterpart of `build_conninfo`. Without `entra_user` the ODBC connection + string is used as-is. With it, `Authentication=ActiveDirectoryDefault` is appended so + mssql-python acquires an Entra ID token itself via azure-identity's + `DefaultAzureCredential` (needs the `azure` extra); the identity is whatever that + credential chain resolves, so `entra_user` only switches Entra auth on and is not + sent to the server. A connection string that already sets `Authentication` is left + to speak for itself and rejected here to avoid ambiguity.""" + if entra_user is None: + return conn_str + if re.search(r"(^|;)\s*Authentication\s*=", conn_str, re.IGNORECASE): + raise ValueError("--entra-user can't be combined with an Authentication= setting in the connection string") + return f"{conn_str.rstrip().rstrip(';')};Authentication=ActiveDirectoryDefault" + + def build_conninfo( url: str, entra_user: str | None = None, diff --git a/tests/test_migrate_mssql.py b/tests/test_migrate_mssql.py index c1a8047..01084d4 100644 --- a/tests/test_migrate_mssql.py +++ b/tests/test_migrate_mssql.py @@ -213,10 +213,35 @@ def test_cli_check_and_apply_with_mssql_dialect(fake: FakeConnection, tmp_path: assert any("insert into [dbo].[schema_migrations]" in s for s in fake.sql()) -def test_cli_rejects_entra_user_and_unknown_dialect(tmp_path: Path): +def test_build_mssql_conninfo(): + from pgdevkit.connection import build_mssql_conninfo + + assert build_mssql_conninfo("Server=x;Database=d") == "Server=x;Database=d" + assert ( + build_mssql_conninfo("Server=x;Database=d;", "a@b.c") + == "Server=x;Database=d;Authentication=ActiveDirectoryDefault" + ) + with pytest.raises(ValueError): + build_mssql_conninfo("Server=x;authentication = ActiveDirectoryMSI", "a@b.c") + + +def test_cli_entra_user_with_mssql_adds_authentication(monkeypatch: pytest.MonkeyPatch, tmp_path: Path): + seen: list[str] = [] + con = FakeConnection(objects={"[dbo].[schema_migrations]"}) + monkeypatch.setattr(migrate_mssql, "connect", lambda conninfo: (seen.append(conninfo), con)[1]) result = runner.invoke( app, ["migrate", "check", str(tmp_path), "--url", "Server=x", "--dialect", "mssql", "--entra-user", "a@b.c"] ) + assert result.exit_code == 0, result.output + assert seen == ["Server=x;Authentication=ActiveDirectoryDefault"] + + +def test_cli_rejects_conflicting_authentication_and_unknown_dialect(tmp_path: Path): + result = runner.invoke( + app, + ["migrate", "check", str(tmp_path), "--url", "Server=x;Authentication=ActiveDirectoryMSI", + "--dialect", "mssql", "--entra-user", "a@b.c"], + ) assert result.exit_code == 2 result = runner.invoke(app, ["migrate", "check", str(tmp_path), "--url", "x", "--dialect", "oracle"]) assert result.exit_code == 2 From fe393dfc771653b1db76faf9e89a198b5329fb27 Mon Sep 17 00:00:00 2001 From: Adrian Ehrsam Date: Mon, 5 Oct 2026 20:23:57 +0000 Subject: [PATCH 3/3] Address review: drain result sets, hide MSSQL creds in prompt, per-dialect tracking table key, strip block comments, mssql extra hint; --entra-user in compare Co-Authored-By: Claude Sonnet 5.5 --- README.md | 12 ++++- pgdevkit/cli.py | 48 +++++++++++++++---- pgdevkit/migrate.py | 6 ++- pgdevkit/migrate_mssql.py | 67 ++++++++++++++++++++++---- tests/test_migrate_mssql.py | 96 ++++++++++++++++++++++++++++++++++++- 5 files changed, 206 insertions(+), 23 deletions(-) diff --git a/README.md b/README.md index bd0f68d..c2ed34a 100644 --- a/README.md +++ b/README.md @@ -52,6 +52,12 @@ Requires the `mssql` extra: `pip install pgdevkit[mssql]` (pulls in own driver — no system ODBC driver install needed). MSSQL has no composite type or native enum equivalent, so those areas of a `database/` tree don't have a direct equivalent on this backend — see `docs/database-layout.md`. + +`pgdb compare --dialect mssql --entra-user ` appends +`Authentication=ActiveDirectoryDefault` to the connection string so mssql-python +gets an Entra ID token via `DefaultAzureCredential` (also install the `azure` +extra). The identity is whatever that credential chain resolves; the flag's +value only switches Entra auth on. Current Azure SQL/SQL Server (2025+) does have a native `json` column type, which parses/introspects/diffs like any other column type; see "`pgdevkit.db` — helpers for application code" below for how JSON values are @@ -324,10 +330,14 @@ pgdb migrate apply path/to/database/_migration_scripts --dialect mssql \ - Migration files are T-SQL: they are split into batches on standalone `GO` lines (needed for e.g. `CREATE VIEW`, which must be first in its batch), and all batches - of one file run in a single transaction that is rolled back if any batch fails. + of one file run in a single transaction that is rolled back if any batch fails + (including a failure in a later statement of a multi-statement batch). `GO ` + is not honored: the batch runs once. - The tracking table needs `filename nvarchar(450) primary key, applied_at datetimeoffset not null default sysdatetimeoffset(), applied_by nvarchar(128) not null default suser_sname()`. As on Postgres, a migration creating it can bootstrap it. + Override its name with `mssql_migrations_table` in `[tool.pgdevkit]` (the Postgres + `migrations_table` key is deliberately ignored for MSSQL) or with `--tracking-table`. - `--ask` auto-detects "already done" for single-statement `GO` batches that create a table, view or schema, or add a column; anything else is asked about. - Post-apply verification checks every `CREATE TABLE` target via `OBJECT_ID`. diff --git a/pgdevkit/cli.py b/pgdevkit/cli.py index fd8c2d2..088f3d1 100644 --- a/pgdevkit/cli.py +++ b/pgdevkit/cli.py @@ -9,11 +9,12 @@ import psycopg import typer from rich.console import Console +from rich.markup import escape from rich.table import Table from rich import box from tqdm import tqdm -from . import migrate, stats, testdb +from . import migrate, migrate_mssql, stats, testdb from .backends import get_backend from .connection import build_conninfo, build_mssql_conninfo from .diff import DiffKind, compute_diff @@ -73,7 +74,12 @@ def _as_set(values: list[str]) -> frozenset[str] | None: @app.command() def compare( url: str = typer.Option(..., "--url", help="PostgreSQL DSN (postgresql://user:pass@host:port/db)"), - entra_user: str | None = typer.Option(None, "--entra-user", help="Azure Entra user (triggers token auth)"), + entra_user: str | None = typer.Option( + None, + "--entra-user", + help="Azure Entra user (triggers token auth); with --dialect mssql this adds " + "Authentication=ActiveDirectoryDefault and the value itself is not used", + ), databricks_workspace_host: str | None = typer.Option( None, "--databricks-workspace-host", @@ -103,18 +109,21 @@ def compare( ) try: - conninfo = build_conninfo( - url, - entra_user, - databricks_workspace_host=databricks_workspace_host, - databricks_instance=databricks_instance, - ) + backend = get_backend(dialect) except ValueError as e: err_console.print(f"[red]Error:[/red] {e}") raise typer.Exit(2) try: - backend = get_backend(dialect) + if backend.dialect.name == "mssql": + conninfo = build_mssql_conninfo(url, entra_user) + else: + conninfo = build_conninfo( + url, + entra_user, + databricks_workspace_host=databricks_workspace_host, + databricks_instance=databricks_instance, + ) except ValueError as e: err_console.print(f"[red]Error:[/red] {e}") raise typer.Exit(2) @@ -379,6 +388,19 @@ def testdb_list_orphaned() -> None: console.print(name) +def _describe_target(url: str, dialect: str) -> str: + """Where a migrate command points, without credentials: host/db for a Postgres URL, + Server/Database only for an MSSQL connection string (never UID/PWD).""" + if dialect == "mssql": + parts = dict( + (k.strip().lower(), v.strip()) for k, _, v in (p.partition("=") for p in url.split(";")) if k.strip() + ) + server = parts.get("server") or parts.get("data source") or "?" + database = parts.get("database") or parts.get("initial catalog") + return f"{server}/{database}" if database else server + return url.rsplit("@", 1)[-1] if "@" in url else url + + def _migrate_target( url: str, entra_user: str | None, dialect: str, tracking_table: str | None, migrations_dir: Path ) -> tuple[str, str, str]: @@ -388,6 +410,12 @@ def _migrate_target( except ValueError as e: err_console.print(f"[red]Error:[/red] {e}") raise typer.Exit(2) + if resolved.name == "mssql": + try: + migrate_mssql.require_driver() + except ImportError as e: + err_console.print(f"[red]Error:[/red] {escape(str(e))}") + raise typer.Exit(2) try: conninfo = (build_mssql_conninfo if resolved.name == "mssql" else build_conninfo)(url, entra_user) except ValueError as e: @@ -504,7 +532,7 @@ def migrate_apply( conninfo, dialect, tracking_table = _migrate_target(url, entra_user, dialect, tracking_table, migrations_dir) areas, exclude_areas = _as_set(area), _as_set(exclude_area) schemas, exclude_schemas = _as_set(schema), _as_set(exclude_schema) - target_desc = url.rsplit("@", 1)[-1] if "@" in url else url + target_desc = _describe_target(url, dialect) if not yes: typer.confirm(f"About to run migrations against {target_desc}. Continue?", abort=True) diff --git a/pgdevkit/migrate.py b/pgdevkit/migrate.py index 15fcbc2..2e67b94 100644 --- a/pgdevkit/migrate.py +++ b/pgdevkit/migrate.py @@ -64,7 +64,8 @@ def _find_pyproject(start: Path) -> Path | None: def default_tracking_table(start: Path | None = None, dialect: str | Dialect = "postgres") -> str: """The project's configured tracking table: `[tool.pgdevkit].migrations_table` in the nearest pyproject.toml at or above `start` (default: cwd), or "public.schema_migrations" - if neither is set ("dbo.schema_migrations" for MSSQL). Lets a project fix its tracking table once instead of passing + if neither is set. For MSSQL the key is `mssql_migrations_table` (default + "dbo.schema_migrations"), so a Postgres-style `migrations_table` never leaks into it. Lets a project fix its tracking table once instead of passing --tracking-table on every `pgdb migrate` invocation.""" default = f"{resolve_dialect(dialect).default_schema}.schema_migrations" pyproject = _find_pyproject((start or Path.cwd()).resolve()) @@ -72,7 +73,8 @@ def default_tracking_table(start: Path | None = None, dialect: str | Dialect = " return default data = tomllib.loads(pyproject.read_text(encoding="utf-8")) section = data.get("tool", {}).get("pgdevkit", {}) - return section.get("migrations_table", default) + key = "mssql_migrations_table" if resolve_dialect(dialect).name == "mssql" else "migrations_table" + return section.get(key, default) def _split_sql(sql: str) -> list[str]: diff --git a/pgdevkit/migrate_mssql.py b/pgdevkit/migrate_mssql.py index 2fc2603..2e67fe1 100644 --- a/pgdevkit/migrate_mssql.py +++ b/pgdevkit/migrate_mssql.py @@ -21,10 +21,8 @@ from datetime import datetime from typing import Any -import mssql_python import sqlglot -from .sql_text import strip_line_comments from .testdb.query import split_tsql_batches _IDENTIFIER = r"[A-Za-z_][A-Za-z0-9_]*" @@ -61,10 +59,60 @@ def split_batches(sql: str) -> list[str]: return split_tsql_batches(sql) +def require_driver() -> None: + """Raise a clear ImportError if the `mssql` extra isn't installed.""" + try: + import mssql_python # noqa: F401 + except ImportError: + raise ImportError("MSSQL support needs the mssql extra: pip install pgdevkit[mssql]") from None + + def connect(conninfo: str) -> Any: + require_driver() + import mssql_python + return mssql_python.connect(conninfo) +def strip_comments(sql: str) -> str: + """Drop `--` line comments and (nestable) `/* ... */` block comments, leaving string + literals and [bracketed identifiers] untouched.""" + out: list[str] = [] + i, n, depth = 0, len(sql), 0 + quote: str | None = None # "'" or "]" while inside a literal / bracketed identifier + while i < n: + c, nxt = sql[i], sql[i + 1 : i + 2] + if depth: + if c == "/" and nxt == "*": + depth, i = depth + 1, i + 2 + elif c == "*" and nxt == "/": + depth, i = depth - 1, i + 2 + else: + i += 1 + elif quote: + out.append(c) + if c == quote: + if nxt == quote: # doubled '' or ]] is an escape + out.append(nxt) + i += 1 + else: + quote = None + i += 1 + elif c == "-" and nxt == "-": + while i < n and sql[i] != "\n": + i += 1 + elif c == "/" and nxt == "*": + depth, i = 1, i + 2 + else: + if c == "'": + quote = "'" + elif c == "[": + quote = "]" + out.append(c) + i += 1 + return "".join(out) + + def _scalar(con: Any, sql: str, params: tuple = ()) -> Any: cur = con.cursor() try: @@ -131,6 +179,10 @@ def execute_batches(conninfo: str, batches: list[str]) -> None: try: for batch in batches: cur.execute(batch) + # A statement-level error after the first statement of a batch is only + # reported once its result set is reached, so drain them all before commit. + while cur.nextset(): + pass finally: cur.close() con.commit() @@ -142,14 +194,11 @@ def execute_batches(conninfo: str, batches: list[str]) -> None: def created_table_names(batches: list[str]) -> list[str]: - """Table names any CREATE TABLE in these batches targets. Temp tables (#x) are skipped - since they don't outlive the batch.""" + """Table names any CREATE TABLE in these batches targets (comments ignored). Temp + tables (#x) never match `_NAME`, so they're skipped.""" names: list[str] = [] for batch in batches: - stripped = strip_line_comments(batch) - names.extend( - m.group(1) for m in _CREATE_TABLE_RE.finditer(stripped) if not m.group(1).startswith("#") - ) + names.extend(m.group(1) for m in _CREATE_TABLE_RE.finditer(strip_comments(batch))) return names @@ -177,7 +226,7 @@ def idempotent_target(batch: str) -> tuple[str, ...] | None: view, ("schema", name), or ("column", table, column) for an ADD column. Anything else (CREATE OR ALTER, indexes, procedures, data changes, a batch holding more than one statement, a multi-column ADD, ...) is None -- never guessed.""" - stripped = strip_line_comments(batch).strip() + stripped = strip_comments(batch).strip() if re.search(r"\bOR\s+ALTER\b", stripped, re.IGNORECASE) or not _single_statement(stripped): return None for regex in (_CREATE_TABLE_RE, _CREATE_VIEW_RE): diff --git a/tests/test_migrate_mssql.py b/tests/test_migrate_mssql.py index 01084d4..0c762d9 100644 --- a/tests/test_migrate_mssql.py +++ b/tests/test_migrate_mssql.py @@ -33,6 +33,12 @@ def execute(self, sql: str, params: tuple = ()) -> None: elif sql.startswith("select filename"): self._rows = self.con.applied_rows + def nextset(self) -> bool: + self.con.log.append(("", ())) + if self.con.fail_on_nextset: + raise RuntimeError("late boom") + return False + def fetchone(self) -> tuple | None: return self._row @@ -51,6 +57,7 @@ def __init__(self, **kw: Any) -> None: self.columns: set[tuple] = kw.get("columns", set()) self.applied_rows: list[tuple] = kw.get("applied_rows", []) self.fail_on: str | None = kw.get("fail_on") + self.fail_on_nextset: bool = kw.get("fail_on_nextset", False) self.commits = 0 self.rollbacks = 0 @@ -123,7 +130,7 @@ def test_apply_migration_splits_on_go_runs_one_transaction_and_records(fake: Fak result = migrate.apply_migration("conn", path, "dbo.schema_migrations", dialect="mssql") assert result.executed and result.verified_tables == ["app.widgets"] - statements = fake.sql() + statements = [q for q in fake.sql() if q != ""] assert statements[0] == "CREATE TABLE app.widgets (id int);" assert statements[1].startswith("CREATE VIEW app.v") assert fake.commits >= 2 # migration transaction + tracking insert @@ -245,3 +252,90 @@ def test_cli_rejects_conflicting_authentication_and_unknown_dialect(tmp_path: Pa assert result.exit_code == 2 result = runner.invoke(app, ["migrate", "check", str(tmp_path), "--url", "x", "--dialect", "oracle"]) assert result.exit_code == 2 + + +def test_late_statement_error_surfaced_by_draining_result_sets_rolls_back(fake: FakeConnection, tmp_path: Path): + fake.fail_on_nextset = True + path = tmp_path / "001_multi.sql" + path.write_text("UPDATE t SET a = 1; SELECT 1/0;\n", encoding="utf-8") + with pytest.raises(RuntimeError, match="late boom"): + migrate.apply_migration("conn", path, "dbo.schema_migrations", dialect="mssql") + assert fake.commits == 0 and fake.rollbacks == 1 + assert not any("insert into" in s for s in fake.sql()) + + +def test_strip_comments_handles_block_nested_line_and_literals(): + sql = "/* CREATE TABLE a (id int) /* nested */ still comment */ CREATE TABLE b (id int) -- CREATE TABLE c\n" \ + "SELECT '/* not a comment */', '--nope', [we--ird]" + out = migrate_mssql.strip_comments(sql) + assert "CREATE TABLE a" not in out and "CREATE TABLE c" not in out + assert "CREATE TABLE b" in out + assert "'/* not a comment */'" in out and "'--nope'" in out and "[we--ird]" in out + + +def test_block_commented_create_table_is_not_a_verification_target(): + batches = ["/* old: CREATE TABLE dbo.legacy (id int) */\nCREATE TABLE dbo.real (id int)"] + assert migrate_mssql.created_table_names(batches) == ["dbo.real"] + assert migrate_mssql.idempotent_target("/* CREATE TABLE x (id int) */ CREATE SCHEMA app") == ("schema", "app") + + +def test_mssql_tracking_table_ignores_postgres_migrations_table_key(tmp_path: Path): + (tmp_path / "pyproject.toml").write_text( + "[tool.pgdevkit]\nmigrations_table = 'public.schema_migrations'\nmssql_migrations_table = 'ops.migrations'\n" + ) + assert migrate.default_tracking_table(tmp_path, dialect="mssql") == "ops.migrations" + assert migrate.default_tracking_table(tmp_path) == "public.schema_migrations" + + (tmp_path / "pyproject.toml").write_text("[tool.pgdevkit]\nmigrations_table = 'public.schema_migrations'\n") + assert migrate.default_tracking_table(tmp_path, dialect="mssql") == "dbo.schema_migrations" + + +def test_confirmation_prompt_never_shows_mssql_credentials(fake: FakeConnection, tmp_path: Path): + (tmp_path / "001_a.sql").write_text("CREATE SCHEMA app\n", encoding="utf-8") + fake.objects = {"[dbo].[schema_migrations]"} + result = runner.invoke( + app, + ["migrate", "apply", str(tmp_path), "--dialect", "mssql", + "--url", "Server=h,1433;Database=d;UID=u;PWD=s3@cret"], + input="y\n", + ) + assert result.exit_code == 0, result.output + assert "h,1433/d" in result.output + assert "s3@cret" not in result.output and "PWD" not in result.output + + +def test_postgres_target_description_unchanged(): + from pgdevkit.cli import _describe_target + + assert _describe_target("postgresql://u:p@host:5432/db", "postgres") == "host:5432/db" + + +def test_missing_mssql_extra_gives_install_hint(monkeypatch: pytest.MonkeyPatch, tmp_path: Path): + def boom() -> None: + raise ImportError("MSSQL support needs the mssql extra: pip install pgdevkit[mssql]") + + monkeypatch.setattr(migrate_mssql, "require_driver", boom) + result = runner.invoke(app, ["migrate", "apply", str(tmp_path), "--url", "Server=x", "--dialect", "mssql", "-y"]) + assert result.exit_code == 2 + assert "pgdevkit[mssql]" in result.output + + +def test_compare_entra_user_with_mssql_uses_authentication_keyword(monkeypatch: pytest.MonkeyPatch, tmp_path: Path): + seen: list[str] = [] + + class StubBackend: + from pgdevkit.dialect import MSSQL as dialect + + def introspect(self, conninfo: str): + seen.append(conninfo) + from pgdevkit.models import DatabaseSchema + + return DatabaseSchema() + + monkeypatch.setattr("pgdevkit.cli.get_backend", lambda d: StubBackend()) + result = runner.invoke( + app, + ["compare", "--url", "Server=x", "--dialect", "mssql", "--entra-user", "a@b.c", str(tmp_path)], + ) + assert result.exit_code == 0, result.output + assert seen == ["Server=x;Authentication=ActiveDirectoryDefault"]