fix: flag an INSERT whose rows come from a parenthesised query

Postgres takes the row source parenthesised, so `INSERT INTO "t" ("a") (SELECT
...)` copies a whole table at boot. Reading only the unparenthesised text let it
through: 777eb8af10 caught it, then e7dea842c3 traded it away to stop a VALUES
list joined to a query by a set operation from bounding nothing.

Read the top level first so set operations still count, then fall back to the
whole statement when no top-level VALUES bounds the insert. `TABLE t` is a row
source as much as a `SELECT` is, and it was passing too
This commit is contained in:
mateo-berri 2026-08-21 18:58:34 -07:00
parent aabaa5151b
commit 64267ebd28
2 changed files with 63 additions and 13 deletions

View file

@ -11,10 +11,12 @@ Flagged, per statement, by its leading keyword:
UPDATE rewrites every matching row, and `WHERE` does not bound the scan
DELETE same scan, and the dead tuples outlive the migration
MERGE both of the above in one statement
INSERT only when it draws rows from a `SELECT`; an insert whose row source is
a leading `VALUES` is bounded by the rows spelled out there and passes,
scalar subqueries in that list included, while a `VALUES` reached
through a subquery or a set operation bounds nothing
INSERT only when its rows come from a query rather than a literal `VALUES`
list. The query counts wherever it sits, since Postgres takes it
parenthesised, and `TABLE t` is one as much as a `SELECT` is. An
insert bounded by a leading `VALUES` passes, scalar subqueries in that
list included, while a `VALUES` reached through a subquery or joined
to a query by a set operation bounds nothing
WITH a CTE-led statement containing any of the above
Referential actions (`ON DELETE CASCADE`, `ON UPDATE CASCADE`) are schema, never a
@ -292,24 +294,38 @@ def offending_keyword(statement: str) -> str | None:
return keyword
if keyword == "INSERT":
return "INSERT ... SELECT" if draws_rows_from_a_select(statement) else None
source = row_source_keyword(statement)
return None if source is None else f"INSERT ... {source}"
if keyword == "WITH":
nested = next((name for name in sorted(REWRITES_ROWS) if contains(statement, name)), None)
if nested is not None:
return f"WITH ... {nested}"
if contains(statement, "INSERT") and draws_rows_from_a_select(statement):
return "WITH ... INSERT ... SELECT"
if contains(statement, "INSERT"):
source = row_source_keyword(statement)
if source is not None:
return f"WITH ... INSERT ... {source}"
return None
def draws_rows_from_a_select(statement: str) -> bool:
"""Whether an `INSERT` takes its rows from a query rather than a literal list. Only a
`SELECT` the insert is built on counts, so the scalar subqueries and helper CTEs that
sit in parentheses around a `VALUES` list do not make it a rewrite, while one reached
through a set operation does."""
return contains(strip_parens(statement), "SELECT")
def row_source_keyword(statement: str) -> str | None:
"""Which keyword supplies an `INSERT` its rows, or `None` when a literal `VALUES` list
does. A query outside every parenthesis is the row source outright, including one a set
operation joins to a `VALUES` list. Failing that, a `VALUES` outside every parenthesis
is itself the row source, so the scalar subqueries and helper CTEs nested within that
list do not make the insert a rewrite. Failing both, the rows come from a parenthesised
query, which Postgres accepts and which reading only the unparenthesised text would let
through: `INSERT INTO "t" ("a") (SELECT ...)` copies a whole table."""
outer = strip_parens(statement)
joined = row_source_in(outer)
if joined is not None:
return joined
return None if contains(outer, "VALUES") else row_source_in(statement)
def row_source_in(text: str) -> str | None:
return next((word for word in ("SELECT", "TABLE") if contains(text, word)), None)
def hands_off_sql(statement: str, executed: frozenset[str]) -> bool:

View file

@ -125,6 +125,32 @@ class TestInsert:
sql = 'INSERT INTO "Foo" ("id") VALUES (1), (2) UNION ALL SELECT "id" FROM "Bar";'
assert _keywords(tmp_path, sql) == ("INSERT ... SELECT",)
def test_a_parenthesised_select_row_source_is_flagged(self, tmp_path):
sql = 'INSERT INTO "Foo" ("id") (SELECT "id" FROM "Bar");'
assert _keywords(tmp_path, sql) == ("INSERT ... SELECT",)
def test_a_parenthesised_select_row_source_without_a_column_list_is_flagged(self, tmp_path):
sql = 'INSERT INTO "Foo" (SELECT "id" FROM "Bar");'
assert _keywords(tmp_path, sql) == ("INSERT ... SELECT",)
def test_a_parenthesised_select_row_source_spanning_lines_is_flagged(self, tmp_path):
sql = 'INSERT INTO "Foo" ("id")\n(\n SELECT "id" FROM "Bar"\n);'
assert _keywords(tmp_path, sql) == ("INSERT ... SELECT",)
def test_a_parenthesised_select_over_a_values_list_is_flagged(self, tmp_path):
sql = 'INSERT INTO "Foo" ("id") (SELECT * FROM (VALUES (1), (2)) AS "v"("id"));'
assert _keywords(tmp_path, sql) == ("INSERT ... SELECT",)
def test_a_set_operation_over_parenthesised_selects_is_flagged(self, tmp_path):
sql = 'INSERT INTO "Foo" ("id") (SELECT 1) UNION (SELECT 2);'
assert _keywords(tmp_path, sql) == ("INSERT ... SELECT",)
def test_a_table_row_source_is_flagged(self, tmp_path):
assert _keywords(tmp_path, 'INSERT INTO "Foo" TABLE "Bar";') == ("INSERT ... TABLE",)
def test_a_table_named_in_the_insert_target_does_not_flag_it(self, tmp_path):
assert _keywords(tmp_path, 'INSERT INTO "audit table" ("id") VALUES (1);') == ()
class TestCommonTableExpressions:
def test_cte_led_update_is_flagged(self, tmp_path):
@ -139,6 +165,14 @@ class TestCommonTableExpressions:
sql = 'WITH batch AS (SELECT "id" FROM "Bar") INSERT INTO "Foo" ("id") SELECT "id" FROM batch;'
assert _keywords(tmp_path, sql) == ("WITH ... INSERT ... SELECT",)
def test_cte_led_insert_from_a_parenthesised_select_is_flagged(self, tmp_path):
sql = 'WITH batch AS (SELECT "id" FROM "Bar") INSERT INTO "Foo" ("id") (SELECT "id" FROM batch);'
assert _keywords(tmp_path, sql) == ("WITH ... INSERT ... SELECT",)
def test_cte_led_insert_into_a_values_list_passes(self, tmp_path):
sql = 'WITH batch AS (SELECT max("id") FROM "Bar") INSERT INTO "Foo" ("id") VALUES (1);'
assert _keywords(tmp_path, sql) == ()
def test_read_only_cte_passes(self, tmp_path):
sql = 'WITH batch AS (SELECT "id" FROM "Foo") SELECT count(*) FROM batch;'
assert _keywords(tmp_path, sql) == ()