diff --git a/pgdevkit/cli.py b/pgdevkit/cli.py index a848677..65c205f 100644 --- a/pgdevkit/cli.py +++ b/pgdevkit/cli.py @@ -227,16 +227,22 @@ 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.""" 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..."): - 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) diff --git a/pgdevkit/stats.py b/pgdevkit/stats.py index b6d01ca..7e2e756 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,132 @@ 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}") + col_defs = _mssql_q(conn, _MSSQL_COLUMNS_SQL, (ident,)) + measured: dict[str, Any] = {} + if exact: + # 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] = { + "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 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 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: + 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..2afe0d5 100644 --- a/tests/test_stats_cli.py +++ b/tests/test_stats_cli.py @@ -73,3 +73,61 @@ 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: + 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 + + 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