Repository navigation
Expand file tree
/
Copy pathrun.py
More file actions
87 lines (70 loc) · 2.91 KB
/
Copy pathrun.py
File metadata and controls
87 lines (70 loc) · 2.91 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
"""Run any query in this folder against the fixture schema.
python sql/run.py 03_cohort_retention.sql
python sql/run.py --all
Uses `sqlite3` from the standard library, so there is nothing to install and
nothing to configure - the whole block is reproducible by anyone who clones the
repo. The queries stick to standard SQL (window functions, CTEs) rather than
SQLite-specific syntax, so they transfer to Postgres, DuckDB or BigQuery with
at most a date-function rename.
"""
import argparse
import sqlite3
import sys
from pathlib import Path
HERE = Path(__file__).parent
def build_database() -> sqlite3.Connection:
"""Fresh in-memory database, loaded from schema.sql."""
connection = sqlite3.connect(":memory:")
connection.executescript((HERE / "schema.sql").read_text(encoding="utf-8"))
return connection
def run_query(connection: sqlite3.Connection, sql_path: Path):
"""Execute a .sql file and return (column_names, rows)."""
cursor = connection.execute(sql_path.read_text(encoding="utf-8"))
columns = [d[0] for d in cursor.description] if cursor.description else []
return columns, cursor.fetchall()
def format_table(columns, rows) -> str:
"""Minimal fixed-width table - enough to read a result set in a terminal."""
if not columns:
return "(no rows)"
widths = [len(c) for c in columns]
printable = [["" if v is None else str(v) for v in row] for row in rows]
for row in printable:
for i, value in enumerate(row):
widths[i] = max(widths[i], len(value))
lines = [
" ".join(c.ljust(widths[i]) for i, c in enumerate(columns)),
" ".join("-" * w for w in widths),
]
lines.extend(" ".join(v.ljust(widths[i]) for i, v in enumerate(row))
for row in printable)
lines.append(f"({len(rows)} rows)")
return "\n".join(lines)
def main() -> int:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("query", nargs="?", help="a .sql file in this folder")
parser.add_argument("--all", action="store_true", help="run every query")
args = parser.parse_args()
queries = sorted(p for p in HERE.glob("*.sql") if p.name != "schema.sql")
if args.all:
targets = queries
elif args.query:
target = HERE / args.query
if not target.exists():
print(f"no such query: {args.query}\n\navailable:", file=sys.stderr)
for q in queries:
print(f" {q.name}", file=sys.stderr)
return 1
targets = [target]
else:
print("available queries:")
for q in queries:
print(f" {q.name}")
return 0
connection = build_database()
for path in targets:
print(f"\n=== {path.name} " + "=" * max(0, 60 - len(path.name)))
columns, rows = run_query(connection, path)
print(format_table(columns, rows))
return 0
if __name__ == "__main__":
raise SystemExit(main())