Skip to content
Open
1 change: 1 addition & 0 deletions sqlmodel/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -145,4 +145,5 @@
from .sql.expression import type_coerce as type_coerce
from .sql.expression import within_group as within_group
from .sql.sqltypes import AutoString as AutoString
from .sql.sqltypes import IntEnum as IntEnum
from .sql.sqltypes import UTCDateTime as UTCDateTime
35 changes: 34 additions & 1 deletion sqlmodel/sql/sqltypes.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,7 @@
import enum
from datetime import datetime, timedelta, timezone
from typing import Any, cast
from enum import IntEnum as _IntEnum
from typing import Any, TypeVar, cast

from sqlalchemy import types
from sqlalchemy.engine.interfaces import Dialect
Expand Down Expand Up @@ -59,3 +61,34 @@ def load_dialect_impl(self, dialect: Dialect) -> "types.TypeEngine[Any]":
if impl.length is None and dialect.name == "mysql":
return dialect.type_descriptor(types.String(self.mysql_default_length))
return super().load_dialect_impl(dialect)


_TIntEnum = TypeVar("_TIntEnum", bound="_IntEnum")


class IntEnum(types.TypeDecorator[_TIntEnum | None]):
impl = types.SmallInteger
cache_ok = True

def __init__(self, enum_type: type[_TIntEnum], *args: Any, **kwargs: Any):
super().__init__(*args, **kwargs)

# validate the input enum type
if not issubclass(enum_type, enum.IntEnum):
raise TypeError("Input must be enum.IntEnum")

self.enum_type = enum_type

def process_result_value(
self,
value: int | None,
dialect: Dialect,
) -> _TIntEnum | None:
return None if (value is None) else self.enum_type(value)

def process_bind_param(
self,
value: _TIntEnum | None,
dialect: Dialect,
) -> int | None:
return None if (value is None) else value.value
14 changes: 10 additions & 4 deletions tests/test_enums.py
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,7 @@ def test_postgres_ddl_sql(clear_sqlmodel, capsys: pytest.CaptureFixture[str]):
captured = capsys.readouterr()
assert "CREATE TYPE myenum1 AS ENUM ('A', 'B');" in captured.out
assert "CREATE TYPE myenum2 AS ENUM ('C', 'D');" in captured.out
assert "int_enum_field SMALLINT NOT NULL" in captured.out


def test_sqlite_ddl_sql(clear_sqlmodel, capsys: pytest.CaptureFixture[str]):
Expand All @@ -51,6 +52,7 @@ def test_sqlite_ddl_sql(clear_sqlmodel, capsys: pytest.CaptureFixture[str]):

captured = capsys.readouterr()
assert "enum_field VARCHAR(1) NOT NULL" in captured.out, captured
assert "int_enum_field SMALLINT NOT NULL" in captured.out, captured
assert "CREATE TYPE" not in captured.out


Expand All @@ -61,10 +63,12 @@ def test_json_schema_flat_model_pydantic_v2():
"properties": {
"id": {"title": "Id", "type": "string", "format": "uuid"},
"enum_field": {"$ref": "#/$defs/MyEnum1"},
"int_enum_field": {"$ref": "#/$defs/MyEnum3"},
},
"required": ["id", "enum_field"],
"required": ["id", "enum_field", "int_enum_field"],
"$defs": {
"MyEnum1": {"enum": ["A", "B"], "title": "MyEnum1", "type": "string"}
"MyEnum1": {"enum": ["A", "B"], "title": "MyEnum1", "type": "string"},
"MyEnum3": {"enum": [1, 2], "title": "MyEnum3", "type": "integer"},
},
}

Expand All @@ -76,9 +80,11 @@ def test_json_schema_inherit_model_pydantic_v2():
"properties": {
"id": {"title": "Id", "type": "string", "format": "uuid"},
"enum_field": {"$ref": "#/$defs/MyEnum2"},
"int_enum_field": {"$ref": "#/$defs/MyEnum3"},
},
"required": ["id", "enum_field"],
"required": ["id", "enum_field", "int_enum_field"],
"$defs": {
"MyEnum2": {"enum": ["C", "D"], "title": "MyEnum2", "type": "string"}
"MyEnum2": {"enum": ["C", "D"], "title": "MyEnum2", "type": "string"},
"MyEnum3": {"enum": [1, 2], "title": "MyEnum3", "type": "integer"},
},
}
9 changes: 8 additions & 1 deletion tests/test_enums_models.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,7 @@
import enum
import uuid

from sqlmodel import Field, SQLModel
from sqlmodel import Field, IntEnum, SQLModel


class MyEnum1(str, enum.Enum):
Expand All @@ -14,14 +14,21 @@ class MyEnum2(str, enum.Enum):
D = "D"


class MyEnum3(enum.IntEnum):
E = 1
F = 2


class BaseModel(SQLModel):
id: uuid.UUID = Field(primary_key=True)
enum_field: MyEnum2
int_enum_field: MyEnum3 = Field(sa_type=IntEnum(MyEnum3))


class FlatModel(SQLModel, table=True):
id: uuid.UUID = Field(primary_key=True)
enum_field: MyEnum1
int_enum_field: MyEnum3 = Field(sa_type=IntEnum(MyEnum3))


class InheritModel(BaseModel, table=True):
Expand Down
Loading