mirror of
https://github.com/Canner/WrenAI.git
synced 2026-09-24 23:29:49 +08:00
fix(duckdb): case-insensitive .duckdb discovery, deterministic attach aliases, and single-statement SQL hardening (#2411)
Co-authored-by: Bartok9 <bartok9@users.noreply.github.com> Co-authored-by: Bartok Rage <bartok@team-of-light.local>
This commit is contained in:
co-authored by
Bartok9
Bartok Rage
parent
3122f4fbed
commit
cfbd209d90
@@ -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:
|
||||
|
||||
@@ -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"
|
||||
Reference in New Issue
Block a user