mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-16 23:41:43 +00:00
fix(migrations-check): read the table name past comments, ignore referential SET DEFAULT
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
bf39aebcf1
commit
4513f78df5
2 changed files with 49 additions and 5 deletions
|
|
@ -562,13 +562,14 @@ def row_source_in(text: str) -> str | None:
|
|||
def rewrites_a_log_table(clause: str, region: str, base: int) -> str | None:
|
||||
"""The keyword to report when an `ALTER TABLE` adds a defaulted column to a request-log
|
||||
table, which Postgres 10 answers by rewriting the whole table. The table is read from the
|
||||
region rather than the masked clause, since masking blanks the quoted name in place, and each
|
||||
action of the statement is read on its own so that a `SET DEFAULT` on one column does not
|
||||
stand in for a default on a column another action adds."""
|
||||
region rather than the masked clause, since masking blanks the quoted name in place, after
|
||||
stepping over any comment sitting between `TABLE` and the name, which masking blanked as
|
||||
well. Each action of the statement is read on its own so that a `SET DEFAULT` on one column
|
||||
does not stand in for a default on a column another action adds."""
|
||||
altered = ALTERS_A_TABLE.search(clause)
|
||||
if altered is None:
|
||||
return None
|
||||
named = TABLE_NAME.match(region, base + altered.end())
|
||||
named = TABLE_NAME.match(region, skip_comments(region, base + altered.end()))
|
||||
if named is None or named.group(1).strip('"') not in REQUEST_LOG_TABLES:
|
||||
return None
|
||||
actions = strip_parens(clause[named.end() - base :]).split(",")
|
||||
|
|
@ -577,9 +578,30 @@ def rewrites_a_log_table(clause: str, region: str, base: int) -> str | None:
|
|||
return f"ADD COLUMN ... DEFAULT on {named.group(1)}"
|
||||
|
||||
|
||||
def skip_comments(sql: str, start: int) -> int:
|
||||
index = start
|
||||
while index < len(sql):
|
||||
pair = sql[index : index + 2]
|
||||
if pair == "--":
|
||||
stop = sql.find("\n", index)
|
||||
index = len(sql) if stop == -1 else stop
|
||||
elif pair == "/*":
|
||||
index = skip_block_comment(sql, index)
|
||||
elif sql[index].isspace():
|
||||
index += 1
|
||||
else:
|
||||
return index
|
||||
return index
|
||||
|
||||
|
||||
def adds_a_defaulted_column(action: str) -> bool:
|
||||
"""Whether an `ALTER TABLE` action is an `ADD COLUMN` carrying a column default. A `DEFAULT`
|
||||
right after `SET` is the referential action of an inline foreign key, which fills nothing
|
||||
in, so it does not count."""
|
||||
words = tuple(word.group().upper() for word in FIRST_WORD.finditer(action))
|
||||
return words[:1] == ("ADD",) and words[1:2] != ("CONSTRAINT",) and "DEFAULT" in words
|
||||
if words[:1] != ("ADD",) or words[1:2] == ("CONSTRAINT",):
|
||||
return False
|
||||
return any(word == "DEFAULT" and previous != "SET" for previous, word in zip(words, words[1:]))
|
||||
|
||||
|
||||
def hands_off_sql(statement: str, executed: frozenset[str]) -> bool:
|
||||
|
|
|
|||
|
|
@ -120,6 +120,28 @@ class TestDefaultedColumnsOnRequestLogTables:
|
|||
sql = 'ALTER TABLE "LiteLLM_SpendLogs" ADD COLUMN "a" TEXT, ALTER COLUMN "b" SET DEFAULT 1;'
|
||||
assert _keywords(tmp_path, sql) == ()
|
||||
|
||||
def test_a_referential_set_default_on_the_new_column_passes(self, tmp_path):
|
||||
sql = (
|
||||
'ALTER TABLE "LiteLLM_SpendLogs" ADD COLUMN "team_id" TEXT '
|
||||
'REFERENCES "LiteLLM_TeamTable"("team_id") ON DELETE SET DEFAULT;'
|
||||
)
|
||||
assert _keywords(tmp_path, sql) == ()
|
||||
|
||||
def test_a_column_default_beside_a_referential_set_default_is_flagged(self, tmp_path):
|
||||
sql = (
|
||||
'ALTER TABLE "LiteLLM_SpendLogs" ADD COLUMN "team_id" TEXT DEFAULT \'t\' '
|
||||
'REFERENCES "LiteLLM_TeamTable"("team_id") ON DELETE SET DEFAULT;'
|
||||
)
|
||||
assert _keywords(tmp_path, sql) == (SPEND_LOGS_DEFAULT,)
|
||||
|
||||
def test_a_block_comment_before_the_table_name_is_flagged(self, tmp_path):
|
||||
sql = 'ALTER TABLE /* audit */ "LiteLLM_SpendLogs" ADD COLUMN "a" TEXT DEFAULT \'x\';'
|
||||
assert _keywords(tmp_path, sql) == (SPEND_LOGS_DEFAULT,)
|
||||
|
||||
def test_a_line_comment_before_the_table_name_is_flagged(self, tmp_path):
|
||||
sql = 'ALTER TABLE IF EXISTS -- audit\n"LiteLLM_SpendLogs" ADD COLUMN "a" TEXT DEFAULT \'x\';'
|
||||
assert _keywords(tmp_path, sql) == (SPEND_LOGS_DEFAULT,)
|
||||
|
||||
def test_a_defaulted_column_among_other_actions_is_flagged(self, tmp_path):
|
||||
sql = 'ALTER TABLE "LiteLLM_SpendLogs" ADD COLUMN "a" TEXT, ADD COLUMN "b" INTEGER DEFAULT 0;'
|
||||
assert _keywords(tmp_path, sql) == (SPEND_LOGS_DEFAULT,)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue