Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
45 changes: 43 additions & 2 deletions sqlite_utils/db.py
Original file line number Diff line number Diff line change
Expand Up @@ -1074,6 +1074,10 @@ def quote_default_value(self, value: str) -> str:
):
return value

if str(value).upper().startswith("X'") and str(value).endswith("'"):
# Blob literal, e.g. X'FF'
return value

if str(value).upper() in ("CURRENT_TIME", "CURRENT_DATE", "CURRENT_TIMESTAMP"):
return value

Expand All @@ -1083,11 +1087,46 @@ def quote_default_value(self, value: str) -> str:
return value

if str(value).endswith(")"):
# Expr
# Expr - already parenthesized values pass through unchanged
if str(value).startswith("("):
return value
return f"({value})"

return self.quote(value)

def transform_default_fragment(self, value: str) -> str:
"""Return a re-emittable DEFAULT fragment for a PRAGMA-returned value.

PRAGMA strips the parentheses from an expression default (``DEFAULT
(1+2)`` is reported as ``1+2``), so expressions must be wrapped again
or SQLite would treat them as string literals. Literals and keywords
are passed through unchanged.
"""
s = str(value)
if s.startswith("'") and s.endswith("'"):
# A single quoted string literal
if "''" not in s[1:-1] and "'" not in s[1:-1]:
return s
if s.upper().startswith("X'") and s.endswith("'"):
# Blob literal, e.g. X'FF'
return s
if s.upper() in (
"TRUE",
"FALSE",
"NULL",
"CURRENT_TIME",
"CURRENT_DATE",
"CURRENT_TIMESTAMP",
):
return s
try:
float(s)
return s
except ValueError:
pass
# Anything else was a parenthesized expression, e.g. "1+2" or "'a'||'b'"
return f"({s})"

def table_names(self, fts4: bool = False, fts5: bool = False) -> list[str]:
"""
List of string table names in this database.
Expand Down Expand Up @@ -3061,7 +3100,9 @@ def fk_with_renamed_columns(fk: ForeignKey) -> ForeignKey:
)
# defaults=
create_table_defaults = {
(rename.get(c.name) or c.name): c.default_value
(rename.get(c.name) or c.name): self.db.transform_default_fragment(
c.default_value
)
for c in self.columns
if c.default_value is not None and c.name not in drop
}
Expand Down
29 changes: 29 additions & 0 deletions tests/test_transform.py
Original file line number Diff line number Diff line change
Expand Up @@ -288,6 +288,35 @@ def test_transform_preserves_keyword_literal_defaults(fresh_db):
assert after == (1, 0, None)


def test_transform_preserves_expression_defaults(fresh_db):
# PRAGMA strips the parentheses from expression defaults, so transform()
# used to re-emit them as string literals: DEFAULT (1+2) became DEFAULT
# '1+2' and every new row silently got the text '1+2' instead of 3.
fresh_db.execute(
"CREATE TABLE t ("
" a TEXT,"
" b INTEGER DEFAULT (1+2),"
" c BLOB DEFAULT (X'ff'),"
" d TEXT DEFAULT ('x' || 'y')"
")"
)
table = fresh_db.table("t")
table.insert({"a": "first"})

# Rebuild via an unrelated change.
table.transform(rename={"a": "aa"})

schema = table.schema
assert "DEFAULT (1+2)" in schema
assert "DEFAULT X'ff'" in schema
assert "DEFAULT ('x' || 'y')" in schema

table.insert({"aa": "second"})
assert fresh_db.execute(
"SELECT b, c, d FROM t WHERE aa = 'second'"
).fetchone() == (3, b"\xff", "xy")


def test_transform_not_null(fresh_db):
dogs = fresh_db.table("dogs")
dogs.insert({"id": 1, "name": "Cleo", "age": "5"}, pk="id")
Expand Down
Loading