Files
DBClient/app/drivers/sqlite_driver.py
T
2026-05-21 15:46:41 -04:00

281 lines
11 KiB
Python

"""SQLite driver implementation using stdlib sqlite3."""
import sqlite3
from contextlib import contextmanager
from typing import Optional
from app.drivers.base import (
BaseDriver, ColumnInfo, IndexInfo, ForeignKeyInfo, TableInfo
)
class SQLiteDriver(BaseDriver):
"""SQLite database driver (uses stdlib sqlite3)."""
def __init__(self, config: dict):
super().__init__(config)
self.db_type = "sqlite"
def connect(self) -> None:
db_path = self.config.get("database", ":memory:")
self._connection = sqlite3.connect(
db_path,
check_same_thread=False,
timeout=int(self.config.get("connection_timeout", 30)),
)
self._connection.execute("PRAGMA journal_mode=WAL")
self._connection.execute("PRAGMA foreign_keys=ON")
def disconnect(self) -> None:
if self._connection:
try:
self._connection.close()
except Exception:
pass
finally:
self._connection = None
def test_connection(self) -> tuple:
try:
db_path = self.config.get("database", "")
conn = sqlite3.connect(db_path, timeout=5)
conn.execute("SELECT 1")
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:
return [self.config.get("database", "main")]
def get_tables(self, database: str = "") -> list:
with self._cur() as c:
c.execute("""
SELECT name FROM sqlite_master
WHERE type='table' AND name NOT LIKE 'sqlite_%'
ORDER BY name
""")
names = [r[0] for r in c.fetchall()]
tables = []
for name in names:
try:
with self._cur() as rc:
rc.execute(f'SELECT COUNT(*) FROM "{name}"')
row_count = rc.fetchone()[0]
except Exception:
row_count = 0
tables.append(TableInfo(name=name, schema="main", row_count=row_count))
return tables
def get_views(self, database: str = "") -> list:
with self._cur() as c:
c.execute("SELECT name FROM sqlite_master WHERE type='view' ORDER BY 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(f'PRAGMA table_info("{table}")')
cols = []
for r in c.fetchall():
# cid, name, type, notnull, dflt_value, pk
cols.append(ColumnInfo(
name=r[1],
data_type=r[2] or "TEXT",
nullable=not bool(r[3]),
default=str(r[4]) if r[4] is not None else None,
is_primary_key=bool(r[5]),
is_foreign_key=False,
))
return cols
def get_indexes(self, database: str, table: str) -> list:
with self._cur() as c:
c.execute(f'PRAGMA index_list("{table}")')
index_rows = c.fetchall()
indexes = []
for r in index_rows:
idx_name = r[1]
is_unique = bool(r[2])
with self._cur() as cc:
cc.execute(f'PRAGMA index_info("{idx_name}")')
cols = [row[2] for row in cc.fetchall()]
indexes.append(IndexInfo(name=idx_name, columns=cols, is_unique=is_unique))
return indexes
def get_foreign_keys(self, database: str, table: str) -> list:
with self._cur() as c:
c.execute(f'PRAGMA foreign_key_list("{table}")')
return [ForeignKeyInfo(
name=f"fk_{r[3]}",
column=r[3], ref_table=r[2], ref_column=r[4],
on_update=r[5] or "", on_delete=r[6] or "",
) for r in c.fetchall()]
def get_table_ddl(self, database: str, table: str) -> str:
with self._cur() as c:
c.execute(
"SELECT sql FROM sqlite_master WHERE name = ? AND type = 'table'",
(table,)
)
row = c.fetchone()
return row[0] if row else ""
def get_functions(self, database: str) -> list:
return []
def get_stored_procedures(self, database: str) -> list:
return []
def get_triggers(self, database: str, table: str = "") -> list:
with self._cur() as c:
if table:
c.execute(
"SELECT name FROM sqlite_master WHERE type='trigger' AND tbl_name=?",
(table,)
)
else:
c.execute("SELECT name FROM sqlite_master WHERE type='trigger' ORDER BY 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 or ())
if c.description:
cols = [d[0] for d in c.description]
rows = c.fetchall()
return cols, rows, len(rows)
self._connection.commit()
return [], [], c.rowcount
def execute_script(self, sql: str) -> list:
results = []
stmts = [s.strip() for s in sql.split(';') if s.strip()]
for stmt in stmts:
try:
with self._cur() as c:
c.execute(stmt)
if c.description:
cols = [d[0] for d in c.description]
rows = c.fetchall()
results.append((cols, rows, len(rows), ""))
else:
self._connection.commit()
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(["?"] * len(data))
with self._cur() as c:
c.execute(f'INSERT INTO "{table}" ({cols}) VALUES ({ph})',
tuple(data.values()))
self._connection.commit()
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}" = ?')
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}" = ?' 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))
self._connection.commit()
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))
self._connection.commit()
return True
# ── Server tools ──────────────────────────────────────────────────────────
def explain_query(self, sql: str) -> tuple:
with self._cur() as c:
c.execute(f"EXPLAIN QUERY PLAN {sql}")
cols = [d[0] for d in c.description]
return cols, c.fetchall()
def get_process_list(self) -> tuple:
return ["Info"], [("SQLite does not support process listing.",)]
def kill_process(self, process_id: int) -> bool:
return False
# ── Table designer ────────────────────────────────────────────────────────
def add_column(self, _database: str, table: str, col_name: str,
col_type: str, nullable: bool = True,
default=None) -> bool:
null_clause = "" if nullable else " NOT NULL"
default_clause = f" DEFAULT {default}" if default is not None and default != "" else ""
sql = (f'ALTER TABLE "{table}" '
f'ADD COLUMN "{col_name}" {col_type}{default_clause}{null_clause}')
with self._cur() as c:
c.execute(sql)
self._connection.commit()
return True
def drop_column(self, _database: str, table: str, col_name: str) -> bool:
# Requires SQLite 3.35.0+
with self._cur() as c:
c.execute(f'ALTER TABLE "{table}" DROP COLUMN "{col_name}"')
self._connection.commit()
return True
def rename_column(self, _database: str, table: str,
old_name: str, new_name: str) -> bool:
# Requires SQLite 3.25.0+
with self._cur() as c:
c.execute(f'ALTER TABLE "{table}" RENAME COLUMN "{old_name}" TO "{new_name}"')
self._connection.commit()
return True