mirror of
https://github.com/Canner/WrenAI.git
synced 2026-09-24 23:29:49 +08:00
fix(policy): block table-valued functions reached via JOIN in strict mode (#2405)
This commit is contained in:
@@ -15,6 +15,21 @@ from sqlglot.errors import SqlglotError
|
||||
from wren.config import WrenConfig
|
||||
from wren.model.error import ErrorCode, ErrorPhase, WrenError
|
||||
|
||||
# Row-expansion operators (UNNEST / FLATTEN / EXPLODE) are not data sources —
|
||||
# they restructure an array/struct expression that is already in query scope
|
||||
# (e.g. a governed model column like ``orders.items``), so they never read data
|
||||
# outside the manifest and must not be treated as a disallowed table-valued
|
||||
# function. The class mapping is stable across dialects: ``UNNEST`` ->
|
||||
# ``exp.Unnest`` (trino/postgres/duckdb), Snowflake ``FLATTEN`` -> ``exp.Explode``,
|
||||
# so no per-dialect name matching is needed.
|
||||
#
|
||||
# Note: ``UNNEST(read_csv(...))`` is therefore also allowed, but that is an
|
||||
# instance of the broader pre-existing gap (data-reading TVFs in non-source
|
||||
# positions — already reachable today via projection/WHERE subqueries, since
|
||||
# strict mode only governs the top-level source position) and is tracked
|
||||
# separately. It is not a regression introduced here.
|
||||
_ROW_EXPANSION_FUNCS: tuple[type[exp.Func], ...] = (exp.Unnest, exp.Explode)
|
||||
|
||||
# Dialects we probe when canonicalising the user's denylist. sqlglot can map
|
||||
# the same function name (e.g. ``version()``) onto different concrete AST
|
||||
# subclasses depending on the dialect — postgres/mysql/duckdb/trino/clickhouse
|
||||
@@ -137,12 +152,28 @@ def _check_tables(
|
||||
phase=ErrorPhase.SQL_POLICY_CHECK,
|
||||
)
|
||||
|
||||
# Func subclasses used as FROM sources (e.g. UNNEST) produce no exp.Table
|
||||
# node at all. Scan for Func nodes inside From clauses.
|
||||
for from_clause in ast.find_all(exp.From):
|
||||
source = from_clause.this
|
||||
# Func subclasses used as query sources (e.g. UNNEST, generate_series)
|
||||
# produce no exp.Table node at all. They can appear both as the FROM
|
||||
# source AND as a JOIN source (e.g. ``orders CROSS JOIN UNNEST(items)``),
|
||||
# so scan both — checking only FROM let a table-valued function slip
|
||||
# through strict mode whenever it was reached via a JOIN.
|
||||
for clause in ast.find_all(exp.From, exp.Join):
|
||||
source = clause.this
|
||||
if isinstance(source, exp.Alias):
|
||||
source = source.this
|
||||
# LATERAL-wrapped TVFs (e.g. ``LATERAL FLATTEN(...)`` /
|
||||
# ``LATERAL generate_series(...)``) parse to an exp.Lateral node that
|
||||
# *wraps* the function rather than being an exp.Func itself, so the
|
||||
# bare-Func check below would miss them. Unwrap to inspect the inner
|
||||
# source — this is the same "TVF reached via JOIN" bug class.
|
||||
if isinstance(source, exp.Lateral):
|
||||
source = source.this
|
||||
if isinstance(source, exp.Alias):
|
||||
source = source.this
|
||||
if isinstance(source, _ROW_EXPANSION_FUNCS):
|
||||
# Row-expansion over an in-scope expression (e.g. a governed model
|
||||
# column) reads nothing outside the manifest — allow it.
|
||||
continue
|
||||
if isinstance(source, exp.Func):
|
||||
raise WrenError(
|
||||
ErrorCode.MODEL_NOT_FOUND,
|
||||
|
||||
@@ -214,16 +214,92 @@ def test_tvf_generate_series_blocked():
|
||||
assert exc_info.value.error_code == ErrorCode.MODEL_NOT_FOUND
|
||||
|
||||
|
||||
def test_tvf_unnest_blocked():
|
||||
def test_unnest_row_expansion_allowed():
|
||||
# UNNEST is a row-expansion operator, not a data source: it restructures an
|
||||
# array/struct already in query scope, so it reads nothing outside the
|
||||
# manifest and must not be blocked in strict mode.
|
||||
sql = "SELECT * FROM unnest(ARRAY[1,2,3]) AS t(x)"
|
||||
ast = parse_one(sql, dialect="duckdb")
|
||||
config = WrenConfig(strict_mode=True)
|
||||
with pytest.raises(WrenError) as exc_info:
|
||||
validate_sql_policy(ast, _MODELS, config)
|
||||
assert exc_info.value.error_code == ErrorCode.MODEL_NOT_FOUND
|
||||
validate_sql_policy(ast, _MODELS, config)
|
||||
|
||||
|
||||
def test_tvf_allowed_when_not_strict():
|
||||
ast = parse_one("SELECT * FROM read_csv('file.csv')", dialect="duckdb")
|
||||
config = WrenConfig(strict_mode=False)
|
||||
validate_sql_policy(ast, _MODELS, config)
|
||||
|
||||
|
||||
def test_unnest_model_column_allowed():
|
||||
# UNNEST over a governed model column reached via JOIN is row-expansion of
|
||||
# an in-scope column (RLAC/CLAC already apply to ``orders.items``); it must
|
||||
# NOT be false-blocked as a non-model source. UNNEST parses to exp.Unnest
|
||||
# (an exp.Func subclass) with no exp.Table node, so the FROM/JOIN func scan
|
||||
# sees it — the row-expansion allow-list lets it through.
|
||||
sql = "SELECT * FROM orders CROSS JOIN UNNEST(orders.items) AS t(item)"
|
||||
ast = parse_one(sql, dialect="trino")
|
||||
config = WrenConfig(strict_mode=True)
|
||||
validate_sql_policy(ast, _MODELS, config)
|
||||
|
||||
|
||||
def test_unnest_over_reader_tvf_allowed_known_gap():
|
||||
# Regression anchor for CodeRabbit's nested-reader question.
|
||||
#
|
||||
# ``UNNEST(read_csv(...))`` is currently ALLOWED, and intentionally so for
|
||||
# this PR. The row-expansion allow-list keys on the wrapper class
|
||||
# (exp.Unnest/exp.Explode), and the inner reader contributes no exp.Table
|
||||
# node, so the data-source guard never sees ``read_csv``. This is NOT a
|
||||
# regression introduced here: strict mode only governs the *top-level*
|
||||
# source position, so readers in non-source positions (projection, WHERE
|
||||
# subqueries, and now a row-expansion argument) were already reachable on
|
||||
# main. The collaborator (goldmedal) confirmed this and filed the broader
|
||||
# non-source-reader gap as a separate issue rather than blocking it here.
|
||||
#
|
||||
# This test pins the agreed behavior so the gap is explicit and any future
|
||||
# tightening is a deliberate, reviewed change rather than a silent flip.
|
||||
sql = "SELECT * FROM orders CROSS JOIN UNNEST(read_csv('s3://b/f.csv')) AS t(c)"
|
||||
ast = parse_one(sql, dialect="duckdb")
|
||||
config = WrenConfig(strict_mode=True)
|
||||
validate_sql_policy(ast, _MODELS, config)
|
||||
|
||||
|
||||
def test_tvf_generate_series_in_join_blocked():
|
||||
sql = "SELECT * FROM orders JOIN generate_series(1, 10) AS g(x) ON true"
|
||||
ast = parse_one(sql, dialect="postgres")
|
||||
config = WrenConfig(strict_mode=True)
|
||||
with pytest.raises(WrenError) as exc_info:
|
||||
validate_sql_policy(ast, _MODELS, config)
|
||||
assert exc_info.value.error_code == ErrorCode.MODEL_NOT_FOUND
|
||||
|
||||
|
||||
def test_lateral_flatten_model_column_allowed():
|
||||
# Snowflake FLATTEN over a governed model column, reached via a LATERAL
|
||||
# comma-join, parses to exp.Lateral wrapping exp.Explode. After unwrapping
|
||||
# the LATERAL, it is a row-expansion operator over an in-scope column — it
|
||||
# restructures ``orders.items`` rather than reading a new source, so it
|
||||
# must be allowed rather than false-blocked.
|
||||
sql = "SELECT * FROM orders, LATERAL FLATTEN(input => orders.items) f"
|
||||
ast = parse_one(sql, dialect="snowflake")
|
||||
config = WrenConfig(strict_mode=True)
|
||||
validate_sql_policy(ast, _MODELS, config)
|
||||
|
||||
|
||||
def test_tvf_lateral_generate_series_in_join_blocked():
|
||||
# generate_series IS a data-generating source (not row-expansion of an
|
||||
# in-scope column), so a LATERAL-wrapped generate_series reached via JOIN
|
||||
# must still be blocked.
|
||||
sql = "SELECT * FROM orders CROSS JOIN LATERAL generate_series(1, 3) g"
|
||||
ast = parse_one(sql, dialect="snowflake")
|
||||
config = WrenConfig(strict_mode=True)
|
||||
with pytest.raises(WrenError) as exc_info:
|
||||
validate_sql_policy(ast, _MODELS, config)
|
||||
assert exc_info.value.error_code == ErrorCode.MODEL_NOT_FOUND
|
||||
|
||||
|
||||
def test_join_between_two_mdl_models_allowed():
|
||||
# Guard against over-blocking: a plain JOIN between two manifest models
|
||||
# must still pass.
|
||||
sql = "SELECT * FROM orders o JOIN customers c ON o.customer_id = c.id"
|
||||
ast = parse_one(sql, dialect="trino")
|
||||
config = WrenConfig(strict_mode=True)
|
||||
validate_sql_policy(ast, _MODELS, config)
|
||||
|
||||
Reference in New Issue
Block a user