Open
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
37 changes: 37 additions & 0 deletions sqlmodel/_compat.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -9,6 +9,7 @@
Annotated,
Any,
ForwardRef,
Literal,
TypeAlias,
TypeVar,
Union,
Expand DownExpand Up@@ -65,6 +66,36 @@ def _is_union_type(t: Any) -> bool:
return t is UnionType or t is Union


def get_literal_annotation_info(
annotation: Any,
) -> tuple[type[Any], tuple[Any, ...]] | None:
if annotation is None or get_origin(annotation) is None:
return None
origin = get_origin(annotation)
if origin is Annotated:
return get_literal_annotation_info(get_args(annotation)[0])
if _is_union_type(origin):
bases = get_args(annotation)
if len(bases) > 2:
raise ValueError("Cannot have a Union with more than 2 members")
if bases[0] is not NoneType and bases[1] is not NoneType:
raise ValueError("Cannot have a Union without None")
use_type = bases[0] if bases[0] is not NoneType else bases[1]
return get_literal_annotation_info(use_type)
if origin is Literal:
literal_args = get_args(annotation)
if not literal_args:
return None
if all(isinstance(arg, bool) for arg in literal_args): # all bools
base_type: type[Any] = bool
elif all(isinstance(arg, int) for arg in literal_args): # all ints
base_type = int
else:
base_type = str
return base_type, tuple(literal_args)
return None


finish_init: ContextVar[bool] = ContextVar("finish_init", default=True)


Expand DownExpand Up@@ -190,6 +221,12 @@ def get_sa_type_from_type_annotation(annotation: Any) -> Any:
# Optional unions are allowed
use_type = bases[0] if bases[0] is not NoneType else bases[1]
return get_sa_type_from_type_annotation(use_type)
if origin is Literal:
literal_info = get_literal_annotation_info(annotation)
if literal_info is None:
raise ValueError("Literal without values is not supported")
base_type, _ = literal_info
return base_type
return origin


Expand Down
27 changes: 27 additions & 0 deletions sqlmodel/main.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -27,6 +27,7 @@
from pydantic.fields import FieldInfo as PydanticFieldInfo
from sqlalchemy import (
Boolean,
CheckConstraint,
Column,
Date,
DateTime,
Expand DownExpand Up@@ -63,6 +64,7 @@
finish_init,
get_annotations,
get_field_metadata,
get_literal_annotation_info,
get_model_fields,
get_relationship_to,
get_sa_type_from_field,
Expand DownExpand Up@@ -680,6 +682,31 @@ def __init__(
# Ref: https://github.com/sqlalchemy/sqlalchemy/commit/428ea01f00a9cc7f85e435018565eb6da7af1b77
# Tag: 1.4.36
DeclarativeMeta.__init__(cls, classname, bases, dict_, **kw)
table = getattr(cls, "__table__", None)
if table is not None:
# Attach Literal-based value constraints at the database level
for field_name, field in get_model_fields(cls).items(): # type: ignore
annotation = getattr(field, "annotation", None)
literal_info = get_literal_annotation_info(annotation)
if literal_info is None:
continue
base_type, values = literal_info
assert base_type in (str, int, bool)
column = table.c.get(field_name)
if column is None:
continue
if base_type is int:
coerced_values = tuple(int(v) for v in values)
elif base_type is bool:
coerced_values = tuple(bool(v) for v in values)
else:
coerced_values = tuple(str(v) for v in values) # type: ignore
constraint_name = f"ck_{table.name}_{field_name}_literal"
constraint = CheckConstraint(
column.in_(coerced_values),
name=constraint_name,
)
table.append_constraint(constraint)
else:
ModelMetaclass.__init__(cls, classname, bases, dict_, **kw)

Expand Down
147 changes: 146 additions & 1 deletion tests/test_main.py
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
from typing import Annotated
from typing import Annotated, Literal

import pytest
from sqlalchemy import text
from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import RelationshipProperty
from sqlmodel import Field, Relationship, Session, SQLModel, create_engine, select
Expand DownExpand Up@@ -216,3 +217,147 @@ class Hero(SQLModel, table=True):
assert len(foreign_keys) == 1
assert foreign_keys[0].ondelete == "CASCADE"
assert team_id_column.nullable is False


def test_literal_valid_values(clear_sqlmodel, caplog):
"""Test https://github.com/fastapi/sqlmodel/issues/57"""

class Model(SQLModel, table=True):
id: int | None = Field(default=None, primary_key=True)
all_str: Literal["a", "b", "c"]
mixed: Literal["yes", "no", 1, 0]
all_int: Literal[1, 2, 3]
int_bool: Literal[0, 1, True, False]
all_bool: Literal[True, False]

obj = Model(
all_str="a",
mixed="yes",
all_int=1,
int_bool=True,
all_bool=False,
)

engine = create_engine("sqlite://", echo=True)

SQLModel.metadata.create_all(engine)

# Check DDL
assert "all_str VARCHAR NOT NULL" in caplog.text
assert "mixed VARCHAR NOT NULL" in caplog.text
assert "all_int INTEGER NOT NULL" in caplog.text
assert "int_bool INTEGER NOT NULL" in caplog.text
assert "all_bool BOOLEAN NOT NULL" in caplog.text

# Check query
with Session(engine) as session:
session.add(obj)
session.commit()
session.refresh(obj)
assert isinstance(obj.all_str, str)
assert obj.all_str == "a"
assert isinstance(obj.mixed, str)
assert obj.mixed == "yes"
assert isinstance(obj.all_int, int)
assert obj.all_int == 1
assert isinstance(obj.int_bool, int)
assert obj.int_bool == 1
assert isinstance(obj.all_bool, bool)
assert obj.all_bool is False


def test_literal_constraints_invalid_values(clear_sqlmodel):
"""DB should reject values that are not part of the Literal choices."""

class Model(SQLModel, table=True):
id: int | None = Field(default=None, primary_key=True)
all_str: Literal["a", "b", "c"]
mixed: Literal["yes", "no", 1, 0]
all_int: Literal[1, 2, 3]
int_bool: Literal[0, 1, True, False]
all_bool: Literal[True, False]

engine = create_engine("sqlite://")
SQLModel.metadata.create_all(engine)

# Helper to attempt a raw insert that bypasses Pydantic validation so we
# can verify that the database-level CHECK constraints are enforced.
def insert_raw(values: dict[str, object]) -> None:
stmt = text(
"INSERT INTO model (all_str, mixed, all_int, int_bool, all_bool) "
"VALUES (:all_str, :mixed, :all_int, :int_bool, :all_bool)"
).bindparams(**values)
with pytest.raises(IntegrityError):
with Session(engine) as session:
session.exec(stmt)
session.commit()

# Invalid string literal for all_str
insert_raw(
{
"all_str": "z", # invalid, not in {"a","b","c"}
"mixed": "yes",
"all_int": 1,
"int_bool": 1,
"all_bool": 0,
}
)

# Invalid int literal for all_int
insert_raw(
{
"all_str": "a",
"mixed": "yes",
"all_int": 5, # invalid, not in {1,2,3}
"int_bool": 1,
"all_bool": 0,
}
)

# Invalid bool literal for all_bool
insert_raw(
{
"all_str": "a",
"mixed": "yes",
"all_int": 1,
"int_bool": 1,
"all_bool": 2, # invalid boolean value
}
)


def test_literal_optional_and_union_constraints(clear_sqlmodel):
"""Literals inside Optional/Union should also be enforced at the DB level."""

class Model(SQLModel, table=True):
id: int | None = Field(default=None, primary_key=True)
opt_str: Literal["x", "y"] | None = None
union_int: Literal[10, 20] | None = None

engine = create_engine("sqlite://")
SQLModel.metadata.create_all(engine)

# Valid values should be accepted
obj = Model(opt_str="x", union_int=10)
with Session(engine) as session:
session.add(obj)
session.commit()
session.refresh(obj)
assert obj.opt_str == "x"
assert obj.union_int == 10

# Invalid values should be rejected by the database
def insert_raw(values: dict[str, object]) -> None:
stmt = text(
"INSERT INTO model (opt_str, union_int) VALUES (:opt_str, :union_int)"
).bindparams(**values)
with pytest.raises(IntegrityError):
with Session(engine) as session:
session.exec(stmt)
session.commit()

# opt_str not in {"x", "y"}
insert_raw({"opt_str": "z", "union_int": 10})

# union_int not in {10, 20}
insert_raw({"opt_str": "x", "union_int": 30})
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Add copy buttons to all
 blocks\n(function() {\n function addCopyButtons() {\n document.querySelectorAll('pre code').forEach(function(codeBlock) {\n if (codeBlock.parentElement.hasAttribute('data-copy-added')) return;\n codeBlock.parentElement.setAttribute('data-copy-added', 'true');\n \n var btn = document.createElement('button');\n btn.textContent = 'Copy';\n btn.style.cssText = 'position:absolute;top:4px;right:4px;padding:2px 8px;font-size:11px;background:#4ecdc4;border:none;border-radius:4px;color:#1a1a2e;cursor:pointer;opacity:0.7;transition:opacity 0.2s;';\n btn.onmouseover = function() { this.style.opacity = '1'; };\n btn.onmouseout = function() { this.style.opacity = '0.7'; };\n btn.onclick = function() {\n navigator.clipboard.writeText(codeBlock.textContent).then(function() {\n btn.textContent = 'Copied!';\n setTimeout(function() { btn.textContent = 'Copy'; }, 1500);\n });\n };\n codeBlock.parentElement.style.position = 'relative';\n codeBlock.parentElement.appendChild(btn);\n });\n }\n \n addCopyButtons();\n \n // Re-run on dynamic content\n var observer = new MutationObserver(addCopyButtons);\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Add Copy Buttons to Code Blocks");
}
} catch(__e) { console.warn('[Userscript:Add Copy Buttons to Code Blocks]', __e); }
})();
(function(){
try {
var __m = "github.com";
var __re = new RegExp('^' + "github\\.com" + '
Skip to content
Open
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
37 changes: 37 additions & 0 deletions sqlmodel/_compat.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -9,6 +9,7 @@
Annotated,
Any,
ForwardRef,
Literal,
TypeAlias,
TypeVar,
Union,
Expand DownExpand Up@@ -65,6 +66,36 @@ def _is_union_type(t: Any) -> bool:
return t is UnionType or t is Union


def get_literal_annotation_info(
annotation: Any,
) -> tuple[type[Any], tuple[Any, ...]] | None:
if annotation is None or get_origin(annotation) is None:
return None
origin = get_origin(annotation)
if origin is Annotated:
return get_literal_annotation_info(get_args(annotation)[0])
if _is_union_type(origin):
bases = get_args(annotation)
if len(bases) > 2:
raise ValueError("Cannot have a Union with more than 2 members")
if bases[0] is not NoneType and bases[1] is not NoneType:
raise ValueError("Cannot have a Union without None")
use_type = bases[0] if bases[0] is not NoneType else bases[1]
return get_literal_annotation_info(use_type)
if origin is Literal:
literal_args = get_args(annotation)
if not literal_args:
return None
if all(isinstance(arg, bool) for arg in literal_args): # all bools
base_type: type[Any] = bool
elif all(isinstance(arg, int) for arg in literal_args): # all ints
base_type = int
else:
base_type = str
return base_type, tuple(literal_args)
return None


finish_init: ContextVar[bool] = ContextVar("finish_init", default=True)


Expand DownExpand Up@@ -190,6 +221,12 @@ def get_sa_type_from_type_annotation(annotation: Any) -> Any:
# Optional unions are allowed
use_type = bases[0] if bases[0] is not NoneType else bases[1]
return get_sa_type_from_type_annotation(use_type)
if origin is Literal:
literal_info = get_literal_annotation_info(annotation)
if literal_info is None:
raise ValueError("Literal without values is not supported")
base_type, _ = literal_info
return base_type
return origin


Expand Down
27 changes: 27 additions & 0 deletions sqlmodel/main.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -27,6 +27,7 @@
from pydantic.fields import FieldInfo as PydanticFieldInfo
from sqlalchemy import (
Boolean,
CheckConstraint,
Column,
Date,
DateTime,
Expand DownExpand Up@@ -63,6 +64,7 @@
finish_init,
get_annotations,
get_field_metadata,
get_literal_annotation_info,
get_model_fields,
get_relationship_to,
get_sa_type_from_field,
Expand DownExpand Up@@ -680,6 +682,31 @@ def __init__(
# Ref: https://github.com/sqlalchemy/sqlalchemy/commit/428ea01f00a9cc7f85e435018565eb6da7af1b77
# Tag: 1.4.36
DeclarativeMeta.__init__(cls, classname, bases, dict_, **kw)
table = getattr(cls, "__table__", None)
if table is not None:
# Attach Literal-based value constraints at the database level
for field_name, field in get_model_fields(cls).items(): # type: ignore
annotation = getattr(field, "annotation", None)
literal_info = get_literal_annotation_info(annotation)
if literal_info is None:
continue
base_type, values = literal_info
assert base_type in (str, int, bool)
column = table.c.get(field_name)
if column is None:
continue
if base_type is int:
coerced_values = tuple(int(v) for v in values)
elif base_type is bool:
coerced_values = tuple(bool(v) for v in values)
else:
coerced_values = tuple(str(v) for v in values) # type: ignore
constraint_name = f"ck_{table.name}_{field_name}_literal"
constraint = CheckConstraint(
column.in_(coerced_values),
name=constraint_name,
)
table.append_constraint(constraint)
else:
ModelMetaclass.__init__(cls, classname, bases, dict_, **kw)

Expand Down
147 changes: 146 additions & 1 deletion tests/test_main.py
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
from typing import Annotated
from typing import Annotated, Literal

import pytest
from sqlalchemy import text
from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import RelationshipProperty
from sqlmodel import Field, Relationship, Session, SQLModel, create_engine, select
Expand DownExpand Up@@ -216,3 +217,147 @@ class Hero(SQLModel, table=True):
assert len(foreign_keys) == 1
assert foreign_keys[0].ondelete == "CASCADE"
assert team_id_column.nullable is False


def test_literal_valid_values(clear_sqlmodel, caplog):
"""Test https://github.com/fastapi/sqlmodel/issues/57"""

class Model(SQLModel, table=True):
id: int | None = Field(default=None, primary_key=True)
all_str: Literal["a", "b", "c"]
mixed: Literal["yes", "no", 1, 0]
all_int: Literal[1, 2, 3]
int_bool: Literal[0, 1, True, False]
all_bool: Literal[True, False]

obj = Model(
all_str="a",
mixed="yes",
all_int=1,
int_bool=True,
all_bool=False,
)

engine = create_engine("sqlite://", echo=True)

SQLModel.metadata.create_all(engine)

# Check DDL
assert "all_str VARCHAR NOT NULL" in caplog.text
assert "mixed VARCHAR NOT NULL" in caplog.text
assert "all_int INTEGER NOT NULL" in caplog.text
assert "int_bool INTEGER NOT NULL" in caplog.text
assert "all_bool BOOLEAN NOT NULL" in caplog.text

# Check query
with Session(engine) as session:
session.add(obj)
session.commit()
session.refresh(obj)
assert isinstance(obj.all_str, str)
assert obj.all_str == "a"
assert isinstance(obj.mixed, str)
assert obj.mixed == "yes"
assert isinstance(obj.all_int, int)
assert obj.all_int == 1
assert isinstance(obj.int_bool, int)
assert obj.int_bool == 1
assert isinstance(obj.all_bool, bool)
assert obj.all_bool is False


def test_literal_constraints_invalid_values(clear_sqlmodel):
"""DB should reject values that are not part of the Literal choices."""

class Model(SQLModel, table=True):
id: int | None = Field(default=None, primary_key=True)
all_str: Literal["a", "b", "c"]
mixed: Literal["yes", "no", 1, 0]
all_int: Literal[1, 2, 3]
int_bool: Literal[0, 1, True, False]
all_bool: Literal[True, False]

engine = create_engine("sqlite://")
SQLModel.metadata.create_all(engine)

# Helper to attempt a raw insert that bypasses Pydantic validation so we
# can verify that the database-level CHECK constraints are enforced.
def insert_raw(values: dict[str, object]) -> None:
stmt = text(
"INSERT INTO model (all_str, mixed, all_int, int_bool, all_bool) "
"VALUES (:all_str, :mixed, :all_int, :int_bool, :all_bool)"
).bindparams(**values)
with pytest.raises(IntegrityError):
with Session(engine) as session:
session.exec(stmt)
session.commit()

# Invalid string literal for all_str
insert_raw(
{
"all_str": "z", # invalid, not in {"a","b","c"}
"mixed": "yes",
"all_int": 1,
"int_bool": 1,
"all_bool": 0,
}
)

# Invalid int literal for all_int
insert_raw(
{
"all_str": "a",
"mixed": "yes",
"all_int": 5, # invalid, not in {1,2,3}
"int_bool": 1,
"all_bool": 0,
}
)

# Invalid bool literal for all_bool
insert_raw(
{
"all_str": "a",
"mixed": "yes",
"all_int": 1,
"int_bool": 1,
"all_bool": 2, # invalid boolean value
}
)


def test_literal_optional_and_union_constraints(clear_sqlmodel):
"""Literals inside Optional/Union should also be enforced at the DB level."""

class Model(SQLModel, table=True):
id: int | None = Field(default=None, primary_key=True)
opt_str: Literal["x", "y"] | None = None
union_int: Literal[10, 20] | None = None

engine = create_engine("sqlite://")
SQLModel.metadata.create_all(engine)

# Valid values should be accepted
obj = Model(opt_str="x", union_int=10)
with Session(engine) as session:
session.add(obj)
session.commit()
session.refresh(obj)
assert obj.opt_str == "x"
assert obj.union_int == 10

# Invalid values should be rejected by the database
def insert_raw(values: dict[str, object]) -> None:
stmt = text(
"INSERT INTO model (opt_str, union_int) VALUES (:opt_str, :union_int)"
).bindparams(**values)
with pytest.raises(IntegrityError):
with Session(engine) as session:
session.exec(stmt)
session.commit()

# opt_str not in {"x", "y"}
insert_raw({"opt_str": "z", "union_int": 10})

# union_int not in {10, 20}
insert_raw({"opt_str": "x", "union_int": 30})
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Force GitHub README to respect dark mode\n(function() {\n var style = document.createElement('style');\n style.textContent = '\n .markdown-body {\n color-scheme: dark light;\n }\n .markdown-body pre { background: #161b22 !important; }\n .markdown-body code { background: rgba(110, 118, 129, 0.4) !important; }\n .markdown-body table th, .markdown-body table td { border-color: #30363d !important; }\n .markdown-body img { background: #0d1117; }\n .markdown-body blockquote { border-left-color: #8b949e; }\n .markdown-body hr { border-color: #30363d; }\n ';\n document.head.appendChild(style);\n})();", "GitHub Dark Mode README Fix"); } } catch(__e) { console.warn('[Userscript:GitHub Dark Mode README Fix]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Open
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
37 changes: 37 additions & 0 deletions sqlmodel/_compat.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -9,6 +9,7 @@
Annotated,
Any,
ForwardRef,
Literal,
TypeAlias,
TypeVar,
Union,
Expand DownExpand Up@@ -65,6 +66,36 @@ def _is_union_type(t: Any) -> bool:
return t is UnionType or t is Union


def get_literal_annotation_info(
annotation: Any,
) -> tuple[type[Any], tuple[Any, ...]] | None:
if annotation is None or get_origin(annotation) is None:
return None
origin = get_origin(annotation)
if origin is Annotated:
return get_literal_annotation_info(get_args(annotation)[0])
if _is_union_type(origin):
bases = get_args(annotation)
if len(bases) > 2:
raise ValueError("Cannot have a Union with more than 2 members")
if bases[0] is not NoneType and bases[1] is not NoneType:
raise ValueError("Cannot have a Union without None")
use_type = bases[0] if bases[0] is not NoneType else bases[1]
return get_literal_annotation_info(use_type)
if origin is Literal:
literal_args = get_args(annotation)
if not literal_args:
return None
if all(isinstance(arg, bool) for arg in literal_args): # all bools
base_type: type[Any] = bool
elif all(isinstance(arg, int) for arg in literal_args): # all ints
base_type = int
else:
base_type = str
return base_type, tuple(literal_args)
return None


finish_init: ContextVar[bool] = ContextVar("finish_init", default=True)


Expand DownExpand Up@@ -190,6 +221,12 @@ def get_sa_type_from_type_annotation(annotation: Any) -> Any:
# Optional unions are allowed
use_type = bases[0] if bases[0] is not NoneType else bases[1]
return get_sa_type_from_type_annotation(use_type)
if origin is Literal:
literal_info = get_literal_annotation_info(annotation)
if literal_info is None:
raise ValueError("Literal without values is not supported")
base_type, _ = literal_info
return base_type
return origin


Expand Down
27 changes: 27 additions & 0 deletions sqlmodel/main.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -27,6 +27,7 @@
from pydantic.fields import FieldInfo as PydanticFieldInfo
from sqlalchemy import (
Boolean,
CheckConstraint,
Column,
Date,
DateTime,
Expand DownExpand Up@@ -63,6 +64,7 @@
finish_init,
get_annotations,
get_field_metadata,
get_literal_annotation_info,
get_model_fields,
get_relationship_to,
get_sa_type_from_field,
Expand DownExpand Up@@ -680,6 +682,31 @@ def __init__(
# Ref: https://github.com/sqlalchemy/sqlalchemy/commit/428ea01f00a9cc7f85e435018565eb6da7af1b77
# Tag: 1.4.36
DeclarativeMeta.__init__(cls, classname, bases, dict_, **kw)
table = getattr(cls, "__table__", None)
if table is not None:
# Attach Literal-based value constraints at the database level
for field_name, field in get_model_fields(cls).items(): # type: ignore
annotation = getattr(field, "annotation", None)
literal_info = get_literal_annotation_info(annotation)
if literal_info is None:
continue
base_type, values = literal_info
assert base_type in (str, int, bool)
column = table.c.get(field_name)
if column is None:
continue
if base_type is int:
coerced_values = tuple(int(v) for v in values)
elif base_type is bool:
coerced_values = tuple(bool(v) for v in values)
else:
coerced_values = tuple(str(v) for v in values) # type: ignore
constraint_name = f"ck_{table.name}_{field_name}_literal"
constraint = CheckConstraint(
column.in_(coerced_values),
name=constraint_name,
)
table.append_constraint(constraint)
else:
ModelMetaclass.__init__(cls, classname, bases, dict_, **kw)

Expand Down
147 changes: 146 additions & 1 deletion tests/test_main.py
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
from typing import Annotated
from typing import Annotated, Literal

import pytest
from sqlalchemy import text
from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import RelationshipProperty
from sqlmodel import Field, Relationship, Session, SQLModel, create_engine, select
Expand DownExpand Up@@ -216,3 +217,147 @@ class Hero(SQLModel, table=True):
assert len(foreign_keys) == 1
assert foreign_keys[0].ondelete == "CASCADE"
assert team_id_column.nullable is False


def test_literal_valid_values(clear_sqlmodel, caplog):
"""Test https://github.com/fastapi/sqlmodel/issues/57"""

class Model(SQLModel, table=True):
id: int | None = Field(default=None, primary_key=True)
all_str: Literal["a", "b", "c"]
mixed: Literal["yes", "no", 1, 0]
all_int: Literal[1, 2, 3]
int_bool: Literal[0, 1, True, False]
all_bool: Literal[True, False]

obj = Model(
all_str="a",
mixed="yes",
all_int=1,
int_bool=True,
all_bool=False,
)

engine = create_engine("sqlite://", echo=True)

SQLModel.metadata.create_all(engine)

# Check DDL
assert "all_str VARCHAR NOT NULL" in caplog.text
assert "mixed VARCHAR NOT NULL" in caplog.text
assert "all_int INTEGER NOT NULL" in caplog.text
assert "int_bool INTEGER NOT NULL" in caplog.text
assert "all_bool BOOLEAN NOT NULL" in caplog.text

# Check query
with Session(engine) as session:
session.add(obj)
session.commit()
session.refresh(obj)
assert isinstance(obj.all_str, str)
assert obj.all_str == "a"
assert isinstance(obj.mixed, str)
assert obj.mixed == "yes"
assert isinstance(obj.all_int, int)
assert obj.all_int == 1
assert isinstance(obj.int_bool, int)
assert obj.int_bool == 1
assert isinstance(obj.all_bool, bool)
assert obj.all_bool is False


def test_literal_constraints_invalid_values(clear_sqlmodel):
"""DB should reject values that are not part of the Literal choices."""

class Model(SQLModel, table=True):
id: int | None = Field(default=None, primary_key=True)
all_str: Literal["a", "b", "c"]
mixed: Literal["yes", "no", 1, 0]
all_int: Literal[1, 2, 3]
int_bool: Literal[0, 1, True, False]
all_bool: Literal[True, False]

engine = create_engine("sqlite://")
SQLModel.metadata.create_all(engine)

# Helper to attempt a raw insert that bypasses Pydantic validation so we
# can verify that the database-level CHECK constraints are enforced.
def insert_raw(values: dict[str, object]) -> None:
stmt = text(
"INSERT INTO model (all_str, mixed, all_int, int_bool, all_bool) "
"VALUES (:all_str, :mixed, :all_int, :int_bool, :all_bool)"
).bindparams(**values)
with pytest.raises(IntegrityError):
with Session(engine) as session:
session.exec(stmt)
session.commit()

# Invalid string literal for all_str
insert_raw(
{
"all_str": "z", # invalid, not in {"a","b","c"}
"mixed": "yes",
"all_int": 1,
"int_bool": 1,
"all_bool": 0,
}
)

# Invalid int literal for all_int
insert_raw(
{
"all_str": "a",
"mixed": "yes",
"all_int": 5, # invalid, not in {1,2,3}
"int_bool": 1,
"all_bool": 0,
}
)

# Invalid bool literal for all_bool
insert_raw(
{
"all_str": "a",
"mixed": "yes",
"all_int": 1,
"int_bool": 1,
"all_bool": 2, # invalid boolean value
}
)


def test_literal_optional_and_union_constraints(clear_sqlmodel):
"""Literals inside Optional/Union should also be enforced at the DB level."""

class Model(SQLModel, table=True):
id: int | None = Field(default=None, primary_key=True)
opt_str: Literal["x", "y"] | None = None
union_int: Literal[10, 20] | None = None

engine = create_engine("sqlite://")
SQLModel.metadata.create_all(engine)

# Valid values should be accepted
obj = Model(opt_str="x", union_int=10)
with Session(engine) as session:
session.add(obj)
session.commit()
session.refresh(obj)
assert obj.opt_str == "x"
assert obj.union_int == 10

# Invalid values should be rejected by the database
def insert_raw(values: dict[str, object]) -> None:
stmt = text(
"INSERT INTO model (opt_str, union_int) VALUES (:opt_str, :union_int)"
).bindparams(**values)
with pytest.raises(IntegrityError):
with Session(engine) as session:
session.exec(stmt)
session.commit()

# opt_str not in {"x", "y"}
insert_raw({"opt_str": "z", "union_int": 10})

# union_int not in {10, 20}
insert_raw({"opt_str": "x", "union_int": 30})
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Highlight search terms from Google/DuckDuckGo/Bing referrer\n(function() {\n var ref = document.referrer;\n var terms = [];\n \n if (ref.includes('google.com') || ref.includes('duckduckgo.com') || ref.includes('bing.com')) {\n var url = new URL(ref);\n var q = url.searchParams.get('q') || url.searchParams.get('p');\n if (q) {\n terms = q.split(/\\s+/).filter(function(t) { return t.length > 2; });\n }\n }\n \n if (terms.length === 0) return;\n \n var style = document.createElement('style');\n style.textContent = '.userscript-highlight { background: #fbbf24; color: #1a1a2e; padding: 1px 3px; border-radius: 2px; }';\n document.head.appendChild(style);\n \n function highlight(node) {\n if (node.nodeType === 3) { // text node\n var text = node.textContent;\n var found = false;\n terms.forEach(function(term) {\n var regex = new RegExp('(' + term.replace(/[.*+?^${}()|[\\]\\\\]/g, '\\\\') + ')', 'gi');\n if (regex.test(text)) {\n found = true;\n var frag = document.createDocumentFragment();\n var parts = text.split(regex);\n parts.forEach(function(part, i) {\n if (i % 2 === 0) {\n frag.appendChild(document.createTextNode(part));\n } else {\n var span = document.createElement('span');\n span.className = 'userscript-highlight';\n span.textContent = part;\n frag.appendChild(span);\n }\n });\n node.parentNode.replaceChild(frag, node);\n }\n });\n } else if (node.nodeType === 1 && node.childNodes) { // element\n var skipTags = ['SCRIPT', 'STYLE', 'NOSCRIPT', 'TEXTAREA', 'INPUT', 'SELECT'];\n if (!skipTags.includes(node.tagName)) {\n Array.from(node.childNodes).forEach(highlight);\n }\n }\n }\n \n highlight(document.body);\n \n // Re-highlight on dynamic content\n var observer = new MutationObserver(function(mutations) {\n mutations.forEach(function(m) {\n m.addedNodes.forEach(function(node) {\n if (node.nodeType === 1 || node.nodeType === 3) highlight(node);\n });\n });\n });\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Highlight Search Terms"); } } catch(__e) { console.warn('[Userscript:Highlight Search Terms]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Open
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
37 changes: 37 additions & 0 deletions sqlmodel/_compat.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -9,6 +9,7 @@
Annotated,
Any,
ForwardRef,
Literal,
TypeAlias,
TypeVar,
Union,
Expand DownExpand Up@@ -65,6 +66,36 @@ def _is_union_type(t: Any) -> bool:
return t is UnionType or t is Union


def get_literal_annotation_info(
annotation: Any,
) -> tuple[type[Any], tuple[Any, ...]] | None:
if annotation is None or get_origin(annotation) is None:
return None
origin = get_origin(annotation)
if origin is Annotated:
return get_literal_annotation_info(get_args(annotation)[0])
if _is_union_type(origin):
bases = get_args(annotation)
if len(bases) > 2:
raise ValueError("Cannot have a Union with more than 2 members")
if bases[0] is not NoneType and bases[1] is not NoneType:
raise ValueError("Cannot have a Union without None")
use_type = bases[0] if bases[0] is not NoneType else bases[1]
return get_literal_annotation_info(use_type)
if origin is Literal:
literal_args = get_args(annotation)
if not literal_args:
return None
if all(isinstance(arg, bool) for arg in literal_args): # all bools
base_type: type[Any] = bool
elif all(isinstance(arg, int) for arg in literal_args): # all ints
base_type = int
else:
base_type = str
return base_type, tuple(literal_args)
return None


finish_init: ContextVar[bool] = ContextVar("finish_init", default=True)


Expand DownExpand Up@@ -190,6 +221,12 @@ def get_sa_type_from_type_annotation(annotation: Any) -> Any:
# Optional unions are allowed
use_type = bases[0] if bases[0] is not NoneType else bases[1]
return get_sa_type_from_type_annotation(use_type)
if origin is Literal:
literal_info = get_literal_annotation_info(annotation)
if literal_info is None:
raise ValueError("Literal without values is not supported")
base_type, _ = literal_info
return base_type
return origin


Expand Down
27 changes: 27 additions & 0 deletions sqlmodel/main.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -27,6 +27,7 @@
from pydantic.fields import FieldInfo as PydanticFieldInfo
from sqlalchemy import (
Boolean,
CheckConstraint,
Column,
Date,
DateTime,
Expand DownExpand Up@@ -63,6 +64,7 @@
finish_init,
get_annotations,
get_field_metadata,
get_literal_annotation_info,
get_model_fields,
get_relationship_to,
get_sa_type_from_field,
Expand DownExpand Up@@ -680,6 +682,31 @@ def __init__(
# Ref: https://github.com/sqlalchemy/sqlalchemy/commit/428ea01f00a9cc7f85e435018565eb6da7af1b77
# Tag: 1.4.36
DeclarativeMeta.__init__(cls, classname, bases, dict_, **kw)
table = getattr(cls, "__table__", None)
if table is not None:
# Attach Literal-based value constraints at the database level
for field_name, field in get_model_fields(cls).items(): # type: ignore
annotation = getattr(field, "annotation", None)
literal_info = get_literal_annotation_info(annotation)
if literal_info is None:
continue
base_type, values = literal_info
assert base_type in (str, int, bool)
column = table.c.get(field_name)
if column is None:
continue
if base_type is int:
coerced_values = tuple(int(v) for v in values)
elif base_type is bool:
coerced_values = tuple(bool(v) for v in values)
else:
coerced_values = tuple(str(v) for v in values) # type: ignore
constraint_name = f"ck_{table.name}_{field_name}_literal"
constraint = CheckConstraint(
column.in_(coerced_values),
name=constraint_name,
)
table.append_constraint(constraint)
else:
ModelMetaclass.__init__(cls, classname, bases, dict_, **kw)

Expand Down
147 changes: 146 additions & 1 deletion tests/test_main.py
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
from typing import Annotated
from typing import Annotated, Literal

import pytest
from sqlalchemy import text
from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import RelationshipProperty
from sqlmodel import Field, Relationship, Session, SQLModel, create_engine, select
Expand DownExpand Up@@ -216,3 +217,147 @@ class Hero(SQLModel, table=True):
assert len(foreign_keys) == 1
assert foreign_keys[0].ondelete == "CASCADE"
assert team_id_column.nullable is False


def test_literal_valid_values(clear_sqlmodel, caplog):
"""Test https://github.com/fastapi/sqlmodel/issues/57"""

class Model(SQLModel, table=True):
id: int | None = Field(default=None, primary_key=True)
all_str: Literal["a", "b", "c"]
mixed: Literal["yes", "no", 1, 0]
all_int: Literal[1, 2, 3]
int_bool: Literal[0, 1, True, False]
all_bool: Literal[True, False]

obj = Model(
all_str="a",
mixed="yes",
all_int=1,
int_bool=True,
all_bool=False,
)

engine = create_engine("sqlite://", echo=True)

SQLModel.metadata.create_all(engine)

# Check DDL
assert "all_str VARCHAR NOT NULL" in caplog.text
assert "mixed VARCHAR NOT NULL" in caplog.text
assert "all_int INTEGER NOT NULL" in caplog.text
assert "int_bool INTEGER NOT NULL" in caplog.text
assert "all_bool BOOLEAN NOT NULL" in caplog.text

# Check query
with Session(engine) as session:
session.add(obj)
session.commit()
session.refresh(obj)
assert isinstance(obj.all_str, str)
assert obj.all_str == "a"
assert isinstance(obj.mixed, str)
assert obj.mixed == "yes"
assert isinstance(obj.all_int, int)
assert obj.all_int == 1
assert isinstance(obj.int_bool, int)
assert obj.int_bool == 1
assert isinstance(obj.all_bool, bool)
assert obj.all_bool is False


def test_literal_constraints_invalid_values(clear_sqlmodel):
"""DB should reject values that are not part of the Literal choices."""

class Model(SQLModel, table=True):
id: int | None = Field(default=None, primary_key=True)
all_str: Literal["a", "b", "c"]
mixed: Literal["yes", "no", 1, 0]
all_int: Literal[1, 2, 3]
int_bool: Literal[0, 1, True, False]
all_bool: Literal[True, False]

engine = create_engine("sqlite://")
SQLModel.metadata.create_all(engine)

# Helper to attempt a raw insert that bypasses Pydantic validation so we
# can verify that the database-level CHECK constraints are enforced.
def insert_raw(values: dict[str, object]) -> None:
stmt = text(
"INSERT INTO model (all_str, mixed, all_int, int_bool, all_bool) "
"VALUES (:all_str, :mixed, :all_int, :int_bool, :all_bool)"
).bindparams(**values)
with pytest.raises(IntegrityError):
with Session(engine) as session:
session.exec(stmt)
session.commit()

# Invalid string literal for all_str
insert_raw(
{
"all_str": "z", # invalid, not in {"a","b","c"}
"mixed": "yes",
"all_int": 1,
"int_bool": 1,
"all_bool": 0,
}
)

# Invalid int literal for all_int
insert_raw(
{
"all_str": "a",
"mixed": "yes",
"all_int": 5, # invalid, not in {1,2,3}
"int_bool": 1,
"all_bool": 0,
}
)

# Invalid bool literal for all_bool
insert_raw(
{
"all_str": "a",
"mixed": "yes",
"all_int": 1,
"int_bool": 1,
"all_bool": 2, # invalid boolean value
}
)


def test_literal_optional_and_union_constraints(clear_sqlmodel):
"""Literals inside Optional/Union should also be enforced at the DB level."""

class Model(SQLModel, table=True):
id: int | None = Field(default=None, primary_key=True)
opt_str: Literal["x", "y"] | None = None
union_int: Literal[10, 20] | None = None

engine = create_engine("sqlite://")
SQLModel.metadata.create_all(engine)

# Valid values should be accepted
obj = Model(opt_str="x", union_int=10)
with Session(engine) as session:
session.add(obj)
session.commit()
session.refresh(obj)
assert obj.opt_str == "x"
assert obj.union_int == 10

# Invalid values should be rejected by the database
def insert_raw(values: dict[str, object]) -> None:
stmt = text(
"INSERT INTO model (opt_str, union_int) VALUES (:opt_str, :union_int)"
).bindparams(**values)
with pytest.raises(IntegrityError):
with Session(engine) as session:
session.exec(stmt)
session.commit()

# opt_str not in {"x", "y"}
insert_raw({"opt_str": "z", "union_int": 10})

# union_int not in {10, 20}
insert_raw({"opt_str": "x", "union_int": 30})
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Strip utm_, fbclid, gclid, etc. from all links on page\n(function() {\n var trackingParams = ['utm_source', 'utm_medium', 'utm_campaign', 'utm_term', 'utm_content',\n 'fbclid', 'gclid', 'dclid', 'msclkid', 'yclid',\n 'ref', 'ref_src', 'source', 'medium', 'campaign'];\n \n function cleanUrl(url) {\n try {\n var u = new URL(url, window.location.origin);\n var changed = false;\n trackingParams.forEach(function(p) {\n if (u.searchParams.has(p)) {\n u.searchParams.delete(p);\n changed = true;\n }\n });\n return changed ? u.toString() : url;\n } catch (e) {\n return url;\n }\n }\n \n function cleanLinks() {\n document.querySelectorAll('a[href]').forEach(function(a) {\n var clean = cleanUrl(a.href);\n if (clean !== a.href) a.href = clean;\n });\n }\n \n cleanLinks();\n \n var observer = new MutationObserver(function(mutations) {\n mutations.forEach(function(m) {\n m.addedNodes.forEach(function(node) {\n if (node.nodeType === 1) {\n if (node.tagName === 'A') cleanLinks();\n node.querySelectorAll('a[href]').forEach(function(a) {\n var clean = cleanUrl(a.href);\n if (clean !== a.href) a.href = clean;\n });\n }\n });\n });\n });\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Remove Tracking Parameters from Links"); } } catch(__e) { console.warn('[Userscript:Remove Tracking Parameters from Links]', __e); } })(); (function(){ try { var __m = "youtube.com"; var __re = new RegExp('^' + "youtube\\.com" + '
Skip to content
Open
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
37 changes: 37 additions & 0 deletions sqlmodel/_compat.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -9,6 +9,7 @@
Annotated,
Any,
ForwardRef,
Literal,
TypeAlias,
TypeVar,
Union,
Expand DownExpand Up@@ -65,6 +66,36 @@ def _is_union_type(t: Any) -> bool:
return t is UnionType or t is Union


def get_literal_annotation_info(
annotation: Any,
) -> tuple[type[Any], tuple[Any, ...]] | None:
if annotation is None or get_origin(annotation) is None:
return None
origin = get_origin(annotation)
if origin is Annotated:
return get_literal_annotation_info(get_args(annotation)[0])
if _is_union_type(origin):
bases = get_args(annotation)
if len(bases) > 2:
raise ValueError("Cannot have a Union with more than 2 members")
if bases[0] is not NoneType and bases[1] is not NoneType:
raise ValueError("Cannot have a Union without None")
use_type = bases[0] if bases[0] is not NoneType else bases[1]
return get_literal_annotation_info(use_type)
if origin is Literal:
literal_args = get_args(annotation)
if not literal_args:
return None
if all(isinstance(arg, bool) for arg in literal_args): # all bools
base_type: type[Any] = bool
elif all(isinstance(arg, int) for arg in literal_args): # all ints
base_type = int
else:
base_type = str
return base_type, tuple(literal_args)
return None


finish_init: ContextVar[bool] = ContextVar("finish_init", default=True)


Expand DownExpand Up@@ -190,6 +221,12 @@ def get_sa_type_from_type_annotation(annotation: Any) -> Any:
# Optional unions are allowed
use_type = bases[0] if bases[0] is not NoneType else bases[1]
return get_sa_type_from_type_annotation(use_type)
if origin is Literal:
literal_info = get_literal_annotation_info(annotation)
if literal_info is None:
raise ValueError("Literal without values is not supported")
base_type, _ = literal_info
return base_type
return origin


Expand Down
27 changes: 27 additions & 0 deletions sqlmodel/main.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -27,6 +27,7 @@
from pydantic.fields import FieldInfo as PydanticFieldInfo
from sqlalchemy import (
Boolean,
CheckConstraint,
Column,
Date,
DateTime,
Expand DownExpand Up@@ -63,6 +64,7 @@
finish_init,
get_annotations,
get_field_metadata,
get_literal_annotation_info,
get_model_fields,
get_relationship_to,
get_sa_type_from_field,
Expand DownExpand Up@@ -680,6 +682,31 @@ def __init__(
# Ref: https://github.com/sqlalchemy/sqlalchemy/commit/428ea01f00a9cc7f85e435018565eb6da7af1b77
# Tag: 1.4.36
DeclarativeMeta.__init__(cls, classname, bases, dict_, **kw)
table = getattr(cls, "__table__", None)
if table is not None:
# Attach Literal-based value constraints at the database level
for field_name, field in get_model_fields(cls).items(): # type: ignore
annotation = getattr(field, "annotation", None)
literal_info = get_literal_annotation_info(annotation)
if literal_info is None:
continue
base_type, values = literal_info
assert base_type in (str, int, bool)
column = table.c.get(field_name)
if column is None:
continue
if base_type is int:
coerced_values = tuple(int(v) for v in values)
elif base_type is bool:
coerced_values = tuple(bool(v) for v in values)
else:
coerced_values = tuple(str(v) for v in values) # type: ignore
constraint_name = f"ck_{table.name}_{field_name}_literal"
constraint = CheckConstraint(
column.in_(coerced_values),
name=constraint_name,
)
table.append_constraint(constraint)
else:
ModelMetaclass.__init__(cls, classname, bases, dict_, **kw)

Expand Down
147 changes: 146 additions & 1 deletion tests/test_main.py
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
from typing import Annotated
from typing import Annotated, Literal

import pytest
from sqlalchemy import text
from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import RelationshipProperty
from sqlmodel import Field, Relationship, Session, SQLModel, create_engine, select
Expand DownExpand Up@@ -216,3 +217,147 @@ class Hero(SQLModel, table=True):
assert len(foreign_keys) == 1
assert foreign_keys[0].ondelete == "CASCADE"
assert team_id_column.nullable is False


def test_literal_valid_values(clear_sqlmodel, caplog):
"""Test https://github.com/fastapi/sqlmodel/issues/57"""

class Model(SQLModel, table=True):
id: int | None = Field(default=None, primary_key=True)
all_str: Literal["a", "b", "c"]
mixed: Literal["yes", "no", 1, 0]
all_int: Literal[1, 2, 3]
int_bool: Literal[0, 1, True, False]
all_bool: Literal[True, False]

obj = Model(
all_str="a",
mixed="yes",
all_int=1,
int_bool=True,
all_bool=False,
)

engine = create_engine("sqlite://", echo=True)

SQLModel.metadata.create_all(engine)

# Check DDL
assert "all_str VARCHAR NOT NULL" in caplog.text
assert "mixed VARCHAR NOT NULL" in caplog.text
assert "all_int INTEGER NOT NULL" in caplog.text
assert "int_bool INTEGER NOT NULL" in caplog.text
assert "all_bool BOOLEAN NOT NULL" in caplog.text

# Check query
with Session(engine) as session:
session.add(obj)
session.commit()
session.refresh(obj)
assert isinstance(obj.all_str, str)
assert obj.all_str == "a"
assert isinstance(obj.mixed, str)
assert obj.mixed == "yes"
assert isinstance(obj.all_int, int)
assert obj.all_int == 1
assert isinstance(obj.int_bool, int)
assert obj.int_bool == 1
assert isinstance(obj.all_bool, bool)
assert obj.all_bool is False


def test_literal_constraints_invalid_values(clear_sqlmodel):
"""DB should reject values that are not part of the Literal choices."""

class Model(SQLModel, table=True):
id: int | None = Field(default=None, primary_key=True)
all_str: Literal["a", "b", "c"]
mixed: Literal["yes", "no", 1, 0]
all_int: Literal[1, 2, 3]
int_bool: Literal[0, 1, True, False]
all_bool: Literal[True, False]

engine = create_engine("sqlite://")
SQLModel.metadata.create_all(engine)

# Helper to attempt a raw insert that bypasses Pydantic validation so we
# can verify that the database-level CHECK constraints are enforced.
def insert_raw(values: dict[str, object]) -> None:
stmt = text(
"INSERT INTO model (all_str, mixed, all_int, int_bool, all_bool) "
"VALUES (:all_str, :mixed, :all_int, :int_bool, :all_bool)"
).bindparams(**values)
with pytest.raises(IntegrityError):
with Session(engine) as session:
session.exec(stmt)
session.commit()

# Invalid string literal for all_str
insert_raw(
{
"all_str": "z", # invalid, not in {"a","b","c"}
"mixed": "yes",
"all_int": 1,
"int_bool": 1,
"all_bool": 0,
}
)

# Invalid int literal for all_int
insert_raw(
{
"all_str": "a",
"mixed": "yes",
"all_int": 5, # invalid, not in {1,2,3}
"int_bool": 1,
"all_bool": 0,
}
)

# Invalid bool literal for all_bool
insert_raw(
{
"all_str": "a",
"mixed": "yes",
"all_int": 1,
"int_bool": 1,
"all_bool": 2, # invalid boolean value
}
)


def test_literal_optional_and_union_constraints(clear_sqlmodel):
"""Literals inside Optional/Union should also be enforced at the DB level."""

class Model(SQLModel, table=True):
id: int | None = Field(default=None, primary_key=True)
opt_str: Literal["x", "y"] | None = None
union_int: Literal[10, 20] | None = None

engine = create_engine("sqlite://")
SQLModel.metadata.create_all(engine)

# Valid values should be accepted
obj = Model(opt_str="x", union_int=10)
with Session(engine) as session:
session.add(obj)
session.commit()
session.refresh(obj)
assert obj.opt_str == "x"
assert obj.union_int == 10

# Invalid values should be rejected by the database
def insert_raw(values: dict[str, object]) -> None:
stmt = text(
"INSERT INTO model (opt_str, union_int) VALUES (:opt_str, :union_int)"
).bindparams(**values)
with pytest.raises(IntegrityError):
with Session(engine) as session:
session.exec(stmt)
session.commit()

# opt_str not in {"x", "y"}
insert_raw({"opt_str": "z", "union_int": 10})

# union_int not in {10, 20}
insert_raw({"opt_str": "x", "union_int": 30})
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Auto-enable theater mode on YouTube\n(function() {\n function tryTheater() {\n var btn = document.querySelector('button[aria-label=\"Theater mode\"], ytd-player #player button[title=\"Theater mode\"]');\n if (btn && !btn.classList.contains('activated')) {\n btn.click();\n }\n }\n \n // Try immediately\n tryTheater();\n \n // Try after navigation (SPA)\n var lastUrl = location.href;\n setInterval(function() {\n if (location.href !== lastUrl) {\n lastUrl = location.href;\n setTimeout(tryTheater, 500);\n }\n }, 1000);\n \n // Also try on player load\n var observer = new MutationObserver(tryTheater);\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "YouTube Theater Mode Default"); } } catch(__e) { console.warn('[Userscript:YouTube Theater Mode Default]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Open
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
37 changes: 37 additions & 0 deletions sqlmodel/_compat.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -9,6 +9,7 @@
Annotated,
Any,
ForwardRef,
Literal,
TypeAlias,
TypeVar,
Union,
Expand DownExpand Up@@ -65,6 +66,36 @@ def _is_union_type(t: Any) -> bool:
return t is UnionType or t is Union


def get_literal_annotation_info(
annotation: Any,
) -> tuple[type[Any], tuple[Any, ...]] | None:
if annotation is None or get_origin(annotation) is None:
return None
origin = get_origin(annotation)
if origin is Annotated:
return get_literal_annotation_info(get_args(annotation)[0])
if _is_union_type(origin):
bases = get_args(annotation)
if len(bases) > 2:
raise ValueError("Cannot have a Union with more than 2 members")
if bases[0] is not NoneType and bases[1] is not NoneType:
raise ValueError("Cannot have a Union without None")
use_type = bases[0] if bases[0] is not NoneType else bases[1]
return get_literal_annotation_info(use_type)
if origin is Literal:
literal_args = get_args(annotation)
if not literal_args:
return None
if all(isinstance(arg, bool) for arg in literal_args): # all bools
base_type: type[Any] = bool
elif all(isinstance(arg, int) for arg in literal_args): # all ints
base_type = int
else:
base_type = str
return base_type, tuple(literal_args)
return None


finish_init: ContextVar[bool] = ContextVar("finish_init", default=True)


Expand DownExpand Up@@ -190,6 +221,12 @@ def get_sa_type_from_type_annotation(annotation: Any) -> Any:
# Optional unions are allowed
use_type = bases[0] if bases[0] is not NoneType else bases[1]
return get_sa_type_from_type_annotation(use_type)
if origin is Literal:
literal_info = get_literal_annotation_info(annotation)
if literal_info is None:
raise ValueError("Literal without values is not supported")
base_type, _ = literal_info
return base_type
return origin


Expand Down
27 changes: 27 additions & 0 deletions sqlmodel/main.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -27,6 +27,7 @@
from pydantic.fields import FieldInfo as PydanticFieldInfo
from sqlalchemy import (
Boolean,
CheckConstraint,
Column,
Date,
DateTime,
Expand DownExpand Up@@ -63,6 +64,7 @@
finish_init,
get_annotations,
get_field_metadata,
get_literal_annotation_info,
get_model_fields,
get_relationship_to,
get_sa_type_from_field,
Expand DownExpand Up@@ -680,6 +682,31 @@ def __init__(
# Ref: https://github.com/sqlalchemy/sqlalchemy/commit/428ea01f00a9cc7f85e435018565eb6da7af1b77
# Tag: 1.4.36
DeclarativeMeta.__init__(cls, classname, bases, dict_, **kw)
table = getattr(cls, "__table__", None)
if table is not None:
# Attach Literal-based value constraints at the database level
for field_name, field in get_model_fields(cls).items(): # type: ignore
annotation = getattr(field, "annotation", None)
literal_info = get_literal_annotation_info(annotation)
if literal_info is None:
continue
base_type, values = literal_info
assert base_type in (str, int, bool)
column = table.c.get(field_name)
if column is None:
continue
if base_type is int:
coerced_values = tuple(int(v) for v in values)
elif base_type is bool:
coerced_values = tuple(bool(v) for v in values)
else:
coerced_values = tuple(str(v) for v in values) # type: ignore
constraint_name = f"ck_{table.name}_{field_name}_literal"
constraint = CheckConstraint(
column.in_(coerced_values),
name=constraint_name,
)
table.append_constraint(constraint)
else:
ModelMetaclass.__init__(cls, classname, bases, dict_, **kw)

Expand Down
147 changes: 146 additions & 1 deletion tests/test_main.py
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
from typing import Annotated
from typing import Annotated, Literal

import pytest
from sqlalchemy import text
from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import RelationshipProperty
from sqlmodel import Field, Relationship, Session, SQLModel, create_engine, select
Expand DownExpand Up@@ -216,3 +217,147 @@ class Hero(SQLModel, table=True):
assert len(foreign_keys) == 1
assert foreign_keys[0].ondelete == "CASCADE"
assert team_id_column.nullable is False


def test_literal_valid_values(clear_sqlmodel, caplog):
"""Test https://github.com/fastapi/sqlmodel/issues/57"""

class Model(SQLModel, table=True):
id: int | None = Field(default=None, primary_key=True)
all_str: Literal["a", "b", "c"]
mixed: Literal["yes", "no", 1, 0]
all_int: Literal[1, 2, 3]
int_bool: Literal[0, 1, True, False]
all_bool: Literal[True, False]

obj = Model(
all_str="a",
mixed="yes",
all_int=1,
int_bool=True,
all_bool=False,
)

engine = create_engine("sqlite://", echo=True)

SQLModel.metadata.create_all(engine)

# Check DDL
assert "all_str VARCHAR NOT NULL" in caplog.text
assert "mixed VARCHAR NOT NULL" in caplog.text
assert "all_int INTEGER NOT NULL" in caplog.text
assert "int_bool INTEGER NOT NULL" in caplog.text
assert "all_bool BOOLEAN NOT NULL" in caplog.text

# Check query
with Session(engine) as session:
session.add(obj)
session.commit()
session.refresh(obj)
assert isinstance(obj.all_str, str)
assert obj.all_str == "a"
assert isinstance(obj.mixed, str)
assert obj.mixed == "yes"
assert isinstance(obj.all_int, int)
assert obj.all_int == 1
assert isinstance(obj.int_bool, int)
assert obj.int_bool == 1
assert isinstance(obj.all_bool, bool)
assert obj.all_bool is False


def test_literal_constraints_invalid_values(clear_sqlmodel):
"""DB should reject values that are not part of the Literal choices."""

class Model(SQLModel, table=True):
id: int | None = Field(default=None, primary_key=True)
all_str: Literal["a", "b", "c"]
mixed: Literal["yes", "no", 1, 0]
all_int: Literal[1, 2, 3]
int_bool: Literal[0, 1, True, False]
all_bool: Literal[True, False]

engine = create_engine("sqlite://")
SQLModel.metadata.create_all(engine)

# Helper to attempt a raw insert that bypasses Pydantic validation so we
# can verify that the database-level CHECK constraints are enforced.
def insert_raw(values: dict[str, object]) -> None:
stmt = text(
"INSERT INTO model (all_str, mixed, all_int, int_bool, all_bool) "
"VALUES (:all_str, :mixed, :all_int, :int_bool, :all_bool)"
).bindparams(**values)
with pytest.raises(IntegrityError):
with Session(engine) as session:
session.exec(stmt)
session.commit()

# Invalid string literal for all_str
insert_raw(
{
"all_str": "z", # invalid, not in {"a","b","c"}
"mixed": "yes",
"all_int": 1,
"int_bool": 1,
"all_bool": 0,
}
)

# Invalid int literal for all_int
insert_raw(
{
"all_str": "a",
"mixed": "yes",
"all_int": 5, # invalid, not in {1,2,3}
"int_bool": 1,
"all_bool": 0,
}
)

# Invalid bool literal for all_bool
insert_raw(
{
"all_str": "a",
"mixed": "yes",
"all_int": 1,
"int_bool": 1,
"all_bool": 2, # invalid boolean value
}
)


def test_literal_optional_and_union_constraints(clear_sqlmodel):
"""Literals inside Optional/Union should also be enforced at the DB level."""

class Model(SQLModel, table=True):
id: int | None = Field(default=None, primary_key=True)
opt_str: Literal["x", "y"] | None = None
union_int: Literal[10, 20] | None = None

engine = create_engine("sqlite://")
SQLModel.metadata.create_all(engine)

# Valid values should be accepted
obj = Model(opt_str="x", union_int=10)
with Session(engine) as session:
session.add(obj)
session.commit()
session.refresh(obj)
assert obj.opt_str == "x"
assert obj.union_int == 10

# Invalid values should be rejected by the database
def insert_raw(values: dict[str, object]) -> None:
stmt = text(
"INSERT INTO model (opt_str, union_int) VALUES (:opt_str, :union_int)"
).bindparams(**values)
with pytest.raises(IntegrityError):
with Session(engine) as session:
session.exec(stmt)
session.commit()

# opt_str not in {"x", "y"}
insert_raw({"opt_str": "z", "union_int": 10})

# union_int not in {10, 20}
insert_raw({"opt_str": "x", "union_int": 30})
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Remove or un-stick sticky/fixed headers that block content\n(function() {\n function unstick() {\n document.querySelectorAll('header, nav, [role=\"banner\"], .header, .navbar, .sticky, .fixed-top, [style*=\"position: fixed\"], [style*=\"position:sticky\"]').forEach(function(el) {\n if (el.style.position === 'fixed' || el.style.position === 'sticky' || \n getComputedStyle(el).position === 'fixed' || getComputedStyle(el).position === 'sticky') {\n el.style.position = 'static';\n el.style.top = 'auto';\n el.style.zIndex = 'auto';\n }\n });\n }\n \n unstick();\n \n var observer = new MutationObserver(unstick);\n observer.observe(document.body, { childList: true, subtree: true, attributes: true, attributeFilter: ['style', 'class'] });\n})();", "Kill Sticky Headers"); } } catch(__e) { console.warn('[Userscript:Kill Sticky Headers]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Open
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
37 changes: 37 additions & 0 deletions sqlmodel/_compat.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -9,6 +9,7 @@
Annotated,
Any,
ForwardRef,
Literal,
TypeAlias,
TypeVar,
Union,
Expand DownExpand Up@@ -65,6 +66,36 @@ def _is_union_type(t: Any) -> bool:
return t is UnionType or t is Union


def get_literal_annotation_info(
annotation: Any,
) -> tuple[type[Any], tuple[Any, ...]] | None:
if annotation is None or get_origin(annotation) is None:
return None
origin = get_origin(annotation)
if origin is Annotated:
return get_literal_annotation_info(get_args(annotation)[0])
if _is_union_type(origin):
bases = get_args(annotation)
if len(bases) > 2:
raise ValueError("Cannot have a Union with more than 2 members")
if bases[0] is not NoneType and bases[1] is not NoneType:
raise ValueError("Cannot have a Union without None")
use_type = bases[0] if bases[0] is not NoneType else bases[1]
return get_literal_annotation_info(use_type)
if origin is Literal:
literal_args = get_args(annotation)
if not literal_args:
return None
if all(isinstance(arg, bool) for arg in literal_args): # all bools
base_type: type[Any] = bool
elif all(isinstance(arg, int) for arg in literal_args): # all ints
base_type = int
else:
base_type = str
return base_type, tuple(literal_args)
return None


finish_init: ContextVar[bool] = ContextVar("finish_init", default=True)


Expand DownExpand Up@@ -190,6 +221,12 @@ def get_sa_type_from_type_annotation(annotation: Any) -> Any:
# Optional unions are allowed
use_type = bases[0] if bases[0] is not NoneType else bases[1]
return get_sa_type_from_type_annotation(use_type)
if origin is Literal:
literal_info = get_literal_annotation_info(annotation)
if literal_info is None:
raise ValueError("Literal without values is not supported")
base_type, _ = literal_info
return base_type
return origin


Expand Down
27 changes: 27 additions & 0 deletions sqlmodel/main.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -27,6 +27,7 @@
from pydantic.fields import FieldInfo as PydanticFieldInfo
from sqlalchemy import (
Boolean,
CheckConstraint,
Column,
Date,
DateTime,
Expand DownExpand Up@@ -63,6 +64,7 @@
finish_init,
get_annotations,
get_field_metadata,
get_literal_annotation_info,
get_model_fields,
get_relationship_to,
get_sa_type_from_field,
Expand DownExpand Up@@ -680,6 +682,31 @@ def __init__(
# Ref: https://github.com/sqlalchemy/sqlalchemy/commit/428ea01f00a9cc7f85e435018565eb6da7af1b77
# Tag: 1.4.36
DeclarativeMeta.__init__(cls, classname, bases, dict_, **kw)
table = getattr(cls, "__table__", None)
if table is not None:
# Attach Literal-based value constraints at the database level
for field_name, field in get_model_fields(cls).items(): # type: ignore
annotation = getattr(field, "annotation", None)
literal_info = get_literal_annotation_info(annotation)
if literal_info is None:
continue
base_type, values = literal_info
assert base_type in (str, int, bool)
column = table.c.get(field_name)
if column is None:
continue
if base_type is int:
coerced_values = tuple(int(v) for v in values)
elif base_type is bool:
coerced_values = tuple(bool(v) for v in values)
else:
coerced_values = tuple(str(v) for v in values) # type: ignore
constraint_name = f"ck_{table.name}_{field_name}_literal"
constraint = CheckConstraint(
column.in_(coerced_values),
name=constraint_name,
)
table.append_constraint(constraint)
else:
ModelMetaclass.__init__(cls, classname, bases, dict_, **kw)

Expand Down
147 changes: 146 additions & 1 deletion tests/test_main.py
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
from typing import Annotated
from typing import Annotated, Literal

import pytest
from sqlalchemy import text
from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import RelationshipProperty
from sqlmodel import Field, Relationship, Session, SQLModel, create_engine, select
Expand DownExpand Up@@ -216,3 +217,147 @@ class Hero(SQLModel, table=True):
assert len(foreign_keys) == 1
assert foreign_keys[0].ondelete == "CASCADE"
assert team_id_column.nullable is False


def test_literal_valid_values(clear_sqlmodel, caplog):
"""Test https://github.com/fastapi/sqlmodel/issues/57"""

class Model(SQLModel, table=True):
id: int | None = Field(default=None, primary_key=True)
all_str: Literal["a", "b", "c"]
mixed: Literal["yes", "no", 1, 0]
all_int: Literal[1, 2, 3]
int_bool: Literal[0, 1, True, False]
all_bool: Literal[True, False]

obj = Model(
all_str="a",
mixed="yes",
all_int=1,
int_bool=True,
all_bool=False,
)

engine = create_engine("sqlite://", echo=True)

SQLModel.metadata.create_all(engine)

# Check DDL
assert "all_str VARCHAR NOT NULL" in caplog.text
assert "mixed VARCHAR NOT NULL" in caplog.text
assert "all_int INTEGER NOT NULL" in caplog.text
assert "int_bool INTEGER NOT NULL" in caplog.text
assert "all_bool BOOLEAN NOT NULL" in caplog.text

# Check query
with Session(engine) as session:
session.add(obj)
session.commit()
session.refresh(obj)
assert isinstance(obj.all_str, str)
assert obj.all_str == "a"
assert isinstance(obj.mixed, str)
assert obj.mixed == "yes"
assert isinstance(obj.all_int, int)
assert obj.all_int == 1
assert isinstance(obj.int_bool, int)
assert obj.int_bool == 1
assert isinstance(obj.all_bool, bool)
assert obj.all_bool is False


def test_literal_constraints_invalid_values(clear_sqlmodel):
"""DB should reject values that are not part of the Literal choices."""

class Model(SQLModel, table=True):
id: int | None = Field(default=None, primary_key=True)
all_str: Literal["a", "b", "c"]
mixed: Literal["yes", "no", 1, 0]
all_int: Literal[1, 2, 3]
int_bool: Literal[0, 1, True, False]
all_bool: Literal[True, False]

engine = create_engine("sqlite://")
SQLModel.metadata.create_all(engine)

# Helper to attempt a raw insert that bypasses Pydantic validation so we
# can verify that the database-level CHECK constraints are enforced.
def insert_raw(values: dict[str, object]) -> None:
stmt = text(
"INSERT INTO model (all_str, mixed, all_int, int_bool, all_bool) "
"VALUES (:all_str, :mixed, :all_int, :int_bool, :all_bool)"
).bindparams(**values)
with pytest.raises(IntegrityError):
with Session(engine) as session:
session.exec(stmt)
session.commit()

# Invalid string literal for all_str
insert_raw(
{
"all_str": "z", # invalid, not in {"a","b","c"}
"mixed": "yes",
"all_int": 1,
"int_bool": 1,
"all_bool": 0,
}
)

# Invalid int literal for all_int
insert_raw(
{
"all_str": "a",
"mixed": "yes",
"all_int": 5, # invalid, not in {1,2,3}
"int_bool": 1,
"all_bool": 0,
}
)

# Invalid bool literal for all_bool
insert_raw(
{
"all_str": "a",
"mixed": "yes",
"all_int": 1,
"int_bool": 1,
"all_bool": 2, # invalid boolean value
}
)


def test_literal_optional_and_union_constraints(clear_sqlmodel):
"""Literals inside Optional/Union should also be enforced at the DB level."""

class Model(SQLModel, table=True):
id: int | None = Field(default=None, primary_key=True)
opt_str: Literal["x", "y"] | None = None
union_int: Literal[10, 20] | None = None

engine = create_engine("sqlite://")
SQLModel.metadata.create_all(engine)

# Valid values should be accepted
obj = Model(opt_str="x", union_int=10)
with Session(engine) as session:
session.add(obj)
session.commit()
session.refresh(obj)
assert obj.opt_str == "x"
assert obj.union_int == 10

# Invalid values should be rejected by the database
def insert_raw(values: dict[str, object]) -> None:
stmt = text(
"INSERT INTO model (opt_str, union_int) VALUES (:opt_str, :union_int)"
).bindparams(**values)
with pytest.raises(IntegrityError):
with Session(engine) as session:
session.exec(stmt)
session.commit()

# opt_str not in {"x", "y"}
insert_raw({"opt_str": "z", "union_int": 10})

# union_int not in {10, 20}
insert_raw({"opt_str": "x", "union_int": 30})
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Universal Dark Mode - works on any site\n(function() {\n var enabled = true;\n \n function applyDarkMode() {\n if (!enabled) return;\n \n // Create style element if it doesn't exist\n var style = document.getElementById('universal-dark-mode-style');\n if (!style) {\n style = document.createElement('style');\n style.id = 'universal-dark-mode-style';\n document.head.appendChild(style);\n }\n \n // Dark mode CSS - inverts colors but preserves images/video\n style.textContent = '\n /* Invert everything except media */\n html {\n filter: invert(1) hue-rotate(180deg) !important;\n background: #1a1a2e !important;\n }\n \n /* Restore images, videos, iframes, canvas */\n img, video, iframe, canvas, svg, picture, [style*=\"background-image\"] {\n filter: invert(1) hue-rotate(180deg) !important;\n }\n \n /* Preserve specific elements that should not be inverted */\n .no-dark-mode, .no-dark-mode *,\n [data-theme=\"light\"], [data-theme=\"light\"],\n .ace_editor, .ace_editor *,\n .CodeMirror, .CodeMirror *,\n .monaco-editor, .monaco-editor *,\n .markdown-body pre, .markdown-body pre *,\n .highlight, .highlight *,\n pre code, pre code * {\n filter: none !important;\n }\n \n /* Fix common UI elements */\n .modal, .popup, .dropdown-menu, .tooltip, .popover {\n filter: invert(1) hue-rotate(180deg) !important;\n background: #2d2d44 !important;\n border-color: #444 !important;\n }\n \n /* Scrollbars */\n ::-webkit-scrollbar { background: #1a1a2e !important; }\n ::-webkit-scrollbar-thumb { background: #444 !important; }\n ::-webkit-scrollbar-thumb:hover { background: #555 !important; }\n \n /* Selection */\n ::selection { background: #4ecdc4 !important; color: #1a1a2e !important; }\n ::-moz-selection { background: #4ecdc4 !important; color: #1a1a2e !important; }\n ';\n }\n \n function removeDarkMode() {\n var style = document.getElementById('universal-dark-mode-style');\n if (style) style.remove();\n }\n \n // Toggle with Alt+Shift+D\n document.addEventListener('keydown', function(e) {\n if (e.altKey && e.shiftKey && e.key === 'D') {\n e.preventDefault();\n enabled = !enabled;\n if (enabled) {\n applyDarkMode();\n console.log('[Universal Dark Mode] Enabled');\n } else {\n removeDarkMode();\n console.log('[Universal Dark Mode] Disabled');\n }\n }\n });\n \n // Apply on load\n applyDarkMode();\n \n // Re-apply on dynamic content\n var observer = new MutationObserver(function(mutations) {\n if (enabled && !document.getElementById('universal-dark-mode-style')) {\n applyDarkMode();\n }\n });\n observer.observe(document.head, { childList: true });\n \n console.log('[Universal Dark Mode] Loaded - Press Alt+Shift+D to toggle');\n})();", "Universal Dark Mode"); } } catch(__e) { console.warn('[Userscript:Universal Dark Mode]', __e); } })(); })();
Skip to content
Open
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
37 changes: 37 additions & 0 deletions sqlmodel/_compat.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -9,6 +9,7 @@
Annotated,
Any,
ForwardRef,
Literal,
TypeAlias,
TypeVar,
Union,
Expand DownExpand Up@@ -65,6 +66,36 @@ def _is_union_type(t: Any) -> bool:
return t is UnionType or t is Union


def get_literal_annotation_info(
annotation: Any,
) -> tuple[type[Any], tuple[Any, ...]] | None:
if annotation is None or get_origin(annotation) is None:
return None
origin = get_origin(annotation)
if origin is Annotated:
return get_literal_annotation_info(get_args(annotation)[0])
if _is_union_type(origin):
bases = get_args(annotation)
if len(bases) > 2:
raise ValueError("Cannot have a Union with more than 2 members")
if bases[0] is not NoneType and bases[1] is not NoneType:
raise ValueError("Cannot have a Union without None")
use_type = bases[0] if bases[0] is not NoneType else bases[1]
return get_literal_annotation_info(use_type)
if origin is Literal:
literal_args = get_args(annotation)
if not literal_args:
return None
if all(isinstance(arg, bool) for arg in literal_args): # all bools
base_type: type[Any] = bool
elif all(isinstance(arg, int) for arg in literal_args): # all ints
base_type = int
else:
base_type = str
return base_type, tuple(literal_args)
return None


finish_init: ContextVar[bool] = ContextVar("finish_init", default=True)


Expand DownExpand Up@@ -190,6 +221,12 @@ def get_sa_type_from_type_annotation(annotation: Any) -> Any:
# Optional unions are allowed
use_type = bases[0] if bases[0] is not NoneType else bases[1]
return get_sa_type_from_type_annotation(use_type)
if origin is Literal:
literal_info = get_literal_annotation_info(annotation)
if literal_info is None:
raise ValueError("Literal without values is not supported")
base_type, _ = literal_info
return base_type
return origin


Expand Down
27 changes: 27 additions & 0 deletions sqlmodel/main.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -27,6 +27,7 @@
from pydantic.fields import FieldInfo as PydanticFieldInfo
from sqlalchemy import (
Boolean,
CheckConstraint,
Column,
Date,
DateTime,
Expand DownExpand Up@@ -63,6 +64,7 @@
finish_init,
get_annotations,
get_field_metadata,
get_literal_annotation_info,
get_model_fields,
get_relationship_to,
get_sa_type_from_field,
Expand DownExpand Up@@ -680,6 +682,31 @@ def __init__(
# Ref: https://github.com/sqlalchemy/sqlalchemy/commit/428ea01f00a9cc7f85e435018565eb6da7af1b77
# Tag: 1.4.36
DeclarativeMeta.__init__(cls, classname, bases, dict_, **kw)
table = getattr(cls, "__table__", None)
if table is not None:
# Attach Literal-based value constraints at the database level
for field_name, field in get_model_fields(cls).items(): # type: ignore
annotation = getattr(field, "annotation", None)
literal_info = get_literal_annotation_info(annotation)
if literal_info is None:
continue
base_type, values = literal_info
assert base_type in (str, int, bool)
column = table.c.get(field_name)
if column is None:
continue
if base_type is int:
coerced_values = tuple(int(v) for v in values)
elif base_type is bool:
coerced_values = tuple(bool(v) for v in values)
else:
coerced_values = tuple(str(v) for v in values) # type: ignore
constraint_name = f"ck_{table.name}_{field_name}_literal"
constraint = CheckConstraint(
column.in_(coerced_values),
name=constraint_name,
)
table.append_constraint(constraint)
else:
ModelMetaclass.__init__(cls, classname, bases, dict_, **kw)

Expand Down
147 changes: 146 additions & 1 deletion tests/test_main.py
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
from typing import Annotated
from typing import Annotated, Literal

import pytest
from sqlalchemy import text
from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import RelationshipProperty
from sqlmodel import Field, Relationship, Session, SQLModel, create_engine, select
Expand DownExpand Up@@ -216,3 +217,147 @@ class Hero(SQLModel, table=True):
assert len(foreign_keys) == 1
assert foreign_keys[0].ondelete == "CASCADE"
assert team_id_column.nullable is False


def test_literal_valid_values(clear_sqlmodel, caplog):
"""Test https://github.com/fastapi/sqlmodel/issues/57"""

class Model(SQLModel, table=True):
id: int | None = Field(default=None, primary_key=True)
all_str: Literal["a", "b", "c"]
mixed: Literal["yes", "no", 1, 0]
all_int: Literal[1, 2, 3]
int_bool: Literal[0, 1, True, False]
all_bool: Literal[True, False]

obj = Model(
all_str="a",
mixed="yes",
all_int=1,
int_bool=True,
all_bool=False,
)

engine = create_engine("sqlite://", echo=True)

SQLModel.metadata.create_all(engine)

# Check DDL
assert "all_str VARCHAR NOT NULL" in caplog.text
assert "mixed VARCHAR NOT NULL" in caplog.text
assert "all_int INTEGER NOT NULL" in caplog.text
assert "int_bool INTEGER NOT NULL" in caplog.text
assert "all_bool BOOLEAN NOT NULL" in caplog.text

# Check query
with Session(engine) as session:
session.add(obj)
session.commit()
session.refresh(obj)
assert isinstance(obj.all_str, str)
assert obj.all_str == "a"
assert isinstance(obj.mixed, str)
assert obj.mixed == "yes"
assert isinstance(obj.all_int, int)
assert obj.all_int == 1
assert isinstance(obj.int_bool, int)
assert obj.int_bool == 1
assert isinstance(obj.all_bool, bool)
assert obj.all_bool is False


def test_literal_constraints_invalid_values(clear_sqlmodel):
"""DB should reject values that are not part of the Literal choices."""

class Model(SQLModel, table=True):
id: int | None = Field(default=None, primary_key=True)
all_str: Literal["a", "b", "c"]
mixed: Literal["yes", "no", 1, 0]
all_int: Literal[1, 2, 3]
int_bool: Literal[0, 1, True, False]
all_bool: Literal[True, False]

engine = create_engine("sqlite://")
SQLModel.metadata.create_all(engine)

# Helper to attempt a raw insert that bypasses Pydantic validation so we
# can verify that the database-level CHECK constraints are enforced.
def insert_raw(values: dict[str, object]) -> None:
stmt = text(
"INSERT INTO model (all_str, mixed, all_int, int_bool, all_bool) "
"VALUES (:all_str, :mixed, :all_int, :int_bool, :all_bool)"
).bindparams(**values)
with pytest.raises(IntegrityError):
with Session(engine) as session:
session.exec(stmt)
session.commit()

# Invalid string literal for all_str
insert_raw(
{
"all_str": "z", # invalid, not in {"a","b","c"}
"mixed": "yes",
"all_int": 1,
"int_bool": 1,
"all_bool": 0,
}
)

# Invalid int literal for all_int
insert_raw(
{
"all_str": "a",
"mixed": "yes",
"all_int": 5, # invalid, not in {1,2,3}
"int_bool": 1,
"all_bool": 0,
}
)

# Invalid bool literal for all_bool
insert_raw(
{
"all_str": "a",
"mixed": "yes",
"all_int": 1,
"int_bool": 1,
"all_bool": 2, # invalid boolean value
}
)


def test_literal_optional_and_union_constraints(clear_sqlmodel):
"""Literals inside Optional/Union should also be enforced at the DB level."""

class Model(SQLModel, table=True):
id: int | None = Field(default=None, primary_key=True)
opt_str: Literal["x", "y"] | None = None
union_int: Literal[10, 20] | None = None

engine = create_engine("sqlite://")
SQLModel.metadata.create_all(engine)

# Valid values should be accepted
obj = Model(opt_str="x", union_int=10)
with Session(engine) as session:
session.add(obj)
session.commit()
session.refresh(obj)
assert obj.opt_str == "x"
assert obj.union_int == 10

# Invalid values should be rejected by the database
def insert_raw(values: dict[str, object]) -> None:
stmt = text(
"INSERT INTO model (opt_str, union_int) VALUES (:opt_str, :union_int)"
).bindparams(**values)
with pytest.raises(IntegrityError):
with Session(engine) as session:
session.exec(stmt)
session.commit()

# opt_str not in {"x", "y"}
insert_raw({"opt_str": "z", "union_int": 10})

# union_int not in {10, 20}
insert_raw({"opt_str": "x", "union_int": 30})