execute_script() semicolon splitter (all drivers)
- Add BaseDriver._split_statements() that tracks single/double-quoted
strings and -- / /* */ comments so semicolons inside procedure bodies
are not treated as statement boundaries.
- Replace the naive sql.split(';') in all four drivers with this helper.
- MSSQL execute_script() also pre-splits on GO (case-insensitive, own
line) so scripts pasted from SSMS work correctly.
MSSQL EXPLAIN
- SET SHOWPLAN_TEXT ON/execute/SET SHOWPLAN_TEXT OFF must be separate
execute() calls; pyodbc rejects multiple statements in one call.
Hold self._lock for the entire sequence to keep session state atomic.
Keyring (connections.py)
- Log warnings (with traceback) on load/save failures instead of
silently returning "".
- save_password() now calls delete_password() when password is empty,
so clearing a saved password actually removes the old keyring entry
rather than leaving a stale one behind.
SQLite get_tables() O(N) lock acquisitions
- Run all COUNT(*) queries inside the same `with self._cur() as c:`
block, reducing N+1 lock acquisitions to 1.
Unused import: remove psycopg2.extras from postgres_driver.py.
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
347 lines
14 KiB
Python
347 lines
14 KiB
Python
"""PostgreSQL driver implementation using psycopg2."""
|
|
import psycopg2
|
|
from contextlib import contextmanager
|
|
from typing import Optional
|
|
from app.drivers.base import (
|
|
BaseDriver, ColumnInfo, IndexInfo, ForeignKeyInfo, TableInfo
|
|
)
|
|
|
|
|
|
class PostgreSQLDriver(BaseDriver):
|
|
"""PostgreSQL database driver."""
|
|
|
|
def __init__(self, config: dict):
|
|
super().__init__(config)
|
|
self.db_type = "postgresql"
|
|
|
|
def _dsn(self) -> dict:
|
|
kw = {
|
|
"host": self.config.get("host", "localhost"),
|
|
"port": int(self.config.get("port", 5432)),
|
|
"user": self.config.get("user", ""),
|
|
"password": self.config.get("password", ""),
|
|
"connect_timeout": int(self.config.get("connection_timeout", 30)),
|
|
}
|
|
db = self.config.get("database", "")
|
|
if db:
|
|
kw["dbname"] = db
|
|
if self.config.get("ssl"):
|
|
kw["sslmode"] = "verify-ca" if self.config.get("ssl_ca") else "require"
|
|
if self.config.get("ssl_ca"): kw["sslrootcert"] = self.config["ssl_ca"]
|
|
if self.config.get("ssl_cert"): kw["sslcert"] = self.config["ssl_cert"]
|
|
if self.config.get("ssl_key"): kw["sslkey"] = self.config["ssl_key"]
|
|
return kw
|
|
|
|
def connect(self) -> None:
|
|
self._connection = psycopg2.connect(**self._dsn())
|
|
self._connection.autocommit = True
|
|
|
|
def disconnect(self) -> None:
|
|
if self._connection:
|
|
try:
|
|
self._connection.close()
|
|
except Exception:
|
|
pass
|
|
finally:
|
|
self._connection = None
|
|
|
|
def test_connection(self) -> tuple:
|
|
try:
|
|
conn = psycopg2.connect(**self._dsn())
|
|
conn.close()
|
|
return True, "Connection successful"
|
|
except Exception as e:
|
|
return False, str(e)
|
|
|
|
@contextmanager
|
|
def _cur(self):
|
|
"""Yield a cursor while holding the driver lock (thread-safe execute→fetch)."""
|
|
with self._lock:
|
|
cursor = self._connection.cursor()
|
|
try:
|
|
yield cursor
|
|
finally:
|
|
try:
|
|
cursor.close()
|
|
except Exception:
|
|
pass
|
|
|
|
# ── Schema introspection ──────────────────────────────────────────────────
|
|
|
|
def get_databases(self) -> list:
|
|
with self._cur() as c:
|
|
c.execute(
|
|
"SELECT datname FROM pg_database "
|
|
"WHERE datistemplate = false ORDER BY datname"
|
|
)
|
|
return [r[0] for r in c.fetchall()]
|
|
|
|
def get_tables(self, database: str) -> list:
|
|
with self._cur() as c:
|
|
c.execute("""
|
|
SELECT t.table_name, t.table_schema,
|
|
COALESCE(s.n_live_tup, 0),
|
|
0,
|
|
'',
|
|
COALESCE(obj_description(
|
|
(quote_ident(t.table_schema)||'.'||quote_ident(t.table_name))::regclass,
|
|
'pg_class'), '')
|
|
FROM information_schema.tables t
|
|
LEFT JOIN pg_stat_user_tables s
|
|
ON s.schemaname = t.table_schema AND s.relname = t.table_name
|
|
WHERE t.table_schema NOT IN ('pg_catalog','information_schema')
|
|
AND t.table_type = 'BASE TABLE'
|
|
ORDER BY t.table_name
|
|
""")
|
|
return [TableInfo(name=r[0], schema=r[1], row_count=r[2],
|
|
size_bytes=r[3], engine=r[4], comment=r[5])
|
|
for r in c.fetchall()]
|
|
|
|
def get_views(self, database: str) -> list:
|
|
with self._cur() as c:
|
|
c.execute("""
|
|
SELECT table_name FROM information_schema.views
|
|
WHERE table_schema NOT IN ('pg_catalog','information_schema')
|
|
ORDER BY table_name
|
|
""")
|
|
return [r[0] for r in c.fetchall()]
|
|
|
|
def get_columns(self, database: str, table: str) -> list:
|
|
with self._cur() as c:
|
|
c.execute("""
|
|
SELECT c.column_name, c.data_type, c.is_nullable,
|
|
c.column_default,
|
|
(SELECT true FROM information_schema.table_constraints tc
|
|
JOIN information_schema.key_column_usage kcu
|
|
ON tc.constraint_name = kcu.constraint_name
|
|
WHERE tc.constraint_type = 'PRIMARY KEY'
|
|
AND kcu.table_name = c.table_name
|
|
AND kcu.column_name = c.column_name
|
|
LIMIT 1) IS NOT NULL,
|
|
false
|
|
FROM information_schema.columns c
|
|
WHERE c.table_name = %s
|
|
ORDER BY c.ordinal_position
|
|
""", (table,))
|
|
return [ColumnInfo(name=r[0], data_type=r[1],
|
|
nullable=(r[2] == "YES"), default=r[3],
|
|
is_primary_key=bool(r[4]),
|
|
is_foreign_key=bool(r[5]))
|
|
for r in c.fetchall()]
|
|
|
|
def get_indexes(self, database: str, table: str) -> list:
|
|
with self._cur() as c:
|
|
c.execute("""
|
|
SELECT i.relname, ix.indisunique,
|
|
array_agg(a.attname ORDER BY k.n) AS cols
|
|
FROM pg_class t
|
|
JOIN pg_index ix ON t.oid = ix.indrelid
|
|
JOIN pg_class i ON i.oid = ix.indexrelid
|
|
JOIN unnest(ix.indkey) WITH ORDINALITY AS k(attnum, n)
|
|
ON TRUE
|
|
JOIN pg_attribute a ON a.attrelid = t.oid AND a.attnum = k.attnum
|
|
WHERE t.relname = %s
|
|
GROUP BY i.relname, ix.indisunique
|
|
ORDER BY i.relname
|
|
""", (table,))
|
|
return [IndexInfo(name=r[0], columns=list(r[2]),
|
|
is_unique=bool(r[1]))
|
|
for r in c.fetchall()]
|
|
|
|
def get_foreign_keys(self, database: str, table: str) -> list:
|
|
with self._cur() as c:
|
|
c.execute("""
|
|
SELECT tc.constraint_name, kcu.column_name,
|
|
ccu.table_name, ccu.column_name,
|
|
rc.update_rule, rc.delete_rule
|
|
FROM information_schema.table_constraints tc
|
|
JOIN information_schema.key_column_usage kcu
|
|
ON tc.constraint_name = kcu.constraint_name
|
|
JOIN information_schema.constraint_column_usage ccu
|
|
ON ccu.constraint_name = tc.constraint_name
|
|
JOIN information_schema.referential_constraints rc
|
|
ON rc.constraint_name = tc.constraint_name
|
|
WHERE tc.constraint_type = 'FOREIGN KEY'
|
|
AND tc.table_name = %s
|
|
""", (table,))
|
|
return [ForeignKeyInfo(name=r[0], column=r[1],
|
|
ref_table=r[2], ref_column=r[3],
|
|
on_update=r[4] or "", on_delete=r[5] or "")
|
|
for r in c.fetchall()]
|
|
|
|
def get_table_ddl(self, database: str, table: str) -> str:
|
|
cols = self.get_columns(database, table)
|
|
lines = [f"CREATE TABLE {table} ("]
|
|
col_defs = []
|
|
for c in cols:
|
|
d = f" {c.name} {c.data_type}"
|
|
if not c.nullable: d += " NOT NULL"
|
|
if c.default: d += f" DEFAULT {c.default}"
|
|
col_defs.append(d)
|
|
lines.append(",\n".join(col_defs))
|
|
lines.append(");")
|
|
return "\n".join(lines)
|
|
|
|
def get_functions(self, database: str) -> list:
|
|
with self._cur() as c:
|
|
c.execute("""
|
|
SELECT routine_name FROM information_schema.routines
|
|
WHERE routine_type = 'FUNCTION'
|
|
AND routine_schema NOT IN ('pg_catalog','information_schema')
|
|
ORDER BY routine_name
|
|
""")
|
|
return [r[0] for r in c.fetchall()]
|
|
|
|
def get_stored_procedures(self, database: str) -> list:
|
|
with self._cur() as c:
|
|
c.execute("""
|
|
SELECT routine_name FROM information_schema.routines
|
|
WHERE routine_type = 'PROCEDURE'
|
|
AND routine_schema NOT IN ('pg_catalog','information_schema')
|
|
ORDER BY routine_name
|
|
""")
|
|
return [r[0] for r in c.fetchall()]
|
|
|
|
def get_triggers(self, database: str, table: str = "") -> list:
|
|
with self._cur() as c:
|
|
if table:
|
|
c.execute("""
|
|
SELECT trigger_name FROM information_schema.triggers
|
|
WHERE event_object_table = %s ORDER BY trigger_name
|
|
""", (table,))
|
|
else:
|
|
c.execute(
|
|
"SELECT trigger_name FROM information_schema.triggers "
|
|
"ORDER BY trigger_name"
|
|
)
|
|
return [r[0] for r in c.fetchall()]
|
|
|
|
# ── Query execution ───────────────────────────────────────────────────────
|
|
|
|
def execute_query(self, sql: str, params: Optional[tuple] = None) -> tuple:
|
|
with self._cur() as c:
|
|
c.execute(sql, params)
|
|
if c.description:
|
|
cols = [d[0] for d in c.description]
|
|
rows = c.fetchall()
|
|
return cols, rows, len(rows)
|
|
return [], [], c.rowcount
|
|
|
|
def execute_script(self, sql: str) -> list:
|
|
results = []
|
|
with self._cur() as c:
|
|
for stmt in self._split_statements(sql):
|
|
try:
|
|
c.execute(stmt)
|
|
if c.description:
|
|
cols = [d[0] for d in c.description]
|
|
rows = c.fetchall()
|
|
results.append((cols, rows, len(rows), ""))
|
|
else:
|
|
results.append(([], [], c.rowcount,
|
|
f"{c.rowcount} row(s) affected"))
|
|
except Exception as e:
|
|
results.append(([], [], 0, f"Error: {e}"))
|
|
return results
|
|
|
|
# ── Table data CRUD ───────────────────────────────────────────────────────
|
|
|
|
def get_table_data(self, database: str, table: str,
|
|
where: str = "", order_by: str = "",
|
|
limit: int = 1000, offset: int = 0) -> tuple:
|
|
sql = f'SELECT * FROM "{table}"'
|
|
if where: sql += f" WHERE {where}"
|
|
if order_by: sql += f" ORDER BY {order_by}"
|
|
sql += f" LIMIT {limit} OFFSET {offset}"
|
|
return self.execute_query(sql)
|
|
|
|
def get_table_row_count(self, database: str, table: str, where: str = "") -> int:
|
|
sql = f'SELECT COUNT(*) FROM "{table}"'
|
|
if where: sql += f" WHERE {where}"
|
|
with self._cur() as c:
|
|
c.execute(sql)
|
|
return c.fetchone()[0]
|
|
|
|
def insert_row(self, database: str, table: str, data: dict) -> bool:
|
|
cols = ", ".join(f'"{c}"' for c in data)
|
|
ph = ", ".join(["%s"] * len(data))
|
|
with self._cur() as c:
|
|
c.execute(f'INSERT INTO "{table}" ({cols}) VALUES ({ph})',
|
|
tuple(data.values()))
|
|
return True
|
|
|
|
@staticmethod
|
|
def _where(where: dict) -> tuple:
|
|
parts, params = [], []
|
|
for col, val in where.items():
|
|
if val is None:
|
|
parts.append(f'"{col}" IS NULL')
|
|
else:
|
|
parts.append(f'"{col}" = %s')
|
|
params.append(val)
|
|
return " AND ".join(parts), params
|
|
|
|
def update_row(self, database: str, table: str, data: dict, where: dict) -> bool:
|
|
set_cl = ", ".join(f'"{c}" = %s' for c in data)
|
|
where_cl, where_params = self._where(where)
|
|
with self._cur() as c:
|
|
c.execute(f'UPDATE "{table}" SET {set_cl} WHERE {where_cl}',
|
|
tuple(data.values()) + tuple(where_params))
|
|
return True
|
|
|
|
def delete_row(self, database: str, table: str, where: dict) -> bool:
|
|
where_cl, where_params = self._where(where)
|
|
with self._cur() as c:
|
|
c.execute(f'DELETE FROM "{table}" WHERE {where_cl}',
|
|
tuple(where_params))
|
|
return True
|
|
|
|
# ── Server tools ──────────────────────────────────────────────────────────
|
|
|
|
def explain_query(self, sql: str) -> tuple:
|
|
with self._cur() as c:
|
|
c.execute(f"EXPLAIN ANALYZE {sql}")
|
|
return ["Plan"], c.fetchall()
|
|
|
|
def get_process_list(self) -> tuple:
|
|
with self._cur() as c:
|
|
c.execute("""
|
|
SELECT pid, usename, application_name, client_addr,
|
|
state, query, query_start
|
|
FROM pg_stat_activity WHERE state IS NOT NULL
|
|
ORDER BY query_start DESC NULLS LAST
|
|
""")
|
|
cols = [d[0] for d in c.description]
|
|
return cols, c.fetchall()
|
|
|
|
def kill_process(self, process_id: int) -> bool:
|
|
with self._cur() as c:
|
|
c.execute("SELECT pg_terminate_backend(%s)", (process_id,))
|
|
return True
|
|
|
|
# ── Table designer ────────────────────────────────────────────────────────
|
|
|
|
def add_column(self, _database: str, table: str, col_name: str,
|
|
col_type: str, nullable: bool = True,
|
|
default=None) -> bool:
|
|
default_clause = f" DEFAULT {default}" if default is not None and default != "" else ""
|
|
null_clause = "" if nullable else " NOT NULL"
|
|
sql = (f'ALTER TABLE "{table}" '
|
|
f'ADD COLUMN "{col_name}" {col_type}{default_clause}{null_clause}')
|
|
with self._cur() as c:
|
|
c.execute(sql)
|
|
return True
|
|
|
|
def drop_column(self, _database: str, table: str, col_name: str) -> bool:
|
|
sql = f'ALTER TABLE "{table}" DROP COLUMN "{col_name}"'
|
|
with self._cur() as c:
|
|
c.execute(sql)
|
|
return True
|
|
|
|
def rename_column(self, _database: str, table: str,
|
|
old_name: str, new_name: str) -> bool:
|
|
sql = f'ALTER TABLE "{table}" RENAME COLUMN "{old_name}" TO "{new_name}"'
|
|
with self._cur() as c:
|
|
c.execute(sql)
|
|
return True
|