Skip to content

Commit 7d1eba7

Browse files
authored
✨ Add support for SQLAlchemy 2.1 (#2112)
1 parent a99fbcd commit 7d1eba7

13 files changed

Lines changed: 388 additions & 103 deletions

‎.pre-commit-config.yaml‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -37,7 +37,7 @@ repos:
3737
- id: local-ty
3838
name: ty check
3939
entry: >-
40-
uv run ty check sqlmodel tests/test_dataclass_transform.py tests/test_field_sa_type.py
40+
uv run ty check sqlmodel tests/test_asyncio.py tests/test_dataclass_transform.py tests/test_field_sa_type.py
4141
tests/test_select_typing.py
4242
require_serial: true
4343
language: unsupported

‎.python-version‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1 +1 @@
1-
3.10
1+
3.11

‎pyproject.toml‎

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -35,11 +35,14 @@ classifiers = [
3535
]
3636

3737
dependencies = [
38-
"SQLAlchemy >=2.0.14,<2.1.0",
38+
"SQLAlchemy >=2.0.14,<2.2.0",
3939
"pydantic>=2.11.0",
4040
"typing-extensions>=4.5.0",
4141
]
4242

43+
[project.optional-dependencies]
44+
asyncio = ["SQLAlchemy[asyncio]"]
45+
4346
[project.urls]
4447
Homepage = "https://github.com/fastapi/sqlmodel"
4548
Documentation = "https://sqlmodel.tiangolo.com"
@@ -75,6 +78,8 @@ github-actions = [
7578
"smokeshow >=0.5.0",
7679
]
7780
tests = [
81+
"SQLAlchemy[asyncio]",
82+
"aiosqlite >=0.17.0",
7883
"alembic >=1.12.0",
7984
"black >=24.1.0",
8085
"coverage[toml] >=6.2",

‎scripts/lint.sh‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,6 @@
33
set -e
44
set -x
55

6-
ty check sqlmodel tests/test_dataclass_transform.py tests/test_field_sa_type.py tests/test_select_typing.py
6+
ty check sqlmodel tests/test_asyncio.py tests/test_dataclass_transform.py tests/test_field_sa_type.py tests/test_select_typing.py
77
ruff check sqlmodel tests docs_src scripts
88
ruff format sqlmodel tests docs_src scripts --check

‎sqlmodel/ext/asyncio/session.py‎

Lines changed: 13 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -17,13 +17,14 @@
1717
from sqlalchemy.sql.base import Executable as _Executable
1818
from sqlalchemy.sql.dml import UpdateBase
1919
from sqlalchemy.util.concurrency import greenlet_spawn
20-
from typing_extensions import deprecated
20+
from typing_extensions import TypeVarTuple, Unpack, deprecated
2121

2222
from ...orm.session import Session
2323
from ...sql.base import Executable
2424
from ...sql.expression import Select, SelectOfScalar
2525

2626
_TSelectParam = TypeVar("_TSelectParam", bound=Any)
27+
_Ts = TypeVarTuple("_Ts")
2728

2829

2930
class AsyncSession(_AsyncSession):
@@ -33,14 +34,14 @@ class AsyncSession(_AsyncSession):
3334
@overload
3435
async def exec(
3536
self,
36-
statement: Select[_TSelectParam],
37+
statement: Select[Unpack[_Ts]],
3738
*,
3839
params: Mapping[str, Any] | Sequence[Mapping[str, Any]] | None = None,
3940
execution_options: Mapping[str, Any] = util.EMPTY_DICT,
4041
bind_arguments: dict[str, Any] | None = None,
4142
_parent_execute_state: Any | None = None,
4243
_add_event: Any | None = None,
43-
) -> TupleResult[_TSelectParam]: ...
44+
) -> TupleResult[tuple[Unpack[_Ts]]]: ...
4445

4546
@overload
4647
async def exec(
@@ -64,11 +65,11 @@ async def exec(
6465
bind_arguments: dict[str, Any] | None = None,
6566
_parent_execute_state: Any | None = None,
6667
_add_event: Any | None = None,
67-
) -> CursorResult[Any]: ...
68+
) -> CursorResult[Unpack[tuple[Any, ...]]]: ...
6869

6970
async def exec(
7071
self,
71-
statement: Select[_TSelectParam]
72+
statement: Select[Unpack[_Ts]]
7273
| SelectOfScalar[_TSelectParam]
7374
| Executable[_TSelectParam]
7475
| UpdateBase,
@@ -78,7 +79,11 @@ async def exec(
7879
bind_arguments: dict[str, Any] | None = None,
7980
_parent_execute_state: Any | None = None,
8081
_add_event: Any | None = None,
81-
) -> TupleResult[_TSelectParam] | ScalarResult[_TSelectParam] | CursorResult[Any]:
82+
) -> (
83+
TupleResult[tuple[Unpack[_Ts]]]
84+
| ScalarResult[_TSelectParam]
85+
| CursorResult[Unpack[tuple[Any, ...]]]
86+
):
8287
if execution_options:
8388
execution_options = util.immutabledict(execution_options).union(
8489
_EXECUTE_OPTIONS
@@ -96,7 +101,7 @@ async def exec(
96101
_add_event=_add_event,
97102
)
98103
result_value = await _ensure_sync_result(
99-
cast(Result[_TSelectParam], result), self.exec
104+
cast(Result[Unpack[tuple[Any, ...]]], result), self.exec
100105
)
101106
return result_value # type: ignore
102107

@@ -131,7 +136,7 @@ async def execute(
131136
bind_arguments: dict[str, Any] | None = None,
132137
_parent_execute_state: Any | None = None,
133138
_add_event: Any | None = None,
134-
) -> Result[Any]:
139+
) -> Result[Unpack[tuple[Any, ...]]]:
135140
"""
136141
🚨 You probably want to use `session.exec()` instead of `session.execute()`.
137142

‎sqlmodel/orm/session.py‎

Lines changed: 12 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -17,23 +17,24 @@
1717
from sqlalchemy.sql.dml import UpdateBase
1818
from sqlmodel.sql.base import Executable
1919
from sqlmodel.sql.expression import Select, SelectOfScalar
20-
from typing_extensions import deprecated
20+
from typing_extensions import TypeVarTuple, Unpack, deprecated
2121

2222
_TSelectParam = TypeVar("_TSelectParam", bound=Any)
23+
_Ts = TypeVarTuple("_Ts")
2324

2425

2526
class Session(_Session):
2627
@overload
2728
def exec(
2829
self,
29-
statement: Select[_TSelectParam],
30+
statement: Select[Unpack[_Ts]],
3031
*,
3132
params: Mapping[str, Any] | Sequence[Mapping[str, Any]] | None = None,
3233
execution_options: Mapping[str, Any] = util.EMPTY_DICT,
3334
bind_arguments: dict[str, Any] | None = None,
3435
_parent_execute_state: Any | None = None,
3536
_add_event: Any | None = None,
36-
) -> TupleResult[_TSelectParam]: ...
37+
) -> TupleResult[tuple[Unpack[_Ts]]]: ...
3738

3839
@overload
3940
def exec(
@@ -57,11 +58,11 @@ def exec(
5758
bind_arguments: dict[str, Any] | None = None,
5859
_parent_execute_state: Any | None = None,
5960
_add_event: Any | None = None,
60-
) -> CursorResult[Any]: ...
61+
) -> CursorResult[Unpack[tuple[Any, ...]]]: ...
6162

6263
def exec(
6364
self,
64-
statement: Select[_TSelectParam]
65+
statement: Select[Unpack[_Ts]]
6566
| SelectOfScalar[_TSelectParam]
6667
| Executable[_TSelectParam]
6768
| UpdateBase,
@@ -71,7 +72,11 @@ def exec(
7172
bind_arguments: dict[str, Any] | None = None,
7273
_parent_execute_state: Any | None = None,
7374
_add_event: Any | None = None,
74-
) -> TupleResult[_TSelectParam] | ScalarResult[_TSelectParam] | CursorResult[Any]:
75+
) -> (
76+
TupleResult[tuple[Unpack[_Ts]]]
77+
| ScalarResult[_TSelectParam]
78+
| CursorResult[Unpack[tuple[Any, ...]]]
79+
):
7580
results = super().execute(
7681
statement,
7782
params=params,
@@ -114,7 +119,7 @@ def execute(
114119
bind_arguments: dict[str, Any] | None = None,
115120
_parent_execute_state: Any | None = None,
116121
_add_event: Any | None = None,
117-
) -> Result[Any]:
122+
) -> Result[Unpack[tuple[Any, ...]]]:
118123
"""
119124
🚨 You probably want to use `session.exec()` instead of `session.execute()`.
120125

‎sqlmodel/sql/_expression_select_cls.py‎

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -6,14 +6,15 @@
66
_ColumnExpressionArgument,
77
)
88
from sqlalchemy.sql.expression import Select as _Select
9-
from typing_extensions import Self
9+
from typing_extensions import Self, TypeVarTuple, Unpack
1010

1111
_T = TypeVar("_T")
12+
_Ts = TypeVarTuple("_Ts")
1213

1314

1415
# Separate this class in SelectBase, Select, and SelectOfScalar so that they can share
1516
# where and having without having type overlap incompatibility in session.exec().
16-
class SelectBase(_Select[tuple[_T]]):
17+
class SelectBase(_Select[Unpack[_Ts]]):
1718
inherit_cache = True
1819

1920
def where(self, *whereclause: _ColumnExpressionArgument[bool] | bool) -> Self:
@@ -29,7 +30,7 @@ def having(self, *having: _ColumnExpressionArgument[bool] | bool) -> Self:
2930
return super().having(*having) # ty: ignore[invalid-argument-type]
3031

3132

32-
class Select(SelectBase[_T]):
33+
class Select(SelectBase[Unpack[_Ts]]):
3334
inherit_cache = True
3435

3536

0 commit comments

Comments
 (0)