From cfbd209d906497d0f2c0b481ffba08c9756d991c Mon Sep 17 00:00:00 2001 From: Bartok Date: Wed, 1 Jul 2026 03:24:29 -0600 Subject: [PATCH] fix(duckdb): case-insensitive .duckdb discovery, deterministic attach aliases, and single-statement SQL hardening (#2411) Co-authored-by: Bartok9 Co-authored-by: Bartok Rage --- core/wren/src/wren/connector/duckdb.py | 75 +++++++++- .../tests/unit/test_duckdb_file_listing.py | 137 ++++++++++++++++++ 2 files changed, 207 insertions(+), 5 deletions(-) create mode 100644 core/wren/tests/unit/test_duckdb_file_listing.py diff --git a/core/wren/src/wren/connector/duckdb.py b/core/wren/src/wren/connector/duckdb.py index 966f702ed..f420342d2 100644 --- a/core/wren/src/wren/connector/duckdb.py +++ b/core/wren/src/wren/connector/duckdb.py @@ -1,4 +1,5 @@ import os +import re import opendal import pyarrow as pa @@ -12,6 +13,20 @@ from wren.model import ( ) from wren.model.error import ErrorCode, WrenError +_TRAILING_SEMICOLONS_RE = re.compile(r"[;\s]+\Z") + + +def _strip_trailing_semicolon(sql: str) -> str: + """Strip the terminating run of ``;`` characters and surrounding whitespace. + + Matches the canner/clickhouse/trino helpers of the same name. Wrapping + user SQL as ``SELECT * FROM ({sql}) AS _q LIMIT N`` breaks when ``sql`` + ends in a semicolon — ``SELECT 1;`` is invalid inside a subquery. Only the + terminating run is removed so semicolons inside string literals + (e.g. ``SELECT ';' AS x``) are preserved. + """ + return _TRAILING_SEMICOLONS_RE.sub("", sql) + def _escape_sql(value: str) -> str: return value.replace("'", "''") @@ -73,22 +88,63 @@ class DuckDBConnector(ConnectorABC): raise def query(self, sql: str, limit: int | None = None) -> pa.Table: + """Execute ``sql`` and return the result as an Arrow table. + + When ``limit`` is provided the query is wrapped in a ``LIMIT`` clause + so only that many rows are fetched. + """ if limit is not None: - sql = f"SELECT * FROM ({sql}) AS _q LIMIT {int(limit)}" + # Strip the terminating run of ``;`` / whitespace before wrapping so + # the subquery stays valid SQL (e.g. ``SELECT 1;`` must not become + # ``SELECT * FROM (SELECT 1;) AS _q LIMIT ...``). Semicolons inside + # string literals are preserved. + stripped = _strip_trailing_semicolon(sql) + sql = f"SELECT * FROM ({stripped}) AS _q LIMIT {int(limit)}" return self.connection.execute(sql).fetch_arrow_table() def dry_run(self, sql: str) -> None: - self.connection.execute(f"EXPLAIN {sql}") + """Validate ``sql`` without returning rows or side effects. + + ``duckdb.execute`` runs semicolon-separated batches, so an ``EXPLAIN`` + prefix on raw input would still execute any trailing statements. Rather + than reject multi-statement input outright (which false-positives on + semicolons inside string literals), we neutralize it the same way the + other connectors do: wrap in a ``LIMIT 0`` subquery. Any trailing + statement then becomes a natural syntax error inside the subquery, and + no rows are materialized. + """ + stripped = _strip_trailing_semicolon(sql) + self.connection.execute(f"SELECT * FROM ({stripped}) AS _q LIMIT 0") def _attach_database(self, connection_info) -> None: + """Attach every discovered DuckDB file as a read-only database. + + Each file is attached under an alias derived from its base name. + Raises ``WrenError`` if no files are found or an attach fails. + """ db_files = self._list_duckdb_files(connection_info) if not db_files: raise WrenError(ErrorCode.DUCKDB_FILE_NOT_FOUND, "No DuckDB files found.") - for file in db_files: + # Sort for deterministic alias assignment: OpenDAL listing order is not + # guaranteed, so without this the bare alias could attach to a different + # file across runs when case-colliding names are present. + used_aliases: set[str] = set() + for file in sorted(db_files): try: escaped_file = file.replace("'", "''") - alias = os.path.splitext(os.path.basename(file))[0].replace('"', '""') + base_alias = os.path.splitext(os.path.basename(file))[0] + # Case-insensitive discovery can surface files whose names differ + # only by case (e.g. warehouse.duckdb / warehouse.DUCKDB), which + # would otherwise derive the same attach alias and collide. Make + # each alias unique deterministically. + unique_alias = base_alias + suffix = 1 + while unique_alias.lower() in used_aliases: + unique_alias = f"{base_alias}_{suffix}" + suffix += 1 + used_aliases.add(unique_alias.lower()) + alias = unique_alias.replace('"', '""') self.connection.execute( f"ATTACH DATABASE '{escaped_file}' AS \"{alias}\" (READ_ONLY);" ) @@ -98,13 +154,21 @@ class DuckDBConnector(ConnectorABC): ) def _list_duckdb_files(self, connection_info) -> list[str]: + """List DuckDB database files in the configured directory. + + Walks the connection's root directory and returns the full paths of + all non-directory entries whose name ends with ``.duckdb``. The + extension comparison is case-insensitive so files exported or renamed + with upper/mixed-case extensions (e.g. ``WAREHOUSE.DUCKDB``) are still + discovered. Raises ``WrenError`` if the directory cannot be listed. + """ op = opendal.Operator("fs", root=connection_info.url) files = [] try: for file in op.list("/"): if file.path != "/": stat = op.stat(file.path) - if not stat.mode.is_dir() and file.path.endswith(".duckdb"): + if not stat.mode.is_dir() and file.path.lower().endswith(".duckdb"): files.append(f"{connection_info.url}/{file.path}") except Exception as e: raise WrenError( @@ -113,6 +177,7 @@ class DuckDBConnector(ConnectorABC): return files def close(self) -> None: + """Close the underlying DuckDB connection, logging any error.""" try: self.connection.close() except Exception as e: diff --git a/core/wren/tests/unit/test_duckdb_file_listing.py b/core/wren/tests/unit/test_duckdb_file_listing.py new file mode 100644 index 000000000..214d1e604 --- /dev/null +++ b/core/wren/tests/unit/test_duckdb_file_listing.py @@ -0,0 +1,137 @@ +"""Regression test: DuckDB file discovery is case-insensitive on extension.""" + +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +import wren.connector.duckdb as duckdb_mod +from wren.connector.duckdb import DuckDBConnector + + +def _entry(path): + """Build a fake opendal list entry exposing only a ``path`` attribute.""" + return SimpleNamespace(path=path) + + +def _stat(is_dir): + """Build a fake opendal stat result whose ``mode.is_dir()`` returns ``is_dir``.""" + mode = MagicMock() + mode.is_dir.return_value = is_dir + return SimpleNamespace(mode=mode) + + +def test_list_duckdb_files_matches_uppercase_extension(): + # DuckDB database files are commonly named *.duckdb but the extension may + # be upper- or mixed-case (e.g. exported "WAREHOUSE.DUCKDB"). The listing + # must not drop them just because of letter case. + entries = [ + _entry("/"), + _entry("data.duckdb"), + _entry("WAREHOUSE.DUCKDB"), + _entry("Mixed.DuckDB"), + _entry("notes.txt"), + _entry("subdir/"), + ] + stat_by_path = { + "data.duckdb": _stat(False), + "WAREHOUSE.DUCKDB": _stat(False), + "Mixed.DuckDB": _stat(False), + "notes.txt": _stat(False), + "subdir/": _stat(True), + } + + fake_op = MagicMock() + fake_op.list.return_value = entries + fake_op.stat.side_effect = lambda p: stat_by_path[p] + + connection_info = SimpleNamespace(url="/tmp/dbs") + + # Bypass __init__ (it would open a real duckdb connection). + connector = DuckDBConnector.__new__(DuckDBConnector) + + with patch.object(duckdb_mod.opendal, "Operator", return_value=fake_op): + files = connector._list_duckdb_files(connection_info) + + assert files == [ + "/tmp/dbs/data.duckdb", + "/tmp/dbs/WAREHOUSE.DUCKDB", + "/tmp/dbs/Mixed.DuckDB", + ] + + +def test_attach_database_uniquifies_case_colliding_aliases(): + # Case-insensitive discovery can surface files whose basenames differ only + # by case (warehouse.duckdb / warehouse.DUCKDB). Each must attach under a + # distinct alias so the second ATTACH does not collide. + connector = DuckDBConnector.__new__(DuckDBConnector) + connector._IOException = RuntimeError + connector._HTTPException = RuntimeError + connector.connection = MagicMock() + + files = [ + "/tmp/dbs/warehouse.duckdb", + "/tmp/dbs/warehouse.DUCKDB", + "/tmp/dbs/data.duckdb", + ] + + with patch.object(connector, "_list_duckdb_files", return_value=files): + connector._attach_database(SimpleNamespace(url="/tmp/dbs")) + + executed = [c.args[0] for c in connector.connection.execute.call_args_list] + aliases = [stmt.split(' AS "')[1].split('"')[0] for stmt in executed] + assert len(aliases) == len(set(aliases)), aliases + # _attach_database sorts files for deterministic alias assignment, so the + # alphabetically-first basename ("data") attaches first; the two + # case-colliding "warehouse" files get "warehouse" and "warehouse_1". + assert aliases == ["data", "warehouse", "warehouse_1"] + + +def test_query_strips_trailing_semicolon_before_limit_wrap(): + # A semicolon-terminated statement must not produce invalid SQL such as + # ``SELECT * FROM (SELECT 1;) AS _q LIMIT 5`` when a limit is applied. + connector = DuckDBConnector.__new__(DuckDBConnector) + connector.connection = MagicMock() + connector.connection.execute.return_value.fetch_arrow_table.return_value = "tbl" + + result = connector.query("SELECT 1;", limit=5) + + executed = connector.connection.execute.call_args.args[0] + assert executed == "SELECT * FROM (SELECT 1) AS _q LIMIT 5" + assert result == "tbl" + + +def test_dry_run_wraps_in_limit_zero_subquery(): + # dry_run neutralizes multi-statement input by wrapping in a LIMIT 0 + # subquery (matching the other connectors) rather than pre-rejecting it, + # so no rows materialize and any trailing statement becomes a natural + # syntax error inside the subquery. + connector = DuckDBConnector.__new__(DuckDBConnector) + connector.connection = MagicMock() + + connector.dry_run("SELECT 1; DROP TABLE t;") + + executed = connector.connection.execute.call_args.args[0] + # The trailing terminator is stripped; the interior ``;`` stays inside the + # subquery where DuckDB rejects it as a syntax error (no side effects). + assert executed == "SELECT * FROM (SELECT 1; DROP TABLE t) AS _q LIMIT 0" + + +def test_dry_run_strips_trailing_semicolon(): + connector = DuckDBConnector.__new__(DuckDBConnector) + connector.connection = MagicMock() + + connector.dry_run("SELECT 1;") + + executed = connector.connection.execute.call_args.args[0] + assert executed == "SELECT * FROM (SELECT 1) AS _q LIMIT 0" + + +def test_dry_run_preserves_semicolon_in_string_literal(): + # A single valid statement with a semicolon inside a string literal must + # not be mangled or falsely rejected. + connector = DuckDBConnector.__new__(DuckDBConnector) + connector.connection = MagicMock() + + connector.dry_run("SELECT ';' AS x") + + executed = connector.connection.execute.call_args.args[0] + assert executed == "SELECT * FROM (SELECT ';' AS x) AS _q LIMIT 0"