fix(policy): block table-valued functions reached via JOIN in strict mode (#2405)

This commit is contained in:
Bartok
2026-06-29 17:52:34 +08:00
committed by GitHub
parent 1be4baaeba
commit a2a37b3955
2 changed files with 115 additions and 8 deletions
+35 -4
View File
@@ -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,
+80 -4
View File
@@ -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)