Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 7 additions & 1 deletion pgdevkit/cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -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/<schema.name>.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)
Expand Down
122 changes: 121 additions & 1 deletion pgdevkit/stats.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"

Expand Down Expand Up @@ -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:
Expand Down
2 changes: 2 additions & 0 deletions skills/pgdevkit/SKILL.md
Original file line number Diff line number Diff line change
Expand Up @@ -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/<schema.table>.json` for columns).

---
Expand Down
58 changes: 58 additions & 0 deletions tests/test_stats_cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Loading