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)