diff --git a/.github/workflows/python-test.yml b/.github/workflows/python-test.yml index a59f3d5..82158eb 100644 --- a/.github/workflows/python-test.yml +++ b/.github/workflows/python-test.yml @@ -31,4 +31,26 @@ jobs: - name: ty check run: uv run ty check pgdevkit - name: Test with pytest - run: uv run -m pytest --capture=tee-sys --maxfail=3 tests + run: uv run -m pytest --capture=tee-sys --maxfail=3 -m "not mssql" tests + + mssql-test: + # Separate job (not folded into `build`) so a live SQL Server -- a much + # larger image and slower cold start than the Postgres container -- + # never slows down or blocks the fast, always-run Postgres suite above. + # mssql-python bundles its own ODBC driver, so unlike pyodbc this needs + # no system driver package install here. + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + with: + submodules: "recursive" + - name: Set up Python 3.14 + uses: actions/setup-python@v5 + with: + python-version: "3.14" + - name: Install uv + run: curl -LsSf https://astral.sh/uv/install.sh | sh + - name: Install project dependencies + run: uv sync --all-extras --all-groups + - name: Run MSSQL-only tests + run: uv run -m pytest --capture=tee-sys --maxfail=3 -m mssql tests diff --git a/README.md b/README.md index 0e21d4f..dc8f20d 100644 --- a/README.md +++ b/README.md @@ -38,6 +38,22 @@ pgdb compare --url postgresql://instance-abc.database.azuredatabricks.net:5432/d (`--url`'s own user/password, if any, are discarded and replaced — `--entra-user` plus the fetched token become the connection's actual credentials.) +### MSSQL + +`pgdb compare`/`pgdb fetch-missing` default to Postgres. Pass `--dialect mssql` +to compare against a SQL Server database instead: + +```bash +pgdb compare --dialect mssql --url "Server=host,1433;Database=db;UID=user;PWD=pass" path/to/database/ +``` + +Requires the `mssql` extra: `pip install pgdevkit[mssql]` (pulls in +[mssql-python](https://github.com/microsoft/mssql-python), which bundles its +own driver — no system ODBC driver install needed). MSSQL has no composite +type, native enum, or first-class JSONB column type, so those areas of a +`database/` tree don't have a direct equivalent on this backend — see +`docs/database-layout.md`. + ## `pgdb testdb` Manages a single shared, Podman-backed Postgres container for local tests @@ -89,12 +105,29 @@ The role named by `PGDEVKIT_TESTDB_USER` must exist and match your OS user (`CREATE ROLE SUPERUSER LOGIN;`) and `pg_hba.conf` must allow `peer` auth for local connections (Debian/Ubuntu Postgres ships this by default). +### MSSQL + +Add `engine = "mssql"` to `[tool.pgdevkit]` (or set +`PGDEVKIT_TESTDB_ENGINE=mssql` for a one-off run) to manage a shared SQL +Server container instead of Postgres — same one-container-per-machine, +one-database-per-workspace model. Requires the `mssql` extra (see above). + +Container defaults (`localhost:14330`, `sa`/a generated complexity-valid +password) can be overridden with `PGDEVKIT_TESTDB_MSSQL_HOST`, `_PORT`, +`_USER`, `_PASSWORD`, `_IMAGE`, `_MEMORY_LIMIT_MB`. The container only +bootstraps the `sa` login — additional logins are a known limitation. +`pgdb testdb shell` execs into +[`sqlcmd`](https://github.com/microsoft/go-sqlcmd) (an external prerequisite, +the same category as `psql` for the Postgres path) rather than a Python +REPL. + ## `pgdevkit.db` — helpers for application code Install with the `db` extra: `pip install pgdevkit[db]`. -- **`PostgresTableModel`** — a `pydantic.BaseModel` base class for models - that map 1:1 to a table row. Implement `get_table_name()` (returns +- **`TableModel`** (formerly `PostgresTableModel`, still importable under + that name) — a `pydantic.BaseModel` base class for models that map 1:1 to + a table row, for either engine. Implement `get_table_name()` (returns `(schema, table)`) and `get_primary_key()` on each model. - **`PgPool`** — an async connection pool keyed off `{env_prefix}HOST/PORT/DB/USER/PASSWORD` env vars. Call `await pool.open()` @@ -107,8 +140,12 @@ Install with the `db` extra: `pip install pgdevkit[db]`. - **CRUD functions** — `pg_retrieve`, `pg_retrieve_many`, `pg_insert`, `pg_insert_many`, `pg_update`, `pg_update_dict`, `pg_upsert`, `pg_upsert_dict`, `pg_upsert_many`, `pg_upsert_many_dict`, `pg_delete`, - `pg_delete_dict` — typed (`PostgresTableModel`-based) or dict-based CRUD - against a table, built on `psycopg` for safe identifier/value handling. + `pg_delete_dict` — typed (`TableModel`-based) or dict-based CRUD against a + table, built on `psycopg` for safe identifier/value handling. The `mssql` + extra provides an `mssql_*`-prefixed mirror of the same functions in + `pgdevkit.db.mssql_crud`, built on `mssql-python` (`MERGE`-based upsert, + `OUTPUT` instead of `RETURNING`) — MSSQL has no composite/enum/JSONB + equivalent, so `complex_helper` is always `None` on that path. - **`SqlLoader`** — loads and caches `.sql` files from `{root}//.sql`, for keeping hand-written queries out of Python source. diff --git a/docs/database-layout.md b/docs/database-layout.md index 68d2887..b27f709 100644 --- a/docs/database-layout.md +++ b/docs/database-layout.md @@ -121,6 +121,21 @@ comment on column dim.user.is_active is 'False once a user is soft-deleted; keep --- +## MSSQL projects (`engine = "mssql"`) + +Everything above is engine-agnostic *as a folder/apply-order convention*, +with two exceptions: + +- `types/` (custom types / enums) has no direct T-SQL equivalent -- MSSQL + has neither a native enum type nor composite types, so a `CREATE TYPE ... + AS ENUM`/composite `.sql` file is a Postgres-only construct. `pgdb compare` + reports every such object as missing on an MSSQL database (correctly -- + it genuinely doesn't exist there), rather than erroring. +- T-SQL scripts conventionally separate batches with a standalone `GO` line + (an `sqlcmd`/SSMS scripting convention, not valid inside a single + driver `execute()` call). `pgdb testdb` splits on these automatically when + applying a file; hand-written `.sql` files may use `GO` freely. + ## Backfilling untracked objects If a table, scalar function, or table function was created directly on the diff --git a/pgdevkit/backends/__init__.py b/pgdevkit/backends/__init__.py new file mode 100644 index 0000000..5b9241b --- /dev/null +++ b/pgdevkit/backends/__init__.py @@ -0,0 +1,21 @@ +from __future__ import annotations + +from ..dialect import Dialect, resolve_dialect +from .base import Backend +from .mssql import MssqlBackend +from .postgres import PostgresBackend + +_REGISTRY: dict[str, Backend] = { + "postgres": PostgresBackend(), + "mssql": MssqlBackend(), +} + + +def get_backend(dialect: str | Dialect = "postgres") -> Backend: + """Look up the `Backend` for a dialect name (or an already-resolved + `Dialect`). Defaults to postgres.""" + resolved = resolve_dialect(dialect) + return _REGISTRY[resolved.name] + + +__all__ = ["Backend", "MssqlBackend", "PostgresBackend", "get_backend"] diff --git a/pgdevkit/backends/base.py b/pgdevkit/backends/base.py new file mode 100644 index 0000000..6c16684 --- /dev/null +++ b/pgdevkit/backends/base.py @@ -0,0 +1,27 @@ +from __future__ import annotations + +from typing import Any, Callable, Protocol + +from ..dialect import Dialect +from ..models import DatabaseSchema + + +class Backend(Protocol): + """Introspection + a couple of engine facts, behind one interface. + + CRUD is deliberately NOT part of this protocol -- psycopg's + `AsyncConnection` and an MSSQL driver's connection type are unrelated, + so a unified `backend.retrieve()`/`backend.insert()` surface would force + existing Postgres callers to go through a new indirection just to keep + working. Callers that want CRUD import `pgdevkit.db.crud`'s `pg_*` + functions or `pgdevkit.db.mssql_crud`'s `mssql_*` functions directly, + exactly as `db/crud.py`'s functions are imported today.""" + + dialect: Dialect + + def introspect(self, conninfo: str) -> DatabaseSchema: ... + + def complex_helper_factory(self) -> Callable[..., Any] | None: + """A `ComplexHelper`-like factory for composite/enum/JSONB columns, + or None when the engine has no equivalent (MSSQL).""" + ... diff --git a/pgdevkit/backends/mssql.py b/pgdevkit/backends/mssql.py new file mode 100644 index 0000000..1be8e6e --- /dev/null +++ b/pgdevkit/backends/mssql.py @@ -0,0 +1,27 @@ +from __future__ import annotations + +from typing import Any, Callable + +from ..dialect import MSSQL, Dialect +from ..models import DatabaseSchema + + +class MssqlBackend: + dialect: Dialect = MSSQL + + def introspect(self, conninfo: str) -> DatabaseSchema: + # Imported lazily so `import pgdevkit.backends` (and thus + # `pgdevkit.cli`) doesn't require mssql-python/the mssql extra to + # be installed unless a caller actually asks for the mssql backend. + from ..mssql_introspect import introspect_mssql_db + + return introspect_mssql_db(conninfo) + + def complex_helper_factory(self) -> Callable[..., Any] | None: + # MSSQL has no composite type, native enum, or first-class JSONB + # column type -- there is nothing for a ComplexHelper to adapt. + # Every `complex_helper` parameter in db/crud.py (and its + # db/mssql_crud.py counterpart) is already Optional, so callers on + # this backend simply pass/receive None and every complex-type + # branch takes its existing no-op path. + return None diff --git a/pgdevkit/backends/postgres.py b/pgdevkit/backends/postgres.py new file mode 100644 index 0000000..5c6452f --- /dev/null +++ b/pgdevkit/backends/postgres.py @@ -0,0 +1,19 @@ +from __future__ import annotations + +from typing import Any, Callable + +from ..dialect import POSTGRES, Dialect +from ..introspect import introspect_db +from ..models import DatabaseSchema + + +class PostgresBackend: + dialect: Dialect = POSTGRES + + def introspect(self, conninfo: str) -> DatabaseSchema: + return introspect_db(conninfo) + + def complex_helper_factory(self) -> Callable[..., Any] | None: + from ..db.complex_types import ComplexHelper + + return ComplexHelper diff --git a/pgdevkit/cli.py b/pgdevkit/cli.py index 4d60856..916e3a4 100644 --- a/pgdevkit/cli.py +++ b/pgdevkit/cli.py @@ -10,10 +10,10 @@ from rich import box from . import testdb +from .backends import get_backend from .connection import build_conninfo from .diff import DiffKind, compute_diff from .fetch_missing import SUBFOLDER, find_missing_objects, layer_folder_for, reconstruct_ddl -from .introspect import introspect_db from .parser import parse_directory app = typer.Typer(name="pgdb", help="PostgreSQL database schema tools") @@ -37,9 +37,10 @@ def compare( None, "--databricks-instance", help="Lakebase instance name (required for Lakebase hosts)" ), report_extra_db: bool = typer.Option(False, "--report-extra-db", help="Report objects in DB but not in scripts"), + dialect: str = typer.Option("postgres", "--dialect", help="postgres (default) or mssql"), scripts_dir: Path = typer.Argument(..., help="Directory containing SQL scripts"), ) -> None: - """Compare SQL scripts to a live PostgreSQL database and report differences.""" + """Compare SQL scripts to a live database and report differences.""" if not scripts_dir.is_dir(): err_console.print(f"[red]Error:[/red] {scripts_dir} is not a directory") raise typer.Exit(2) @@ -55,13 +56,19 @@ def compare( err_console.print(f"[red]Error:[/red] {e}") raise typer.Exit(2) + try: + backend = get_backend(dialect) + except ValueError as e: + err_console.print(f"[red]Error:[/red] {e}") + raise typer.Exit(2) + with console.status("Parsing SQL scripts..."): - scripts_schema = parse_directory(scripts_dir) + scripts_schema = parse_directory(scripts_dir, dialect=backend.dialect) with console.status("Introspecting database..."): - db_schema = introspect_db(conninfo) + db_schema = backend.introspect(conninfo) - diffs = compute_diff(scripts_schema, db_schema, report_extra_db=report_extra_db) + diffs = compute_diff(scripts_schema, db_schema, report_extra_db=report_extra_db, dialect=backend.dialect) if not diffs: console.print("[green]No differences found.[/green]") @@ -204,8 +211,10 @@ def testdb_status() -> None: @testdb_app.command("shell") def testdb_shell() -> None: - """Drop into psql against this workspace's database.""" - os.execvp("psql", ["psql", testdb.dsn_for()]) + """Drop into an interactive shell (psql, or sqlcmd for MSSQL) against + this workspace's database.""" + binary, argv = testdb.shell_argv() + os.execvp(binary, argv) @testdb_app.command("clean") diff --git a/pgdevkit/db/__init__.py b/pgdevkit/db/__init__.py index d8ae8b0..d41aaa7 100644 --- a/pgdevkit/db/__init__.py +++ b/pgdevkit/db/__init__.py @@ -17,13 +17,14 @@ pg_upsert_many_dict, ) from .loader import SqlLoader -from .model import PostgresTableModel +from .model import PostgresTableModel, TableModel __all__ = [ "ComplexHelper", "PgPool", "PostgresTableModel", "SqlLoader", + "TableModel", "pg_delete", "pg_delete_dict", "pg_insert", diff --git a/pgdevkit/db/model.py b/pgdevkit/db/model.py index c04e16d..286c8e8 100644 --- a/pgdevkit/db/model.py +++ b/pgdevkit/db/model.py @@ -6,8 +6,10 @@ from pydantic import BaseModel -class PostgresTableModel(BaseModel, ABC): - """Base class for models that map 1:1 to a database table/row. +class TableModel(BaseModel, ABC): + """Base class for models that map 1:1 to a database table/row (any + engine -- schema/table naming is equally meaningful for Postgres and + MSSQL, this base class was never actually Postgres-specific). Models representing partial results (joins, aggregations, projections) should extend `pydantic.BaseModel` directly instead.""" @@ -21,3 +23,8 @@ def get_table_name() -> tuple[str, str]: @abstractmethod def get_primary_key() -> Sequence[str]: """Return the primary key column name(s).""" + + +# Backward-compat alias -- this class was named PostgresTableModel before +# MSSQL support existed; kept so existing imports keep working unchanged. +PostgresTableModel = TableModel diff --git a/pgdevkit/db/mssql_crud.py b/pgdevkit/db/mssql_crud.py new file mode 100644 index 0000000..fc37d47 --- /dev/null +++ b/pgdevkit/db/mssql_crud.py @@ -0,0 +1,285 @@ +from __future__ import annotations + +import asyncio +from typing import Any, Callable, Mapping, Optional, Sequence, Type, TypeVar + +from .mssql_sql import ident, qualified +from .model import TableModel + +T = TypeVar("T", bound=TableModel) + +# `con` below is an mssql-python (github.com/microsoft/mssql-python) +# connection, but typed as `Any` rather than `mssql_python.Connection` so +# this module -- and, importantly, the pure `_build_*` query builders +# below, which have no driver dependency at all -- stays importable (and +# unit-testable) without the `mssql` extra installed. mssql-python bundles +# its own ODBC driver, so unlike pyodbc it needs no system driver install; +# its Connection/Cursor API otherwise mirrors pyodbc's (cursor(), execute(), +# executemany(), fetchone()/fetchall(), qmark `?` placeholders via a +# positional params list), which is what `_execute_returning`/`_execute_many` +# below rely on. `complex_helper` is likewise typed loosely: MSSQL has no +# ComplexHelper equivalent (see backends/mssql.py), so every caller on this +# backend passes/receives None here -- the parameter exists purely for +# signature symmetry with db/crud.py's `pg_*` functions. + + +def _build_retrieve(table_name: tuple[str, str], pks: dict) -> tuple[str, list]: + where = " AND ".join(f"{ident(k)} = ?" for k in pks) + sql = f"SELECT * FROM {qualified(*table_name)} WHERE {where}" + return sql, list(pks.values()) + + +def _build_retrieve_many(table_name: tuple[str, str], filters: dict) -> tuple[str, list]: + if not filters: + return f"SELECT * FROM {qualified(*table_name)}", [] + where = " AND ".join(f"{ident(k)} = ?" for k in filters) + sql = f"SELECT * FROM {qualified(*table_name)} WHERE {where}" + return sql, list(filters.values()) + + +def _build_insert(table_name: tuple[str, str], data: dict) -> tuple[str, list]: + fields = list(data) + cols = ", ".join(ident(k) for k in fields) + placeholders = ", ".join("?" for _ in fields) + sql = f"INSERT INTO {qualified(*table_name)} ({cols}) OUTPUT INSERTED.* VALUES ({placeholders})" + return sql, [data[k] for k in fields] + + +def _build_insert_many(table_name: tuple[str, str], fields: Sequence[str]) -> str: + cols = ", ".join(ident(k) for k in fields) + placeholders = ", ".join("?" for _ in fields) + return f"INSERT INTO {qualified(*table_name)} ({cols}) VALUES ({placeholders})" + + +def _build_update(table_name: tuple[str, str], data: dict, primary_keys: Sequence[str]) -> tuple[str, list]: + set_fields = [k for k in data if k not in primary_keys] + set_clause = ", ".join(f"{ident(k)} = ?" for k in set_fields) + where_clause = " AND ".join(f"{ident(pk)} = ?" for pk in primary_keys) + sql = f"UPDATE {qualified(*table_name)} SET {set_clause} OUTPUT INSERTED.* WHERE {where_clause}" + params = [data[k] for k in set_fields] + [data[pk] for pk in primary_keys] + return sql, params + + +def _build_update_many(table_name: tuple[str, str], fields: Sequence[str], primary_keys: Sequence[str]) -> str: + set_clause = ", ".join(f"{ident(k)} = ?" for k in fields if k not in primary_keys) + where_clause = " AND ".join(f"t.{ident(pk)} = ?" for pk in primary_keys) + return f"UPDATE {qualified(*table_name)} SET {set_clause} WHERE {where_clause}" + + +def _build_upsert_merge(table_name: tuple[str, str], data: dict, primary_keys: Sequence[str]) -> tuple[str, list]: + """MERGE INTO ... USING (SELECT ? AS col, ...) AS s ON pk = pk WHEN + MATCHED THEN UPDATE ... WHEN NOT MATCHED THEN INSERT ... OUTPUT + INSERTED.* -- the MSSQL replacement for Postgres's `INSERT ... ON + CONFLICT ... DO UPDATE ... EXCLUDED.col`. Structurally different from + an upsert-by-string-swap: MERGE is its own statement shape.""" + fields = list(data) + src_cols = ", ".join(f"? AS {ident(k)}" for k in fields) + on_clause = " AND ".join(f"t.{ident(pk)} = s.{ident(pk)}" for pk in primary_keys) + update_fields = [k for k in fields if k not in primary_keys] + insert_cols = ", ".join(ident(k) for k in fields) + insert_vals = ", ".join(f"s.{ident(k)}" for k in fields) + + sql = f"MERGE INTO {qualified(*table_name)} AS t USING (SELECT {src_cols}) AS s ON {on_clause} " + if update_fields: + update_clause = ", ".join(f"t.{ident(k)} = s.{ident(k)}" for k in update_fields) + sql += f"WHEN MATCHED THEN UPDATE SET {update_clause} " + sql += f"WHEN NOT MATCHED THEN INSERT ({insert_cols}) VALUES ({insert_vals}) OUTPUT INSERTED.*;" + return sql, [data[k] for k in fields] + + +def _build_delete(table_name: tuple[str, str], data: dict) -> tuple[str, list]: + where_clause = " AND ".join(f"{ident(k)} = ?" for k in data) + sql = f"DELETE FROM {qualified(*table_name)} OUTPUT DELETED.* WHERE {where_clause}" + return sql, list(data.values()) + + +async def _execute_returning(con: Any, sql: str, params: list) -> dict | None: + def _run() -> dict | None: + cur = con.cursor() + try: + cur.execute(sql, params) + cols = [c[0] for c in cur.description] + row = cur.fetchone() + return dict(zip(cols, row)) if row is not None else None + finally: + cur.close() + + return await asyncio.to_thread(_run) + + +async def _execute_many(con: Any, sql: str, param_rows: Sequence[Sequence]) -> None: + def _run() -> None: + cur = con.cursor() + try: + cur.executemany(sql, list(param_rows)) + finally: + cur.close() + + await asyncio.to_thread(_run) + + +async def mssql_retrieve( + con: Any, + data_type: Type[T], + pks: dict, + *, + complex_helper: Any | None = None, +) -> T | None: + """Fetch a single row by primary key(s). MSSQL has no ComplexHelper + equivalent (see backends/mssql.py) -- `complex_helper` exists only for + signature symmetry with `db.crud.pg_retrieve` and is otherwise unused.""" + sql, params = _build_retrieve(data_type.get_table_name(), pks) + row = await _execute_returning(con, sql, params) + return data_type(**row) if row else None + + +async def mssql_retrieve_many( + con: Any, + data_type: Type[T], + filters: dict, + *, + from_dict: Optional[Callable[[Mapping], T]] = None, + complex_helper: Any | None = None, +) -> Sequence[T]: + """Fetch multiple rows matching all filter key=value pairs.""" + sql, params = _build_retrieve_many(data_type.get_table_name(), filters) + + def _run() -> list[dict]: + cur = con.cursor() + try: + cur.execute(sql, params) + cols = [c[0] for c in cur.description] + return [dict(zip(cols, row)) for row in cur.fetchall()] + finally: + cur.close() + + rows = await asyncio.to_thread(_run) + fn = from_dict or (lambda d: data_type(**d)) + return [fn(r) for r in rows] + + +async def mssql_insert( + con: Any, + table_name: tuple[str, str], + data: dict, + *, + complex_helper: Any | None = None, +) -> dict[str, Any]: + """Insert one row and return the full row (`OUTPUT INSERTED.*`).""" + sql, params = _build_insert(table_name, data) + row = await _execute_returning(con, sql, params) + assert row is not None + return row + + +async def mssql_insert_many( + con: Any, + table_name: tuple[str, str], + data: Sequence[dict], + *, + complex_helper: Any | None = None, +) -> None: + """Batch insert -- no OUTPUT, one round-trip via executemany.""" + if not data: + return + fields = list(data[0]) + sql = _build_insert_many(table_name, fields) + await _execute_many(con, sql, [[row[k] for k in fields] for row in data]) + + +async def mssql_update_dict( + con: Any, + table_name: tuple[str, str], + data: dict, + primary_keys: Sequence[str], +) -> dict | None: + """Update a row identified by primary_keys. Returns the updated row.""" + sql, params = _build_update(table_name, data, primary_keys) + return await _execute_returning(con, sql, params) + + +async def mssql_update(con: Any, data: T, data_type: type[T]) -> dict | None: + """Update a typed model instance.""" + return await mssql_update_dict(con, data_type.get_table_name(), data.model_dump(), data_type.get_primary_key()) + + +async def mssql_upsert_dict( + con: Any, + table_name: tuple[str, str], + data: dict, + primary_keys: Sequence[str], + *, + complex_helper: Any | None = None, +) -> dict: + """MERGE-based upsert, returns the row as a dict.""" + sql, params = _build_upsert_merge(table_name, data, primary_keys) + row = await _execute_returning(con, sql, params) + assert row is not None + return row + + +async def mssql_upsert( + con: Any, data: T, data_type: type[T], *, complex_helper: Any | None = None +) -> dict: + """Upsert a typed model instance.""" + return await mssql_upsert_dict(con, data_type.get_table_name(), data.model_dump(), data_type.get_primary_key()) + + +async def mssql_upsert_many_dict( + con: Any, + table_name: tuple[str, str], + data: Sequence[dict], + primary_keys: Sequence[str], + *, + must_exist: bool = False, + complex_helper: Any | None = None, +) -> None: + """Batch upsert. + + `must_exist=True` switches to a plain UPDATE (no INSERT) matched on + `primary_keys` -- for callers that only ever update pre-existing rows + and want a missing row to be a silent no-op rather than create one.""" + if not data: + return + fields = list(data[0]) + if must_exist: + sql = _build_update_many(table_name, fields, primary_keys) + non_pk = [k for k in fields if k not in primary_keys] + rows = [[row[k] for k in non_pk] + [row[pk] for pk in primary_keys] for row in data] + await _execute_many(con, sql, rows) + else: + # MERGE's USING clause is per-row here (first cut) -- a set-based + # multi-row MERGE ... USING (VALUES (...), (...)) is more efficient + # but adds real complexity (a dynamic column-count VALUES list); + # row-by-row via executemany matches how the must_exist branch above + # already works. + for row in data: + sql, params = _build_upsert_merge(table_name, row, primary_keys) + + def _run() -> None: + cur = con.cursor() + try: + cur.execute(sql, params) + finally: + cur.close() + + await asyncio.to_thread(_run) + + +async def mssql_upsert_many( + con: Any, data: Sequence[T], data_type: type[T], *, complex_helper: Any | None = None +) -> None: + await mssql_upsert_many_dict(con, data_type.get_table_name(), [d.model_dump() for d in data], data_type.get_primary_key()) + + +async def mssql_delete_dict(con: Any, table_name: tuple[str, str], data: dict) -> dict | None: + """Delete by arbitrary key dict, returns the deleted row.""" + sql, params = _build_delete(table_name, data) + return await _execute_returning(con, sql, params) + + +async def mssql_delete(con: Any, data: T, data_type: type[T]) -> T | None: + """Delete a typed model instance by its primary key(s).""" + pk_dict = {pk: getattr(data, pk) for pk in data_type.get_primary_key()} + row = await mssql_delete_dict(con, data_type.get_table_name(), pk_dict) + return data_type.model_validate(row) if row else None diff --git a/pgdevkit/db/mssql_sql.py b/pgdevkit/db/mssql_sql.py new file mode 100644 index 0000000..e3c8a2e --- /dev/null +++ b/pgdevkit/db/mssql_sql.py @@ -0,0 +1,14 @@ +from __future__ import annotations + + +def ident(name: str) -> str: + """Bracket-quote a single identifier, doubling any embedded `]` + (T-SQL's escaping rule) -- the mssql-python driver has no + `psycopg.sql.Identifier` equivalent, so this is the composable-SQL + builder Postgres gets for free, hand-rolled for the one thing it's + actually needed for here.""" + return f"[{name.replace(']', ']]')}]" + + +def qualified(schema: str, table: str) -> str: + return f"{ident(schema)}.{ident(table)}" diff --git a/pgdevkit/dialect.py b/pgdevkit/dialect.py new file mode 100644 index 0000000..b414695 --- /dev/null +++ b/pgdevkit/dialect.py @@ -0,0 +1,93 @@ +from __future__ import annotations + +from dataclasses import dataclass, field + + +# Postgres type-name synonyms so scripts-vs-db type comparisons in diff.py +# are spelling-insensitive (e.g. a script written as "int4" matching a +# catalog-reported "integer"). Moved here (unchanged) from diff.py so both +# dialects' tables live next to the Dialect they belong to. +_POSTGRES_TYPE_SYNONYMS = { + "int": "integer", "int4": "integer", + "int2": "smallint", + "int8": "bigint", + "float4": "real", + "float8": "double precision", + "bool": "boolean", + "decimal": "numeric", + "varchar": "character varying", + "char": "character", "bpchar": "character", + "timestamptz": "timestamp with time zone", + "timestamp": "timestamp without time zone", + "timetz": "time with time zone", + "time": "time without time zone", + "varbit": "bit varying", + "serial": "integer", "serial4": "integer", + "smallserial": "smallint", "serial2": "smallint", + "bigserial": "bigint", "serial8": "bigint", +} + +# T-SQL's ISO/ODBC synonyms (per Microsoft's documented list) plus the one +# genuinely deprecated pair (timestamp/rowversion) that scripts still use. +_MSSQL_TYPE_SYNONYMS = { + "integer": "int", + "double precision": "float", + "national character": "nchar", + "national char": "nchar", + "national character varying": "nvarchar", + "national char varying": "nvarchar", + "char varying": "varchar", + "binary varying": "varbinary", + "numeric": "decimal", + "timestamp": "rowversion", +} + + +@dataclass(frozen=True) +class Dialect: + """A thin wrapper around a sqlglot dialect name plus the handful of + other facts that vary between engines and were previously hardcoded + throughout parser.py/diff.py/schema.py (default schema, type-name + synonyms, enum/composite-type support). Intentionally NOT a + reimplementation of anything sqlglot already does — `sqlglot_name` is + passed straight through to `sqlglot.parse()`/`.sql(dialect=...)`.""" + + name: str + sqlglot_name: str + default_schema: str + type_synonyms: dict[str, str] = field(default_factory=dict) + supports_enums: bool = True + supports_composites: bool = True + + +POSTGRES = Dialect( + name="postgres", + sqlglot_name="postgres", + default_schema="public", + type_synonyms=_POSTGRES_TYPE_SYNONYMS, + supports_enums=True, + supports_composites=True, +) + +MSSQL = Dialect( + name="mssql", + sqlglot_name="tsql", + default_schema="dbo", + type_synonyms=_MSSQL_TYPE_SYNONYMS, + supports_enums=False, + supports_composites=False, +) + +_REGISTRY = {"postgres": POSTGRES, "mssql": MSSQL} + + +def resolve_dialect(dialect: str | Dialect = "postgres") -> Dialect: + """Resolve a dialect name (or an already-resolved `Dialect`) to a + `Dialect` instance. Defaults to postgres, matching every caller's + default before this module existed.""" + if isinstance(dialect, Dialect): + return dialect + try: + return _REGISTRY[dialect] + except KeyError: + raise ValueError(f"Unknown dialect {dialect!r}; expected one of {sorted(_REGISTRY)}") from None diff --git a/pgdevkit/diff.py b/pgdevkit/diff.py index 390e753..a3e04cc 100644 --- a/pgdevkit/diff.py +++ b/pgdevkit/diff.py @@ -6,6 +6,7 @@ import sqlglot import sqlglot.expressions as exp +from .dialect import Dialect, POSTGRES, resolve_dialect from .models import DatabaseSchema, FunctionDef, IndexDef, TableDef @@ -23,7 +24,14 @@ class DiffEntry: detail: str = "" -def compute_diff(scripts: DatabaseSchema, db: DatabaseSchema, report_extra_db: bool = False) -> list[DiffEntry]: +def compute_diff( + scripts: DatabaseSchema, + db: DatabaseSchema, + report_extra_db: bool = False, + *, + dialect: str | Dialect = "postgres", +) -> list[DiffEntry]: + resolved = resolve_dialect(dialect) diffs: list[DiffEntry] = [] _diff_set("schema", scripts.schemas, db.schemas, diffs, report_extra_db) @@ -36,7 +44,7 @@ def compute_diff(scripts: DatabaseSchema, db: DatabaseSchema, report_extra_db: b tables_missing_in_db.add(name) diffs.append(DiffEntry(DiffKind.MISSING_IN_DB, "table", name)) else: - _diff_table(name, obj, db.tables[name], diffs) + _diff_table(name, obj, db.tables[name], diffs, resolved) if report_extra_db: for name in db.tables: if name not in scripts.tables: @@ -60,7 +68,7 @@ def compute_diff(scripts: DatabaseSchema, db: DatabaseSchema, report_extra_db: b if name not in db.functions: diffs.append(DiffEntry(DiffKind.MISSING_IN_DB, "function", name)) else: - _diff_function(name, obj, db.functions[name], diffs) + _diff_function(name, obj, db.functions[name], diffs, resolved) if report_extra_db: for name in db.functions: if name not in scripts.functions: @@ -91,7 +99,7 @@ def compute_diff(scripts: DatabaseSchema, db: DatabaseSchema, report_extra_db: b if f"{obj.schema}.{obj.table}" not in tables_missing_in_db: diffs.append(DiffEntry(DiffKind.MISSING_IN_DB, "index", name)) else: - _diff_index(name, obj, db.indexes[name], diffs) + _diff_index(name, obj, db.indexes[name], diffs, resolved) if report_extra_db: for name, obj in db.indexes.items(): if name not in scripts.indexes: @@ -111,7 +119,7 @@ def _diff_set(obj_type: str, scripts_set: set, db_set: set, diffs: list[DiffEntr diffs.append(DiffEntry(DiffKind.MISSING_IN_SCRIPTS, obj_type, item)) -def _diff_table(name: str, s: TableDef, d: TableDef, diffs: list[DiffEntry]) -> None: +def _diff_table(name: str, s: TableDef, d: TableDef, diffs: list[DiffEntry], dialect: Dialect) -> None: if s.is_partition or d.is_partition: return @@ -123,7 +131,7 @@ def _diff_table(name: str, s: TableDef, d: TableDef, diffs: list[DiffEntry]) -> else: dc = dcols[cname] issues = [] - if _norm_type(sc.data_type) != _norm_type(dc.data_type): + if _norm_type(sc.data_type, dialect) != _norm_type(dc.data_type, dialect): issues.append(f"type: {sc.data_type!r} vs {dc.data_type!r}") if sc.is_nullable != dc.is_nullable: issues.append(f"nullable: {sc.is_nullable} vs {dc.is_nullable}") @@ -134,9 +142,9 @@ def _diff_table(name: str, s: TableDef, d: TableDef, diffs: list[DiffEntry]) -> diffs.append(DiffEntry(DiffKind.MISSING_IN_SCRIPTS, "column", f"{name}.{cname}")) -def _diff_function(name: str, s: FunctionDef, d: FunctionDef, diffs: list[DiffEntry]) -> None: +def _diff_function(name: str, s: FunctionDef, d: FunctionDef, diffs: list[DiffEntry], dialect: Dialect) -> None: issues = [] - if _norm_type(s.return_type) != _norm_type(d.return_type): + if _norm_type(s.return_type, dialect) != _norm_type(d.return_type, dialect): issues.append(f"return_type: {s.return_type!r} vs {d.return_type!r}") if _norm_body(s.body) != _norm_body(d.body): issues.append("body differs") @@ -144,9 +152,9 @@ def _diff_function(name: str, s: FunctionDef, d: FunctionDef, diffs: list[DiffEn diffs.append(DiffEntry(DiffKind.MISMATCH, "function", name, "; ".join(issues))) -def _diff_index(name: str, s: IndexDef, d: IndexDef, diffs: list[DiffEntry]) -> None: - s_info = _parse_index_def(s.definition) - d_info = _parse_index_def(d.definition) +def _diff_index(name: str, s: IndexDef, d: IndexDef, diffs: list[DiffEntry], dialect: Dialect) -> None: + s_info = _parse_index_def(s.definition, dialect) + d_info = _parse_index_def(d.definition, dialect) if s_info is None or d_info is None: if _norm_sql(s.definition) != _norm_sql(d.definition): @@ -167,9 +175,9 @@ def _diff_index(name: str, s: IndexDef, d: IndexDef, diffs: list[DiffEntry]) -> diffs.append(DiffEntry(DiffKind.MISMATCH, "index", name, "; ".join(issues))) -def _parse_index_def(definition: str) -> dict | None: +def _parse_index_def(definition: str, dialect: Dialect = POSTGRES) -> dict | None: try: - parsed = sqlglot.parse_one(definition, dialect="postgres") + parsed = sqlglot.parse_one(definition, dialect=dialect.sqlglot_name) except Exception: return None if not isinstance(parsed, exp.Create): @@ -184,11 +192,11 @@ def _parse_index_def(definition: str) -> dict | None: columns = [] for col in (params.args.get("columns") if params else None) or []: - columns.append(_norm_sql(col.sql(dialect="postgres"))) + columns.append(_norm_sql(col.sql(dialect=dialect.sqlglot_name))) where_node = params.args.get("where") if params else None where_ast = _unwrap_paren(where_node.this) if where_node else None - where_display = _norm_sql(where_node.this.sql(dialect="postgres")) if where_node else None + where_display = _norm_sql(where_node.this.sql(dialect=dialect.sqlglot_name)) if where_node else None return { "unique": bool(parsed.args.get("unique")), @@ -215,28 +223,7 @@ def _norm_body(s: str) -> str: return "\n".join(l.lower() for l in lines if l) -_TYPE_SYNONYMS = { - "int": "integer", "int4": "integer", - "int2": "smallint", - "int8": "bigint", - "float4": "real", - "float8": "double precision", - "bool": "boolean", - "decimal": "numeric", - "varchar": "character varying", - "char": "character", "bpchar": "character", - "timestamptz": "timestamp with time zone", - "timestamp": "timestamp without time zone", - "timetz": "time with time zone", - "time": "time without time zone", - "varbit": "bit varying", - "serial": "integer", "serial4": "integer", - "smallserial": "smallint", "serial2": "smallint", - "bigserial": "bigint", "serial8": "bigint", -} - - -def _norm_type(t: str) -> str: +def _norm_type(t: str, dialect: Dialect = POSTGRES) -> str: s = " ".join(t.lower().split()) array_suffix = "" @@ -249,5 +236,5 @@ def _norm_type(t: str) -> str: base = (s[: match.start()] + s[match.end() :]).strip() if match else s base = " ".join(base.split()) - base = _TYPE_SYNONYMS.get(base, base) + base = dialect.type_synonyms.get(base, base) return f"{base}{params}{array_suffix}" diff --git a/pgdevkit/mssql_introspect.py b/pgdevkit/mssql_introspect.py new file mode 100644 index 0000000..506de50 --- /dev/null +++ b/pgdevkit/mssql_introspect.py @@ -0,0 +1,275 @@ +from __future__ import annotations + +from typing import Any + +import mssql_python +import sqlglot +import sqlglot.expressions as exp + +from .models import ( + ColumnDef, ConstraintDef, DatabaseSchema, FunctionDef, IndexDef, TableDef, ViewDef, +) +from .parser import _parse_function_details_tsql + +# Schemas that are SQL Server system/fixed-role schemas, not user schemas -- +# the equivalent of introspect.py's "pg_catalog"/"information_schema"/"pg_%" +# exclusion. +_SYSTEM_SCHEMAS = {"sys", "INFORMATION_SCHEMA", "guest"} + + +def _is_system_schema(name: str) -> bool: + return name in _SYSTEM_SCHEMAS or name.startswith("db_") + + +def _q(conn: Any, sql: str, params: tuple = ()) -> list[dict[str, Any]]: + cur = conn.cursor() + try: + cur.execute(sql, params) + cols = [c[0] for c in cur.description] + return [dict(zip(cols, row)) for row in cur.fetchall()] + finally: + cur.close() + + +def introspect_mssql_db(conninfo: str) -> DatabaseSchema: + conn = mssql_python.connect(conninfo) + try: + db = DatabaseSchema() + _load_schemas(conn, db) + _load_tables(conn, db) + _load_views(conn, db) + _load_functions(conn, db) + # MSSQL has no native enum or composite type -- a scripts.sql file + # that declares `CREATE TYPE ... AS ENUM`/a composite is a + # Postgres-only construct on this backend. Leaving these empty + # (rather than raising) means compute_diff reports every such + # object as MISSING_IN_DB, which is the honest answer: it genuinely + # doesn't exist as a first-class DB object here. + db.enums = {} + db.composites = {} + _load_indexes(conn, db) + return db + finally: + conn.close() + + +def _load_schemas(conn: Any, db: DatabaseSchema) -> None: + rows = _q(conn, "SELECT schema_name FROM information_schema.schemata") + db.schemas = {r["schema_name"] for r in rows if not _is_system_schema(r["schema_name"])} + + +_WCHAR_TYPES = {"nvarchar", "nchar"} +_CHAR_TYPES = {"varchar", "char", "varbinary", "binary"} +_DECIMAL_TYPES = {"decimal", "numeric"} + + +def _format_type(type_name: str, max_length: int, precision: int, scale: int) -> str: + """Render a sys.columns/sys.types row as a type string comparable to + what parser.py produces from a script's column definition (e.g. + "nvarchar(50)", "decimal(18,2)") -- the MSSQL analog of Postgres's + format_type(). A first cut: covers the character/decimal/float cases + that actually carry a meaningful length/precision; anything else is + rendered bare (int, bigint, bit, date, datetime2, uniqueidentifier, ...).""" + tn = type_name.lower() + if tn in _WCHAR_TYPES: + return f"{tn}(max)" if max_length == -1 else f"{tn}({max_length // 2})" + if tn in _CHAR_TYPES: + return f"{tn}(max)" if max_length == -1 else f"{tn}({max_length})" + if tn in _DECIMAL_TYPES: + return f"{tn}({precision},{scale})" + if tn == "float" and precision and precision != 53: + return f"{tn}({precision})" + return tn + + +def _load_tables(conn: Any, db: DatabaseSchema) -> None: + tables = _q(conn, """ + SELECT s.name AS [schema], t.name AS name, t.object_id AS object_id + FROM sys.tables t + JOIN sys.schemas s ON s.schema_id = t.schema_id + """) + + for row in tables: + tschema, tname, object_id = row["schema"], row["name"], row["object_id"] + if _is_system_schema(tschema): + continue + # SQL Server has no equivalent to Postgres's declarative + # partitioning (a table that IS a partition of another table) -- + # its own table partitioning is an internal storage detail of one + # table, not a distinct child-table relationship, so there is + # nothing to set here besides False. + table = TableDef(schema=tschema, name=tname, is_partition=False) + + for c in _q(conn, """ + SELECT c.name AS name, + ty.name AS base_type, + c.max_length AS max_length, + c.precision AS precision, + c.scale AS scale, + c.is_nullable AS is_nullable, + c.is_identity AS is_identity, + cc.definition AS is_computed_def, + dc.definition AS col_default + FROM sys.columns c + JOIN sys.types ty ON ty.user_type_id = c.user_type_id + LEFT JOIN sys.default_constraints dc + ON dc.parent_object_id = c.object_id AND dc.parent_column_id = c.column_id + LEFT JOIN sys.computed_columns cc + ON cc.object_id = c.object_id AND cc.column_id = c.column_id + WHERE c.object_id = ? + ORDER BY c.column_id + """, (object_id,)): + table.columns.append(ColumnDef( + name=c["name"], + data_type=_format_type(c["base_type"], c["max_length"], c["precision"], c["scale"]), + is_nullable=bool(c["is_nullable"]), + default=c["col_default"], + is_generated=bool(c["is_identity"]) or c["is_computed_def"] is not None, + )) + + table.constraints.extend(_load_key_constraints(conn, object_id)) + table.constraints.extend(_load_foreign_keys(conn, object_id)) + table.constraints.extend(_load_check_constraints(conn, object_id)) + + db.tables[table.qualified_name] = table + + +def _load_key_constraints(conn: Any, object_id: int) -> list[ConstraintDef]: + constraints = [] + for r in _q(conn, """ + SELECT kc.name AS name, kc.type AS type, i.index_id AS index_id + FROM sys.key_constraints kc + JOIN sys.indexes i ON i.object_id = kc.parent_object_id AND i.index_id = kc.unique_index_id + WHERE kc.parent_object_id = ? + """, (object_id,)): + cols = _q(conn, """ + SELECT c.name AS name + FROM sys.index_columns ic + JOIN sys.columns c ON c.object_id = ic.object_id AND c.column_id = ic.column_id + WHERE ic.object_id = ? AND ic.index_id = ? + ORDER BY ic.key_ordinal + """, (object_id, r["index_id"])) + col_list = ", ".join(c["name"] for c in cols) + kind = "PRIMARY KEY" if r["type"] == "PK" else "UNIQUE" + constraints.append(ConstraintDef(name=r["name"], kind=kind, definition=f"{kind.lower()} ({col_list})")) + return constraints + + +def _load_foreign_keys(conn: Any, object_id: int) -> list[ConstraintDef]: + constraints = [] + for r in _q(conn, "SELECT name, object_id FROM sys.foreign_keys WHERE parent_object_id = ?", (object_id,)): + cols = _q(conn, """ + SELECT pc.name AS col, rc.name AS ref_col, rt.name AS ref_table, rs.name AS ref_schema + FROM sys.foreign_key_columns fkc + JOIN sys.columns pc ON pc.object_id = fkc.parent_object_id AND pc.column_id = fkc.parent_column_id + JOIN sys.columns rc ON rc.object_id = fkc.referenced_object_id AND rc.column_id = fkc.referenced_column_id + JOIN sys.tables rt ON rt.object_id = fkc.referenced_object_id + JOIN sys.schemas rs ON rs.schema_id = rt.schema_id + WHERE fkc.constraint_object_id = ? + ORDER BY fkc.constraint_column_id + """, (r["object_id"],)) + if not cols: + continue + col_list = ", ".join(c["col"] for c in cols) + ref_list = ", ".join(c["ref_col"] for c in cols) + ref_table = f"{cols[0]['ref_schema']}.{cols[0]['ref_table']}" + definition = f"foreign key ({col_list}) references {ref_table} ({ref_list})" + constraints.append(ConstraintDef(name=r["name"], kind="FOREIGN KEY", definition=definition)) + return constraints + + +def _load_check_constraints(conn: Any, object_id: int) -> list[ConstraintDef]: + return [ + ConstraintDef(name=r["name"], kind="CHECK", definition=(r["definition"] or "").lower()) + for r in _q(conn, "SELECT name, definition FROM sys.check_constraints WHERE parent_object_id = ?", (object_id,)) + ] + + +def _load_views(conn: Any, db: DatabaseSchema) -> None: + for r in _q(conn, """ + SELECT s.name AS [schema], v.name AS name, m.definition AS definition + FROM sys.views v + JOIN sys.schemas s ON s.schema_id = v.schema_id + JOIN sys.sql_modules m ON m.object_id = v.object_id + """): + if _is_system_schema(r["schema"]): + continue + definition = _extract_view_query(r["definition"] or "").lower() + view = ViewDef(schema=r["schema"], name=r["name"], definition=definition) + db.views[view.qualified_name] = view + + +def _extract_view_query(definition: str) -> str: + """`sys.sql_modules.definition` is the verbatim `CREATE [OR ALTER] VIEW + ... AS ` statement text -- unlike Postgres's `pg_get_viewdef()`, + which returns only the query body. Parse it back out so `ViewDef.definition` + means the same thing on both backends and compares equal to parser.py's + script-side definition (also query-only).""" + try: + parsed = sqlglot.parse_one(definition, dialect="tsql") + except Exception: # noqa: BLE001 + return definition + if isinstance(parsed, exp.Create) and parsed.expression is not None: + return parsed.expression.sql(dialect="tsql") + return definition + + +_FUNCTION_KINDS = {"FN": "function", "IF": "function", "TF": "function", "P": "procedure"} + + +def _load_functions(conn: Any, db: DatabaseSchema) -> None: + for r in _q(conn, """ + SELECT s.name AS [schema], o.name AS name, + o.type AS type_code, m.definition AS definition + FROM sys.objects o + JOIN sys.schemas s ON s.schema_id = o.schema_id + JOIN sys.sql_modules m ON m.object_id = o.object_id + WHERE o.type IN ('FN', 'IF', 'TF', 'P') + """): + if _is_system_schema(r["schema"]): + continue + # Same verbatim-statement-text situation as views (see + # _extract_view_query) -- reuse parser.py's own T-SQL signature/body + # extraction on the catalog's stored definition text, so the + # introspected side is parsed exactly the same way the script side + # is, rather than maintaining two separate extraction paths that can + # drift out of sync. + args, return_type, language, body = _parse_function_details_tsql(r["definition"] or "") + + func = FunctionDef( + schema=r["schema"], name=r["name"], + args=args, return_type=return_type, + language=language, body=body, + kind=_FUNCTION_KINDS[r["type_code"]], + ) + db.functions[func.qualified_name] = func + + +def _load_indexes(conn: Any, db: DatabaseSchema) -> None: + for r in _q(conn, """ + SELECT s.name AS [schema], t.name AS table_name, i.name AS index_name, + i.index_id AS index_id, i.is_unique AS is_unique, t.object_id AS object_id, + i.filter_definition AS filter_definition + FROM sys.indexes i + JOIN sys.tables t ON t.object_id = i.object_id + JOIN sys.schemas s ON s.schema_id = t.schema_id + WHERE i.is_primary_key = 0 AND i.name IS NOT NULL + """): + if _is_system_schema(r["schema"]): + continue + cols = _q(conn, """ + SELECT c.name AS name + FROM sys.index_columns ic + JOIN sys.columns c ON c.object_id = ic.object_id AND c.column_id = ic.column_id + WHERE ic.object_id = ? AND ic.index_id = ? AND ic.is_included_column = 0 + ORDER BY ic.key_ordinal + """, (r["object_id"], r["index_id"])) + col_list = ", ".join(c["name"] for c in cols) + unique = "UNIQUE " if r["is_unique"] else "" + where_clause = f" WHERE {r['filter_definition']}" if r["filter_definition"] else "" + definition = ( + f"create {unique}index {r['index_name']} " + f"on {r['schema']}.{r['table_name']} ({col_list}){where_clause}" + ).lower() + idx = IndexDef(schema=r["schema"], table=r["table_name"], name=r["index_name"], definition=definition) + db.indexes[idx.qualified_name] = idx diff --git a/pgdevkit/parser.py b/pgdevkit/parser.py index 39b61f7..e4251e4 100644 --- a/pgdevkit/parser.py +++ b/pgdevkit/parser.py @@ -7,6 +7,7 @@ import sqlglot import sqlglot.expressions as exp +from .dialect import Dialect, POSTGRES, resolve_dialect from .models import ( ColumnDef, ConstraintDef, CompositeTypeDef, DatabaseSchema, EnumDef, FunctionDef, IndexDef, TableDef, ViewDef, @@ -14,33 +15,34 @@ logger = logging.getLogger(__name__) -# Regex to extract dollar-quoted body +# Regex to extract dollar-quoted body (Postgres-only construct) _DOLLAR_BODY = re.compile(r'\$(\w*)\$(.*?)\$\1\$', re.DOTALL | re.IGNORECASE) -# Regex for CREATE TYPE AS ENUM inside DO blocks +# Regex for CREATE TYPE AS ENUM inside DO blocks (Postgres-only construct) _DO_ENUM = re.compile( r'CREATE\s+TYPE\s+(\w+(?:\.\w+)?)\s+AS\s+ENUM\s*\(([^)]+)\)', re.IGNORECASE | re.DOTALL, ) -# Regex for CREATE TYPE AS composite inside DO blocks +# Regex for CREATE TYPE AS composite inside DO blocks (Postgres-only construct) _DO_COMPOSITE = re.compile( r'CREATE\s+TYPE\s+(\w+(?:\.\w+)?)\s+AS\s*\(([^)]+)\)', re.IGNORECASE | re.DOTALL, ) -def parse_directory(scripts_dir: Path) -> DatabaseSchema: +def parse_directory(scripts_dir: Path, *, dialect: str | Dialect = "postgres") -> DatabaseSchema: + resolved = resolve_dialect(dialect) db_schema = DatabaseSchema() for sql_file in sorted(scripts_dir.rglob("*.sql")): - _parse_file(sql_file, db_schema) + _parse_file(sql_file, db_schema, resolved) return db_schema -def _parse_file(path: Path, db_schema: DatabaseSchema) -> None: +def _parse_file(path: Path, db_schema: DatabaseSchema, dialect: Dialect = POSTGRES) -> None: content = path.read_text(encoding="utf-8") try: - exprs = sqlglot.parse(content, dialect="postgres", error_level=sqlglot.ErrorLevel.WARN) + exprs = sqlglot.parse(content, dialect=dialect.sqlglot_name, error_level=sqlglot.ErrorLevel.WARN) except Exception as e: logger.warning("sqlglot failed on %s: %s", path.name, e) exprs = [] @@ -49,32 +51,35 @@ def _parse_file(path: Path, db_schema: DatabaseSchema) -> None: if expr is None: continue try: - _handle_expr(expr, content, db_schema) + _handle_expr(expr, content, db_schema, dialect) except Exception as e: logger.debug("Skipping expression in %s: %s", path.name, e) - _extract_do_block_objects(content, db_schema) + # DO $$ ... $$ blocks are a Postgres-only construct; T-SQL has no + # equivalent, so this scan simply doesn't apply to other dialects. + if dialect.name == "postgres": + _extract_do_block_objects(content, db_schema) -def _handle_expr(expr: exp.Expression, raw: str, db_schema: DatabaseSchema) -> None: +def _handle_expr(expr: exp.Expression, raw: str, db_schema: DatabaseSchema, dialect: Dialect) -> None: if not isinstance(expr, exp.Create): return kind = (expr.args.get("kind") or "").upper() if kind == "TABLE": - _handle_table(expr, db_schema) + _handle_table(expr, db_schema, dialect) elif kind == "VIEW": - _handle_view(expr, db_schema) + _handle_view(expr, db_schema, dialect) elif kind in ("FUNCTION", "PROCEDURE"): - _handle_function(expr, raw, db_schema, kind.lower()) + _handle_function(expr, raw, db_schema, kind.lower(), dialect) elif kind == "TYPE": - _handle_type(expr, db_schema) + _handle_type(expr, db_schema, dialect) elif kind == "SCHEMA": _handle_schema_create(expr, db_schema) elif kind == "INDEX": - _handle_index(expr, db_schema) + _handle_index(expr, db_schema, dialect) -def _resolve_name(expr: exp.Create) -> tuple[str, str] | None: +def _resolve_name(expr: exp.Create, dialect: Dialect) -> tuple[str, str] | None: """Return (schema, name) from a CREATE expression.""" this = expr.this if isinstance(this, exp.Schema): @@ -84,13 +89,13 @@ def _resolve_name(expr: exp.Create) -> tuple[str, str] | None: if isinstance(table_node, exp.Table): db_node = table_node.args.get("db") - schema = db_node.name if db_node else "public" + schema = db_node.name if db_node else dialect.default_schema return schema, table_node.name return None -def _handle_table(expr: exp.Create, db_schema: DatabaseSchema) -> None: - result = _resolve_name(expr) +def _handle_table(expr: exp.Create, db_schema: DatabaseSchema, dialect: Dialect) -> None: + result = _resolve_name(expr, dialect) if not result: return tschema, tname = result @@ -106,11 +111,11 @@ def _handle_table(expr: exp.Create, db_schema: DatabaseSchema) -> None: pk_columns: set[str] = set() for item in items: if isinstance(item, exp.ColumnDef): - col = _parse_column_def(item) + col = _parse_column_def(item, dialect) if col: table.columns.append(col) else: - constr = _parse_table_constraint(item) + constr = _parse_table_constraint(item, dialect) if constr: table.constraints.append(constr) pk_columns |= _extract_primary_key_columns(item) @@ -131,11 +136,11 @@ def _extract_primary_key_columns(item: exp.Expression) -> set[str]: return set() -def _parse_column_def(col: exp.ColumnDef) -> ColumnDef | None: +def _parse_column_def(col: exp.ColumnDef, dialect: Dialect) -> ColumnDef | None: name = col.name if not name or col.kind is None: return None - data_type = col.kind.sql(dialect="postgres").lower() + data_type = col.kind.sql(dialect=dialect.sqlglot_name).lower() is_serial = col.kind.this in ( exp.DataType.Type.SERIAL, exp.DataType.Type.SMALLSERIAL, exp.DataType.Type.BIGSERIAL, ) @@ -145,10 +150,17 @@ def _parse_column_def(col: exp.ColumnDef) -> ColumnDef | None: for c in col.constraints: ck = c.kind - if isinstance(ck, (exp.NotNullColumnConstraint, exp.PrimaryKeyColumnConstraint)): + if isinstance(ck, exp.PrimaryKeyColumnConstraint): is_nullable = False + elif isinstance(ck, exp.NotNullColumnConstraint): + # sqlglot represents both "NOT NULL" and an explicit "NULL" + # (common T-SQL style) as this same node, distinguished only by + # allow_null -- true for the latter, which must NOT mark the + # column non-nullable. + if not ck.args.get("allow_null"): + is_nullable = False elif isinstance(ck, exp.DefaultColumnConstraint): - default = ck.this.sql(dialect="postgres") if ck.this else None + default = ck.this.sql(dialect=dialect.sqlglot_name) if ck.this else None elif isinstance(ck, exp.GeneratedAsIdentityColumnConstraint): is_generated = True is_nullable = False @@ -158,7 +170,7 @@ def _parse_column_def(col: exp.ColumnDef) -> ColumnDef | None: return ColumnDef(name=name, data_type=data_type, is_nullable=is_nullable, default=default, is_generated=is_generated) -def _parse_table_constraint(item: exp.Expression) -> ConstraintDef | None: +def _parse_table_constraint(item: exp.Expression, dialect: Dialect) -> ConstraintDef | None: name = None kind = "UNKNOWN" @@ -179,28 +191,28 @@ def _parse_table_constraint(item: exp.Expression) -> ConstraintDef | None: else: return None - definition = item.sql(dialect="postgres").lower() + definition = item.sql(dialect=dialect.sqlglot_name).lower() return ConstraintDef(name=name, kind=kind, definition=definition) -def _handle_view(expr: exp.Create, db_schema: DatabaseSchema) -> None: - result = _resolve_name(expr) +def _handle_view(expr: exp.Create, db_schema: DatabaseSchema, dialect: Dialect) -> None: + result = _resolve_name(expr, dialect) if not result: return vschema, vname = result query = expr.expression - definition = query.sql(dialect="postgres").lower() if query else "" + definition = query.sql(dialect=dialect.sqlglot_name).lower() if query else "" view = ViewDef(schema=vschema, name=vname, definition=definition) db_schema.views[view.qualified_name] = view -def _handle_function(expr: exp.Create, raw: str, db_schema: DatabaseSchema, kind: str) -> None: +def _handle_function(expr: exp.Create, raw: str, db_schema: DatabaseSchema, kind: str, dialect: Dialect) -> None: # Get name/schema from sqlglot func_node = expr.this fname = func_node.name if hasattr(func_node, "name") else "" if fname: db_node = func_node.args.get("db") if hasattr(func_node, "args") else None - fschema = db_node.name if db_node else "public" + fschema = db_node.name if db_node else dialect.default_schema else: # sqlglot (30.11.0) parses "CREATE FUNCTION myapp.greet(...)" as a # UserDefinedFunction wrapping a Table (this=Identifier(greet), @@ -210,15 +222,20 @@ def _handle_function(expr: exp.Create, raw: str, db_schema: DatabaseSchema, kind if isinstance(inner, exp.Table) and inner.name: fname = inner.name db_node = inner.args.get("db") - fschema = db_node.name if db_node else "public" + fschema = db_node.name if db_node else dialect.default_schema else: - result = _resolve_name(expr) + result = _resolve_name(expr, dialect) if not result: return fschema, fname = result - # Extract args, return type, language, body with regex on raw SQL - args, return_type, language, body = _parse_function_details(raw) + # Extract args, return type, language, body with regex on raw SQL -- + # dialect-specific since Postgres (dollar-quoted body, LANGUAGE clause) + # and T-SQL (AS BEGIN...END, no LANGUAGE clause) use different syntax. + if dialect.name == "mssql": + args, return_type, language, body = _parse_function_details_tsql(raw) + else: + args, return_type, language, body = _parse_function_details(raw) func = FunctionDef( schema=fschema, @@ -259,8 +276,52 @@ def _parse_function_details(sql: str) -> tuple[str, str, str, str]: return args, return_type, language, body -def _handle_type(expr: exp.Create, db_schema: DatabaseSchema) -> None: - result = _resolve_name(expr) +# T-SQL has no LANGUAGE clause (it's always effectively "sql") and no +# dollar-quoting -- a function/procedure body is just "AS [BEGIN] ... [END]" +# running to the end of the statement. This is a first-cut regex covering +# the common single-object-per-file convention this project uses; it isn't +# meant to handle every T-SQL corner case (nested BEGIN/END blocks with +# their own trailing semicolons, WITH ENCRYPTION/SCHEMABINDING options +# between the signature and AS, etc). +_FUNC_SIG_TSQL = re.compile( + r'CREATE\s+(?:OR\s+ALTER\s+)?(?:FUNCTION|PROCEDURE|PROC)\s+' + r'(?:\[?\w+\]?\.)?\[?\w+\]?\s*' + r'(\(([^)]*)\))?\s*' + r'(?:RETURNS\s+([^\s(]+(?:\s*\([^)]*\))?))?\s*' + r'AS\b', + re.IGNORECASE | re.DOTALL, +) + +_BEGIN_END_WRAPPER = re.compile(r'^\s*BEGIN\b(.*)\bEND\s*;?\s*$', re.IGNORECASE | re.DOTALL) + + +def _parse_function_details_tsql(sql: str) -> tuple[str, str, str, str]: + args, return_type, body = "", "", "" + + m = _FUNC_SIG_TSQL.search(sql) + if m: + args = re.sub(r'\s+', ' ', m.group(2) or "").strip().lower() + return_type = (m.group(3) or "").strip().lower() + raw_body = sql[m.end():].strip() + wrapper = _BEGIN_END_WRAPPER.match(raw_body) + if wrapper: + raw_body = wrapper.group(1) + lines = [l.strip() for l in raw_body.splitlines()] + body = "\n".join(l.lower() for l in lines if l) + + return args, return_type, "sql", body + + +def _handle_type(expr: exp.Create, db_schema: DatabaseSchema, dialect: Dialect) -> None: + # CREATE TYPE ... AS ENUM / AS (composite fields) are Postgres-only + # constructs. T-SQL's CREATE TYPE forms (table types, alias types) parse + # to different AST shapes that simply won't match the isinstance checks + # below, so this naturally no-ops for dialects without enum/composite + # support rather than needing an explicit dialect branch. + if not dialect.supports_enums and not dialect.supports_composites: + return + + result = _resolve_name(expr, dialect) if not result: return tschema, tname = result @@ -278,7 +339,7 @@ def _handle_type(expr: exp.Create, db_schema: DatabaseSchema) -> None: fields = [] for col in expression.expressions: if isinstance(col, exp.ColumnDef) and col.kind: - fields.append((col.name, col.kind.sql(dialect="postgres").lower())) + fields.append((col.name, col.kind.sql(dialect=dialect.sqlglot_name).lower())) comp = CompositeTypeDef(schema=tschema, name=tname, fields=fields) db_schema.composites[comp.qualified_name] = comp @@ -296,16 +357,16 @@ def _handle_schema_create(expr: exp.Create, db_schema: DatabaseSchema) -> None: db_schema.schemas.add(name) -def _handle_index(expr: exp.Create, db_schema: DatabaseSchema) -> None: +def _handle_index(expr: exp.Create, db_schema: DatabaseSchema, dialect: Dialect) -> None: this = expr.this index_name = this.name if hasattr(this, "name") else "" table_node = expr.find(exp.Table) if not table_node: return db_node = table_node.args.get("db") - tschema = db_node.name if db_node else "public" + tschema = db_node.name if db_node else dialect.default_schema tname = table_node.name - definition = expr.sql(dialect="postgres").lower() + definition = expr.sql(dialect=dialect.sqlglot_name).lower() idx = IndexDef(schema=tschema, table=tname, name=index_name, definition=definition) db_schema.indexes[idx.qualified_name] = idx diff --git a/pgdevkit/testdb/__init__.py b/pgdevkit/testdb/__init__.py index 102da2a..9c56f9e 100644 --- a/pgdevkit/testdb/__init__.py +++ b/pgdevkit/testdb/__init__.py @@ -1,3 +1,3 @@ -from .api import clean_testdb, dsn_for, ensure_testdb, reset_testdb, run_sql, status +from .api import clean_testdb, dsn_for, ensure_testdb, reset_testdb, run_sql, shell_argv, status -__all__ = ["clean_testdb", "dsn_for", "ensure_testdb", "reset_testdb", "run_sql", "status"] +__all__ = ["clean_testdb", "dsn_for", "ensure_testdb", "reset_testdb", "run_sql", "shell_argv", "status"] diff --git a/pgdevkit/testdb/_docker.py b/pgdevkit/testdb/_docker.py new file mode 100644 index 0000000..41cfcbf --- /dev/null +++ b/pgdevkit/testdb/_docker.py @@ -0,0 +1,41 @@ +from __future__ import annotations + +import os + +import docker + +# Candidate Docker-API-compatible socket URLs tried after plain +# docker.from_env() (which only looks at DOCKER_HOST / the default Docker +# socket) fails to connect -- covers rootful and rootless Podman, which +# speaks the same API but doesn't always advertise itself via DOCKER_HOST. +_FALLBACK_SOCKET_URLS = [ + f"unix://{os.environ['XDG_RUNTIME_DIR']}/podman/podman.sock" if os.environ.get("XDG_RUNTIME_DIR") else None, + "unix:///run/podman/podman.sock", +] + + +def client() -> docker.DockerClient: + """A Docker-API client, working against a real Docker daemon or a + Podman one (Podman exposes the same API over its own socket) -- callers + never need to know or care which one is actually running. Shared by + both the Postgres and MSSQL container modules -- a container is a + container, regardless of which image runs inside it.""" + try: + c = docker.from_env() + c.ping() + return c + except Exception: # noqa: BLE001 + pass + for base_url in _FALLBACK_SOCKET_URLS: + if base_url is None: + continue + try: + c = docker.DockerClient(base_url=base_url) + c.ping() + return c + except Exception: # noqa: BLE001 + continue + raise RuntimeError( + "Could not reach a Docker-compatible API. Set DOCKER_HOST, or make sure " + "Docker or Podman's API socket is running." + ) diff --git a/pgdevkit/testdb/api.py b/pgdevkit/testdb/api.py index 9b1fb0c..2e0e85a 100644 --- a/pgdevkit/testdb/api.py +++ b/pgdevkit/testdb/api.py @@ -28,6 +28,15 @@ def _resolve(project_root: Path | None) -> tuple[ProjectConfig, str]: return config, db_name +def _mssql_api(): + # Imported lazily so importing pgdevkit.testdb (and thus pgdevkit.cli) + # doesn't require the mssql extra unless a project actually opts into + # `engine = "mssql"`. + from .mssql import api as mssql_api + + return mssql_api + + def _env_for(config: ProjectConfig, db_name: str) -> dict[str, str]: prefix = config.env_prefix return { @@ -72,9 +81,13 @@ async def _apply(config: ProjectConfig, db_name: str, force_reset: bool) -> None def ensure_testdb(project_root: Path | None = None, force_reset: bool = False) -> dict[str, str]: """Ensure the shared container is running, this workspace's database exists, and its schema is applied. Returns the {PREFIX}POSTGRES_* env - vars for this workspace.""" - ensure_container() + vars for this workspace (or the mssql equivalent's env vars, per + `config.engine`).""" config, db_name = _resolve(project_root) + if config.engine == "mssql": + return _mssql_api().ensure_testdb(config, db_name, force_reset) + + ensure_container() async def _run() -> None: if force_reset: @@ -97,6 +110,9 @@ def clean_testdb(project_root: Path | None = None, all: bool = False) -> None: belonging to this project (matched by its name-slug prefix), across every worktree/branch.""" config, db_name = _resolve(project_root) + if config.engine == "mssql": + _mssql_api().clean_testdb(config, db_name, all) + return async def _run() -> None: if not all: @@ -118,7 +134,10 @@ async def _run() -> None: def status(project_root: Path | None = None) -> dict[str, str]: config, db_name = _resolve(project_root) + if config.engine == "mssql": + return _mssql_api().status(config, db_name) return { + "engine": config.engine, "container": constants.CONTAINER_NAME, "host": constants.HOST, "port": str(constants.PORT), @@ -128,10 +147,23 @@ def status(project_root: Path | None = None) -> dict[str, str]: def run_sql(sql: str, project_root: Path | None = None) -> list[dict] | None: - _, db_name = _resolve(project_root) + config, db_name = _resolve(project_root) + if config.engine == "mssql": + return _mssql_api().run_sql(config, db_name, sql) return asyncio.run(query.execute(_db_dsn(db_name), sql)) def dsn_for(project_root: Path | None = None) -> str: - _, db_name = _resolve(project_root) + config, db_name = _resolve(project_root) + if config.engine == "mssql": + return _mssql_api().dsn_for(config, db_name) return _db_dsn(db_name) + + +def shell_argv(project_root: Path | None = None) -> tuple[str, list[str]]: + """The (binary, argv) to `os.execvp` for an interactive shell against + this workspace's database -- `psql` for Postgres, `sqlcmd` for MSSQL.""" + config, db_name = _resolve(project_root) + if config.engine == "mssql": + return _mssql_api().shell_argv(config, db_name) + return "psql", ["psql", _db_dsn(db_name)] diff --git a/pgdevkit/testdb/config.py b/pgdevkit/testdb/config.py index e8f0e1c..c8e9349 100644 --- a/pgdevkit/testdb/config.py +++ b/pgdevkit/testdb/config.py @@ -1,9 +1,12 @@ from __future__ import annotations +import os import tomllib from dataclasses import dataclass, field from pathlib import Path +_ENGINES = ("postgres", "mssql") + @dataclass(frozen=True) class ProjectConfig: @@ -11,11 +14,14 @@ class ProjectConfig: database_dir: str = "database" env_prefix: str = "" extensions: tuple[str, ...] = () + engine: str = "postgres" root: Path = field(default_factory=Path) def __post_init__(self) -> None: if not self.env_prefix: object.__setattr__(self, "env_prefix", f"{self.name.upper()}_") + if self.engine not in _ENGINES: + raise ValueError(f"[tool.pgdevkit].engine must be one of {_ENGINES}, got {self.engine!r}") def _find_pyproject(start: Path) -> Path | None: @@ -42,10 +48,17 @@ def load_config(start: Path | None = None) -> ProjectConfig: f"[tool.pgdevkit].extensions in {pyproject} must be a list, got {type(extensions).__name__}" ) + # PGDEVKIT_TESTDB_ENGINE lets CI/ad-hoc runs flip engines without + # editing pyproject.toml; the toml value is the durable, per-project + # default (a project's database/ tree is written in one dialect, so + # this isn't meant to vary per-invocation the way a CLI flag would). + engine = os.environ.get("PGDEVKIT_TESTDB_ENGINE") or section.get("engine", "postgres") + return ProjectConfig( name=section.get("name") or root.name, database_dir=section.get("database_dir", "database"), env_prefix=section.get("env_prefix", ""), extensions=tuple(extensions), + engine=engine, root=root, ) diff --git a/pgdevkit/testdb/container.py b/pgdevkit/testdb/container.py index cad0e93..c3fec47 100644 --- a/pgdevkit/testdb/container.py +++ b/pgdevkit/testdb/container.py @@ -8,40 +8,7 @@ import psycopg from . import constants - -# Candidate Docker-API-compatible socket URLs tried after plain -# docker.from_env() (which only looks at DOCKER_HOST / the default Docker -# socket) fails to connect -- covers rootful and rootless Podman, which -# speaks the same API but doesn't always advertise itself via DOCKER_HOST. -_FALLBACK_SOCKET_URLS = [ - f"unix://{os.environ['XDG_RUNTIME_DIR']}/podman/podman.sock" if os.environ.get("XDG_RUNTIME_DIR") else None, - "unix:///run/podman/podman.sock", -] - - -def _client() -> docker.DockerClient: - """A Docker-API client, working against a real Docker daemon or a - Podman one (Podman exposes the same API over its own socket) -- callers - never need to know or care which one is actually running.""" - try: - client = docker.from_env() - client.ping() - return client - except Exception: # noqa: BLE001 - pass - for base_url in _FALLBACK_SOCKET_URLS: - if base_url is None: - continue - try: - client = docker.DockerClient(base_url=base_url) - client.ping() - return client - except Exception: # noqa: BLE001 - continue - raise RuntimeError( - "Could not reach a Docker-compatible API. Set DOCKER_HOST, or make sure " - "Docker or Podman's API socket is running." - ) +from ._docker import client as _client def _available(timeout: float = 3.0) -> bool: diff --git a/pgdevkit/testdb/mssql/__init__.py b/pgdevkit/testdb/mssql/__init__.py new file mode 100644 index 0000000..9d48db4 --- /dev/null +++ b/pgdevkit/testdb/mssql/__init__.py @@ -0,0 +1 @@ +from __future__ import annotations diff --git a/pgdevkit/testdb/mssql/api.py b/pgdevkit/testdb/mssql/api.py new file mode 100644 index 0000000..2284233 --- /dev/null +++ b/pgdevkit/testdb/mssql/api.py @@ -0,0 +1,219 @@ +from __future__ import annotations + +import asyncio +import json +from pathlib import Path +from typing import Any + +import mssql_python + +from ...db.mssql_sql import ident +from ...dialect import MSSQL +from .. import query +from ..config import ProjectConfig +from ..schema import _iter_sql_files, _strip_layer_prefix +from . import constants +from .container import ensure_mssql_container + + +def _admin_dsn() -> str: + return constants.conninfo("master") + + +def _db_dsn(db_name: str) -> str: + return constants.conninfo(db_name) + + +def _env_for(config: ProjectConfig, db_name: str) -> dict[str, str]: + prefix = config.env_prefix + return { + f"{prefix}MSSQL_HOST": constants.HOST, + f"{prefix}MSSQL_PORT": str(constants.PORT), + f"{prefix}MSSQL_DB": db_name, + f"{prefix}MSSQL_USER": constants.USER, + f"{prefix}MSSQL_PASSWORD": constants.PASSWORD, + } + + +async def _ensure_database(db_name: str) -> None: + def _run() -> None: + conn = mssql_python.connect(_admin_dsn(), autocommit=True) + try: + cur = conn.cursor() + cur.execute("SELECT 1 FROM sys.databases WHERE name = ?", [db_name]) + if cur.fetchone(): + return + cur.execute(f"CREATE DATABASE {ident(db_name)}") + finally: + conn.close() + + await asyncio.to_thread(_run) + + +async def _drop_database(db_name: str) -> None: + def _run() -> None: + conn = mssql_python.connect(_admin_dsn(), autocommit=True) + try: + cur = conn.cursor() + cur.execute("SELECT 1 FROM sys.databases WHERE name = ?", [db_name]) + if not cur.fetchone(): + return + # One statement kills other sessions and drops -- SQL Server's + # equivalent of Postgres's pg_terminate_backend()+DROP DATABASE. + cur.execute(f"ALTER DATABASE {ident(db_name)} SET SINGLE_USER WITH ROLLBACK IMMEDIATE") + cur.execute(f"DROP DATABASE IF EXISTS {ident(db_name)}") + finally: + conn.close() + + await asyncio.to_thread(_run) + + +async def _insert_test_data(json_file: Path, table: str, force_reset: bool, conn: Any) -> None: + if not json_file.exists(): + return + rows: list[dict[str, Any]] = json.loads(json_file.read_text(encoding="utf-8")) + if not rows: + return + schema, table_name = table.split(".") + qualified = f"{ident(schema)}.{ident(table_name)}" + + def _run() -> None: + cur = conn.cursor() + if not force_reset: + cur.execute(f"SELECT count(*) FROM {qualified}") + (count,) = cur.fetchone() + if count == len(rows): + return + # Unlike the Postgres path (ComplexHelper-driven composite/enum/JSONB + # conversion), MSSQL fixtures are limited to flat scalar columns for + # this first cut -- there is no MSSQL equivalent to convert into, so + # dict/list values are serialized as JSON text (matching how a + # NVARCHAR(MAX)-typed "JSON column" is conventionally stored on this + # engine) rather than silently dropped or erroring. + col_names = list(rows[0]) + cur.execute(f"DELETE FROM {qualified}") + cols = ", ".join(ident(c) for c in col_names) + placeholders = ", ".join("?" for _ in col_names) + insert_sql = f"INSERT INTO {qualified} ({cols}) VALUES ({placeholders})" + param_rows = [ + [json.dumps(row[c]) if isinstance(row[c], (dict, list)) else row[c] for c in col_names] + for row in rows + ] + cur.executemany(insert_sql, param_rows) + + await asyncio.to_thread(_run) + + +async def _apply(config: ProjectConfig, db_name: str, force_reset: bool) -> None: + def _connect() -> Any: + return mssql_python.connect(_db_dsn(db_name), autocommit=True) + + conn = await asyncio.to_thread(_connect) + try: + database_dir = config.root / config.database_dir + if not database_dir.is_dir(): + return + for file, sql in _iter_sql_files(database_dir, MSSQL): + for batch in query.split_tsql_batches(sql): + + def _exec(batch: str = batch) -> None: + conn.cursor().execute(batch) + + await asyncio.to_thread(_exec) + json_file = file.with_suffix(".test_data.json") + if json_file.exists(): + schema_name = _strip_layer_prefix(file.parent.parent.name) + table_stem = _strip_layer_prefix(file.stem) + await _insert_test_data(json_file, f"{schema_name}.{table_stem}", force_reset, conn) + finally: + await asyncio.to_thread(conn.close) + + +def ensure_testdb(config: ProjectConfig, db_name: str, force_reset: bool) -> dict[str, str]: + ensure_mssql_container() + + async def _run() -> None: + if force_reset: + await _drop_database(db_name) + await _ensure_database(db_name) + await _apply(config, db_name, force_reset) + + asyncio.run(_run()) + return _env_for(config, db_name) + + +def clean_testdb(config: ProjectConfig, db_name: str, all: bool) -> None: + from ..naming import slugify + + async def _run() -> None: + if not all: + await _drop_database(db_name) + return + prefix = f"{slugify(config.name)}_" + escaped_prefix = prefix.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_") + + def _list_names() -> list[str]: + conn = mssql_python.connect(_admin_dsn(), autocommit=True) + try: + cur = conn.cursor() + cur.execute("SELECT name FROM sys.databases WHERE name LIKE ? ESCAPE '\\'", [f"{escaped_prefix}%"]) + return [row[0] for row in cur.fetchall()] + finally: + conn.close() + + names = await asyncio.to_thread(_list_names) + for name in names: + await _drop_database(name) + + asyncio.run(_run()) + + +def status(config: ProjectConfig, db_name: str) -> dict[str, str]: + return { + "engine": config.engine, + "container": constants.CONTAINER_NAME, + "host": constants.HOST, + "port": str(constants.PORT), + "database": db_name, + "dsn": _db_dsn(db_name), + } + + +def run_sql(config: ProjectConfig, db_name: str, sql: str) -> list[dict] | None: + def _run() -> list[dict] | None: + conn = mssql_python.connect(_db_dsn(db_name), autocommit=True) + try: + last_rows: list[dict] | None = None + for batch in query.split_tsql_batches(sql): + cur = conn.cursor() + cur.execute(batch) + if cur.description: + cols = [c[0] for c in cur.description] + last_rows = [dict(zip(cols, row)) for row in cur.fetchall()] + else: + last_rows = None + return last_rows + finally: + conn.close() + + return _run() + + +def dsn_for(config: ProjectConfig, db_name: str) -> str: + return _db_dsn(db_name) + + +def shell_argv(config: ProjectConfig, db_name: str) -> tuple[str, list[str]]: + """`sqlcmd` (the modern standalone github.com/microsoft/go-sqlcmd build, + not the legacy mssql-tools18 package) is the documented external + prerequisite here -- the same category as `psql` being assumed on PATH + for the Postgres path. `-C` trusts the container's self-signed cert, + required since sqlcmd v18+ defaults to encrypted+verified connections.""" + return "sqlcmd", [ + "sqlcmd", + "-S", f"{constants.HOST},{constants.PORT}", + "-U", constants.USER, + "-P", constants.PASSWORD, + "-d", db_name, + "-C", + ] diff --git a/pgdevkit/testdb/mssql/constants.py b/pgdevkit/testdb/mssql/constants.py new file mode 100644 index 0000000..44fad9f --- /dev/null +++ b/pgdevkit/testdb/mssql/constants.py @@ -0,0 +1,60 @@ +from __future__ import annotations + +import os +import string + +CONTAINER_NAME = "pgdevkit-mssql" +# amd64-only -- overridable for Apple Silicon dev machines, where this runs +# emulated under Docker Desktop (slow, occasionally flaky startup). The +# documented escape hatch is mcr.microsoft.com/azure-sql-edge (multi-arch, +# missing a few full-SQL-Server features). +IMAGE = os.environ.get("PGDEVKIT_TESTDB_MSSQL_IMAGE", "mcr.microsoft.com/mssql/server:2022-latest") +HOST = os.environ.get("PGDEVKIT_TESTDB_MSSQL_HOST", "localhost") +# Deliberately not SQL Server's native 1433 -- avoids colliding with a +# locally-installed instance, mirroring how the Postgres container's own +# default (54322) differs from Postgres's native 5432. +PORT = int(os.environ.get("PGDEVKIT_TESTDB_MSSQL_PORT", "14330")) +# The container only bootstraps the `sa` login -- additional logins are a +# known limitation of this first cut. +USER = os.environ.get("PGDEVKIT_TESTDB_MSSQL_USER", "sa") +# SQL Server's password-complexity rule rejects the Postgres container's +# plain default ("testpwd") outright, so this needs its own complexity-valid +# default -- see validate_sa_password(). +PASSWORD = os.environ.get("PGDEVKIT_TESTDB_MSSQL_PASSWORD", "TestPwd!2026") +MEMORY_LIMIT_MB = int(os.environ.get("PGDEVKIT_TESTDB_MSSQL_MEMORY_LIMIT_MB", "2048")) + + +def validate_sa_password(password: str) -> None: + """Check SQL Server's SA-password complexity rule before starting a + container with it, so a weak custom password fails fast with a clear + message instead of surfacing as an opaque container crash-loop. + + Rule (per Microsoft's documented policy): at least 8 characters, and at + least 3 of {uppercase, lowercase, digit, symbol}; must not contain the + login name "sa".""" + if len(password) < 8: + raise ValueError("MSSQL_SA_PASSWORD must be at least 8 characters long") + if "sa" in password.lower(): + raise ValueError("MSSQL_SA_PASSWORD must not contain the login name 'sa'") + classes_met = sum([ + any(c.islower() for c in password), + any(c.isupper() for c in password), + any(c.isdigit() for c in password), + any(c in string.punctuation for c in password), + ]) + if classes_met < 3: + raise ValueError( + "MSSQL_SA_PASSWORD must contain at least 3 of: uppercase letter, " + "lowercase letter, digit, symbol" + ) + + +def conninfo(dbname: str) -> str: + """Build an mssql-python connection string from HOST/PORT/USER/PASSWORD. + `Driver`/`APP` are deliberately omitted -- mssql-python controls those + itself (it bundles its own driver, so no system ODBC driver install is + needed) and raises if a caller tries to set them.""" + return ( + f"Server={HOST},{PORT};Database={dbname};UID={USER};PWD={PASSWORD};" + "Encrypt=yes;TrustServerCertificate=yes" + ) diff --git a/pgdevkit/testdb/mssql/container.py b/pgdevkit/testdb/mssql/container.py new file mode 100644 index 0000000..343d3d7 --- /dev/null +++ b/pgdevkit/testdb/mssql/container.py @@ -0,0 +1,90 @@ +from __future__ import annotations + +import os +import time + +import docker +import docker.errors +import mssql_python + +from .. import _docker +from . import constants + + +def _available(timeout: float = 3.0) -> bool: + """Quick check (short timeout) for whether SQL Server is already + reachable at HOST:PORT, so a container started outside pgdevkit's + control doesn't trigger another Docker API call.""" + try: + conn = mssql_python.connect(constants.conninfo("master"), timeout=max(1, round(timeout))) + conn.close() + return True + except Exception: # noqa: BLE001 + return False + + +def _create_container(client: docker.DockerClient) -> None: + constants.validate_sa_password(constants.PASSWORD) + try: + client.containers.run( + constants.IMAGE, + name=constants.CONTAINER_NAME, + detach=True, + ports={"1433/tcp": constants.PORT}, + environment={ + "ACCEPT_EULA": "Y", + "MSSQL_SA_PASSWORD": constants.PASSWORD, + # Unset defaults to the Evaluation edition, which stops + # working after 180 days -- a real footgun for a long-lived + # shared dev container. + "MSSQL_PID": "Developer", + "MSSQL_MEMORY_LIMIT_MB": str(constants.MEMORY_LIMIT_MB), + }, + ) + except docker.errors.APIError as e: + if getattr(e, "status_code", None) == 409 or "already in use" in str(e): + client.containers.get(constants.CONTAINER_NAME).start() + return + raise RuntimeError(f"Starting the {constants.CONTAINER_NAME} container failed: {e}") from e + + +def _wait_ready(timeout: float = 90.0) -> None: + # SQL Server accepts TCP connections before its internal init finishes, + # surfacing as transient "Login failed"/"server is not currently + # accepting connections" errors -- treated as "not ready yet," same as + # the Postgres container's broad `except Exception`. Cold start is + # slower than Postgres's, hence the longer default timeout. + deadline = time.monotonic() + timeout + last_error: Exception | None = None + while time.monotonic() < deadline: + try: + conn = mssql_python.connect(constants.conninfo("master"), timeout=2) + conn.close() + return + except Exception as e: # noqa: BLE001 + last_error = e + time.sleep(0.5) + raise RuntimeError(f"SQL Server did not become ready within {timeout}s: {last_error}") + + +def ensure_mssql_container() -> None: + """Idempotently ensure the shared pgdevkit-mssql container is running + and accepting connections. Never touches the Docker API if SQL Server + is already reachable, or if PGDEVKIT_SKIP_CONTAINER says to assume it + is (the same on/off switch used for the Postgres container -- one + project uses one engine, so there's no need for a second env var).""" + if os.environ.get("PGDEVKIT_SKIP_CONTAINER"): + return + if _available(): + return + client = _docker.client() + try: + container = client.containers.get(constants.CONTAINER_NAME) + except docker.errors.NotFound: + container = None + if container is not None: + if container.status != "running": + container.start() + else: + _create_container(client) + _wait_ready() diff --git a/pgdevkit/testdb/query.py b/pgdevkit/testdb/query.py index 7fb7a75..dfc7eb3 100644 --- a/pgdevkit/testdb/query.py +++ b/pgdevkit/testdb/query.py @@ -7,6 +7,44 @@ from psycopg.rows import dict_row _DOLLAR_TAG = re.compile(r"\$[A-Za-z_]*\$") +_GO_LINE = re.compile(r"^\s*GO\s*(\d+)?\s*$", re.IGNORECASE) + + +def split_tsql_batches(sql: str) -> list[str]: + """Split T-SQL script text on standalone `GO` batch-separator lines (the + sqlcmd/SSMS convention T-SQL scripts conventionally use) -- `GO` is not + valid inside a single driver `execute()` call, unlike Postgres's `;` + statement separator which the driver handles natively. Tracks `/* ... */` + block comments as opaque so a `GO`-looking line inside a comment doesn't + split; does not attempt full tokenization of string literals spanning a + `GO` line, which is exceedingly rare in practice for schema/DDL scripts. + A script with no `GO` lines at all (any Postgres script, or a T-SQL one + that just doesn't use them) returns as a single batch, unchanged.""" + batches: list[str] = [] + buf: list[str] = [] + in_block_comment = False + for line in sql.splitlines(): + stripped = line.strip() + if in_block_comment: + buf.append(line) + if "*/" in line: + in_block_comment = False + continue + if stripped.startswith("/*") and "*/" not in stripped: + in_block_comment = True + buf.append(line) + continue + if _GO_LINE.match(line): + batch = "\n".join(buf).strip() + if batch: + batches.append(batch) + buf = [] + continue + buf.append(line) + tail = "\n".join(buf).strip() + if tail: + batches.append(tail) + return batches def _split_statements(sql: str) -> list[str]: diff --git a/pgdevkit/testdb/schema.py b/pgdevkit/testdb/schema.py index 3afe255..e958cbf 100644 --- a/pgdevkit/testdb/schema.py +++ b/pgdevkit/testdb/schema.py @@ -14,6 +14,7 @@ from psycopg.sql import SQL, Identifier, Placeholder from ..db.complex_types import ComplexHelper +from ..dialect import Dialect, POSTGRES logger = logging.getLogger(__name__) logging.getLogger("sqlglot").setLevel(logging.ERROR) @@ -63,6 +64,17 @@ def _strip_layer_prefix(schema_name: str) -> str: } +# Schemas that hold system catalog views/tables, never a file this project +# manages -- a reference to one is never a real cross-file dependency to +# wait for. Matters most for T-SQL, where "IF NOT EXISTS (SELECT ... FROM +# sys.schemas/sys.tables/sys.objects ...) BEGIN CREATE ... END" is the +# idiomatic idempotency-guard pattern (T-SQL has no native "CREATE TABLE IF +# NOT EXISTS"/"CREATE SCHEMA IF NOT EXISTS"), so without this exclusion +# nearly every T-SQL file would pick up a spurious, never-resolvable +# dependency on "sys.*" and get shuffled into the delayed-retry path, whose +# reverse-order resolution can then apply files out of their intended order. +_SYSTEM_SCHEMAS = {"pg_catalog", "information_schema", "sys"} + _DECLARE_RE = re.compile( r"CREATE\s+(?:OR\s+REPLACE\s+)?(?:TABLE|VIEW|FUNCTION|PROCEDURE|TYPE|SCHEMA)\s+(\w+\.\w+)", re.IGNORECASE ) @@ -76,16 +88,17 @@ def _get_sql_deps_regex_fallback(sql: str) -> set[str]: is fine.""" declares = set(_DECLARE_RE.findall(sql)) deps = set(_DEPEND_RE.findall(sql)) + deps = {d for d in deps if d.split(".", 1)[0].lower() not in _SYSTEM_SCHEMAS} return deps - declares -def _get_sql_deps(sql: str) -> set[str]: +def _get_sql_deps(sql: str, dialect: Dialect = POSTGRES) -> set[str]: try: # error_level=IGNORE lets sqlglot recover from statements it can't # fully parse (e.g. a schema-qualified `DROP TRIGGER ... ON # schema.table`) and keep going, instead of raising and losing every # other statement's dependency info in the same file. - exprs = sqlglot.parse(sql, dialect="postgres", error_level=sqlglot.ErrorLevel.IGNORE) + exprs = sqlglot.parse(sql, dialect=dialect.sqlglot_name, error_level=sqlglot.ErrorLevel.IGNORE) except Exception: # noqa: BLE001 return _get_sql_deps_regex_fallback(sql) deps: set[str] = set() @@ -93,7 +106,10 @@ def _get_sql_deps(sql: str) -> set[str]: if e is None: continue for t in e.find_all(exp.Table): - if t.args.get("this") is not None and t.args.get("db") is not None: + db_node = t.args.get("db") + if t.args.get("this") is not None and db_node is not None: + if db_node.name.lower() in _SYSTEM_SCHEMAS: + continue # exp.table_name(), not str(t): str() includes " AS alias" for # an aliased reference (e.g. "FROM editing.visit v"), which # would never match the plain declared name any dependent @@ -102,7 +118,7 @@ def _get_sql_deps(sql: str) -> set[str]: return deps -def _iter_sql_files(database_dir: Path): +def _iter_sql_files(database_dir: Path, dialect: Dialect = POSTGRES): """Yield (Path, sql_content) pairs in dependency-safe execution order.""" files: list[Path] = [] for root, _, dbfiles in os.walk(database_dir): @@ -120,7 +136,7 @@ def _iter_sql_files(database_dir: Path): for file in sorted(files, key=lambda p: (_get_type_order(p), p.name)): content = file.read_text(encoding="utf-8") - deps = _get_sql_deps(content) + deps = _get_sql_deps(content, dialect) if file.parent.name in _SCHEMA_QUALIFIED_TYPES: schema = _strip_layer_prefix(file.parent.parent.name) full_name = f"{schema}.{_strip_layer_prefix(file.stem)}" @@ -141,7 +157,7 @@ def _iter_sql_files(database_dir: Path): progressed = False for i in range(len(delayed) - 1, -1, -1): declared_name, file, content = delayed[i] - deps = _get_sql_deps(content) + deps = _get_sql_deps(content, dialect) if declared_name: deps.discard(declared_name) if all(d in delivered or d not in all_declared for d in deps): @@ -204,6 +220,8 @@ async def apply_schema( database_dir: Path, extensions: tuple[str, ...] = (), force_reset: bool = False, + *, + dialect: Dialect = POSTGRES, ) -> None: """Apply every .sql file under database_dir (in dependency-safe order) and seed any matching .test_data.json files. Safe to call repeatedly. @@ -227,7 +245,7 @@ async def _apply(file: Path, sql: str) -> None: await _insert_test_data(json_file, f"{schema_name}.{table_stem}", force_reset, con, complex_helper) failures: list[tuple[Path, str]] = [] - for file, sql in _iter_sql_files(database_dir): + for file, sql in _iter_sql_files(database_dir, dialect): try: await _apply(file, sql) except Exception as e: # noqa: BLE001 diff --git a/pyproject.toml b/pyproject.toml index e2308a1..58c3222 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -11,7 +11,7 @@ packages = ["pgdevkit"] [project] name = "pgdevkit" -version = "0.2.4" +version = "0.3.0" description = "A helper for developing with Postgres" readme = "README.md" requires-python = ">=3.14" @@ -31,6 +31,7 @@ db = [ "psycopg-pool>=3.3.0", "pydantic>=2.0", ] +mssql = ["mssql-python>=1.0.0"] [project.scripts] pgdb = "pgdevkit.cli:app" @@ -52,3 +53,4 @@ test = [ [tool.pytest.ini_options] pythonpath = ["."] asyncio_mode = "auto" +markers = ["mssql: requires a live MSSQL testdb container and the mssql extra (msodbcsql driver)"] diff --git a/tests/db/test_mssql_crud_live.py b/tests/db/test_mssql_crud_live.py new file mode 100644 index 0000000..22adea4 --- /dev/null +++ b/tests/db/test_mssql_crud_live.py @@ -0,0 +1,73 @@ +from __future__ import annotations + +import asyncio +from pathlib import Path + +import mssql_python +import pytest + +from pgdevkit.db.mssql_crud import mssql_delete_dict, mssql_insert, mssql_retrieve, mssql_update_dict, mssql_upsert_dict +from pgdevkit.db.model import TableModel +from pgdevkit.testdb.api import clean_testdb, ensure_testdb, status +# _make_project (not the project_factory fixture) -- project_factory lives +# in tests/testdb/conftest.py, whose fixtures pytest only injects into tests +# under tests/testdb/ itself. This file sits outside that directory, so it +# calls the same helper directly with its own tmp_path instead. +from tests.testdb.conftest import _make_project, requires_mssql + +pytestmark = pytest.mark.mssql + + +class Widget(TableModel): + id: int + name: str + + @staticmethod + def get_table_name() -> tuple[str, str]: + return ("app", "widget") + + @staticmethod + def get_primary_key() -> list[str]: + return ["id"] + + +@requires_mssql +def test_crud_round_trip(tmp_path: Path): + # A plain (not `async def`) test function -- `ensure_testdb`/ + # `clean_testdb` are sync wrappers that call `asyncio.run()` internally, + # which raises "cannot be called from a running event loop" if this test + # were itself async (pytest-asyncio already runs async tests inside a + # loop). The actual async CRUD calls get their own separate, sequential + # `asyncio.run()` below instead. + project = _make_project(tmp_path, "mssqlcrudlive", "main", engine="mssql") + try: + ensure_testdb(project) + conn = mssql_python.connect(status(project)["dsn"], autocommit=True) + + async def _run() -> None: + inserted = await mssql_insert(conn, ("app", "widget"), {"id": 2, "name": "cog"}) + assert inserted["name"] == "cog" + + fetched = await mssql_retrieve(conn, Widget, {"id": 2}) + assert fetched is not None and fetched.name == "cog" + + updated = await mssql_update_dict(conn, ("app", "widget"), {"id": 2, "name": "cog2"}, ["id"]) + assert updated is not None and updated["name"] == "cog2" + + upserted = await mssql_upsert_dict(conn, ("app", "widget"), {"id": 3, "name": "sprocket3"}, ["id"]) + assert upserted["name"] == "sprocket3" + upserted_again = await mssql_upsert_dict( + conn, ("app", "widget"), {"id": 3, "name": "sprocket3-updated"}, ["id"] + ) + assert upserted_again["name"] == "sprocket3-updated" + + deleted = await mssql_delete_dict(conn, ("app", "widget"), {"id": 2}) + assert deleted is not None and deleted["name"] == "cog2" + assert await mssql_retrieve(conn, Widget, {"id": 2}) is None + + try: + asyncio.run(_run()) + finally: + conn.close() + finally: + clean_testdb(project) diff --git a/tests/db/test_mssql_crud_sql.py b/tests/db/test_mssql_crud_sql.py new file mode 100644 index 0000000..3066c53 --- /dev/null +++ b/tests/db/test_mssql_crud_sql.py @@ -0,0 +1,85 @@ +from __future__ import annotations + +from pgdevkit.db.mssql_crud import ( + _build_delete, + _build_insert, + _build_insert_many, + _build_retrieve, + _build_retrieve_many, + _build_update, + _build_update_many, + _build_upsert_merge, +) +from pgdevkit.db.mssql_sql import ident, qualified + + +def test_ident_brackets_and_doubles_embedded_bracket(): + assert ident("widget") == "[widget]" + assert ident("weird]name") == "[weird]]name]" + + +def test_qualified_brackets_both_parts(): + assert qualified("dbo", "widget") == "[dbo].[widget]" + + +def test_build_retrieve_uses_qmark_placeholders(): + sql, params = _build_retrieve(("dbo", "widget"), {"id": 1}) + assert sql == "SELECT * FROM [dbo].[widget] WHERE [id] = ?" + assert params == [1] + + +def test_build_retrieve_many_without_filters_selects_all(): + sql, params = _build_retrieve_many(("dbo", "widget"), {}) + assert sql == "SELECT * FROM [dbo].[widget]" + assert params == [] + + +def test_build_retrieve_many_with_filters(): + sql, params = _build_retrieve_many(("dbo", "widget"), {"status": "active"}) + assert sql == "SELECT * FROM [dbo].[widget] WHERE [status] = ?" + assert params == ["active"] + + +def test_build_insert_uses_output_inserted_before_values(): + sql, params = _build_insert(("dbo", "widget"), {"name": "sprocket", "price": 9.99}) + assert sql == "INSERT INTO [dbo].[widget] ([name], [price]) OUTPUT INSERTED.* VALUES (?, ?)" + assert params == ["sprocket", 9.99] + + +def test_build_insert_many_has_no_output_clause(): + sql = _build_insert_many(("dbo", "widget"), ["name", "price"]) + assert sql == "INSERT INTO [dbo].[widget] ([name], [price]) VALUES (?, ?)" + + +def test_build_update_excludes_pk_from_set_clause_and_orders_params(): + sql, params = _build_update(("dbo", "widget"), {"id": 1, "name": "sprocket"}, ["id"]) + assert sql == "UPDATE [dbo].[widget] SET [name] = ? OUTPUT INSERTED.* WHERE [id] = ?" + assert params == ["sprocket", 1] + + +def test_build_update_many_has_no_output_and_qualifies_where_with_alias(): + sql = _build_update_many(("dbo", "widget"), ["id", "name"], ["id"]) + assert sql == "UPDATE [dbo].[widget] SET [name] = ? WHERE t.[id] = ?" + + +def test_build_upsert_merge_shape(): + sql, params = _build_upsert_merge(("dbo", "widget"), {"id": 1, "name": "sprocket"}, ["id"]) + assert sql.startswith("MERGE INTO [dbo].[widget] AS t USING (SELECT ? AS [id], ? AS [name]) AS s ") + assert "ON t.[id] = s.[id]" in sql + assert "WHEN MATCHED THEN UPDATE SET t.[name] = s.[name]" in sql + assert "WHEN NOT MATCHED THEN INSERT ([id], [name]) VALUES (s.[id], s.[name])" in sql + assert sql.rstrip().endswith("OUTPUT INSERTED.*;") + assert params == [1, "sprocket"] + + +def test_build_upsert_merge_with_only_pk_columns_has_no_matched_update_clause(): + sql, params = _build_upsert_merge(("dbo", "widget"), {"id": 1}, ["id"]) + assert "WHEN MATCHED THEN UPDATE" not in sql + assert "WHEN NOT MATCHED THEN INSERT ([id]) VALUES (s.[id])" in sql + assert params == [1] + + +def test_build_delete_uses_output_deleted_before_where(): + sql, params = _build_delete(("dbo", "widget"), {"id": 1}) + assert sql == "DELETE FROM [dbo].[widget] OUTPUT DELETED.* WHERE [id] = ?" + assert params == [1] diff --git a/tests/test_cli_compare.py b/tests/test_cli_compare.py index 46d4b8c..c2bf258 100644 --- a/tests/test_cli_compare.py +++ b/tests/test_cli_compare.py @@ -39,11 +39,11 @@ def fake_get_lakebase_password(workspace_host, instance_name): monkeypatch.setattr("pgdevkit.lakebase.get_lakebase_password", fake_get_lakebase_password) - def fake_introspect_db(conninfo): + def fake_introspect(self, conninfo): assert "LAKEBASE_TOKEN" in conninfo raise RuntimeError("stop after conninfo built — introspection itself isn't under test here") - monkeypatch.setattr("pgdevkit.cli.introspect_db", fake_introspect_db) + monkeypatch.setattr("pgdevkit.backends.postgres.PostgresBackend.introspect", fake_introspect) result = runner.invoke( app, diff --git a/tests/test_compare_mssql_live.py b/tests/test_compare_mssql_live.py new file mode 100644 index 0000000..ba7a90d --- /dev/null +++ b/tests/test_compare_mssql_live.py @@ -0,0 +1,39 @@ +from __future__ import annotations + +from pathlib import Path + +import pytest + +from pgdevkit.backends import get_backend +from pgdevkit.diff import compute_diff +from pgdevkit.parser import parse_directory +from pgdevkit.testdb.api import clean_testdb, ensure_testdb, status +# _make_project (not the project_factory fixture) -- project_factory lives +# in tests/testdb/conftest.py, whose fixtures pytest only injects into tests +# under tests/testdb/ itself. This file sits outside that directory, so it +# calls the same helper directly with its own tmp_path instead. +from tests.testdb.conftest import _make_project, requires_mssql + +# See tests/testdb/test_api_mssql_live.py's module docstring for why this is +# marked `mssql` (selected only by the dedicated CI job) rather than run +# everywhere. This test in particular is the main way the `sys.*` catalog +# queries in pgdevkit/mssql_introspect.py get validated against a real SQL +# Server at all -- there's no way to check their syntactic correctness +# offline. +pytestmark = pytest.mark.mssql + + +@requires_mssql +def test_introspection_matches_applied_scripts_with_no_diff(tmp_path: Path): + project = _make_project(tmp_path, "mssqlcompare", "main", engine="mssql") + try: + ensure_testdb(project) + conninfo = status(project)["dsn"] + + scripts = parse_directory(project / "database", dialect="mssql") + db = get_backend("mssql").introspect(conninfo) + diffs = compute_diff(scripts, db, dialect="mssql") + + assert diffs == [], diffs + finally: + clean_testdb(project) diff --git a/tests/test_dialect.py b/tests/test_dialect.py new file mode 100644 index 0000000..ea452dc --- /dev/null +++ b/tests/test_dialect.py @@ -0,0 +1,44 @@ +from __future__ import annotations + +import pytest + +from pgdevkit.dialect import MSSQL, POSTGRES, Dialect, resolve_dialect + + +def test_resolve_dialect_defaults_to_postgres(): + assert resolve_dialect() is POSTGRES + + +def test_resolve_dialect_by_name(): + assert resolve_dialect("postgres") is POSTGRES + assert resolve_dialect("mssql") is MSSQL + + +def test_resolve_dialect_passes_through_a_dialect_instance(): + assert resolve_dialect(MSSQL) is MSSQL + + +def test_resolve_dialect_rejects_unknown_name(): + with pytest.raises(ValueError, match="Unknown dialect"): + resolve_dialect("oracle") + + +def test_postgres_dialect_fields(): + assert POSTGRES.sqlglot_name == "postgres" + assert POSTGRES.default_schema == "public" + assert POSTGRES.supports_enums + assert POSTGRES.supports_composites + assert POSTGRES.type_synonyms["int4"] == "integer" + + +def test_mssql_dialect_fields(): + assert MSSQL.sqlglot_name == "tsql" + assert MSSQL.default_schema == "dbo" + assert not MSSQL.supports_enums + assert not MSSQL.supports_composites + assert MSSQL.type_synonyms["integer"] == "int" + + +def test_dialect_is_frozen(): + with pytest.raises(Exception): + POSTGRES.name = "mssql" # type: ignore[misc] diff --git a/tests/test_diff_mssql.py b/tests/test_diff_mssql.py new file mode 100644 index 0000000..14309aa --- /dev/null +++ b/tests/test_diff_mssql.py @@ -0,0 +1,60 @@ +from __future__ import annotations + +from pgdevkit.diff import DiffKind, _norm_type, _parse_index_def, compute_diff +from pgdevkit.dialect import MSSQL +from pgdevkit.models import ColumnDef, DatabaseSchema, IndexDef, TableDef + + +def _table(schema: str, name: str, **cols: tuple[str, bool]) -> TableDef: + t = TableDef(schema=schema, name=name) + for cname, (dtype, nullable) in cols.items(): + t.columns.append(ColumnDef(name=cname, data_type=dtype, is_nullable=nullable, default=None)) + return t + + +def test_norm_type_applies_mssql_synonyms(): + assert _norm_type("integer", MSSQL) == _norm_type("int", MSSQL) + assert _norm_type("numeric(10,2)", MSSQL) == _norm_type("decimal(10,2)", MSSQL) + + +def test_norm_type_mssql_synonyms_dont_leak_into_postgres_normalization(): + # "integer" is already canonical for Postgres and must not be rewritten + # to "int" there -- the two dialects' synonym tables are independent. + assert _norm_type("integer") == "integer" + + +def test_compute_diff_reports_missing_table_for_mssql(): + scripts = DatabaseSchema(tables={"dbo.widget": _table("dbo", "widget", id=("int", False))}) + db = DatabaseSchema() + diffs = compute_diff(scripts, db, dialect="mssql") + assert any(d.kind == DiffKind.MISSING_IN_DB and d.object_type == "table" and d.object_name == "dbo.widget" for d in diffs) + + +def test_compute_diff_column_type_synonym_insensitive_for_mssql(): + scripts = DatabaseSchema(tables={"dbo.widget": _table("dbo", "widget", n=("numeric(10,2)", False))}) + db = DatabaseSchema(tables={"dbo.widget": _table("dbo", "widget", n=("decimal(10,2)", False))}) + diffs = compute_diff(scripts, db, dialect="mssql") + assert diffs == [] + + +def test_compute_diff_column_type_mismatch_for_mssql(): + scripts = DatabaseSchema(tables={"dbo.widget": _table("dbo", "widget", n=("int", False))}) + db = DatabaseSchema(tables={"dbo.widget": _table("dbo", "widget", n=("nvarchar(50)", False))}) + diffs = compute_diff(scripts, db, dialect="mssql") + assert any(d.kind == DiffKind.MISMATCH and d.object_type == "column" for d in diffs) + + +def test_parse_index_def_parses_tsql_filtered_index(): + definition = "create unique index ix_widget_name on dbo.widget (name) where name is not null" + info = _parse_index_def(definition, MSSQL) + assert info is not None + assert info["unique"] is True + assert info["columns"] == ["name"] + + +def test_diff_index_matches_equivalent_tsql_definitions(): + idx = IndexDef(schema="dbo", table="widget", name="ix_widget_name", definition="create index ix_widget_name on dbo.widget (name)") + scripts = DatabaseSchema(indexes={idx.qualified_name: idx}) + db = DatabaseSchema(indexes={idx.qualified_name: idx}) + diffs = compute_diff(scripts, db, dialect="mssql") + assert diffs == [] diff --git a/tests/test_mssql_introspect.py b/tests/test_mssql_introspect.py new file mode 100644 index 0000000..cdb59b9 --- /dev/null +++ b/tests/test_mssql_introspect.py @@ -0,0 +1,18 @@ +from __future__ import annotations + +from pgdevkit.mssql_introspect import _extract_view_query + + +def test_extract_view_query_strips_create_or_alter_view_prefix(): + # Regression test for a real bug caught in CI: sys.sql_modules.definition + # is the verbatim CREATE VIEW statement text, unlike Postgres's + # pg_get_viewdef() which returns only the query body -- storing it + # as-is produced a spurious "definition differs" diff against every + # script-defined view, since parser.py's side is query-only. + definition = "CREATE OR ALTER VIEW app.b_base_view AS\nSELECT id, name FROM app.widget;" + assert _extract_view_query(definition) == "SELECT id, name FROM app.widget" + + +def test_extract_view_query_falls_back_to_raw_text_on_unparseable_input(): + garbage = "not a valid create view statement !!!" + assert _extract_view_query(garbage) == garbage diff --git a/tests/test_parser_mssql.py b/tests/test_parser_mssql.py new file mode 100644 index 0000000..dd1d8e9 --- /dev/null +++ b/tests/test_parser_mssql.py @@ -0,0 +1,92 @@ +from __future__ import annotations + +from pathlib import Path + +from pgdevkit.parser import _parse_function_details_tsql, parse_directory + + +def _write(tmp_path: Path, name: str, sql: str) -> None: + (tmp_path / name).write_text(sql, encoding="utf-8") + + +def test_parses_tsql_table_with_identity_and_types(tmp_path: Path): + _write( + tmp_path, + "widget.sql", + """ + CREATE TABLE dbo.widget ( + id INT IDENTITY(1,1) PRIMARY KEY, + name NVARCHAR(100) NOT NULL, + price DECIMAL(10,2) NULL + ); + """, + ) + schema = parse_directory(tmp_path, dialect="mssql") + table = schema.tables["dbo.widget"] + cols = {c.name: c for c in table.columns} + assert cols["id"].is_generated + assert not cols["id"].is_nullable + assert not cols["name"].is_nullable + assert cols["price"].is_nullable + + +def test_unqualified_table_falls_back_to_dbo_schema(tmp_path: Path): + _write(tmp_path, "widget.sql", "CREATE TABLE widget (id INT PRIMARY KEY);") + schema = parse_directory(tmp_path, dialect="mssql") + assert "dbo.widget" in schema.tables + + +def test_unqualified_table_falls_back_to_public_schema_for_postgres(tmp_path: Path): + _write(tmp_path, "widget.sql", "CREATE TABLE widget (id int PRIMARY KEY);") + schema = parse_directory(tmp_path) # default dialect + assert "public.widget" in schema.tables + + +def test_parses_tsql_view(tmp_path: Path): + _write( + tmp_path, + "widget_view.sql", + "CREATE VIEW dbo.widget_view AS SELECT id, name FROM dbo.widget;", + ) + schema = parse_directory(tmp_path, dialect="mssql") + assert "dbo.widget_view" in schema.views + + +def test_parses_tsql_function_with_returns_and_body(tmp_path: Path): + _write( + tmp_path, + "get_price.sql", + """ + CREATE FUNCTION dbo.get_price (@id INT) + RETURNS DECIMAL(10,2) + AS + BEGIN + RETURN 9.99; + END; + """, + ) + schema = parse_directory(tmp_path, dialect="mssql") + func = schema.functions["dbo.get_price"] + assert func.language == "sql" + assert "@id" in func.args + assert func.return_type == "decimal(10,2)" + assert "return 9.99" in func.body + + +def test_parse_function_details_tsql_extracts_args_and_body_directly(): + sql = "CREATE FUNCTION dbo.f (@a INT, @b INT) RETURNS INT AS BEGIN RETURN @a + @b; END;" + args, return_type, language, body = _parse_function_details_tsql(sql) + assert args == "@a int, @b int" + assert return_type == "int" + assert language == "sql" + assert "return @a + @b" in body + + +def test_mssql_enum_style_type_is_ignored_not_erroring(tmp_path: Path): + # A Postgres-style `CREATE TYPE ... AS ENUM` has no T-SQL equivalent -- + # parsing it under dialect="mssql" must not raise, and (since sqlglot's + # tsql dialect won't produce the ENUM/Schema AST shapes _handle_type + # matches on) it also shouldn't populate db.enums. + _write(tmp_path, "mood.sql", "CREATE TYPE mood AS ENUM ('happy', 'sad');") + schema = parse_directory(tmp_path, dialect="mssql") + assert schema.enums == {} diff --git a/tests/testdb/conftest.py b/tests/testdb/conftest.py index 790bf5c..ff75818 100644 --- a/tests/testdb/conftest.py +++ b/tests/testdb/conftest.py @@ -22,7 +22,22 @@ def _has_container_runtime() -> bool: not _has_container_runtime(), reason="no Docker-compatible API reachable" ) + +def _has_mssql_support() -> bool: + try: + import mssql_python # noqa: F401 + except ImportError: + return False + return _has_container_runtime() + + +requires_mssql = pytest.mark.skipif( + not _has_mssql_support(), + reason="no Docker-compatible API reachable, or the mssql extra isn't installed", +) + FIXTURES = Path(__file__).parent / "fixtures" / "database" +MSSQL_FIXTURES = Path(__file__).parent / "fixtures" / "database_mssql" # Appended to every test project's [tool.pgdevkit].name so that two pytest # processes (e.g. from separate git worktrees) running against the shared @@ -31,12 +46,13 @@ def _has_container_runtime() -> bool: RUN_SUFFIX = f"pid{os.getpid()}" -def _make_project(base: Path, name: str, branch: str) -> Path: +def _make_project(base: Path, name: str, branch: str, engine: str = "postgres") -> Path: project = base / f"{name}-{branch}" project.mkdir() - (project / "database").symlink_to(FIXTURES) + (project / "database").symlink_to(MSSQL_FIXTURES if engine == "mssql" else FIXTURES) + engine_line = f'engine = "{engine}"\n' if engine != "postgres" else "" (project / "pyproject.toml").write_text( - f'[tool.pgdevkit]\nname = "{name}_{RUN_SUFFIX}"\nenv_prefix = "{name.upper()}_"\n', + f'[tool.pgdevkit]\nname = "{name}_{RUN_SUFFIX}"\nenv_prefix = "{name.upper()}_"\n{engine_line}', encoding="utf-8", ) for cmd in ( @@ -53,8 +69,8 @@ def _make_project(base: Path, name: str, branch: str) -> Path: @pytest.fixture -def project_factory(tmp_path: Path) -> Callable[[str, str], Path]: - def _factory(name: str, branch: str) -> Path: - return _make_project(tmp_path, name, branch) +def project_factory(tmp_path: Path) -> Callable[..., Path]: + def _factory(name: str, branch: str, engine: str = "postgres") -> Path: + return _make_project(tmp_path, name, branch, engine) return _factory diff --git a/tests/testdb/fixtures/database_mssql/app/tables/widget.sql b/tests/testdb/fixtures/database_mssql/app/tables/widget.sql new file mode 100644 index 0000000..1857bf5 --- /dev/null +++ b/tests/testdb/fixtures/database_mssql/app/tables/widget.sql @@ -0,0 +1,7 @@ +IF NOT EXISTS (SELECT 1 FROM sys.tables t JOIN sys.schemas s ON s.schema_id = t.schema_id WHERE s.name = 'app' AND t.name = 'widget') +BEGIN + CREATE TABLE app.widget ( + id INT PRIMARY KEY, + name NVARCHAR(100) NOT NULL + ); +END diff --git a/tests/testdb/fixtures/database_mssql/app/tables/widget.test_data.json b/tests/testdb/fixtures/database_mssql/app/tables/widget.test_data.json new file mode 100644 index 0000000..dbd8dc2 --- /dev/null +++ b/tests/testdb/fixtures/database_mssql/app/tables/widget.test_data.json @@ -0,0 +1 @@ +[{"id": 1, "name": "sprocket"}] diff --git a/tests/testdb/fixtures/database_mssql/app/views/a_wrapper_view.sql b/tests/testdb/fixtures/database_mssql/app/views/a_wrapper_view.sql new file mode 100644 index 0000000..5cc0f72 --- /dev/null +++ b/tests/testdb/fixtures/database_mssql/app/views/a_wrapper_view.sql @@ -0,0 +1,2 @@ +CREATE OR ALTER VIEW app.a_wrapper_view AS +SELECT id, name FROM app.b_base_view; diff --git a/tests/testdb/fixtures/database_mssql/app/views/b_base_view.sql b/tests/testdb/fixtures/database_mssql/app/views/b_base_view.sql new file mode 100644 index 0000000..d7387af --- /dev/null +++ b/tests/testdb/fixtures/database_mssql/app/views/b_base_view.sql @@ -0,0 +1,2 @@ +CREATE OR ALTER VIEW app.b_base_view AS +SELECT id, name FROM app.widget; diff --git a/tests/testdb/fixtures/database_mssql/schema/app.sql b/tests/testdb/fixtures/database_mssql/schema/app.sql new file mode 100644 index 0000000..79904a8 --- /dev/null +++ b/tests/testdb/fixtures/database_mssql/schema/app.sql @@ -0,0 +1,4 @@ +IF NOT EXISTS (SELECT 1 FROM sys.schemas WHERE name = 'app') +BEGIN + EXEC('CREATE SCHEMA app'); +END diff --git a/tests/testdb/test_api_mssql.py b/tests/testdb/test_api_mssql.py new file mode 100644 index 0000000..48cc268 --- /dev/null +++ b/tests/testdb/test_api_mssql.py @@ -0,0 +1,58 @@ +from __future__ import annotations + +import subprocess +from pathlib import Path + +from pgdevkit.testdb.api import shell_argv, status + + +def _make_mssql_project(base: Path, name: str) -> Path: + project = base / name + project.mkdir() + (project / "pyproject.toml").write_text( + f'[tool.pgdevkit]\nname = "{name}"\nengine = "mssql"\n', encoding="utf-8" + ) + for cmd in ( + ["git", "init", "-q"], + ["git", "config", "user.email", "test@example.com"], + ["git", "config", "user.name", "test"], + ): + subprocess.run(cmd, cwd=project, check=True) + (project / ".gitkeep").write_text("", encoding="utf-8") + subprocess.run(["git", "add", "."], cwd=project, check=True) + subprocess.run(["git", "commit", "-q", "-m", "init"], cwd=project, check=True) + return project + + +def test_status_dispatches_to_mssql_and_reports_engine(tmp_path: Path): + project = _make_mssql_project(tmp_path, "mssqlproj") + info = status(project) + assert info["engine"] == "mssql" + assert "container" in info and "dsn" in info + + +def test_shell_argv_uses_sqlcmd_for_mssql(tmp_path: Path): + project = _make_mssql_project(tmp_path, "mssqlproj2") + binary, argv = shell_argv(project) + assert binary == "sqlcmd" + assert argv[0] == "sqlcmd" + assert "-S" in argv and "-d" in argv + + +def test_shell_argv_uses_psql_for_postgres(tmp_path: Path): + project = tmp_path / "pgproj" + project.mkdir() + (project / "pyproject.toml").write_text('[tool.pgdevkit]\nname = "pgproj"\n', encoding="utf-8") + for cmd in ( + ["git", "init", "-q"], + ["git", "config", "user.email", "test@example.com"], + ["git", "config", "user.name", "test"], + ): + subprocess.run(cmd, cwd=project, check=True) + (project / ".gitkeep").write_text("", encoding="utf-8") + subprocess.run(["git", "add", "."], cwd=project, check=True) + subprocess.run(["git", "commit", "-q", "-m", "init"], cwd=project, check=True) + + binary, argv = shell_argv(project) + assert binary == "psql" + assert argv[0] == "psql" diff --git a/tests/testdb/test_api_mssql_live.py b/tests/testdb/test_api_mssql_live.py new file mode 100644 index 0000000..377268b --- /dev/null +++ b/tests/testdb/test_api_mssql_live.py @@ -0,0 +1,66 @@ +from __future__ import annotations + +from typing import Callable + +import mssql_python +import pytest +from pathlib import Path + +from pgdevkit.testdb.api import clean_testdb, ensure_testdb, status +from tests.testdb.conftest import requires_mssql + +# Selects these (live, container-requiring) tests in the dedicated GitHub +# Actions job -- see .github/workflows/python-test.yml's `mssql-test` job, +# which runs `pytest -m mssql`; the main `build` job runs `-m "not mssql"` +# so a live SQL Server never blocks the fast Postgres-only suite. +pytestmark = pytest.mark.mssql + + +def _query(dsn: str, sql: str) -> list[tuple]: + conn = mssql_python.connect(dsn) + try: + cur = conn.cursor() + cur.execute(sql) + return [tuple(r) for r in cur.fetchall()] + finally: + conn.close() + + +@requires_mssql +def test_ensure_testdb_applies_schema_and_seeds_data(project_factory: Callable[..., Path]): + project = project_factory("mssqllive", "main", engine="mssql") + try: + env = ensure_testdb(project) + assert any(k.endswith("_MSSQL_DB") for k in env) + + rows = _query(status(project)["dsn"], "SELECT id, name FROM app.widget ORDER BY id") + assert rows == [(1, "sprocket")] + finally: + clean_testdb(project) + + +@requires_mssql +def test_ensure_testdb_is_idempotent(project_factory: Callable[..., Path]): + project = project_factory("mssqllive2", "main", engine="mssql") + try: + ensure_testdb(project) + ensure_testdb(project) # must not raise + + rows = _query(status(project)["dsn"], "SELECT count(*) FROM app.widget") + assert rows == [(1,)] + finally: + clean_testdb(project) + + +@requires_mssql +def test_ensure_testdb_resolves_view_to_view_dependency(project_factory: Callable[..., Path]): + # a_wrapper_view selects from b_base_view -- only passes if the + # dependency-ordering logic in schema.py (shared with the Postgres path, + # dialect-parametrized in this PR) also works for T-SQL scripts. + project = project_factory("mssqllive3", "main", engine="mssql") + try: + ensure_testdb(project) + rows = _query(status(project)["dsn"], "SELECT id, name FROM app.a_wrapper_view ORDER BY id") + assert rows == [(1, "sprocket")] + finally: + clean_testdb(project) diff --git a/tests/testdb/test_config_mssql.py b/tests/testdb/test_config_mssql.py new file mode 100644 index 0000000..ed419bc --- /dev/null +++ b/tests/testdb/test_config_mssql.py @@ -0,0 +1,43 @@ +from __future__ import annotations + +from pathlib import Path + +import pytest + +from pgdevkit.testdb.config import load_config + + +def test_engine_defaults_to_postgres_when_absent(tmp_path: Path): + config = load_config(tmp_path) + assert config.engine == "postgres" + + +def test_engine_read_from_pyproject(tmp_path: Path): + (tmp_path / "pyproject.toml").write_text( + '[tool.pgdevkit]\nname = "x"\nengine = "mssql"\n', encoding="utf-8" + ) + config = load_config(tmp_path) + assert config.engine == "mssql" + + +def test_engine_rejects_unknown_value(tmp_path: Path): + (tmp_path / "pyproject.toml").write_text( + '[tool.pgdevkit]\nname = "x"\nengine = "oracle"\n', encoding="utf-8" + ) + with pytest.raises(ValueError, match="engine"): + load_config(tmp_path) + + +def test_env_var_overrides_pyproject_engine(tmp_path: Path, monkeypatch): + (tmp_path / "pyproject.toml").write_text( + '[tool.pgdevkit]\nname = "x"\nengine = "postgres"\n', encoding="utf-8" + ) + monkeypatch.setenv("PGDEVKIT_TESTDB_ENGINE", "mssql") + config = load_config(tmp_path) + assert config.engine == "mssql" + + +def test_env_var_used_when_no_pyproject_value(tmp_path: Path, monkeypatch): + monkeypatch.setenv("PGDEVKIT_TESTDB_ENGINE", "mssql") + config = load_config(tmp_path) + assert config.engine == "mssql" diff --git a/tests/testdb/test_mssql_constants.py b/tests/testdb/test_mssql_constants.py new file mode 100644 index 0000000..f258056 --- /dev/null +++ b/tests/testdb/test_mssql_constants.py @@ -0,0 +1,31 @@ +from __future__ import annotations + +import pytest + +from pgdevkit.testdb.mssql.constants import conninfo, validate_sa_password + + +def test_validate_sa_password_accepts_complexity_valid_password(): + validate_sa_password("TestPwd!2026") # must not raise + + +def test_validate_sa_password_rejects_too_short(): + with pytest.raises(ValueError, match="8 characters"): + validate_sa_password("Ab1!") + + +def test_validate_sa_password_rejects_containing_login_name(): + with pytest.raises(ValueError, match="'sa'"): + validate_sa_password("Sup3rSaPass!") + + +def test_validate_sa_password_rejects_insufficient_character_classes(): + with pytest.raises(ValueError, match="3 of"): + validate_sa_password("alllowercase") + + +def test_conninfo_omits_driver_and_app_keys(): + cs = conninfo("mydb") + assert "Driver=" not in cs + assert "APP=" not in cs + assert "Database=mydb" in cs diff --git a/tests/testdb/test_query.py b/tests/testdb/test_query.py index a317657..30d908b 100644 --- a/tests/testdb/test_query.py +++ b/tests/testdb/test_query.py @@ -5,7 +5,7 @@ from pgdevkit.testdb import constants, query from pgdevkit.testdb.container import ensure_container -from pgdevkit.testdb.query import _split_statements +from pgdevkit.testdb.query import _split_statements, split_tsql_batches from tests.testdb.conftest import RUN_SUFFIX, requires_podman @@ -33,6 +33,32 @@ def test_split_statements_handles_tagged_dollar_quotes(): statements = _split_statements("SELECT $tag$a;b$tag$ AS x; SELECT 2") assert statements == ["SELECT $tag$a;b$tag$ AS x", "SELECT 2"] +def test_split_tsql_batches_splits_on_standalone_go_line(): + sql = "CREATE TABLE dbo.widget (id INT);\nGO\nINSERT INTO dbo.widget VALUES (1);\nGO\n" + batches = split_tsql_batches(sql) + assert len(batches) == 2 + assert batches[0] == "CREATE TABLE dbo.widget (id INT);" + assert batches[1] == "INSERT INTO dbo.widget VALUES (1);" + + +def test_split_tsql_batches_with_no_go_returns_single_batch(): + sql = "SELECT 1;\nSELECT 2;" + assert split_tsql_batches(sql) == [sql] + + +def test_split_tsql_batches_ignores_go_inside_block_comment(): + sql = "SELECT 1;\n/* remember:\nGO\ndo something */\nSELECT 2;" + batches = split_tsql_batches(sql) + assert len(batches) == 1 + assert "GO" in batches[0] + + +def test_split_tsql_batches_accepts_go_with_repeat_count(): + sql = "SELECT 1;\nGO 3\nSELECT 2;" + batches = split_tsql_batches(sql) + assert batches == ["SELECT 1;", "SELECT 2;"] + + TEST_DB = f"pgdevkit_query_selftest_{RUN_SUFFIX}" diff --git a/tests/testdb/test_schema_mssql.py b/tests/testdb/test_schema_mssql.py new file mode 100644 index 0000000..a0a0173 --- /dev/null +++ b/tests/testdb/test_schema_mssql.py @@ -0,0 +1,35 @@ +from __future__ import annotations + +from pathlib import Path + +from pgdevkit.dialect import MSSQL, POSTGRES +from pgdevkit.testdb.schema import _get_sql_deps, _iter_sql_files + +FIXTURES = Path(__file__).parent / "fixtures" / "database_mssql" + + +def test_get_sql_deps_excludes_sys_schema_references(): + # T-SQL has no native "CREATE SCHEMA IF NOT EXISTS", so scripts commonly + # guard with "IF NOT EXISTS (SELECT ... FROM sys.schemas ...)" -- a + # reference to sys.schemas/sys.tables must never count as a real + # cross-file dependency, since no file ever "delivers" it. Regression + # test for a real bug caught in CI: this previously delayed such files + # into the retry loop, whose reverse-order resolution then applied a + # table file before the schema.sql file it actually depended on. + sql = "IF NOT EXISTS (SELECT 1 FROM sys.schemas WHERE name = 'app') BEGIN EXEC('CREATE SCHEMA app'); END" + assert _get_sql_deps(sql, MSSQL) == set() + + +def test_get_sql_deps_excludes_pg_catalog_references_for_postgres_too(): + sql = "SELECT 1 FROM pg_catalog.pg_class WHERE relname = 'widget'" + assert _get_sql_deps(sql, POSTGRES) == set() + + +def test_iter_sql_files_applies_schema_before_dependent_table_and_views(): + order = [f.relative_to(FIXTURES) for f, _ in _iter_sql_files(FIXTURES, MSSQL)] + assert order == [ + Path("schema/app.sql"), + Path("app/tables/widget.sql"), + Path("app/views/b_base_view.sql"), + Path("app/views/a_wrapper_view.sql"), + ] diff --git a/uv.lock b/uv.lock index 0fd4b48..c610b87 100644 --- a/uv.lock +++ b/uv.lock @@ -270,6 +270,38 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/5e/75/bd9b7bb966668920f06b200e84454c8f3566b102183bc55c5473d96cb2b9/msal_extensions-1.3.1-py3-none-any.whl", hash = "sha256:96d3de4d034504e969ac5e85bae8106c8373b5c6568e4c8fa7af2eca9dbe6bca", size = 20583, upload-time = "2025-03-14T23:51:03.016Z" }, ] +[[package]] +name = "mssql-python" +version = "1.12.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "azure-identity" }, + { name = "mssql-python-odbc" }, +] +wheels = [ + { url = "https://files.pythonhosted.org/packages/fa/cb/44cc2da0b4351da84692cec02c90ad6df5ee1bf8652e18fb15b37330968b/mssql_python-1.12.0-cp314-cp314-macosx_15_0_universal2.whl", hash = "sha256:6ba31f455c572abe897bdfc85b54a71645aaf2e53f4bda9e3945b53c2e4f4c7a", size = 28347319, upload-time = "2026-07-24T14:22:15.425Z" }, + { url = "https://files.pythonhosted.org/packages/42/a3/202b2dd5ce5e9ce235caee2c9d61df1030671282cceeba984b63fbb4b23f/mssql_python-1.12.0-cp314-cp314-manylinux_2_28_aarch64.whl", hash = "sha256:6c7aeacb4d38b400e30c78d50e9c1caff40eb6ee1373972822e81763b91ed0b2", size = 27348887, upload-time = "2026-07-24T14:22:18.802Z" }, + { url = "https://files.pythonhosted.org/packages/a1/58/b496e185f2b0978c1ef5e6646eaaa6c203cbd26918820d57d6362c3d2241/mssql_python-1.12.0-cp314-cp314-manylinux_2_28_x86_64.whl", hash = "sha256:c18a1cfdb3d919a840bfe0afb166b7637891f7e7cc303b4f7237baa3a8b2b2c7", size = 28278517, upload-time = "2026-07-24T14:22:21.704Z" }, + { url = "https://files.pythonhosted.org/packages/03/a0/ea25a02928b5723c9e0755d763c3fc72969801a6a7f982ae363ec3bfd25a/mssql_python-1.12.0-cp314-cp314-musllinux_1_2_aarch64.whl", hash = "sha256:f13f0a5aae382b47e04ab0a2a0a32d45ad36656d60bfe363789bdbf11d3acdd6", size = 27031078, upload-time = "2026-07-24T14:22:24.533Z" }, + { url = "https://files.pythonhosted.org/packages/08/04/cc355a09333706b40c76bb24389f180c12151b55ad7f28b11b66e0b79f5e/mssql_python-1.12.0-cp314-cp314-musllinux_1_2_x86_64.whl", hash = "sha256:2f78fc4e937c0addba591f24cd8e6dc4c725304bf5a4d4da478173c18ed2e826", size = 27891294, upload-time = "2026-07-24T14:22:27.212Z" }, + { url = "https://files.pythonhosted.org/packages/73/5b/8bd2d295422f88989d186d443b2392fedd32b58e212f07fa5d6da6fbda32/mssql_python-1.12.0-cp314-cp314-win_amd64.whl", hash = "sha256:670d877e3ba7cacb286ccfa9acf0fb06e518caf99573468d7a5476a70c5d97dd", size = 16113876, upload-time = "2026-07-24T14:22:29.604Z" }, + { url = "https://files.pythonhosted.org/packages/76/b5/bebc8f46a96565693fc622872e6350ecfea98a890330d23371111c59be8e/mssql_python-1.12.0-cp314-cp314-win_arm64.whl", hash = "sha256:b5b0eead5c00e2626cdbcec485779d15700bcf61871b030a91530526ce533b42", size = 19429847, upload-time = "2026-07-24T14:22:31.962Z" }, +] + +[[package]] +name = "mssql-python-odbc" +version = "18.6.2" +source = { registry = "https://pypi.org/simple" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/9e/52/871699a1e6162d1afb4630fc9630ef5c59abeb10d14762c14618b51ffc35/mssql_python_odbc-18.6.2-py3-none-macosx_15_0_universal2.whl", hash = "sha256:518d2c290a6edf89e6cbfded4bf598a0f5ec67f3ce171cf1a755e463397e957a", size = 2016976, upload-time = "2026-07-22T10:08:10.154Z" }, + { url = "https://files.pythonhosted.org/packages/64/f1/8ca3965c3b5488ce6e7a72079ae9f3477c431782ea2efa9f33b65fc03deb/mssql_python_odbc-18.6.2-py3-none-manylinux_2_28_aarch64.whl", hash = "sha256:26f34eeb0c199e82f14c89e1e697b717861cc64971e05855b2bd10efd39fe949", size = 2711629, upload-time = "2026-07-22T10:08:11.781Z" }, + { url = "https://files.pythonhosted.org/packages/75/2d/602d7d194f62cbeb78782efce5c431c168ad557dbd570a1bf6dca25d8d46/mssql_python_odbc-18.6.2-py3-none-manylinux_2_28_x86_64.whl", hash = "sha256:508e4185d486191aad5e659ec37413a3811e52e5d59250d57e6382e787f6652d", size = 3879478, upload-time = "2026-07-22T10:08:13.704Z" }, + { url = "https://files.pythonhosted.org/packages/96/3c/87e27dc3b5454d548c19ba906ae207378ad358fa82298ce423d3bcb7ede3/mssql_python_odbc-18.6.2-py3-none-musllinux_1_2_aarch64.whl", hash = "sha256:6f687645a80aedda238eee7996cbda0527f93f3d7948453057766ed4e199de60", size = 2711629, upload-time = "2026-07-22T10:08:15.146Z" }, + { url = "https://files.pythonhosted.org/packages/0f/24/c7d1a8da1a88fa4bcd0c1c36c7f259c7d0417d63731d69e358c1f5338c80/mssql_python_odbc-18.6.2-py3-none-musllinux_1_2_x86_64.whl", hash = "sha256:12ea164280e1b5492ac8f90b8780fbd0e3a95c1a0d77f6c1dd92f08b29a63a16", size = 3879474, upload-time = "2026-07-22T10:08:16.579Z" }, + { url = "https://files.pythonhosted.org/packages/1f/36/9e23362a0c4a24b000c954273900b1d722e1106e94f9e1224c730f454e62/mssql_python_odbc-18.6.2-py3-none-win_amd64.whl", hash = "sha256:99b734f0479f8536b185f7f74a4c7ab594931546225f9bf6a39d4c0bf8d837b0", size = 3693913, upload-time = "2026-07-22T10:08:17.962Z" }, + { url = "https://files.pythonhosted.org/packages/18/c0/bff06bc2ef7647fddc2f568239383a259d85b5b236d57522fb389780ffd5/mssql_python_odbc-18.6.2-py3-none-win_arm64.whl", hash = "sha256:2ab8057e1565a155ae740b5232c161cb0d9439ce48001a8998e4422f6a02c601", size = 6998918, upload-time = "2026-07-22T10:08:19.503Z" }, +] + [[package]] name = "packaging" version = "26.2" @@ -281,7 +313,7 @@ wheels = [ [[package]] name = "pgdevkit" -version = "0.2.4" +version = "0.3.0" source = { editable = "." } dependencies = [ { name = "docker" }, @@ -301,6 +333,9 @@ db = [ { name = "psycopg-pool" }, { name = "pydantic" }, ] +mssql = [ + { name = "mssql-python" }, +] [package.dev-dependencies] dev = [ @@ -317,6 +352,7 @@ test = [ requires-dist = [ { name = "azure-identity", marker = "extra == 'azure'", specifier = ">=1.19.0" }, { name = "docker", specifier = ">=7.1.0" }, + { name = "mssql-python", marker = "extra == 'mssql'", specifier = ">=1.0.0" }, { name = "psycopg", extras = ["binary"], specifier = ">=3.2.0" }, { name = "psycopg-pool", marker = "extra == 'db'", specifier = ">=3.3.0" }, { name = "pydantic", marker = "extra == 'db'", specifier = ">=2.0" }, @@ -324,7 +360,7 @@ requires-dist = [ { name = "sqlglot", specifier = ">=30.11.0" }, { name = "typer", marker = "extra == 'cli'", specifier = ">=0.26.7" }, ] -provides-extras = ["azure", "cli", "db"] +provides-extras = ["azure", "cli", "db", "mssql"] [package.metadata.requires-dev] dev = [{ name = "ty", specifier = ">=0.0.59" }]