Database 1.0.0
_db.py
4.7 KB · raw
"""Shared by the database tools: saved connections, drivers, and opening one.
Connections are kept by name in connections.json in this folder, as URLs:
sqlite:///C:/path/to/file.db (or a path relative to the workspace)
postgresql://user:password@host:5432/dbname
mysql://user:password@host:3306/dbname
SQLite needs nothing. PostgreSQL (psycopg) and MySQL (PyMySQL) drivers are
installed into _lib/ in this folder by db_install_driver, so removing the
capability removes them. The file holds passwords in plain text: it is the
owner's machine, and there is no secret store a capability can use yet.
"""
from __future__ import annotations
import json
import sqlite3
import sys
from pathlib import Path
from urllib.parse import unquote, urlparse
from aworg.tools.base import ToolError
HERE = Path(__file__).parent
CONNECTIONS = HERE / "connections.json"
LIB = HERE / "_lib"
if LIB.is_dir() and str(LIB) not in sys.path:
sys.path.insert(0, str(LIB))
DRIVERS = {"postgresql": "psycopg[binary]", "mysql": "PyMySQL"}
KINDS = {"sqlite": "sqlite", "postgres": "postgresql", "postgresql": "postgresql", "mysql": "mysql"}
def saved() -> dict[str, str]:
try:
return json.loads(CONNECTIONS.read_text(encoding="utf-8"))
except (OSError, ValueError):
return {}
def save(connections: dict[str, str]) -> None:
CONNECTIONS.write_text(json.dumps(connections, indent=2), encoding="utf-8")
def kind_of(url: str) -> str:
scheme = urlparse(url).scheme.split("+")[0].lower()
if scheme not in KINDS:
raise ToolError(f"Unsupported database {scheme!r}: use sqlite://, postgresql:// or mysql://.")
return KINDS[scheme]
def hide_password(url: str) -> str:
p = urlparse(url)
if p.password:
return url.replace(f":{p.password}@", ":***@", 1)
return url
def resolve(name: str) -> str:
connections = saved()
if name not in connections:
known = ", ".join(connections) or "none saved yet"
raise ToolError(f"No connection called {name!r}. Saved: {known}. Save one with db_connect.")
return connections[name]
def sqlite_path(url: str, workspace: Path | None) -> Path:
"""The file a sqlite URL names. After `sqlite:///`: a drive letter or a
leading `/` is absolute; anything else is relative to the workspace."""
if url.startswith("sqlite:////"):
rest = "/" + url[len("sqlite:////"):]
elif url.startswith("sqlite:///"):
rest = url[len("sqlite:///"):]
else:
rest = url.split("://", 1)[1]
path = Path(unquote(rest))
if not path.is_absolute() and workspace is not None:
path = workspace / path
return path
def open_connection(url: str, workspace: Path | None, read_only: bool):
"""A DB-API connection, read-only when asked, for the URL's kind."""
kind = kind_of(url)
if kind == "sqlite":
path = sqlite_path(url, workspace)
if read_only:
if not path.exists():
raise ToolError(f"No SQLite file at {path}.")
return sqlite3.connect(f"file:{path.as_posix()}?mode=ro", uri=True, timeout=10)
return sqlite3.connect(str(path), timeout=10)
p = urlparse(url)
if kind == "postgresql":
try:
import psycopg
except ImportError:
raise ToolError("The PostgreSQL driver is not installed. Run db_install_driver with kind 'postgresql'.") from None
conn = psycopg.connect(url.replace("postgres://", "postgresql://", 1), connect_timeout=10)
if read_only:
conn.read_only = True
return conn
try:
import pymysql
except ImportError:
raise ToolError("The MySQL driver is not installed. Run db_install_driver with kind 'mysql'.") from None
conn = pymysql.connect(
host=p.hostname or "localhost", port=p.port or 3306,
user=unquote(p.username or ""), password=unquote(p.password or ""),
database=p.path.lstrip("/") or None, connect_timeout=10, autocommit=False,
)
if read_only:
with conn.cursor() as cur:
cur.execute("SET SESSION TRANSACTION READ ONLY")
return conn
def table(columns: list[str], rows: list[tuple], limit: int) -> str:
"""Rows as a plain text table, cells cut at 80 characters."""
def cell(v):
s = "NULL" if v is None else str(v)
s = s.replace("\n", "\\n")
return s if len(s) <= 80 else s[:77] + "..."
body = [[cell(v) for v in r] for r in rows[:limit]]
widths = [max([len(c)] + [len(r[i]) for r in body]) for i, c in enumerate(columns)]
line = lambda vals: " | ".join(v.ljust(w) for v, w in zip(vals, widths))
out = [line(columns), "-+-".join("-" * w for w in widths)] + [line(r) for r in body]
return "\n".join(out)