diff --git a/app/drivers/base.py b/app/drivers/base.py index 92e4df4..d062ce0 100644 --- a/app/drivers/base.py +++ b/app/drivers/base.py @@ -87,6 +87,10 @@ class BaseDriver(ABC): """Returns list of database name strings.""" pass + def quote_identifier(self, name: str) -> str: + """Wrap an identifier in the DB-appropriate quote characters.""" + return f'"{name}"' + @abstractmethod def get_tables(self, database: str) -> list: """Returns list of TableInfo for the given database.""" diff --git a/app/drivers/mssql_driver.py b/app/drivers/mssql_driver.py index 0340a3e..55d23a6 100644 --- a/app/drivers/mssql_driver.py +++ b/app/drivers/mssql_driver.py @@ -19,6 +19,9 @@ class MSSQLDriver(BaseDriver): super().__init__(config) self.db_type = "mssql" + def quote_identifier(self, name: str) -> str: + return f"[{name}]" + def _conn_str(self) -> str: host = self.config.get("host", "localhost") port = int(self.config.get("port", 1433)) diff --git a/app/drivers/mysql_driver.py b/app/drivers/mysql_driver.py index 517770c..2f627a5 100644 --- a/app/drivers/mysql_driver.py +++ b/app/drivers/mysql_driver.py @@ -18,6 +18,9 @@ class MySQLDriver(BaseDriver): super().__init__(config) self.db_type = "mysql" + def quote_identifier(self, name: str) -> str: + return f"`{name}`" + def _connect_kwargs(self) -> dict: kw = { "host": self.config.get("host", "localhost"), diff --git a/app/ui/column_stats_dialog.py b/app/ui/column_stats_dialog.py index b4ca521..ba8ed4e 100644 --- a/app/ui/column_stats_dialog.py +++ b/app/ui/column_stats_dialog.py @@ -19,12 +19,13 @@ class _StatsWorker(QThread): self._column = column def run(self): - col = self._column - tbl = self._table + q = self._driver.quote_identifier + col = q(self._column) + tbl = q(self._table) try: _, rows, _ = self._driver.execute_query( - f'SELECT COUNT(*), COUNT("{col}"), COUNT(DISTINCT "{col}"), ' - f'MIN("{col}"), MAX("{col}") FROM "{tbl}"' + f'SELECT COUNT(*), COUNT({col}), COUNT(DISTINCT {col}), ' + f'MIN({col}), MAX({col}) FROM {tbl}' ) total, non_null, distinct, min_val, max_val = rows[0] null_count = (total or 0) - (non_null or 0) @@ -35,7 +36,7 @@ class _StatsWorker(QThread): avg_val = None try: _, avg_rows, _ = self._driver.execute_query( - f'SELECT AVG("{col}") FROM "{tbl}"' + f'SELECT AVG({col}) FROM {tbl}' ) avg_val = avg_rows[0][0] except Exception: