Skip to content
Merged
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
2 changes: 1 addition & 1 deletion .pre-commit-config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -37,7 +37,7 @@ repos:
- id: local-ty
name: ty check
entry: >-
uv run ty check sqlmodel tests/test_dataclass_transform.py tests/test_field_sa_type.py
uv run ty check sqlmodel tests/test_asyncio.py tests/test_dataclass_transform.py tests/test_field_sa_type.py
tests/test_select_typing.py
require_serial: true
language: unsupported
Expand Down
2 changes: 1 addition & 1 deletion .python-version
Original file line number Diff line number Diff line change
@@ -1 +1 @@
3.10
3.11
7 changes: 6 additions & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -35,11 +35,14 @@ classifiers = [
]

dependencies = [
"SQLAlchemy >=2.0.14,<2.1.0",
"SQLAlchemy >=2.0.14,<2.2.0",
"pydantic>=2.11.0",
"typing-extensions>=4.5.0",
]

[project.optional-dependencies]
asyncio = ["SQLAlchemy[asyncio]"]

[project.urls]
Homepage = "https://github.com/fastapi/sqlmodel"
Documentation = "https://sqlmodel.tiangolo.com"
Expand Down Expand Up @@ -75,6 +78,8 @@ github-actions = [
"smokeshow >=0.5.0",
]
tests = [
"SQLAlchemy[asyncio]",
"aiosqlite >=0.17.0",
"alembic >=1.12.0",
"black >=24.1.0",
"coverage[toml] >=6.2",
Expand Down
2 changes: 1 addition & 1 deletion scripts/lint.sh
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,6 @@
set -e
set -x

ty check sqlmodel tests/test_dataclass_transform.py tests/test_field_sa_type.py tests/test_select_typing.py
ty check sqlmodel tests/test_asyncio.py tests/test_dataclass_transform.py tests/test_field_sa_type.py tests/test_select_typing.py
ruff check sqlmodel tests docs_src scripts
ruff format sqlmodel tests docs_src scripts --check
21 changes: 13 additions & 8 deletions sqlmodel/ext/asyncio/session.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,13 +17,14 @@
from sqlalchemy.sql.base import Executable as _Executable
from sqlalchemy.sql.dml import UpdateBase
from sqlalchemy.util.concurrency import greenlet_spawn
from typing_extensions import deprecated
from typing_extensions import TypeVarTuple, Unpack, deprecated

from ...orm.session import Session
from ...sql.base import Executable
from ...sql.expression import Select, SelectOfScalar

_TSelectParam = TypeVar("_TSelectParam", bound=Any)
_Ts = TypeVarTuple("_Ts")


class AsyncSession(_AsyncSession):
Expand All @@ -33,14 +34,14 @@ class AsyncSession(_AsyncSession):
@overload
async def exec(
self,
statement: Select[_TSelectParam],
statement: Select[Unpack[_Ts]],
*,
params: Mapping[str, Any] | Sequence[Mapping[str, Any]] | None = None,
execution_options: Mapping[str, Any] = util.EMPTY_DICT,
bind_arguments: dict[str, Any] | None = None,
_parent_execute_state: Any | None = None,
_add_event: Any | None = None,
) -> TupleResult[_TSelectParam]: ...
) -> TupleResult[tuple[Unpack[_Ts]]]: ...

@overload
async def exec(
Expand All @@ -64,11 +65,11 @@ async def exec(
bind_arguments: dict[str, Any] | None = None,
_parent_execute_state: Any | None = None,
_add_event: Any | None = None,
) -> CursorResult[Any]: ...
) -> CursorResult[Unpack[tuple[Any, ...]]]: ...

async def exec(
self,
statement: Select[_TSelectParam]
statement: Select[Unpack[_Ts]]
| SelectOfScalar[_TSelectParam]
| Executable[_TSelectParam]
| UpdateBase,
Expand All @@ -78,7 +79,11 @@ async def exec(
bind_arguments: dict[str, Any] | None = None,
_parent_execute_state: Any | None = None,
_add_event: Any | None = None,
) -> TupleResult[_TSelectParam] | ScalarResult[_TSelectParam] | CursorResult[Any]:
) -> (
TupleResult[tuple[Unpack[_Ts]]]
| ScalarResult[_TSelectParam]
| CursorResult[Unpack[tuple[Any, ...]]]
):
if execution_options:
execution_options = util.immutabledict(execution_options).union(
_EXECUTE_OPTIONS
Expand All @@ -96,7 +101,7 @@ async def exec(
_add_event=_add_event,
)
result_value = await _ensure_sync_result(
cast(Result[_TSelectParam], result), self.exec
cast(Result[Unpack[tuple[Any, ...]]], result), self.exec
)
return result_value # type: ignore

Expand Down Expand Up @@ -131,7 +136,7 @@ async def execute(
bind_arguments: dict[str, Any] | None = None,
_parent_execute_state: Any | None = None,
_add_event: Any | None = None,
) -> Result[Any]:
) -> Result[Unpack[tuple[Any, ...]]]:
"""
🚨 You probably want to use `session.exec()` instead of `session.execute()`.
Expand Down
19 changes: 12 additions & 7 deletions sqlmodel/orm/session.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,23 +17,24 @@
from sqlalchemy.sql.dml import UpdateBase
from sqlmodel.sql.base import Executable
from sqlmodel.sql.expression import Select, SelectOfScalar
from typing_extensions import deprecated
from typing_extensions import TypeVarTuple, Unpack, deprecated

_TSelectParam = TypeVar("_TSelectParam", bound=Any)
_Ts = TypeVarTuple("_Ts")


class Session(_Session):
@overload
def exec(
self,
statement: Select[_TSelectParam],
statement: Select[Unpack[_Ts]],
*,
params: Mapping[str, Any] | Sequence[Mapping[str, Any]] | None = None,
execution_options: Mapping[str, Any] = util.EMPTY_DICT,
bind_arguments: dict[str, Any] | None = None,
_parent_execute_state: Any | None = None,
_add_event: Any | None = None,
) -> TupleResult[_TSelectParam]: ...
) -> TupleResult[tuple[Unpack[_Ts]]]: ...

@overload
def exec(
Expand All @@ -57,11 +58,11 @@ def exec(
bind_arguments: dict[str, Any] | None = None,
_parent_execute_state: Any | None = None,
_add_event: Any | None = None,
) -> CursorResult[Any]: ...
) -> CursorResult[Unpack[tuple[Any, ...]]]: ...

def exec(
self,
statement: Select[_TSelectParam]
statement: Select[Unpack[_Ts]]
| SelectOfScalar[_TSelectParam]
| Executable[_TSelectParam]
| UpdateBase,
Expand All @@ -71,7 +72,11 @@ def exec(
bind_arguments: dict[str, Any] | None = None,
_parent_execute_state: Any | None = None,
_add_event: Any | None = None,
) -> TupleResult[_TSelectParam] | ScalarResult[_TSelectParam] | CursorResult[Any]:
) -> (
TupleResult[tuple[Unpack[_Ts]]]
| ScalarResult[_TSelectParam]
| CursorResult[Unpack[tuple[Any, ...]]]
):
results = super().execute(
statement,
params=params,
Expand Down Expand Up @@ -114,7 +119,7 @@ def execute(
bind_arguments: dict[str, Any] | None = None,
_parent_execute_state: Any | None = None,
_add_event: Any | None = None,
) -> Result[Any]:
) -> Result[Unpack[tuple[Any, ...]]]:
"""
🚨 You probably want to use `session.exec()` instead of `session.execute()`.
Expand Down
7 changes: 4 additions & 3 deletions sqlmodel/sql/_expression_select_cls.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,14 +6,15 @@
_ColumnExpressionArgument,
)
from sqlalchemy.sql.expression import Select as _Select
from typing_extensions import Self
from typing_extensions import Self, TypeVarTuple, Unpack

_T = TypeVar("_T")
_Ts = TypeVarTuple("_Ts")


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

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


class Select(SelectBase[_T]):
class Select(SelectBase[Unpack[_Ts]]):
inherit_cache = True


Expand Down
Loading
Loading