fix: catch a prisma routine called through a double-quoted identifier

A migration that defines a row-rewriting routine and calls it as
"backfill"() at the top level slipped past the checker, since masking
blanks double-quoted identifiers before the routine-call search runs, so
the call could not be found by name and the body read as uncalled. mask()
now returns those identifier spans and outside_definition puts them back,
so a call written through a quoted identifier reads as the call it is and
the routine's body gets scanned the same as a bare call
This commit is contained in:
mateo-berri 2026-08-24 13:34:50 -07:00
parent 9a0c541b49
commit 2a4e5fc342
2 changed files with 35 additions and 11 deletions

View file

@ -259,12 +259,17 @@ def defuse_escapes(literal: str) -> str:
return literal.replace("''", "' ")
def mask(sql: str) -> tuple[str, tuple[tuple[int, int], ...], tuple[tuple[int, int], ...]]:
def mask(
sql: str,
) -> tuple[str, tuple[tuple[int, int], ...], tuple[tuple[int, int], ...], tuple[tuple[int, int], ...]]:
"""Blank comments and quoted text, keeping offsets, and locate the spans that can still
hold SQL: dollar-quoted bodies, and the single-quoted literals `EXECUTE` runs."""
hold SQL: dollar-quoted bodies, and the single-quoted literals `EXECUTE` runs. Also locate
the double-quoted identifiers, so a routine called through one can be found by name even
though the call is blanked here the way every other quoted run of text is."""
chunks: list[str] = []
bodies: list[tuple[int, int]] = []
literals: list[tuple[int, int]] = []
identifiers: list[tuple[int, int]] = []
index = 0
length = len(sql)
@ -291,6 +296,8 @@ def mask(sql: str) -> tuple[str, tuple[tuple[int, int], ...], tuple[tuple[int, i
if character == "'":
closed = sql[stop - 1 : stop] == character
literals.append((index + 1, max(index + 1, stop - 1 if closed else stop)))
else:
identifiers.append((index, stop))
chunks.append(blank(sql[index:stop]))
index = stop
continue
@ -309,7 +316,7 @@ def mask(sql: str) -> tuple[str, tuple[tuple[int, int], ...], tuple[tuple[int, i
chunks.append(character)
index += 1
return "".join(chunks), tuple(bodies), tuple(literals)
return "".join(chunks), tuple(bodies), tuple(literals), tuple(identifiers)
def skip_block_comment(sql: str, start: int) -> int:
@ -646,7 +653,7 @@ def scan_region(
before it is scanned, so a `--` or `/*` in one of its nested strings blanks nothing and the
statement after it stays visible, and since that keeps every character on its offset, the
statement reports its true file line and lines up with the markers."""
masked, bodies, literals = mask(region)
masked, bodies, literals, identifiers = mask(region)
executed = executed_names(masked)
runnable = executed_literals(masked, literals, executed)
@ -672,7 +679,7 @@ def scan_region(
yield Violation(migration, line_of(document, offset + keyword_start(clause, base)), keyword)
for body in bodies:
if not runs_when_applied(masked, region, bodies, runnable, body):
if not runs_when_applied(masked, region, bodies, runnable, identifiers, body):
continue
start, end = body
yield from scan_region(document, region[start:end], migration, markers, offset + start)
@ -702,6 +709,7 @@ def runs_when_applied(
region: str,
bodies: tuple[tuple[int, int], ...],
runnable: tuple[tuple[int, int], ...],
identifiers: tuple[tuple[int, int], ...],
body: tuple[int, int],
) -> bool:
"""Whether a dollar-quoted body runs while the migration is being applied. A `DO` block runs
@ -713,9 +721,11 @@ def runs_when_applied(
the same migration names the routine anywhere outside the definition. The definition is
found in the masked text, where one written inside a comment has already been blanked, and
the name is read from the region at those same offsets, since masking blanks a quoted
identifier in place. A name that needed those quotes is blanked at its call sites too and
so can never be found there, which would read as uncalled however the migration runs it,
and the body is read rather than trusted."""
identifier in place. A call written as a quoted identifier is blanked there too, and
`\"backfill\"()` is the same call as `backfill()` in Postgres, so the double-quoted identifiers
are put back before the search and a routine invoked through one is found. A definition whose
own name needs those quotes is read rather than trusted, since matching such a name once it is
put back in the open would be unreliable."""
start, end = body
opens = masked.rfind(";", 0, start) + 1
defined = DEFINES_A_ROUTINE.search(masked, opens, start)
@ -724,7 +734,8 @@ def runs_when_applied(
named = ROUTINE_NAME.match(region, defined.end(), start)
if named is None or named.group(1).startswith('"'):
return True
return contains(outside_definition(masked, region, bodies, runnable, opens, end), re.escape(named.group(1)))
restored = outside_definition(masked, region, bodies, runnable, identifiers, opens, end)
return contains(restored, re.escape(named.group(1)))
def outside_definition(
@ -732,6 +743,7 @@ def outside_definition(
region: str,
bodies: tuple[tuple[int, int], ...],
runnable: tuple[tuple[int, int], ...],
identifiers: tuple[tuple[int, int], ...],
opens: int,
closes: int,
) -> str:
@ -744,13 +756,17 @@ def outside_definition(
runs one as SQL and the call can be written inside it. A single-quoted payload is undoubled as
it goes back, so a `--` or `/*` in one of its nested strings blanks nothing and the call after
it stays visible, and it is padded to the span it fills so the later offsets still land. The
definition is blanked after they are restored, which takes its own body with it, so a routine
that names itself recursively does not thereby count as called."""
double-quoted identifiers come back verbatim, so a routine invoked as `\"backfill\"()` reads as
the call it is. The definition is blanked after they are restored, which takes its own body and
any identifier standing inside it with it, so a routine that names itself recursively does not
thereby count as called."""
text = list(masked)
for start, end in bodies:
text[start:end] = without_comments(region[start:end])
for start, end in runnable:
text[start:end] = without_comments(undouble(region[start:end])).ljust(end - start)
for start, end in identifiers:
text[start:end] = region[start:end]
text[opens:closes] = blank(region[opens:closes])
return "".join(text)

View file

@ -417,6 +417,14 @@ class TestStoredRoutines:
sql = self.DEFINITION + "DO $$ BEGIN PERFORM backfill(); END; $$;\n"
assert _keywords(tmp_path, sql) == ("UPDATE",)
def test_a_call_through_a_quoted_identifier_still_counts(self, tmp_path):
sql = self.DEFINITION + 'SELECT "backfill"();\n'
assert _keywords(tmp_path, sql) == ("UPDATE",)
def test_an_unrelated_quoted_identifier_is_not_a_call(self, tmp_path):
sql = self.DEFINITION + 'SELECT "other"();\n'
assert _keywords(tmp_path, sql) == ()
def test_a_trigger_wiring_the_function_up_counts_as_a_call(self, tmp_path):
sql = self.DEFINITION + 'CREATE TRIGGER t AFTER INSERT ON "Foo" EXECUTE FUNCTION backfill();\n'
assert _keywords(tmp_path, sql) == ("UPDATE",)