From a5a5eb2bc4f37187ce40b3ee5ef6f3a696f13bab Mon Sep 17 00:00:00 2001 From: Adrian Ehrsam Date: Mon, 5 Oct 2026 09:03:34 +0200 Subject: [PATCH 1/3] WIP: MSSQL support for stats (#43) Co-Authored-By: Claude Sonnet 5.5 Claude-Session: https://claude.ai/code/session_01SwHJR2uryShYQnsySXZT6G From cf96e5e2672c4c87053b3ddadad2fbc24f65363a Mon Sep 17 00:00:00 2001 From: Adrian Ehrsam Date: Mon, 5 Oct 2026 09:05:00 +0200 Subject: [PATCH 2/3] Support MSSQL for update-stats via --dialect mssql (#43) Co-Authored-By: Claude Sonnet 5.5 Claude-Session: https://claude.ai/code/session_01SwHJR2uryShYQnsySXZT6G --- pgdevkit/cli.py | 6 +- pgdevkit/stats.py | 115 ++++++++++++++++++++++++++++++++++++++- skills/pgdevkit/SKILL.md | 2 + tests/test_stats_cli.py | 61 +++++++++++++++++++++ 4 files changed, 182 insertions(+), 2 deletions(-) diff --git a/pgdevkit/cli.py b/pgdevkit/cli.py index a848677..821e818 100644 --- a/pgdevkit/cli.py +++ b/pgdevkit/cli.py @@ -227,6 +227,7 @@ def update_stats( table: list[str] = typer.Option([], "--table", help="Only update schema.name (repeatable); default is every table"), exact: bool = typer.Option(False, "--exact", help="Use count(*) for row counts instead of the planner estimate"), analyze: bool = typer.Option(False, "--analyze", help="Run ANALYZE first so column stats are fresh"), + dialect: str = typer.Option("postgres", "--dialect", help="postgres (default) or mssql"), ) -> None: """Store table stats in _stats/_tables.json (keyed by schema.name, sorted) and column stats in _stats/.json.""" @@ -236,10 +237,13 @@ def update_stats( conninfo = build_conninfo(url, entra_user) try: with console.status("Collecting stats..."): - tables, columns = stats.collect_stats(conninfo, set(table) or None, exact=exact, analyze=analyze) + tables, columns = stats.collect_stats(conninfo, set(table) or None, exact=exact, analyze=analyze, dialect=dialect) except KeyError as e: err_console.print(f"[red]Error:[/red] unknown table(s): {e.args[0]}") raise typer.Exit(2) + except ValueError as e: + err_console.print(f"[red]Error:[/red] {e}") + raise typer.Exit(2) path = stats.write_stats(scripts_dir, tables, columns, prune=not table) console.print(f"Updated stats for {len(tables)} table(s) in {path.parent}") diff --git a/pgdevkit/stats.py b/pgdevkit/stats.py index b6d01ca..c761f0d 100644 --- a/pgdevkit/stats.py +++ b/pgdevkit/stats.py @@ -8,6 +8,8 @@ import psycopg.sql from psycopg.rows import dict_row +from .dialect import resolve_dialect + STATS_DIRNAME = "_stats" TABLES_FILE = "_tables.json" @@ -55,14 +57,125 @@ def _write_json(path: Path, data: Any) -> None: path.write_text(json.dumps(data, indent=2, sort_keys=True) + "\n", encoding="utf-8") +_MSSQL_TABLES_SQL = """ +SELECT s.name AS [schema], t.name AS name, + SUM(CASE WHEN ps.index_id IN (0, 1) THEN ps.row_count ELSE 0 END) AS estimated_rows, + SUM(CASE WHEN ps.index_id IN (0, 1) THEN ps.used_page_count ELSE 0 END) * 8192 AS table_bytes, + SUM(CASE WHEN ps.index_id > 1 THEN ps.used_page_count ELSE 0 END) * 8192 AS index_bytes, + SUM(ps.used_page_count) * 8192 AS total_bytes +FROM sys.tables t +JOIN sys.schemas s ON s.schema_id = t.schema_id +LEFT JOIN sys.dm_db_partition_stats ps ON ps.object_id = t.object_id +GROUP BY s.name, t.name +ORDER BY 1, 2 +""" + +_MSSQL_COLUMNS_SQL = """ +SELECT c.name AS name, ty.name AS base_type, c.max_length AS max_length, + c.precision AS precision, c.scale AS scale +FROM sys.columns c +JOIN sys.types ty ON ty.user_type_id = c.user_type_id +WHERE c.object_id = OBJECT_ID(?) +ORDER BY c.column_id +""" + +# Types that can't be COUNT(DISTINCT)ed / measured with DATALENGTH. +_MSSQL_UNMEASURABLE = {"text", "ntext", "image", "xml", "geography", "geometry", "hierarchyid", "sql_variant"} + + +def _mssql_q(conn: Any, sql: str, params: tuple = ()) -> list[dict[str, Any]]: + cur = conn.cursor() + try: + cur.execute(sql, params) + names = [c[0] for c in cur.description] + return [dict(zip(names, row)) for row in cur.fetchall()] + finally: + cur.close() + + +def _mssql_ident(name: str) -> str: + return "[" + name.replace("]", "]]") + "]" + + +def _collect_mssql_stats( + conninfo: str, only: set[str] | None, exact: bool, analyze: bool +) -> tuple[dict[str, dict[str, Any]], dict[str, dict[str, Any]]]: + """MSSQL flavour of collect_stats. Row counts/sizes come from + sys.dm_db_partition_stats. SQL Server has no pg_stats equivalent, so column + null_fraction/n_distinct/avg_width are only filled in with exact=True + (computed by scanning each table); otherwise they are null. + analyze=True runs UPDATE STATISTICS.""" + # Lazy: the mssql extra isn't required for postgres-only use. + import mssql_python + + from .mssql_introspect import _format_type, _is_system_schema + + tables: dict[str, dict[str, Any]] = {} + columns: dict[str, dict[str, Any]] = {} + conn = mssql_python.connect(conninfo, autocommit=True) + try: + rows = [r for r in _mssql_q(conn, _MSSQL_TABLES_SQL) if not _is_system_schema(r["schema"])] + if only is not None: + unknown = only - {f"{r['schema']}.{r['name']}" for r in rows} + if unknown: + raise KeyError(", ".join(sorted(unknown))) + for r in rows: + qn = f"{r['schema']}.{r['name']}" + if only is not None and qn not in only: + continue + ident = f"{_mssql_ident(r['schema'])}.{_mssql_ident(r['name'])}" + if analyze: + conn.execute(f"UPDATE STATISTICS {ident}") + if exact: + row_count = _mssql_q(conn, f"SELECT COUNT_BIG(*) AS n FROM {ident}")[0]["n"] + else: + row_count = r["estimated_rows"] + tables[qn] = { + "row_count": row_count, + "row_count_exact": exact, + "table_bytes": r["table_bytes"], + "index_bytes": r["index_bytes"], + "total_bytes": r["total_bytes"], + } + col_stats: dict[str, Any] = {} + for c in _mssql_q(conn, _MSSQL_COLUMNS_SQL, (ident,)): + entry: dict[str, Any] = { + "data_type": _format_type(c["base_type"], c["max_length"], c["precision"], c["scale"]), + "null_fraction": None, + "n_distinct": None, + "avg_width": None, + } + if exact and row_count and c["base_type"].lower() not in _MSSQL_UNMEASURABLE: + col = _mssql_ident(c["name"]) + m = _mssql_q( + conn, + f"SELECT COUNT_BIG({col}) AS nn, COUNT_BIG(DISTINCT {col}) AS nd, " + f"AVG(CAST(DATALENGTH({col}) AS float)) AS w FROM {ident}", + )[0] + entry["null_fraction"] = 1 - m["nn"] / row_count + entry["n_distinct"] = m["nd"] + entry["avg_width"] = None if m["w"] is None else round(m["w"]) + col_stats[c["name"]] = entry + columns[qn] = col_stats + finally: + conn.close() + return tables, columns + + def collect_stats( - conninfo: str, only: set[str] | None = None, exact: bool = False, analyze: bool = False + conninfo: str, + only: set[str] | None = None, + exact: bool = False, + analyze: bool = False, + dialect: str = "postgres", ) -> tuple[dict[str, dict[str, Any]], dict[str, dict[str, Any]]]: """Returns (table_stats, column_stats), both keyed by "schema.table". Row counts come from pg_class.reltuples (null if the table was never analyzed) unless exact=True, which runs count(*). analyze=True runs ANALYZE first so column stats (null_frac/n_distinct/avg_width) are populated.""" + if resolve_dialect(dialect).name == "mssql": + return _collect_mssql_stats(conninfo, only, exact, analyze) tables: dict[str, dict[str, Any]] = {} columns: dict[str, dict[str, Any]] = {} with psycopg.connect(conninfo, autocommit=True) as conn: diff --git a/skills/pgdevkit/SKILL.md b/skills/pgdevkit/SKILL.md index f270c11..3721b4e 100644 --- a/skills/pgdevkit/SKILL.md +++ b/skills/pgdevkit/SKILL.md @@ -293,6 +293,8 @@ pgdb get-stats database/ public.users public.orders # JSON t pgdb get-stats database/ --no-columns # all tables, table-level stats only ``` +MSSQL: add `--dialect mssql` (needs the `mssql` extra). Row counts and sizes come from `sys.dm_db_partition_stats`; SQL Server has no `pg_stats`, so column `null_fraction`/`n_distinct`/`avg_width` are `null` unless `--exact` is given, which scans each table to compute them. `--analyze` runs `UPDATE STATISTICS`. + To read the stats from code or a script, just `json.load` `database/_stats/_tables.json` (and `database/_stats/.json` for columns). --- diff --git a/tests/test_stats_cli.py b/tests/test_stats_cli.py index 204a6da..1d5f9d1 100644 --- a/tests/test_stats_cli.py +++ b/tests/test_stats_cli.py @@ -73,3 +73,64 @@ def test_update_stats_then_get_stats(tmp_path: Path): finally: with psycopg.connect(admin, autocommit=True) as con: con.execute(f'DROP DATABASE IF EXISTS "{TEST_DB}"') + + +class _FakeMssqlCursor: + def __init__(self, conn): + self.conn, self.description, self.rows = conn, None, [] + + def execute(self, sql, params=()): + self.conn.executed.append(sql) + if "sys.dm_db_partition_stats" in sql: + self.description = [(n,) for n in ("schema", "name", "estimated_rows", "table_bytes", "index_bytes", "total_bytes")] + self.rows = [("dbo", "zeta", 10, 8192, 16384, 24576), ("sys", "junk", 1, 0, 0, 0)] + elif "FROM sys.columns" in sql: + self.description = [(n,) for n in ("name", "base_type", "max_length", "precision", "scale")] + self.rows = [("id", "int", 4, 10, 0), ("note", "nvarchar", 100, 0, 0)] + elif "COUNT_BIG(*)" in sql: + self.description, self.rows = [("n",)], [(10,)] + elif "COUNT_BIG(DISTINCT [id])" in sql: + self.description, self.rows = [("nn",), ("nd",), ("w",)], [(10, 10, 4.0)] + elif "COUNT_BIG(DISTINCT [note])" in sql: + self.description, self.rows = [("nn",), ("nd",), ("w",)], [(0, 0, None)] + + def fetchall(self): + return self.rows + + def close(self): + pass + + +class _FakeMssqlConn: + def __init__(self): + self.executed: list[str] = [] + + def cursor(self): + return _FakeMssqlCursor(self) + + def execute(self, sql): + self.executed.append(sql) + + def close(self): + pass + + +def test_update_stats_mssql(tmp_path: Path, monkeypatch): + import mssql_python + + conn = _FakeMssqlConn() + monkeypatch.setattr(mssql_python, "connect", lambda *a, **k: conn) + + r = runner.invoke(app, ["update-stats", str(tmp_path), "--url", "x", "--dialect", "mssql", "--exact", "--analyze"]) + assert r.exit_code == 0, r.output + tables = json.loads((tmp_path / "_stats" / "_tables.json").read_text()) + assert list(tables) == ["dbo.zeta"] # system schema skipped + assert tables["dbo.zeta"]["row_count"] == 10 and tables["dbo.zeta"]["row_count_exact"] is True + assert tables["dbo.zeta"]["total_bytes"] == 24576 + cols = json.loads((tmp_path / "_stats" / "dbo.zeta.json").read_text()) + assert cols["note"]["data_type"] == "nvarchar(50)" and cols["note"]["null_fraction"] == 1 + assert cols["id"]["n_distinct"] == 10 and cols["id"]["avg_width"] == 4 + assert "UPDATE STATISTICS [dbo].[zeta]" in conn.executed + + assert runner.invoke(app, ["update-stats", str(tmp_path), "--url", "x", "--dialect", "mssql", "--table", "dbo.nope"]).exit_code == 2 + assert runner.invoke(app, ["update-stats", str(tmp_path), "--url", "x", "--dialect", "oracle"]).exit_code == 2 From d99a9b62c06ff632648ab32073846faf2fb385e4 Mon Sep 17 00:00:00 2001 From: Adrian Ehrsam Date: Mon, 5 Oct 2026 09:05:45 +0200 Subject: [PATCH 3/3] Stats MSSQL: one scan per table, validate dialect up front (#43) Co-Authored-By: Claude Sonnet 5.5 Claude-Session: https://claude.ai/code/session_01SwHJR2uryShYQnsySXZT6G --- pgdevkit/cli.py | 8 +++++--- pgdevkit/stats.py | 31 +++++++++++++++++++------------ tests/test_stats_cli.py | 7 ++----- 3 files changed, 26 insertions(+), 20 deletions(-) diff --git a/pgdevkit/cli.py b/pgdevkit/cli.py index 821e818..65c205f 100644 --- a/pgdevkit/cli.py +++ b/pgdevkit/cli.py @@ -234,6 +234,11 @@ def update_stats( if not scripts_dir.is_dir(): err_console.print(f"[red]Error:[/red] {scripts_dir} is not a directory") raise typer.Exit(2) + try: + dialect = get_backend(dialect).dialect.name + except ValueError as e: + err_console.print(f"[red]Error:[/red] {e}") + raise typer.Exit(2) conninfo = build_conninfo(url, entra_user) try: with console.status("Collecting stats..."): @@ -241,9 +246,6 @@ def update_stats( except KeyError as e: err_console.print(f"[red]Error:[/red] unknown table(s): {e.args[0]}") raise typer.Exit(2) - except ValueError as e: - err_console.print(f"[red]Error:[/red] {e}") - raise typer.Exit(2) path = stats.write_stats(scripts_dir, tables, columns, prune=not table) console.print(f"Updated stats for {len(tables)} table(s) in {path.parent}") diff --git a/pgdevkit/stats.py b/pgdevkit/stats.py index c761f0d..7e2e756 100644 --- a/pgdevkit/stats.py +++ b/pgdevkit/stats.py @@ -126,8 +126,20 @@ def _collect_mssql_stats( ident = f"{_mssql_ident(r['schema'])}.{_mssql_ident(r['name'])}" if analyze: conn.execute(f"UPDATE STATISTICS {ident}") + col_defs = _mssql_q(conn, _MSSQL_COLUMNS_SQL, (ident,)) + measured: dict[str, Any] = {} if exact: - row_count = _mssql_q(conn, f"SELECT COUNT_BIG(*) AS n FROM {ident}")[0]["n"] + # One scan per table, so row and column counts share a snapshot. + parts = ["COUNT_BIG(*) AS n"] + for i, c in enumerate(col_defs): + if c["base_type"].lower() not in _MSSQL_UNMEASURABLE: + col = _mssql_ident(c["name"]) + parts.append( + f"COUNT_BIG({col}) AS nn{i}, COUNT_BIG(DISTINCT {col}) AS nd{i}, " + f"AVG(CAST(DATALENGTH({col}) AS float)) AS w{i}" + ) + measured = _mssql_q(conn, f"SELECT {', '.join(parts)} FROM {ident}")[0] + row_count = measured["n"] else: row_count = r["estimated_rows"] tables[qn] = { @@ -138,23 +150,18 @@ def _collect_mssql_stats( "total_bytes": r["total_bytes"], } col_stats: dict[str, Any] = {} - for c in _mssql_q(conn, _MSSQL_COLUMNS_SQL, (ident,)): + for i, c in enumerate(col_defs): entry: dict[str, Any] = { "data_type": _format_type(c["base_type"], c["max_length"], c["precision"], c["scale"]), "null_fraction": None, "n_distinct": None, "avg_width": None, } - if exact and row_count and c["base_type"].lower() not in _MSSQL_UNMEASURABLE: - col = _mssql_ident(c["name"]) - m = _mssql_q( - conn, - f"SELECT COUNT_BIG({col}) AS nn, COUNT_BIG(DISTINCT {col}) AS nd, " - f"AVG(CAST(DATALENGTH({col}) AS float)) AS w FROM {ident}", - )[0] - entry["null_fraction"] = 1 - m["nn"] / row_count - entry["n_distinct"] = m["nd"] - entry["avg_width"] = None if m["w"] is None else round(m["w"]) + if row_count and f"nn{i}" in measured: + entry["null_fraction"] = 1 - measured[f"nn{i}"] / row_count + entry["n_distinct"] = measured[f"nd{i}"] + w = measured[f"w{i}"] + entry["avg_width"] = None if w is None else round(w) col_stats[c["name"]] = entry columns[qn] = col_stats finally: diff --git a/tests/test_stats_cli.py b/tests/test_stats_cli.py index 1d5f9d1..2afe0d5 100644 --- a/tests/test_stats_cli.py +++ b/tests/test_stats_cli.py @@ -88,11 +88,8 @@ def execute(self, sql, params=()): self.description = [(n,) for n in ("name", "base_type", "max_length", "precision", "scale")] self.rows = [("id", "int", 4, 10, 0), ("note", "nvarchar", 100, 0, 0)] elif "COUNT_BIG(*)" in sql: - self.description, self.rows = [("n",)], [(10,)] - elif "COUNT_BIG(DISTINCT [id])" in sql: - self.description, self.rows = [("nn",), ("nd",), ("w",)], [(10, 10, 4.0)] - elif "COUNT_BIG(DISTINCT [note])" in sql: - self.description, self.rows = [("nn",), ("nd",), ("w",)], [(0, 0, None)] + names = ("n", "nn0", "nd0", "w0", "nn1", "nd1", "w1") + self.description, self.rows = [(n,) for n in names], [(10, 10, 10, 4.0, 0, 0, None)] def fetchall(self): return self.rows