From 4c856964e11e129fcd6e9920603e550a5ae55b3e Mon Sep 17 00:00:00 2001 From: BrianLusina <12752833+BrianLusina@users.noreply.github.com> Date: Sun, 12 Jul 2026 19:12:11 +0300 Subject: [PATCH 01/12] feat(sql, session, repository): add async types to session and repository --- sanctumlabs_dbkit/sql/callbacks.py | 15 +- sanctumlabs_dbkit/sql/repository/__init__.py | 7 + .../sql/repository/async_repository.py | 130 ++++++++++++++++++ .../sql/{ => repository}/repository.py | 12 +- sanctumlabs_dbkit/sql/repository/types.py | 10 ++ sanctumlabs_dbkit/sql/session/__init__.py | 15 ++ .../sql/session/async_session.py | 103 ++++++++++++++ .../sql/{ => session}/session.py | 4 +- sanctumlabs_dbkit/sql/session/types.py | 3 + sanctumlabs_dbkit/sql/types.py | 2 + 10 files changed, 289 insertions(+), 12 deletions(-) create mode 100644 sanctumlabs_dbkit/sql/repository/__init__.py create mode 100644 sanctumlabs_dbkit/sql/repository/async_repository.py rename sanctumlabs_dbkit/sql/{ => repository}/repository.py (95%) create mode 100644 sanctumlabs_dbkit/sql/repository/types.py create mode 100644 sanctumlabs_dbkit/sql/session/__init__.py create mode 100644 sanctumlabs_dbkit/sql/session/async_session.py rename sanctumlabs_dbkit/sql/{ => session}/session.py (96%) create mode 100644 sanctumlabs_dbkit/sql/session/types.py diff --git a/sanctumlabs_dbkit/sql/callbacks.py b/sanctumlabs_dbkit/sql/callbacks.py index d444a0b..d219834 100644 --- a/sanctumlabs_dbkit/sql/callbacks.py +++ b/sanctumlabs_dbkit/sql/callbacks.py @@ -4,8 +4,8 @@ from typing import List, cast -from sanctumlabs_dbkit.sql.session import Session -from sanctumlabs_dbkit.sql.types import CommitCallback +from sanctumlabs_dbkit.sql.session import Session, AsyncSession +from sanctumlabs_dbkit.sql.types import CommitCallback, CommitCallbackAsync def on_commit(current_session: Session, callback: CommitCallback) -> None: @@ -14,3 +14,14 @@ def on_commit(current_session: Session, callback: CommitCallback) -> None: List[CommitCallback], current_session.info.setdefault("on_commit_hooks", []) ) commit_hooks.append(callback) + + +def on_commit_async( + current_session: AsyncSession, callback: CommitCallbackAsync +) -> None: + """Sets an async commit callback to the current session""" + commit_hooks = cast( + List[CommitCallbackAsync], + current_session.info.setdefault("on_commit_hooks", []), + ) + commit_hooks.append(callback) diff --git a/sanctumlabs_dbkit/sql/repository/__init__.py b/sanctumlabs_dbkit/sql/repository/__init__.py new file mode 100644 index 0000000..2c7f64b --- /dev/null +++ b/sanctumlabs_dbkit/sql/repository/__init__.py @@ -0,0 +1,7 @@ +from sanctumlabs_dbkit.sql.repository.repository import Repository +from sanctumlabs_dbkit.sql.repository.async_repository import AsyncRepository + +__all__ = [ + "Repository", + "AsyncRepository", +] diff --git a/sanctumlabs_dbkit/sql/repository/async_repository.py b/sanctumlabs_dbkit/sql/repository/async_repository.py new file mode 100644 index 0000000..e868e35 --- /dev/null +++ b/sanctumlabs_dbkit/sql/repository/async_repository.py @@ -0,0 +1,130 @@ +""" +Contains a generic base repository or DAO for access patterns to data for a given database model +""" + +from datetime import datetime, UTC +from typing import ( + Generic, + Any, + Optional, + Sequence, + Type, + cast, + TypeGuard, +) + +from sqlalchemy import ColumnElement, Select, select + +from sanctumlabs_dbkit.exceptions import ModelNotFoundError +from sanctumlabs_dbkit.exceptions import UnsupportedModelOperationError +from sanctumlabs_dbkit.sql.session.async_session import AsyncSession +from sanctumlabs_dbkit.sql.models import AbstractBaseModel +from sanctumlabs_dbkit.sql.repository.types import T + + +class AsyncRepository(Generic[T]): + """ + A base class for implementing an async Repository or DAO. + + ```python + job_dao = AsyncRepository(model=Job, session=async_session) + + job = job_dao.find("123") + """ + + def __init__(self, model: Type[T], session: AsyncSession) -> None: + """Creates an instance of the Repository""" + self.model = model + self.session = session + + @staticmethod + def _supports_soft_deletion(model: Type[T]) -> TypeGuard[Type[AbstractBaseModel]]: + """ + Indicates if the provided model supports soft deletion (has a 'deleted_at' column). This function + takes in an argument due to mypy typeguarding requirements, and is thus static. + """ + return issubclass(model, AbstractBaseModel) + + def create(self, refresh: bool = False, **kwargs: Any) -> T: + """Creates a new entity + + Args: + refresh (bool, optional): whether to refresh the model with the data in the return. Defaults to False. + + Returns: + T: The created model instance + """ + model_instance = self.model(**kwargs) + self.session.add(model_instance) + + if refresh: + self.session.flush() + self.session.refresh(model_instance) + + return cast(T, model_instance) + + def query(self, include_deleted: bool = False) -> Select: + """Returns a select query with the model including deleted records if the include_deleted is set to True""" + selectable = select(self.model) + + if not include_deleted and self._supports_soft_deletion(self.model): + selectable = selectable.where( + self.model.deleted_at == self.model.not_deleted_value() + ) + + return selectable + + async def find(self, pk: Any, include_deleted: bool = False) -> Optional[T]: + """Retrieve a given model given its primary key""" + pk_column = cast(ColumnElement, getattr(self.model, self.model.pk)) + + statement = self.query(include_deleted).where(pk_column == pk).limit(1) + scalars = await self.session.scalars(statement) + + return scalars.first() + + async def find_or_raise(self, pk: Any, include_deleted: bool = False) -> T: + """Finds the given entity or raises an exception if the entity can not be found""" + entity = await self.find(pk, include_deleted) + + if not entity: + raise ModelNotFoundError( + f"The model {self.model.__name__} {pk} does not exist" + ) + + return entity + + async def all(self, include_deleted: bool = False) -> Sequence[T]: + """Retrieves all records for the given model""" + statement = self.query(include_deleted) + scalars = await self.session.scalars(statement) + + return scalars.all() + + async def delete(self, pk: Any) -> None: + """Deletes a given record with the given primary key""" + if not self._supports_soft_deletion(self.model): + raise UnsupportedModelOperationError( + f"The model {self.model.__name__} {pk} does not support soft deletion." + ) + + # Cast here as mypy type narrowing doesn't infer the type of entity + # correctly + entity = cast(AbstractBaseModel, await self.find(pk)) + + if entity: + entity.deleted_at = datetime.now(UTC) + + async def list( + self, limit: int = 20, offset: int = 0, include_deleted: bool = False + ) -> Sequence[T]: + """Returns a list of records for the given database record""" + statement = ( + self.query(include_deleted) + .order_by(self.model.created_at.desc()) + .limit(limit) + .offset(offset) + ) + scalars = await self.session.scalars(statement) + + return scalars.all() diff --git a/sanctumlabs_dbkit/sql/repository.py b/sanctumlabs_dbkit/sql/repository/repository.py similarity index 95% rename from sanctumlabs_dbkit/sql/repository.py rename to sanctumlabs_dbkit/sql/repository/repository.py index 481776f..9613e1af 100644 --- a/sanctumlabs_dbkit/sql/repository.py +++ b/sanctumlabs_dbkit/sql/repository/repository.py @@ -9,21 +9,17 @@ Optional, Sequence, Type, - TypeVar, cast, TypeGuard, - Union, ) + from sqlalchemy import ColumnElement, Select, select from sanctumlabs_dbkit.exceptions import ModelNotFoundError -from sanctumlabs_dbkit.sql.models import AbstractBaseModel, BaseOutboxEvent -from sanctumlabs_dbkit.sql.session import Session from sanctumlabs_dbkit.exceptions import UnsupportedModelOperationError - -RepositoryBaseModel = Union[AbstractBaseModel, BaseOutboxEvent] - -T = TypeVar("T", bound=RepositoryBaseModel) +from sanctumlabs_dbkit.sql.models import AbstractBaseModel +from sanctumlabs_dbkit.sql.repository.types import T +from sanctumlabs_dbkit.sql.session import Session class Repository(Generic[T]): diff --git a/sanctumlabs_dbkit/sql/repository/types.py b/sanctumlabs_dbkit/sql/repository/types.py new file mode 100644 index 0000000..89e6476 --- /dev/null +++ b/sanctumlabs_dbkit/sql/repository/types.py @@ -0,0 +1,10 @@ +from typing import ( + TypeVar, + Union, +) + +from sanctumlabs_dbkit.sql.models import AbstractBaseModel, BaseOutboxEvent + +RepositoryBaseModel = Union[AbstractBaseModel, BaseOutboxEvent] + +T = TypeVar("T", bound=RepositoryBaseModel) diff --git a/sanctumlabs_dbkit/sql/session/__init__.py b/sanctumlabs_dbkit/sql/session/__init__.py new file mode 100644 index 0000000..9bda53a --- /dev/null +++ b/sanctumlabs_dbkit/sql/session/__init__.py @@ -0,0 +1,15 @@ +from sanctumlabs_dbkit.sql.session.async_session import ( + AsyncSession, + AsyncSessionLocal, + async_transaction, +) +from sanctumlabs_dbkit.sql.session.session import Session, SessionLocal, transaction + +__all__ = [ + "Session", + "AsyncSession", + "SessionLocal", + "transaction", + "AsyncSessionLocal", + "async_transaction", +] diff --git a/sanctumlabs_dbkit/sql/session/async_session.py b/sanctumlabs_dbkit/sql/session/async_session.py new file mode 100644 index 0000000..c64240d --- /dev/null +++ b/sanctumlabs_dbkit/sql/session/async_session.py @@ -0,0 +1,103 @@ +""" +Async Session module contains implementation logic for a database session +""" + +import functools +from typing import Any + +from sqlalchemy.ext.asyncio import ( + AsyncSession as BaseAsyncSession, + AsyncSessionTransaction, + async_sessionmaker, +) + +from sanctumlabs_dbkit.sql.session.types import FuncT + + +class AsyncSession(BaseAsyncSession): + """ + Session that subclasses SQLAlchemy Base Session class adding more functionality around a database session + """ + + def begin(self, nested: bool = False) -> AsyncSessionTransaction: + """Begins an async session transaction""" + if nested: + return super().begin_nested() + return super().begin() + + def transaction(self, func: FuncT) -> FuncT: + """ + A decorator to wrap a function within a transaction. + + If we are already within a transaction, a nested transaction will be started. + + Example: + + ```python + from sanctumlabs_dbkit.sql import AsyncSessionLocal + + session = SessionLocal() + + @session.transaction + def create_user(payload) -> User: + user = User(**payload) + session.add(user) + + return user + + create_user({"first_name": "Bob"}) + """ + + @functools.wraps(func) + async def wrapper(*args: Any, **kwargs: Any) -> Any: + async with self.begin(): + return func(*args, **kwargs) + + return wrapper + + +async def async_transaction(func: FuncT) -> FuncT: + """ + A decorator to wrap an instance method within a transaction. + + If we are already within a transaction, a nested transaction will be started. + + Example: + + ```python + from sanctumlabs_dbkit.sql import AsyncSessionLocal + from sanctumlabs_dbkit.sql.async_session import transaction + + class UserService(): + def __init__(session: AsyncSession): + self.session = session + + @transaction + def create(payload) -> User: + user = User(**payload) + self.session.add(user) + + return user + + session = AsyncSessionLocal() + + user_service = UserService(session) + user_service.create({"first_name": "Bob"}) + """ + + @functools.wraps(func) + async def wrapper(self: Any, *args: Any, **kwargs: Any) -> Any: + if not self.session or not isinstance(self.session, AsyncSession): + # pylint: disable=broad-exception-raised + raise Exception( + "The @transaction decorator requires that an instance variable `session` be set to an instance of a " + "`Session`." + ) + + async with self.session.begin(): + return func(self, *args, **kwargs) + + return wrapper + + +AsyncSessionLocal = async_sessionmaker(class_=AsyncSession) diff --git a/sanctumlabs_dbkit/sql/session.py b/sanctumlabs_dbkit/sql/session/session.py similarity index 96% rename from sanctumlabs_dbkit/sql/session.py rename to sanctumlabs_dbkit/sql/session/session.py index d7d46de..f14623f 100644 --- a/sanctumlabs_dbkit/sql/session.py +++ b/sanctumlabs_dbkit/sql/session/session.py @@ -3,11 +3,11 @@ """ import functools -from typing import Any, Callable, TypeVar, cast +from typing import Any, cast from sqlalchemy.orm import SessionTransaction, Session as BaseSession, sessionmaker -FuncT = TypeVar("FuncT", bound=Callable[..., Any]) +from sanctumlabs_dbkit.sql.session.types import FuncT class Session(BaseSession): diff --git a/sanctumlabs_dbkit/sql/session/types.py b/sanctumlabs_dbkit/sql/session/types.py new file mode 100644 index 0000000..c2b83c0 --- /dev/null +++ b/sanctumlabs_dbkit/sql/session/types.py @@ -0,0 +1,3 @@ +from typing import Any, Callable, TypeVar + +FuncT = TypeVar("FuncT", bound=Callable[..., Any]) diff --git a/sanctumlabs_dbkit/sql/types.py b/sanctumlabs_dbkit/sql/types.py index ccb496c..c00f1f0 100644 --- a/sanctumlabs_dbkit/sql/types.py +++ b/sanctumlabs_dbkit/sql/types.py @@ -27,8 +27,10 @@ from wrapt import ObjectProxy from sanctumlabs_dbkit.sql.session import Session +from sanctumlabs_dbkit.sql.session.async_session import AsyncSession CommitCallback = Callable[[Session], None] +CommitCallbackAsync = Callable[[AsyncSession], None] _T = TypeVar("_T", bound=BaseModel) From 2d7e4fb6604ac7458acd78f7f9ea789741254985 Mon Sep 17 00:00:00 2001 From: BrianLusina <12752833+BrianLusina@users.noreply.github.com> Date: Sun, 12 Jul 2026 19:14:50 +0300 Subject: [PATCH 02/12] docs: update doc --- sanctumlabs_dbkit/sql/session/async_session.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/sanctumlabs_dbkit/sql/session/async_session.py b/sanctumlabs_dbkit/sql/session/async_session.py index c64240d..c64f9ce 100644 --- a/sanctumlabs_dbkit/sql/session/async_session.py +++ b/sanctumlabs_dbkit/sql/session/async_session.py @@ -34,7 +34,7 @@ def transaction(self, func: FuncT) -> FuncT: Example: ```python - from sanctumlabs_dbkit.sql import AsyncSessionLocal + from sanctumlabs_dbkit.sql.session import AsyncSessionLocal session = SessionLocal() @@ -65,14 +65,14 @@ async def async_transaction(func: FuncT) -> FuncT: Example: ```python - from sanctumlabs_dbkit.sql import AsyncSessionLocal - from sanctumlabs_dbkit.sql.async_session import transaction + from sanctumlabs_dbkit.sql.session import AsyncSessionLocal + from sanctumlabs_dbkit.sql.sesison import async_transaction class UserService(): def __init__(session: AsyncSession): self.session = session - @transaction + @async_transaction def create(payload) -> User: user = User(**payload) self.session.add(user) From 608fae86ec4b8ac3cd86612531769795efac0ea4 Mon Sep 17 00:00:00 2001 From: Lusina <12752833+BrianLusina@users.noreply.github.com> Date: Sun, 12 Jul 2026 19:42:01 +0300 Subject: [PATCH 03/12] Update sanctumlabs_dbkit/sql/repository/__init__.py Co-authored-by: coderabbitai[bot] <136622811+coderabbitai[bot]@users.noreply.github.com> Signed-off-by: Lusina <12752833+BrianLusina@users.noreply.github.com> --- sanctumlabs_dbkit/sql/repository/__init__.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/sanctumlabs_dbkit/sql/repository/__init__.py b/sanctumlabs_dbkit/sql/repository/__init__.py index 2c7f64b..d276c67 100644 --- a/sanctumlabs_dbkit/sql/repository/__init__.py +++ b/sanctumlabs_dbkit/sql/repository/__init__.py @@ -1,3 +1,5 @@ +"""Repository package exposing sync and async repository implementations.""" + from sanctumlabs_dbkit.sql.repository.repository import Repository from sanctumlabs_dbkit.sql.repository.async_repository import AsyncRepository From 4d39810b4f7621cac785d56c1a233171b083e665 Mon Sep 17 00:00:00 2001 From: Lusina <12752833+BrianLusina@users.noreply.github.com> Date: Sun, 12 Jul 2026 19:44:30 +0300 Subject: [PATCH 04/12] Update sanctumlabs_dbkit/sql/repository/async_repository.py Co-authored-by: coderabbitai[bot] <136622811+coderabbitai[bot]@users.noreply.github.com> Signed-off-by: Lusina <12752833+BrianLusina@users.noreply.github.com> --- sanctumlabs_dbkit/sql/repository/async_repository.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/sanctumlabs_dbkit/sql/repository/async_repository.py b/sanctumlabs_dbkit/sql/repository/async_repository.py index e868e35..5c5a1ae 100644 --- a/sanctumlabs_dbkit/sql/repository/async_repository.py +++ b/sanctumlabs_dbkit/sql/repository/async_repository.py @@ -45,7 +45,7 @@ def _supports_soft_deletion(model: Type[T]) -> TypeGuard[Type[AbstractBaseModel] """ return issubclass(model, AbstractBaseModel) - def create(self, refresh: bool = False, **kwargs: Any) -> T: + async def create(self, refresh: bool = False, **kwargs: Any) -> T: """Creates a new entity Args: @@ -58,10 +58,10 @@ def create(self, refresh: bool = False, **kwargs: Any) -> T: self.session.add(model_instance) if refresh: - self.session.flush() - self.session.refresh(model_instance) + await self.session.flush() + await self.session.refresh(model_instance) - return cast(T, model_instance) + return model_instance def query(self, include_deleted: bool = False) -> Select: """Returns a select query with the model including deleted records if the include_deleted is set to True""" From e2d538b2379b9a9450cdf016b39c9f0f60326c19 Mon Sep 17 00:00:00 2001 From: Lusina <12752833+BrianLusina@users.noreply.github.com> Date: Sun, 12 Jul 2026 19:45:07 +0300 Subject: [PATCH 05/12] Update sanctumlabs_dbkit/sql/session/__init__.py Co-authored-by: coderabbitai[bot] <136622811+coderabbitai[bot]@users.noreply.github.com> Signed-off-by: Lusina <12752833+BrianLusina@users.noreply.github.com> --- sanctumlabs_dbkit/sql/session/__init__.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/sanctumlabs_dbkit/sql/session/__init__.py b/sanctumlabs_dbkit/sql/session/__init__.py index 9bda53a..bd33b52 100644 --- a/sanctumlabs_dbkit/sql/session/__init__.py +++ b/sanctumlabs_dbkit/sql/session/__init__.py @@ -1,3 +1,5 @@ +"""Session package exposing sync and async session implementations.""" + from sanctumlabs_dbkit.sql.session.async_session import ( AsyncSession, AsyncSessionLocal, From 0d27b05a9fcd2837f0ba857a7b0b9faa11ccd866 Mon Sep 17 00:00:00 2001 From: Lusina <12752833+BrianLusina@users.noreply.github.com> Date: Sun, 12 Jul 2026 19:45:54 +0300 Subject: [PATCH 06/12] Update sanctumlabs_dbkit/sql/session/async_session.py Co-authored-by: coderabbitai[bot] <136622811+coderabbitai[bot]@users.noreply.github.com> Signed-off-by: Lusina <12752833+BrianLusina@users.noreply.github.com> --- sanctumlabs_dbkit/sql/session/async_session.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/sanctumlabs_dbkit/sql/session/async_session.py b/sanctumlabs_dbkit/sql/session/async_session.py index c64f9ce..50738ff 100644 --- a/sanctumlabs_dbkit/sql/session/async_session.py +++ b/sanctumlabs_dbkit/sql/session/async_session.py @@ -48,12 +48,12 @@ def create_user(payload) -> User: create_user({"first_name": "Bob"}) """ - @functools.wraps(func) + `@functools.wraps`(func) async def wrapper(*args: Any, **kwargs: Any) -> Any: async with self.begin(): - return func(*args, **kwargs) + return await func(*args, **kwargs) - return wrapper + return cast(FuncT, wrapper) async def async_transaction(func: FuncT) -> FuncT: From 7abdf50b835b8f40437453d21d9ceeccd8b9cdde Mon Sep 17 00:00:00 2001 From: Lusina <12752833+BrianLusina@users.noreply.github.com> Date: Sun, 12 Jul 2026 19:46:46 +0300 Subject: [PATCH 07/12] Update sanctumlabs_dbkit/sql/session/async_session.py Co-authored-by: coderabbitai[bot] <136622811+coderabbitai[bot]@users.noreply.github.com> Signed-off-by: Lusina <12752833+BrianLusina@users.noreply.github.com> --- sanctumlabs_dbkit/sql/session/async_session.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/sanctumlabs_dbkit/sql/session/async_session.py b/sanctumlabs_dbkit/sql/session/async_session.py index 50738ff..f33f82a 100644 --- a/sanctumlabs_dbkit/sql/session/async_session.py +++ b/sanctumlabs_dbkit/sql/session/async_session.py @@ -21,7 +21,7 @@ class AsyncSession(BaseAsyncSession): def begin(self, nested: bool = False) -> AsyncSessionTransaction: """Begins an async session transaction""" - if nested: + if nested or self.in_transaction(): return super().begin_nested() return super().begin() From fafad0502a54cb3e0292e13a09156174cb61c9fa Mon Sep 17 00:00:00 2001 From: Lusina <12752833+BrianLusina@users.noreply.github.com> Date: Sun, 12 Jul 2026 19:49:56 +0300 Subject: [PATCH 08/12] Update sanctumlabs_dbkit/sql/session/async_session.py Co-authored-by: coderabbitai[bot] <136622811+coderabbitai[bot]@users.noreply.github.com> Signed-off-by: Lusina <12752833+BrianLusina@users.noreply.github.com> --- .../sql/session/async_session.py | 37 +------------------ 1 file changed, 1 insertion(+), 36 deletions(-) diff --git a/sanctumlabs_dbkit/sql/session/async_session.py b/sanctumlabs_dbkit/sql/session/async_session.py index f33f82a..1767bf5 100644 --- a/sanctumlabs_dbkit/sql/session/async_session.py +++ b/sanctumlabs_dbkit/sql/session/async_session.py @@ -56,7 +56,7 @@ async def wrapper(*args: Any, **kwargs: Any) -> Any: return cast(FuncT, wrapper) -async def async_transaction(func: FuncT) -> FuncT: +def async_transaction(func: FuncT) -> FuncT: """ A decorator to wrap an instance method within a transaction. @@ -64,40 +64,5 @@ async def async_transaction(func: FuncT) -> FuncT: Example: - ```python - from sanctumlabs_dbkit.sql.session import AsyncSessionLocal - from sanctumlabs_dbkit.sql.sesison import async_transaction - - class UserService(): - def __init__(session: AsyncSession): - self.session = session - - @async_transaction - def create(payload) -> User: - user = User(**payload) - self.session.add(user) - - return user - - session = AsyncSessionLocal() - - user_service = UserService(session) - user_service.create({"first_name": "Bob"}) - """ - - @functools.wraps(func) - async def wrapper(self: Any, *args: Any, **kwargs: Any) -> Any: - if not self.session or not isinstance(self.session, AsyncSession): - # pylint: disable=broad-exception-raised - raise Exception( - "The @transaction decorator requires that an instance variable `session` be set to an instance of a " - "`Session`." - ) - - async with self.session.begin(): - return func(self, *args, **kwargs) - - return wrapper - AsyncSessionLocal = async_sessionmaker(class_=AsyncSession) From ee6e43c90c01bcc4d4cfb75ddfae2824354bc7e3 Mon Sep 17 00:00:00 2001 From: BrianLusina <12752833+BrianLusina@users.noreply.github.com> Date: Sun, 12 Jul 2026 20:13:48 +0300 Subject: [PATCH 09/12] chore: lint fixes --- .../sql/session/async_session.py | 20 +++++++++++++++++-- 1 file changed, 18 insertions(+), 2 deletions(-) diff --git a/sanctumlabs_dbkit/sql/session/async_session.py b/sanctumlabs_dbkit/sql/session/async_session.py index 1767bf5..5ffb6fe 100644 --- a/sanctumlabs_dbkit/sql/session/async_session.py +++ b/sanctumlabs_dbkit/sql/session/async_session.py @@ -3,7 +3,7 @@ """ import functools -from typing import Any +from typing import Any, cast from sqlalchemy.ext.asyncio import ( AsyncSession as BaseAsyncSession, @@ -48,7 +48,7 @@ def create_user(payload) -> User: create_user({"first_name": "Bob"}) """ - `@functools.wraps`(func) + @functools.wraps(func) async def wrapper(*args: Any, **kwargs: Any) -> Any: async with self.begin(): return await func(*args, **kwargs) @@ -63,6 +63,22 @@ def async_transaction(func: FuncT) -> FuncT: If we are already within a transaction, a nested transaction will be started. Example: + AsyncSessionLocal = async_sessionmaker(class_=AsyncSession) + """ + + @functools.wraps(func) + async def wrapper(self: Any, *args: Any, **kwargs: Any) -> Any: + if not self.session or not isinstance(self.session, AsyncSession): + # pylint: disable=broad-exception-raised + raise Exception( + "The @transaction decorator requires that an instance variable `session` be set to an instance of a " + "`Session`." + ) + + async with self.session.begin(): + return func(self, *args, **kwargs) + + return cast(FuncT, wrapper) AsyncSessionLocal = async_sessionmaker(class_=AsyncSession) From 9f1602cfd6ed69db50a89fffba26459f480b93bd Mon Sep 17 00:00:00 2001 From: BrianLusina <12752833+BrianLusina@users.noreply.github.com> Date: Sun, 12 Jul 2026 21:46:38 +0300 Subject: [PATCH 10/12] test: fix tests --- sanctumlabs_dbkit/sql/types.py | 26 ++++++++++++++++++++++++-- tests/sql/conftest.py | 21 ++++++++++++++++++++- 2 files changed, 44 insertions(+), 3 deletions(-) diff --git a/sanctumlabs_dbkit/sql/types.py b/sanctumlabs_dbkit/sql/types.py index c00f1f0..c7481f9 100644 --- a/sanctumlabs_dbkit/sql/types.py +++ b/sanctumlabs_dbkit/sql/types.py @@ -4,6 +4,7 @@ from __future__ import annotations +from decimal import Decimal from typing import ( Any, Callable, @@ -83,10 +84,31 @@ def load_dialect_impl(self, dialect: Dialect) -> TypeEngine[Any]: return dialect.type_descriptor(sa.JSON(none_as_null=True)) def _model_to_dict(self, value: _T) -> Dict[str, Any]: - return value.model_dump( - exclude_defaults=self.serialization_options.exclude_defaults + model_data = value.model_dump( + exclude_defaults=self.serialization_options.exclude_defaults, ) + return cast(Dict[str, Any], _normalise_json_compatible_value(model_data)) + + +def _normalise_json_compatible_value(value: Any) -> Any: + if isinstance(value, Decimal): + if value == value.to_integral_value(): + return int(value) + + return float(value) + + if isinstance(value, dict): + return {k: _normalise_json_compatible_value(v) for k, v in value.items()} + + if isinstance(value, list): + return [_normalise_json_compatible_value(v) for v in value] + + if isinstance(value, tuple): + return tuple(_normalise_json_compatible_value(v) for v in value) + + return value + # pylint: disable=abstract-method, too-many-ancestors class PydanticModel(ColumnUsesPydanticModelsMixin): diff --git a/tests/sql/conftest.py b/tests/sql/conftest.py index 470c4a9..b69bd7c 100644 --- a/tests/sql/conftest.py +++ b/tests/sql/conftest.py @@ -1,9 +1,12 @@ import os +import json from datetime import datetime, UTC +from decimal import Decimal from typing import Any, Generator from uuid import UUID import pytest +from pydantic import BaseModel as PydanticBaseModel from sqlalchemy import create_engine from tests.sql import Business, Card, User @@ -12,13 +15,29 @@ from sanctumlabs_dbkit.sql.session import Session +def pydantic_json_serializer(value: Any) -> str: + def default(obj: Any) -> Any: + if isinstance(obj, PydanticBaseModel): + return obj.model_dump(mode="json") + + if isinstance(obj, Decimal): + if obj == obj.to_integral_value(): + return int(obj) + + return float(obj) + + raise TypeError(f"Object of type {obj.__class__.__name__} is not JSON serializable") + + return json.dumps(value, default=default) + + @pytest.fixture def database_session() -> Generator[Session, Any, None]: database_url = os.environ.get( "DATABASE_URL", "postgresql://sanctumlabs:sanctumlabs@localhost:5432/dbkit-sql" ) - engine = create_engine(database_url) + engine = create_engine(database_url, json_serializer=pydantic_json_serializer) Base.metadata.drop_all(bind=engine) Base.metadata.create_all(bind=engine) From 3e94b3b49f5a5bf3b1bff4044ee06af2ad82e3d9 Mon Sep 17 00:00:00 2001 From: BrianLusina <12752833+BrianLusina@users.noreply.github.com> Date: Mon, 13 Jul 2026 10:10:17 +0300 Subject: [PATCH 11/12] feat(sql, repositories): read write repositories --- sanctumlabs_dbkit/sql/repository/__init__.py | 20 ++- .../sql/repository/async_repository.py | 142 ++++++++++++++++++ .../sql/repository/repository.py | 140 +++++++++++++++++ tests/sql/conftest.py | 4 +- 4 files changed, 303 insertions(+), 3 deletions(-) diff --git a/sanctumlabs_dbkit/sql/repository/__init__.py b/sanctumlabs_dbkit/sql/repository/__init__.py index d276c67..eae52f7 100644 --- a/sanctumlabs_dbkit/sql/repository/__init__.py +++ b/sanctumlabs_dbkit/sql/repository/__init__.py @@ -1,9 +1,25 @@ """Repository package exposing sync and async repository implementations.""" -from sanctumlabs_dbkit.sql.repository.repository import Repository -from sanctumlabs_dbkit.sql.repository.async_repository import AsyncRepository +from sanctumlabs_dbkit.sql.repository.repository import ( + Repository, + BaseRepository, + WriteRepository, + ReadRepository, +) +from sanctumlabs_dbkit.sql.repository.async_repository import ( + AsyncRepository, + AsyncBaseRepository, + AsyncWriteRepository, + AsyncReadRepository, +) __all__ = [ "Repository", + "BaseRepository", + "WriteRepository", + "ReadRepository", "AsyncRepository", + "AsyncBaseRepository", + "AsyncWriteRepository", + "AsyncReadRepository", ] diff --git a/sanctumlabs_dbkit/sql/repository/async_repository.py b/sanctumlabs_dbkit/sql/repository/async_repository.py index 5c5a1ae..f456a04 100644 --- a/sanctumlabs_dbkit/sql/repository/async_repository.py +++ b/sanctumlabs_dbkit/sql/repository/async_repository.py @@ -128,3 +128,145 @@ async def list( scalars = await self.session.scalars(statement) return scalars.all() + + +class AsyncBaseRepository(Generic[T]): + """ + A base class for implementing an async Repository or DAO. + + ```python + job_dao = AsyncRepository(model=Job, session=async_session) + + job = job_dao.find("123") + """ + + def __init__(self, model: Type[T], session: AsyncSession) -> None: + """Creates an instance of the Repository""" + self.model = model + self.session = session + + @staticmethod + def _supports_soft_deletion(model: Type[T]) -> TypeGuard[Type[AbstractBaseModel]]: + """ + Indicates if the provided model supports soft deletion (has a 'deleted_at' column). This function + takes in an argument due to mypy typeguarding requirements, and is thus static. + """ + return issubclass(model, AbstractBaseModel) + + def query(self, include_deleted: bool = False) -> Select: + """Returns a select query with the model including deleted records if the include_deleted is set to True""" + selectable = select(self.model) + + if not include_deleted and self._supports_soft_deletion(self.model): + selectable = selectable.where( + self.model.deleted_at == self.model.not_deleted_value() + ) + + return selectable + + +class AsyncWriteRepository(AsyncBaseRepository, Generic[T]): + """ + A base class for implementing an async write Repository or DAO for performing write operations. + + ```python + job_dao = AsyncWriteRepository(model=Job, session=async_session) + + job = job_dao.find("123") + """ + + def __init__(self, model: Type[T], session: AsyncSession) -> None: + """Creates an instance of the Repository""" + super().__init__(model, session) + self.model = model + self.session = session + + async def create(self, refresh: bool = False, **kwargs: Any) -> T: + """Creates a new entity + + Args: + refresh (bool, optional): whether to refresh the model with the data in the return. Defaults to False. + + Returns: + T: The created model instance + """ + model_instance = self.model(**kwargs) + self.session.add(model_instance) + + if refresh: + await self.session.flush() + await self.session.refresh(model_instance) + + return model_instance + + async def delete(self, pk: Any) -> None: + """Deletes a given record with the given primary key""" + if not self._supports_soft_deletion(self.model): + raise UnsupportedModelOperationError( + f"The model {self.model.__name__} {pk} does not support soft deletion." + ) + + # Cast here as mypy type narrowing doesn't infer the type of entity + # correctly + entity = cast(AbstractBaseModel, await self.find(pk)) + + if entity: + entity.deleted_at = datetime.now(UTC) + + +class AsyncReadRepository(AsyncBaseRepository, Generic[T]): + """ + A base class for implementing an async read Repository or DAO for performing read operations. + + ```python + job_dao = AsyncReadRepository(model=Job, session=async_session) + + job = job_dao.find("123") + """ + + def __init__(self, model: Type[T], session: AsyncSession) -> None: + """Creates an instance of the Repository""" + super().__init__(model, session) + self.model = model + self.session = session + + async def find(self, pk: Any, include_deleted: bool = False) -> Optional[T]: + """Retrieve a given model given its primary key""" + pk_column = cast(ColumnElement, getattr(self.model, self.model.pk)) + + statement = self.query(include_deleted).where(pk_column == pk).limit(1) + scalars = await self.session.scalars(statement) + + return scalars.first() + + async def find_or_raise(self, pk: Any, include_deleted: bool = False) -> T: + """Finds the given entity or raises an exception if the entity can not be found""" + entity = await self.find(pk, include_deleted) + + if not entity: + raise ModelNotFoundError( + f"The model {self.model.__name__} {pk} does not exist" + ) + + return entity + + async def all(self, include_deleted: bool = False) -> Sequence[T]: + """Retrieves all records for the given model""" + statement = self.query(include_deleted) + scalars = await self.session.scalars(statement) + + return scalars.all() + + async def list( + self, limit: int = 20, offset: int = 0, include_deleted: bool = False + ) -> Sequence[T]: + """Returns a list of records for the given database record""" + statement = ( + self.query(include_deleted) + .order_by(self.model.created_at.desc()) + .limit(limit) + .offset(offset) + ) + scalars = await self.session.scalars(statement) + + return scalars.all() diff --git a/sanctumlabs_dbkit/sql/repository/repository.py b/sanctumlabs_dbkit/sql/repository/repository.py index 9613e1af..b79f669 100644 --- a/sanctumlabs_dbkit/sql/repository/repository.py +++ b/sanctumlabs_dbkit/sql/repository/repository.py @@ -2,6 +2,7 @@ Contains a generic base repository or DAO for access patterns to data for a given database model """ +from abc import ABCMeta from datetime import datetime, UTC from typing import ( Generic, @@ -125,3 +126,142 @@ def list( ) return self.session.scalars(statement).all() + + +class BaseRepository(Generic[T], metaclass=ABCMeta): + """ + A base class for implementing a Repository or DAO. This + + ```python + job_dao = Repository(model=Job, session=session) + + job = job_dao.find("123") + """ + + def __init__(self, model: Type[T], session: Session) -> None: + """Creates an instance of the Repository""" + self.model = model + self.session = session + + @staticmethod + def _supports_soft_deletion(model: Type[T]) -> TypeGuard[Type[AbstractBaseModel]]: + """ + Indicates if the provided model supports soft deletion (has a 'deleted_at' column). This function + takes in an argument due to mypy type guarding requirements, and is thus static. + """ + return issubclass(model, AbstractBaseModel) + + def query(self, include_deleted: bool = False) -> Select: + """Returns a select query with the model including deleted records if the include_deleted is set to True""" + selectable = select(self.model) + + if not include_deleted and self._supports_soft_deletion(self.model): + selectable = selectable.where( + self.model.deleted_at == self.model.not_deleted_value() + ) + + return selectable + + +class WriteRepository(BaseRepository, Generic[T]): + """ + A base class for implementing write operations on a Repository or DAO. + + ```python + job_dao = WriteRepository(model=Job, session=session) + + job = job_dao.create("123") + """ + + def __init__(self, model: Type[T], session: Session) -> None: + """Creates an instance of the WriteRepository""" + super().__init__(model, session) + self.model = model + self.session = session + + def create(self, refresh: bool = False, **kwargs: Any) -> T: + """Creates a new entity + + Args: + refresh (bool, optional): whether to refresh the model with the data in the return. Defaults to False. + + Returns: + T: The created model instance + """ + model_instance = self.model(**kwargs) + self.session.add(model_instance) + + if refresh: + self.session.flush() + self.session.refresh(model_instance) + + return cast(T, model_instance) + + def delete(self, pk: Any) -> None: + """Deletes a given record with the given primary key""" + if not self._supports_soft_deletion(self.model): + raise UnsupportedModelOperationError( + f"The model {self.model.__name__} {pk} does not support soft deletion." + ) + + # Cast here as mypy type narrowing doesn't infer the type of entity + # correctly + entity = cast(AbstractBaseModel, self.find(pk)) + + if entity: + entity.deleted_at = datetime.now(UTC) + + +class ReadRepository(BaseRepository, Generic[T]): + """ + A base class for implementing read operations on a Repository or DAO. + + ```python + job_dao = WriteRepository(model=Job, session=session) + + job = job_dao.find("123") + """ + + def __init__(self, model: Type[T], session: Session) -> None: + """Creates an instance of the ReadRepository""" + super().__init__(model, session) + self.model = model + self.session = session + + def find(self, pk: Any, include_deleted: bool = False) -> Optional[T]: + """Retrieve a given model given its primary key""" + pk_column = cast(ColumnElement, getattr(self.model, self.model.pk)) + + statement = self.query(include_deleted).where(pk_column == pk).limit(1) + + return self.session.scalars(statement).first() + + def find_or_raise(self, pk: Any, include_deleted: bool = False) -> T: + """Finds the given entity or raises an exception if the entity can not be found""" + entity = self.find(pk, include_deleted) + + if not entity: + raise ModelNotFoundError( + f"The model {self.model.__name__} {pk} does not exist" + ) + + return entity + + def all(self, include_deleted: bool = False) -> Sequence[T]: + """Retrieves all records for the given model""" + statement = self.query(include_deleted) + + return self.session.scalars(statement).all() + + def list( + self, limit: int = 20, offset: int = 0, include_deleted: bool = False + ) -> Sequence[T]: + """Returns a list of records for the given database record""" + statement = ( + self.query(include_deleted) + .order_by(self.model.created_at.desc()) + .limit(limit) + .offset(offset) + ) + + return self.session.scalars(statement).all() diff --git a/tests/sql/conftest.py b/tests/sql/conftest.py index b69bd7c..5e6eda7 100644 --- a/tests/sql/conftest.py +++ b/tests/sql/conftest.py @@ -26,7 +26,9 @@ def default(obj: Any) -> Any: return float(obj) - raise TypeError(f"Object of type {obj.__class__.__name__} is not JSON serializable") + raise TypeError( + f"Object of type {obj.__class__.__name__} is not JSON serializable" + ) return json.dumps(value, default=default) From 6776b363caff5f621a2c7ecb4d28aec9a54d7102 Mon Sep 17 00:00:00 2001 From: BrianLusina <12752833+BrianLusina@users.noreply.github.com> Date: Tue, 14 Jul 2026 10:44:32 +0300 Subject: [PATCH 12/12] chore: add ty type checker --- Makefile | 4 +++ poetry.lock | 32 +++++++++++++++++-- pyproject.toml | 1 + sanctumlabs_dbkit/sql/mixins.py | 4 +-- .../sql/repository/async_repository.py | 22 ++++++------- .../sql/repository/repository.py | 18 +++++------ sanctumlabs_dbkit/sql/types.py | 2 +- 7 files changed, 58 insertions(+), 25 deletions(-) diff --git a/Makefile b/Makefile index 617e93d..168d9b4 100644 --- a/Makefile +++ b/Makefile @@ -66,6 +66,10 @@ lint-pylint: ## Runs linting with pylint lint-ruff: ## Runs linting with ruff poetry run ruff check sanctumlabs_dbkit +.PHONY: lint-ty +lint-ty: ## Runs type checking with ty + poetry run ty check + .PHONY: lint lint: format-black lint-flake8 lint-mypy lint-pylint diff --git a/poetry.lock b/poetry.lock index 7a17518..7588d37 100644 --- a/poetry.lock +++ b/poetry.lock @@ -378,7 +378,7 @@ description = "Cross-platform colored terminal text." optional = false python-versions = "!=3.0.*,!=3.1.*,!=3.2.*,!=3.3.*,!=3.4.*,!=3.5.*,!=3.6.*,>=2.7" groups = ["dev"] -markers = "platform_system == \"Windows\" or sys_platform == \"win32\"" +markers = "sys_platform == \"win32\" or platform_system == \"Windows\"" files = [ {file = "colorama-0.4.6-py2.py3-none-any.whl", hash = "sha256:4f1d9991f5acc0ca119f9d443620b77f9d6b33703e51011c16baf57afb285fc6"}, {file = "colorama-0.4.6.tar.gz", hash = "sha256:08695f5cb7ed6e0531a20572697297273c47b8cae5a63ffc6d6ed5c201be6e44"}, @@ -2103,6 +2103,34 @@ urllib3 = ">=1.26.0" [package.extras] keyring = ["keyring (>=21.2.0)"] +[[package]] +name = "ty" +version = "0.0.59" +description = "An extremely fast Python type checker, written in Rust." +optional = false +python-versions = ">=3.8" +groups = ["main"] +files = [ + {file = "ty-0.0.59-py3-none-linux_armv6l.whl", hash = "sha256:f8fb08a767ef8f11ea3c537b9d77860726cc2bc39e6f77ad13c02d5b289f20a7"}, + {file = "ty-0.0.59-py3-none-macosx_10_12_x86_64.whl", hash = "sha256:c7f4d5630836c8a0ba13dd4ac7bdae080a7d6ebe965b817ff642dc961bcf2a53"}, + {file = "ty-0.0.59-py3-none-macosx_11_0_arm64.whl", hash = "sha256:872f6fb02c6db5553c4d5fb283b3d50f0985fb9a29a910e4fda4793a775c1926"}, + {file = "ty-0.0.59-py3-none-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:2af8eefbfe806337770eec12c0c819c5f1b8f5b85f8369cb1cc9fa25234a2208"}, + {file = "ty-0.0.59-py3-none-manylinux_2_17_armv7l.manylinux2014_armv7l.whl", hash = "sha256:0acf8b76a1c9a7ddef460b42475f6c76193164426ab080783af1c3175b4b999b"}, + {file = "ty-0.0.59-py3-none-manylinux_2_17_i686.manylinux2014_i686.whl", hash = "sha256:043c2e00eb1d7475f928af7dedd71f69b64e69bfca55e36f4c968479e1373fc4"}, + {file = "ty-0.0.59-py3-none-manylinux_2_17_ppc64le.manylinux2014_ppc64le.whl", hash = "sha256:f0d688d857441df57f48fca66c029d85cf737c510e7be1d01144cdad1e58d968"}, + {file = "ty-0.0.59-py3-none-manylinux_2_17_s390x.manylinux2014_s390x.whl", hash = "sha256:a96c9f88394a3b42c737e2125b2330543f0d90a43b49761f377d96f8c3ee0d62"}, + {file = "ty-0.0.59-py3-none-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:f08dbcb268edcafcb152e59475b5b495ce28d0b340a395c09943557678f4d5a6"}, + {file = "ty-0.0.59-py3-none-manylinux_2_31_riscv64.whl", hash = "sha256:8812764b9a40fdc98df1272826e73a298ef56b06681135e643bcf90aad1896f7"}, + {file = "ty-0.0.59-py3-none-musllinux_1_2_aarch64.whl", hash = "sha256:fd53b8581641d8dad7bfac6d5ea589e91a883d6837e0b9a286fdae30722b7c69"}, + {file = "ty-0.0.59-py3-none-musllinux_1_2_armv7l.whl", hash = "sha256:86da5872124a41877d95058bc17d33ddcff034b587eb5f1e2917ab88ba227dac"}, + {file = "ty-0.0.59-py3-none-musllinux_1_2_i686.whl", hash = "sha256:6a233eef5f2fd4d894881e4a0aec83c9f172bfae1d787d6596ee1939fcc7723e"}, + {file = "ty-0.0.59-py3-none-musllinux_1_2_x86_64.whl", hash = "sha256:7ff678c18b5f1e3128b75a35e50dee7908dea55155baa31cd790619d5014cbf5"}, + {file = "ty-0.0.59-py3-none-win32.whl", hash = "sha256:cf8abb4b8095c5fe39102b8127f5886db308c8d4600909ddbc905512ce9c8163"}, + {file = "ty-0.0.59-py3-none-win_amd64.whl", hash = "sha256:1dde20a82243d24407869e5a608c2f15efddd5cefc662aef461a5af84bfb3f8b"}, + {file = "ty-0.0.59-py3-none-win_arm64.whl", hash = "sha256:987043ee9e021f49493d9135891ac69c1affeee0d4ad4480c5fa4d9c975fc91b"}, + {file = "ty-0.0.59.tar.gz", hash = "sha256:53e53ffeed78ad59cd237fa8ea1316d2b94e13efdea9a945698acab549e005aa"}, +] + [[package]] name = "types-sqlalchemy-utils" version = "1.1.0" @@ -2301,4 +2329,4 @@ dev = ["pytest", "setuptools"] [metadata] lock-version = "2.1" python-versions = "^3.12.0" -content-hash = "f8041ada28abc94c6c5a8056c56853dc5c59914b15a4f11c4badef6a4d898f98" +content-hash = "9e4c1c7fe65e1efc6de8ed0c65a9718655a71ad3ee747bdbae7bcb0623c01b3a" diff --git a/pyproject.toml b/pyproject.toml index be04363..fc49924 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -32,6 +32,7 @@ autoflake = "^2.3.1" types-sqlalchemy-utils = "^1.0.1" psycopg2-binary = "^2.9.9" ruff = ">=0.5.0,<0.16.0" +ty = "^0.0.59" [tool.mypy] disallow_incomplete_defs = true diff --git a/sanctumlabs_dbkit/sql/mixins.py b/sanctumlabs_dbkit/sql/mixins.py index 9bf7862..337d158 100644 --- a/sanctumlabs_dbkit/sql/mixins.py +++ b/sanctumlabs_dbkit/sql/mixins.py @@ -81,10 +81,10 @@ class TableNameMixin: Mixin that creates the table names of a database """ - @declared_attr # type: ignore[arg-type] + @declared_attr.directive # type: ignore[arg-type] def __tablename__(self) -> str: """Table names are snake case plural, for example shipping_records""" - return inflection.pluralize(inflection.underscore(self.__name__)) # type: ignore[attr-defined] + return inflection.pluralize(inflection.underscore(self.__name__)) # type: ignore[attr-defined] # ty:ignore[unresolved-attribute] class BigIntIdentityMixin: diff --git a/sanctumlabs_dbkit/sql/repository/async_repository.py b/sanctumlabs_dbkit/sql/repository/async_repository.py index f456a04..35e1b29 100644 --- a/sanctumlabs_dbkit/sql/repository/async_repository.py +++ b/sanctumlabs_dbkit/sql/repository/async_repository.py @@ -164,15 +164,24 @@ def query(self, include_deleted: bool = False) -> Select: return selectable + async def find(self, pk: Any, include_deleted: bool = False) -> Optional[T]: + """Retrieve a given model given its primary key""" + pk_column = cast(ColumnElement, getattr(self.model, self.model.pk)) + + statement = self.query(include_deleted).where(pk_column == pk).limit(1) + scalars = await self.session.scalars(statement) + + return scalars.first() + class AsyncWriteRepository(AsyncBaseRepository, Generic[T]): """ A base class for implementing an async write Repository or DAO for performing write operations. ```python - job_dao = AsyncWriteRepository(model=Job, session=async_session) + job_repo = AsyncWriteRepository(model=Job, session=async_session) - job = job_dao.find("123") + job = job_repo.create(name="123") """ def __init__(self, model: Type[T], session: AsyncSession) -> None: @@ -230,15 +239,6 @@ def __init__(self, model: Type[T], session: AsyncSession) -> None: self.model = model self.session = session - async def find(self, pk: Any, include_deleted: bool = False) -> Optional[T]: - """Retrieve a given model given its primary key""" - pk_column = cast(ColumnElement, getattr(self.model, self.model.pk)) - - statement = self.query(include_deleted).where(pk_column == pk).limit(1) - scalars = await self.session.scalars(statement) - - return scalars.first() - async def find_or_raise(self, pk: Any, include_deleted: bool = False) -> T: """Finds the given entity or raises an exception if the entity can not be found""" entity = await self.find(pk, include_deleted) diff --git a/sanctumlabs_dbkit/sql/repository/repository.py b/sanctumlabs_dbkit/sql/repository/repository.py index b79f669..838c6e0 100644 --- a/sanctumlabs_dbkit/sql/repository/repository.py +++ b/sanctumlabs_dbkit/sql/repository/repository.py @@ -62,7 +62,7 @@ def create(self, refresh: bool = False, **kwargs: Any) -> T: self.session.flush() self.session.refresh(model_instance) - return cast(T, model_instance) + return model_instance def query(self, include_deleted: bool = False) -> Select: """Returns a select query with the model including deleted records if the include_deleted is set to True""" @@ -162,6 +162,14 @@ def query(self, include_deleted: bool = False) -> Select: return selectable + def find(self, pk: Any, include_deleted: bool = False) -> Optional[T]: + """Retrieve a given model given its primary key""" + pk_column = cast(ColumnElement, getattr(self.model, self.model.pk)) + + statement = self.query(include_deleted).where(pk_column == pk).limit(1) + + return self.session.scalars(statement).first() + class WriteRepository(BaseRepository, Generic[T]): """ @@ -228,14 +236,6 @@ def __init__(self, model: Type[T], session: Session) -> None: self.model = model self.session = session - def find(self, pk: Any, include_deleted: bool = False) -> Optional[T]: - """Retrieve a given model given its primary key""" - pk_column = cast(ColumnElement, getattr(self.model, self.model.pk)) - - statement = self.query(include_deleted).where(pk_column == pk).limit(1) - - return self.session.scalars(statement).first() - def find_or_raise(self, pk: Any, include_deleted: bool = False) -> T: """Finds the given entity or raises an exception if the entity can not be found""" entity = self.find(pk, include_deleted) diff --git a/sanctumlabs_dbkit/sql/types.py b/sanctumlabs_dbkit/sql/types.py index c7481f9..03119d5 100644 --- a/sanctumlabs_dbkit/sql/types.py +++ b/sanctumlabs_dbkit/sql/types.py @@ -80,7 +80,7 @@ def __init__( def load_dialect_impl(self, dialect: Dialect) -> TypeEngine[Any]: # Use JSONB for PostgreSQL and JSON for other databases. if dialect.name == "postgresql": - return dialect.type_descriptor(JSONB(none_as_null=True)) # type: ignore + return dialect.type_descriptor(JSONB(none_as_null=True)) return dialect.type_descriptor(sa.JSON(none_as_null=True)) def _model_to_dict(self, value: _T) -> Dict[str, Any]: