diff --git a/.github/workflows/tests.yml b/.github/workflows/tests.yml index 437ccb4..8065f7c 100644 --- a/.github/workflows/tests.yml +++ b/.github/workflows/tests.yml @@ -50,7 +50,7 @@ jobs: uses: codecov/codecov-action@v7 with: env_vars: PYTHON - report-type: test_results + report_type: test_results token: ${{ secrets.CODECOV_TOKEN }} flags: unittests,python-${{ matrix.python-version }} @@ -58,7 +58,7 @@ jobs: uses: codecov/codecov-action@v7 with: env_vars: PYTHON - report-type: coverage + report_type: coverage token: ${{ secrets.CODECOV_TOKEN }} flags: unittests,python-${{ matrix.python-version }} @@ -69,4 +69,4 @@ jobs: file: coverage.xml language: Python label: code-coverage/pytest - + fail-on-error: false diff --git a/.gitignore b/.gitignore index 92878e8..683075f 100644 --- a/.gitignore +++ b/.gitignore @@ -23,6 +23,7 @@ venv .venv .pdm-build build +typings # vim temporary files *~ diff --git a/pyrannic/__init__.py b/pyrannic/__init__.py index 50d909a..44eeb8b 100644 --- a/pyrannic/__init__.py +++ b/pyrannic/__init__.py @@ -1,12 +1,14 @@ -__version__ = "0.5.6" +__version__ = "0.5.7" from .application import Application as Application from .bootstrap.service_provider import ServiceProvider as ServiceProvider from .config.configuration import Configuration as Configuration from .container.param_functions import Resolves as Resolves -from .ioc import Resolve as Resolve from .database.migration import Migration as Migration from .database.provider import DatabaseServiceProvider as DatabaseServiceProvider +from .http.exceptions.resource_not_found import ( + ResourceNotFoundException as ResourceNotFoundException, +) from .http.providers import ( ExceptionHandlersServiceProvider as ExceptionHandlersServiceProvider, ) @@ -14,10 +16,6 @@ from .http.providers import RoutersServiceProvider as RoutersServiceProvider from .http.resources.collection import ResourceCollection as ResourceCollection from .http.resources.resource import Resource as Resource +from .ioc import Resolve as Resolve from .pagination.meta import PaginationMeta as PaginationMeta from .pagination.paginator import Paginator as Paginator -from .support.facades.config import Config as Config - -from .http.exceptions.resource_not_found import ( - ResourceNotFoundException as ResourceNotFoundException, -) diff --git a/pyrannic/config/provider.py b/pyrannic/config/provider.py index 2838eda..ded6986 100644 --- a/pyrannic/config/provider.py +++ b/pyrannic/config/provider.py @@ -3,7 +3,7 @@ from pyrannic.bootstrap.instance_service_provider import InstanceServiceProvider from pyrannic.config.repository import ConfigRepository from pyrannic.contracts.config.configuration import ConfigurationInterface -from pyrannic.contracts.config.respository import ConfigRepositoryInterface +from pyrannic.contracts.config.repository import ConfigRepositoryInterface from pyrannic.support.path import get_module_paths from pyrannic.support.reflection import get_classes diff --git a/pyrannic/config/repository.py b/pyrannic/config/repository.py index 9e73fc3..14353a5 100644 --- a/pyrannic/config/repository.py +++ b/pyrannic/config/repository.py @@ -2,7 +2,7 @@ from annotated_types import T -from pyrannic.contracts.config.respository import ConfigRepositoryInterface +from pyrannic.contracts.config.repository import ConfigRepositoryInterface from pyrannic.support.collections.dot_dict import get, has, set diff --git a/pyrannic/container/container.py b/pyrannic/container/container.py index 9e6f490..e4cdba5 100644 --- a/pyrannic/container/container.py +++ b/pyrannic/container/container.py @@ -1,11 +1,17 @@ +import inspect from collections.abc import Callable from contextlib import AsyncExitStack -import inspect from types import FunctionType from typing import Any, Awaitable, TypeVar, cast, get_args, get_origin from fastapi import Request from fastapi.concurrency import run_in_threadpool +from fastapi.dependencies.models import Dependant +from fastapi.dependencies.utils import ( + SolvedDependency, + get_dependant, + solve_dependencies, +) from fastapi.exceptions import RequestValidationError from fastapi.types import DependencyCacheKey @@ -17,13 +23,6 @@ ) from pyrannic.support.reflection import is_interface -from fastapi.dependencies.utils import ( - get_dependant, - solve_dependencies, - SolvedDependency, -) -from fastapi.dependencies.models import Dependant - T = TypeVar("T") @@ -200,10 +199,21 @@ def is_bound(self, abstract: str | type) -> bool: or self.is_alias(abstract) ) + async def make( + self, + abstract: str | type[T], + *args: Any, + request: Request | None = None, + **kwargs: Any, + ) -> T: + return await self.resolve(abstract, *args, request=request, **kwargs) + async def resolve( self, abstract: str | type[T], + *args: Any, request: Request | None = None, + **kwargs: Any, ) -> T: binding_key = self.get_alias(abstract) concrete = self._get_contextual_concrete(binding_key) @@ -218,7 +228,7 @@ async def resolve( if not concrete: concrete = self._get_concrete(abstract) - instance = concrete(self._app, request) + instance = concrete(self._app, request, *args, **kwargs) if inspect.isawaitable(instance): instance = await instance @@ -291,8 +301,13 @@ def _get_concrete(self, abstract: str | type[T]) -> Callable[..., Any]: return concrete - async def call(self, callback: type[T] | Callable[..., Any]) -> T: - return await self._resolve(callback, self._app) + async def call( + self, + callback: type[T] | Callable[..., Any], + *args: Any, + **kwargs: Any, + ) -> T: + return await self._resolve(callback, self._app, *args, **kwargs) def resolved(self, abstract: str | type[T]) -> bool: abstract = self.get_alias(abstract) @@ -350,13 +365,25 @@ def _get_closure( self, concrete: type[T], ) -> Callable[[ApplicationInterface, Request], Awaitable[T]]: - async def closure(app: ApplicationInterface, request: Request) -> T: + async def closure( + app: ApplicationInterface, + request: Request, + *args: Any, + **kwargs: Any, + ) -> T: origin = get_origin(concrete) if origin is not None and inspect.isclass(origin): - return await self._resolve_generic(concrete, origin, app, request) + return await self._resolve_generic( + concrete, + origin, + app, + request, + *args, + **kwargs, + ) - return await self._resolve(concrete, app, request) + return await self._resolve(concrete, app, request, *args, **kwargs) return closure @@ -380,12 +407,14 @@ async def _resolve_generic( origin: type[T], app: ApplicationInterface, request: Request | None = None, + *args: Any, + **kwargs: Any, ) -> T: # Add the __orig_class__ attribute to the origin class so that we can retrieve the generic type later. # For example, if we have a generic class Repository[Model], we can retrieve the Model type later by accessing the __orig_class__ attribute. # This is necessary because FastAPI's dependency injection system does not support generic types out of the box. setattr(origin, "__orig_class__", generic) - instance = await self._resolve(origin, app, request) + instance = await self._resolve(origin, app, request, *args, **kwargs) try: setattr(instance, "__orig_class__", generic) @@ -399,25 +428,31 @@ async def _resolve( callback: type[T] | Callable[..., Any], app: ApplicationInterface, request: Request | None = None, + *args: Any, **kwargs: Any, ) -> T: dependant = get_dependant(path="/", call=callback, scope="function") if not request: async with AsyncExitStack() as context_manager: - dependencies = await self._solve_dependencies( - self._fallback_request(app, context_manager), - dependant, - ) + request = self._fallback_request(app, context_manager) + dependencies = await self._solve_dependencies(request, dependant) return await self._resolve_dependant( dependant, dependencies, + *args, **kwargs, ) else: dependencies = await self._solve_dependencies(request, dependant) - return await self._resolve_dependant(dependant, dependencies, **kwargs) + + return await self._resolve_dependant( + dependant, + dependencies, + *args, + **kwargs, + ) # https://stackoverflow.com/a/78279023 # https://github.com/fastapi/fastapi/discussions/7720 @@ -425,37 +460,53 @@ async def _resolve_dependant( self, dependant: Dependant, dependencies: SolvedDependency, + *args: Any, **kwargs: Any, ) -> Any: assert dependant.call # For types - errors = self._validate_errors(dependencies.errors) + errors = self._validate_errors(dependencies.errors, *args, **kwargs) if bool(errors): raise RequestValidationError(errors) if inspect.iscoroutinefunction(dependant.call): - result = await dependant.call(**dependencies.values, **kwargs) + result = await dependant.call(*args, **dependencies.values, **kwargs) else: result = await run_in_threadpool( dependant.call, + *args, **dependencies.values, **kwargs, ) return result - def _validate_errors(self, errors: list[Any]) -> list[Any]: + def _validate_errors( + self, + errors: list[Any], + *args: Any, + **kwargs: Any, + ) -> list[Any]: if not bool(errors): return [] + error_counter = 0 valid_errors: list[Any] = [] + length = len(args) for error in errors: loc = error.get("loc", []) - - if "kwargs" not in loc and "args" not in loc: + attr = loc[-1] if loc else None + + if ( + attr != "kwargs" + and attr != "args" + and attr not in kwargs + and length <= error_counter + ): valid_errors.append(error) + error_counter += 1 return valid_errors diff --git a/pyrannic/container/params.py b/pyrannic/container/params.py index 6b8188c..27f1e01 100644 --- a/pyrannic/container/params.py +++ b/pyrannic/container/params.py @@ -4,8 +4,8 @@ from fastapi.params import Depends -from pyrannic.container.decorators.singleton import singleton from pyrannic.container.decorators.scoped import scoped +from pyrannic.container.decorators.singleton import singleton from pyrannic.contracts.http.request import RequestInterface @@ -28,7 +28,9 @@ def __init__( @classmethod def wrap_dependency(cls, abstract: str | type) -> Callable[..., Any]: async def dependency(request: RequestInterface) -> Any: - return cast(Any, await request.app.container.resolve(abstract, request)) + return cast( + Any, await request.app.container.resolve(abstract, request=request) + ) return dependency diff --git a/pyrannic/contracts/__init__.py b/pyrannic/contracts/__init__.py index 249cc2a..20fc69c 100644 --- a/pyrannic/contracts/__init__.py +++ b/pyrannic/contracts/__init__.py @@ -1,10 +1,13 @@ +from .application import ApplicationInterface as ApplicationInterface from .config.configuration import ConfigurationInterface as ConfigurationInterface -from .config.respository import ConfigRepositoryInterface as ConfigRepositoryInterface +from .config.repository import ConfigRepositoryInterface as ConfigRepositoryInterface +from .container.container import ContainerInterface as ContainerInterface from .database.connector import ConnectorInterface as ConnectorInterface from .database.manager import DatabaseManagerInterface as DatabaseManagerInterface from .http.resources.collection import ( ResourceCollectionInterface as ResourceCollectionInterface, ) from .http.resources.resource import ResourceInterface as ResourceInterface +from .orm.async_repository import AsyncRepositoryInterface as AsyncRepositoryInterface +from .orm.repository import RepositoryInterface as RepositoryInterface from .pagination.paginator import PaginatorInterface as PaginatorInterface -from .container.container import ContainerInterface as ContainerInterface diff --git a/pyrannic/contracts/config/respository.py b/pyrannic/contracts/config/repository.py similarity index 100% rename from pyrannic/contracts/config/respository.py rename to pyrannic/contracts/config/repository.py diff --git a/pyrannic/contracts/container/__init__.py b/pyrannic/contracts/container/__init__.py index e69de29..9332c88 100644 --- a/pyrannic/contracts/container/__init__.py +++ b/pyrannic/contracts/container/__init__.py @@ -0,0 +1,4 @@ +from .container import ContainerInterface as ContainerInterface +from .contextual_binding_builder import ( + ContextualBindingBuilderInterface as ContextualBindingBuilderInterface, +) diff --git a/pyrannic/contracts/container/container.py b/pyrannic/contracts/container/container.py index 82252ca..ae5ba9e 100644 --- a/pyrannic/contracts/container/container.py +++ b/pyrannic/contracts/container/container.py @@ -19,7 +19,6 @@ def bind( shared: bool = False, ) -> None: """Register a binding with the container.""" - pass @abstractmethod def bind_if( @@ -29,7 +28,6 @@ def bind_if( shared: bool = False, ) -> None: """Register a binding if it hasn't already been registered.""" - pass @abstractmethod def scoped( @@ -38,7 +36,6 @@ def scoped( concrete: type[Any] | Callable[..., Any], ) -> None: """Register a scoped binding in the container.""" - pass @abstractmethod def scoped_if( @@ -47,7 +44,6 @@ def scoped_if( concrete: type[Any] | Callable[..., Any], ) -> None: """Register a scoped binding if it hasn't already been registered.""" - pass @abstractmethod def singleton( @@ -56,7 +52,6 @@ def singleton( concrete: type[Any] | Callable[..., Any], ) -> None: """Register a shared binding in the container.""" - pass @abstractmethod def singleton_if( @@ -65,12 +60,10 @@ def singleton_if( concrete: type[Any] | Callable[..., Any], ) -> None: """Register a shared binding if it hasn't already been registered.""" - pass @abstractmethod def instance(self, abstract: str | type[T], instance: T | None = None) -> T: """Register an existing instance as shared in the container or retrieve it from the container.""" - pass @abstractmethod def add_contextual_binding( @@ -80,48 +73,51 @@ def add_contextual_binding( implementation: type | Callable[..., Any], ) -> None: """Add a contextual binding to the container.""" - pass @abstractmethod def when(self, concrete: type | list[type]) -> ContextualBindingBuilderInterface: """Define a contextual binding.""" - pass @abstractmethod def is_bound(self, abstract: str | type) -> bool: """Determine if the given abstract type has been bound.""" - pass + + @abstractmethod + async def make( + self, + abstract: str | type[T], + *args: Any, + request: Request | None = None, + **kwargs: Any, + ) -> T: + """Resolve the given type from the container. Alias for `resolve`.""" @abstractmethod async def resolve( self, abstract: str | type[T], + *args: Any, request: Request | None = None, + **kwargs: Any, ) -> T: """Resolve the given type from the container.""" - pass @abstractmethod async def call(self, callback: type[T] | Callable[..., Any]) -> T: """Call the given callback (Closure, class@method...) and inject its dependencies.""" - pass @abstractmethod def resolved(self, abstract: str | type) -> bool: """Determine if the given abstract type has been resolved.""" - pass @abstractmethod def set_alias(self, abstract: str | type, alias: str | type) -> None: """Alias an abstract type to a different key or type.""" - pass @abstractmethod def is_alias(self, alias: str | type) -> bool: """Determine if a given key or type is an alias.""" - pass @abstractmethod def flush(self) -> None: """Flush the container of all bindings and resolved instances.""" - pass diff --git a/pyrannic/contracts/database/__init__.py b/pyrannic/contracts/database/__init__.py index e69de29..f856f73 100644 --- a/pyrannic/contracts/database/__init__.py +++ b/pyrannic/contracts/database/__init__.py @@ -0,0 +1,4 @@ +from .connector import ConnectorInterface as ConnectorInterface +from .manager import DatabaseManagerInterface as DatabaseManagerInterface +from .migration import MigrationInterface as MigrationInterface +from .schema import SchemaInterface as SchemaInterface diff --git a/pyrannic/contracts/database/connector.py b/pyrannic/contracts/database/connector.py index ea175e2..c2c1cc9 100644 --- a/pyrannic/contracts/database/connector.py +++ b/pyrannic/contracts/database/connector.py @@ -11,14 +11,12 @@ def connection(self) -> Any: """ Establishes and returns a connection to the database. """ - pass @abstractmethod async def disconnect(self) -> None: """ Closes the connection to the database. """ - pass @abstractmethod async def migrate( @@ -28,4 +26,3 @@ async def migrate( """ Runs the provided migrations against the database. """ - pass diff --git a/pyrannic/contracts/orm/__init__.py b/pyrannic/contracts/orm/__init__.py index e69de29..934effc 100644 --- a/pyrannic/contracts/orm/__init__.py +++ b/pyrannic/contracts/orm/__init__.py @@ -0,0 +1 @@ +from .repository import RepositoryInterface as RepositoryInterface diff --git a/pyrannic/contracts/orm/async_repository.py b/pyrannic/contracts/orm/async_repository.py index 784fca8..2a81137 100644 --- a/pyrannic/contracts/orm/async_repository.py +++ b/pyrannic/contracts/orm/async_repository.py @@ -5,7 +5,7 @@ from pyrannic.contracts.pagination.paginator import PaginatorInterface -class RepositoryInterface(ABC, Generic[T]): +class AsyncRepositoryInterface(ABC, Generic[T]): @abstractmethod async def create(self, model: T) -> T: """Insert a new record into the database.""" diff --git a/pyrannic/contracts/orm/query_builder.py b/pyrannic/contracts/orm/query_builder.py index 5397b61..64ae15c 100644 --- a/pyrannic/contracts/orm/query_builder.py +++ b/pyrannic/contracts/orm/query_builder.py @@ -50,12 +50,14 @@ def where_none(self, column_name: str) -> Self: def where_not_none(self, column_name: str) -> Self: """Add a where condition to check if the column is not None.""" + @overload @abstractmethod - def filter(self, *filters: Any | None) -> Self: + def filter(self, *filters: Any) -> Self: """Add filtering conditions to the current query.""" + @overload @abstractmethod - def filter_by(self, **kwargs: Any) -> Self: + def filter(self, **kwargs: Any) -> Self: """Add filtering conditions to the current query.""" @abstractmethod diff --git a/pyrannic/database/provider.py b/pyrannic/database/provider.py index dea2c45..381945c 100644 --- a/pyrannic/database/provider.py +++ b/pyrannic/database/provider.py @@ -5,12 +5,12 @@ from pyrannic.contracts.database.connector import ConnectorInterface from pyrannic.contracts.database.manager import DatabaseManagerInterface from pyrannic.database.manager import DatabaseManager -from pyrannic.orm.sqlalchemy.connector import SqlAlchemyConnector +from pyrannic.orm.sqlalchemy.connector import Connector class DatabaseServiceProvider(ServiceProvider): __singletons__ = { - ConnectorInterface: SqlAlchemyConnector, + ConnectorInterface: Connector, DatabaseManagerInterface: DatabaseManager, } diff --git a/pyrannic/http/resources/collection.py b/pyrannic/http/resources/collection.py index d4b8096..8a1de4f 100644 --- a/pyrannic/http/resources/collection.py +++ b/pyrannic/http/resources/collection.py @@ -51,15 +51,25 @@ def __init__( ) else: if not is_optional(self.__class__, "meta"): + class_name = self.__class__.__name__ + raise RuntimeError( "\n\n" - " The 'meta' attribute is defined as required in your ResourceCollection subclass.\n" - " To fix this exception you have three options:\n" + f"The 'meta' attribute is defined as required in your {class_name} class.\n" + "To fix this exception you have three options:\n" + "\n" " 1. Instead of providing your items collection as a list, use a PaginatorInterface to provide the items.\n" - " 2. Make the 'meta' attribute optional in your ResourceCollection subclass. E.g.:\n" - " class MyResourceCollection(ResourceCollection[MyResource]):\n" - " meta: Optional[PaginationMeta] # or meta: PaginationMeta | None\n" - " 3. Remove the 'meta' attribute from your ResourceCollection subclass if it's not needed.\n" + " If you are using a repository to fetch your items, you can use the 'paginate' method of your repository\n" + f" to get a paginator and pass it to your {class_name} directly:\n" + "\n" + f" {class_name}(your_repository.paginate())\n" + "\n" + f" 2. Make the 'meta' attribute optional in your {class_name} class. E.g.:\n" + "\n" + f" class {class_name}(ResourceCollection[MyResource]):\n" + f" meta: Optional[PaginationMeta] # or meta: PaginationMeta | None\n" + "\n" + f" 3. Remove the 'meta' attribute from your {class_name} class if it's not needed.\n" "\n\n" ) diff --git a/pyrannic/orm/__init__.py b/pyrannic/orm/__init__.py index e69de29..1638edb 100644 --- a/pyrannic/orm/__init__.py +++ b/pyrannic/orm/__init__.py @@ -0,0 +1 @@ +from .abstract_model import AbstractModel as AbstractModel diff --git a/pyrannic/orm/abstract_model.py b/pyrannic/orm/abstract_model.py index 8e83497..eac1cf1 100644 --- a/pyrannic/orm/abstract_model.py +++ b/pyrannic/orm/abstract_model.py @@ -6,14 +6,32 @@ class AbstractModel(ModelInterface, SerializableInterface): - __abstract__ = True + @classmethod + def is_intermediate_table(cls) -> bool: + return False @classmethod def tablename(cls) -> str: + """ + Return the table name for the model. + The table name is generated by removing common suffixes (like "Model", "Entity", "Schema", "Table") + from the class name and converting it to snake_case and pluralizing it. + """ name = cls.__name__ for suffix in _SUFFIXES_TO_REMOVE: if name.endswith(suffix): name = name.replace(suffix, "", 1) - return inflect.pluralize(string.to_snake_case(name)) + snake_case_name = string.to_snake_case(name) + + if cls.is_intermediate_table(): + # For intermediate tables, we want to pluralize the each part of the name and separate it again with underscores. + # For example, if the class name is "UserRole", we want to return "users_roles". + parts = snake_case_name.split("_") + pluralized_parts = [inflect.pluralize(part) for part in parts] + tablename = "_".join(pluralized_parts) + else: + tablename = inflect.pluralize(snake_case_name) + + return tablename diff --git a/pyrannic/orm/sqlalchemy/__init__.py b/pyrannic/orm/sqlalchemy/__init__.py index 7e8f39f..8d76337 100644 --- a/pyrannic/orm/sqlalchemy/__init__.py +++ b/pyrannic/orm/sqlalchemy/__init__.py @@ -1,10 +1,12 @@ +from .async_query_builder import AsyncQueryBuilder as AsyncQueryBuilder from .async_repository import AsyncRepository as AsyncRepository +from .connector import AsyncConnector as AsyncConnector +from .connector import Connector as Connector from .mixins.has_timestamps import HasTimestamps as HasTimestamps from .mixins.soft_deletes import SoftDeletes as SoftDeletes from .model import Model as Model -from .repository import Repository as Repository from .query_builder import QueryBuilder as QueryBuilder -from .async_query_builder import AsyncQueryBuilder as AsyncQueryBuilder -from .connector import SqlAlchemyAsyncConnector as SqlAlchemyAsyncConnector -from .session import Session as Session +from .repository import Repository as Repository +from .schema import Schema as Schema from .session import AsyncSession as AsyncSession +from .session import Session as Session diff --git a/pyrannic/orm/sqlalchemy/abstract_query_builder.py b/pyrannic/orm/sqlalchemy/abstract_query_builder.py index b926dfd..373b3d0 100644 --- a/pyrannic/orm/sqlalchemy/abstract_query_builder.py +++ b/pyrannic/orm/sqlalchemy/abstract_query_builder.py @@ -1,16 +1,9 @@ -from abc import abstractmethod import math +from abc import abstractmethod from datetime import datetime from logging import Logger from typing import Any, Self, overload -from pyrannic.contracts.orm.mixins.soft_deletes import SoftDeletesInterface -from pyrannic.contracts.orm.query_builder import QueryBuilderInterface -from pyrannic.contracts.orm.repository import T -from pyrannic.contracts.orm.scope import ScopeInterface -from pyrannic.orm.sqlalchemy.scopes.soft_deleting_scope import SoftDeletingScope -from pyrannic.support.datetime import get_current_utc_datetime -from pyrannic.support.reflection import get_generic_type from sqlalchemy import ( ColumnExpressionArgument, CompoundSelect, @@ -24,6 +17,14 @@ ) from sqlalchemy.orm import InstrumentedAttribute +from pyrannic.contracts.orm.mixins.soft_deletes import SoftDeletesInterface +from pyrannic.contracts.orm.query_builder import QueryBuilderInterface +from pyrannic.contracts.orm.repository import T +from pyrannic.contracts.orm.scope import ScopeInterface +from pyrannic.orm.sqlalchemy.scopes.soft_deleting_scope import SoftDeletingScope +from pyrannic.support.datetime import get_current_utc_datetime +from pyrannic.support.reflection import get_generic_type + class AbstractQueryBuilder(QueryBuilderInterface[T]): __model__: type[T] @@ -133,20 +134,34 @@ def where_none(self, column_name: str) -> Self: def where_not_none(self, column_name: str) -> Self: return self.where(column(column_name).isnot(None)) - def filter(self, *filters: ColumnExpressionArgument[Any] | None) -> Self: + @overload + def filter(self, *filters: ColumnExpressionArgument[Any]) -> Self: + """ + Apply filtering conditions to the query using SQLAlchemy expressions. + + :param filters: SQLAlchemy expressions for filtering. + :return: The current instance of the query builder. + """ + + @overload + def filter(self, **kwargs: Any) -> Self: self._prepare_query() if isinstance(self._query, (Select, Delete)): - filters = tuple(v for v in filters if v is not None) - self._query = self._query.where(*filters) + self._query = self._query.filter_by(**kwargs) return self - def filter_by(self, **kwargs: Any) -> Self: + def filter(self, *filters: ColumnExpressionArgument[Any], **kwargs: Any) -> Self: self._prepare_query() if isinstance(self._query, (Select, Delete)): - self._query = self._query.filter_by(**kwargs) + if filters: + filters = tuple(v for v in filters if v is not None) + self._query = self._query.where(*filters) + elif kwargs: + kwargs = {k: v for k, v in kwargs.items() if v is not None} + self._query = self._query.filter_by(**kwargs) return self diff --git a/pyrannic/orm/sqlalchemy/async_repository.py b/pyrannic/orm/sqlalchemy/async_repository.py index 15b9cfc..efbc096 100644 --- a/pyrannic/orm/sqlalchemy/async_repository.py +++ b/pyrannic/orm/sqlalchemy/async_repository.py @@ -1,14 +1,14 @@ from typing import Any, Tuple, cast -from pyrannic.contracts.orm.async_repository import RepositoryInterface, T +from sqlalchemy.sql.selectable import TypedReturnsRows + +from pyrannic.contracts.orm.async_repository import AsyncRepositoryInterface, T from pyrannic.contracts.pagination.paginator import PaginatorInterface from pyrannic.orm.sqlalchemy.async_query_builder import AsyncQueryBuilder from pyrannic.pagination.paginator import Paginator -from sqlalchemy.sql.selectable import TypedReturnsRows - -class AsyncRepository(AsyncQueryBuilder[T], RepositoryInterface[T]): +class AsyncRepository(AsyncQueryBuilder[T], AsyncRepositoryInterface[T]): async def create(self, model: T) -> T: try: self.session.add(model) @@ -64,6 +64,7 @@ async def first(self) -> T | None: return model async def all(self) -> list[T]: + self._prepare_query() return await self.get() async def get(self) -> list[T]: diff --git a/pyrannic/orm/sqlalchemy/connector.py b/pyrannic/orm/sqlalchemy/connector.py index c2b7601..c0c7335 100644 --- a/pyrannic/orm/sqlalchemy/connector.py +++ b/pyrannic/orm/sqlalchemy/connector.py @@ -1,6 +1,7 @@ import asyncio from abc import ABC, abstractmethod from logging import Logger +from os import path from typing import Annotated, Any, Generic, TypeVar from sqlalchemy import URL, Engine, create_engine @@ -13,7 +14,8 @@ from sqlalchemy.orm import Session, sessionmaker from pyrannic.container.params import Resolves -from pyrannic.contracts.config.respository import ConfigRepositoryInterface +from pyrannic.contracts.application import ApplicationInterface +from pyrannic.contracts.config.repository import ConfigRepositoryInterface from pyrannic.contracts.database.connector import ConnectorInterface from pyrannic.contracts.database.migration import MigrationInterface from pyrannic.orm.sqlalchemy.schema import Schema @@ -25,13 +27,12 @@ ) -class AbstractSqlAlchemyConnector( - ConnectorInterface, ABC, Generic[EngineType, SessionType] -): +class AbstractConnector(ConnectorInterface, ABC, Generic[EngineType, SessionType]): """ Handles interactions with an SQL Database using SQLAlchemy. """ + _application: ApplicationInterface _logger: Logger _config: ConfigRepositoryInterface _engine: EngineType | None @@ -39,9 +40,11 @@ class AbstractSqlAlchemyConnector( def __init__( self, + application: Annotated[ApplicationInterface, Resolves()], logger: Annotated[Logger, Resolves()], config: Annotated[ConfigRepositoryInterface, Resolves()], ): + self._application = application self._logger = logger self._config = config @@ -57,7 +60,9 @@ def __init__( @property @abstractmethod def engine(self) -> EngineType: - pass + """ + Returns the SQLAlchemy engine instance. + """ async def migrate( self, @@ -102,7 +107,7 @@ def alembic_config(self): alembic_cfg = Config() alembic_cfg.set_main_option( "script_location", - "%(here)s/database/migrations", + path.join("%(here)s", self._application.base_path, "database/migrations"), ) alembic_cfg.set_main_option( "sqlalchemy.url", self.url.render_as_string(hide_password=False) @@ -116,7 +121,8 @@ def alembic_config(self): return alembic_cfg async def _run_migrations( - self, migrations: list[type[MigrationInterface]] | None = None + self, + migrations: list[type[MigrationInterface]] | None = None, ) -> None: if migrations is not None: schema = Schema(self.engine, self._logger) @@ -137,7 +143,7 @@ async def _run_alembic_migrations(self): self._logger.info("|- ✅ Applied Alembic migrations") -class SqlAlchemyConnector(AbstractSqlAlchemyConnector[Engine, sessionmaker[Session]]): +class Connector(AbstractConnector[Engine, sessionmaker[Session]]): """ Handles synchronous interactions with an SQL Database using SQLAlchemy. """ @@ -162,14 +168,14 @@ def engine(self) -> Engine: Returns the SQLAlchemy engine instance. """ - # TODO: Use config.services.sqlalchemy settings for pool size, echo, etc. + # TODO: Use config.services.sqlalchemy settings for pool size, echo, max_overflow, etc. if not self._engine: self._engine = create_engine( self.url, echo=False, - pool_size=5, - max_overflow=5, + # TODO pool_size=5, + # TODO max_overflow=5, pool_pre_ping=True, future=True, # lazy connections ) @@ -177,9 +183,7 @@ def engine(self) -> Engine: return self._engine -class SqlAlchemyAsyncConnector( - AbstractSqlAlchemyConnector[AsyncEngine, async_sessionmaker[AsyncSession]] -): +class AsyncConnector(AbstractConnector[AsyncEngine, async_sessionmaker[AsyncSession]]): """ Handles asynchronous interactions with an SQL Database using SQLAlchemy. """ @@ -210,8 +214,8 @@ def engine(self) -> AsyncEngine: self._engine = create_async_engine( self.url, echo=False, - pool_size=5, - max_overflow=5, + # TODO pool_size=5, + # TODO max_overflow=5, pool_pre_ping=True, future=True, # lazy connections ) diff --git a/pyrannic/orm/sqlalchemy/repository.py b/pyrannic/orm/sqlalchemy/repository.py index 89a291c..0d534a1 100644 --- a/pyrannic/orm/sqlalchemy/repository.py +++ b/pyrannic/orm/sqlalchemy/repository.py @@ -1,12 +1,12 @@ from typing import Any, Tuple, cast +from sqlalchemy.sql.selectable import TypedReturnsRows + from pyrannic.contracts.orm.repository import RepositoryInterface, T from pyrannic.contracts.pagination.paginator import PaginatorInterface from pyrannic.orm.sqlalchemy.query_builder import QueryBuilder from pyrannic.pagination.paginator import Paginator -from sqlalchemy.sql.selectable import TypedReturnsRows - class Repository(QueryBuilder[T], RepositoryInterface[T]): def create(self, model: T) -> T: @@ -64,6 +64,7 @@ def first(self) -> T | None: return model def all(self) -> list[T]: + self._prepare_query() return self.get() def get(self) -> list[T]: diff --git a/pyrannic/orm/sqlalchemy/schema.py b/pyrannic/orm/sqlalchemy/schema.py index a0179b8..4abb88a 100644 --- a/pyrannic/orm/sqlalchemy/schema.py +++ b/pyrannic/orm/sqlalchemy/schema.py @@ -16,9 +16,9 @@ def __init__(self, engine: AsyncEngine | Engine, logger: Logger): self._engine = engine self._logger = logger - async def create(self, blueprint: DeclarativeBase) -> None: + async def create(self, blueprint: type[DeclarativeBase]) -> None: """ - Creates the conversation table in the database if it does not exist. + Creates the blueprint in the database if it does not exist. """ await self._run( blueprint.__table__, @@ -27,9 +27,9 @@ async def create(self, blueprint: DeclarativeBase) -> None: "Failed to create {} table: {}", ) - async def drop(self, blueprint: DeclarativeBase) -> None: + async def drop(self, blueprint: type[DeclarativeBase]) -> None: """ - Drops the conversation table from the database if it exists. + Drops the blueprint from the database if it exists. """ await self._run( blueprint.__table__, diff --git a/pyrannic/support/facades/config.py b/pyrannic/support/facades/config.py index d7a3997..36cf762 100644 --- a/pyrannic/support/facades/config.py +++ b/pyrannic/support/facades/config.py @@ -1,4 +1,4 @@ -from pyrannic.contracts.config.respository import ConfigRepositoryInterface +from pyrannic.contracts.config.repository import ConfigRepositoryInterface from pyrannic.support.facades.facade import facade diff --git a/tests/conftest.py b/tests/conftest.py index a34db55..b947e17 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -3,7 +3,7 @@ import pytest from pyrannic.application import Application -from pyrannic.contracts.application import ApplicationInterface +from pyrannic.contracts import ApplicationInterface @pytest.fixture(scope="module") diff --git a/tests/unit/bootstrap/manager/providers.py b/tests/unit/bootstrap/manager/providers.py index 7b25d64..2811e6b 100644 --- a/tests/unit/bootstrap/manager/providers.py +++ b/tests/unit/bootstrap/manager/providers.py @@ -5,7 +5,7 @@ from pyrannic.bootstrap.service_provider import ServiceProvider from pyrannic.container.param_functions import Resolves from pyrannic.contracts.application import ApplicationInterface -from pyrannic.contracts.config.respository import ConfigRepositoryInterface +from pyrannic.contracts.config.repository import ConfigRepositoryInterface from pyrannic.support.string import to_snake_case diff --git a/tests/unit/bootstrap/manager/test_lifespan_method.py b/tests/unit/bootstrap/manager/test_lifespan_method.py index 16d9ff5..e45db1f 100644 --- a/tests/unit/bootstrap/manager/test_lifespan_method.py +++ b/tests/unit/bootstrap/manager/test_lifespan_method.py @@ -5,13 +5,13 @@ from pyrannic.application import Application from pyrannic.bootstrap.manager import BootstrapManager from pyrannic.bootstrap.service_provider import ServiceProvider -from pyrannic.contracts.config.respository import ConfigRepositoryInterface +from pyrannic.contracts.config.repository import ConfigRepositoryInterface from tests.unit.bootstrap.manager.providers import ( FooServiceProvider, UnbootableCriticalServiceProvider, UnbootableServiceProvider, - UninitializableServiceProvider, UninitializableCriticalServiceProvider, + UninitializableServiceProvider, ) diff --git a/tests/unit/bootstrap/manager/test_start_critical_services.py b/tests/unit/bootstrap/manager/test_start_critical_services.py index 7ef8d85..e921a21 100644 --- a/tests/unit/bootstrap/manager/test_start_critical_services.py +++ b/tests/unit/bootstrap/manager/test_start_critical_services.py @@ -6,7 +6,7 @@ from pyrannic.bootstrap.manager import BootstrapManager from pyrannic.bootstrap.service_provider import ServiceProvider from pyrannic.config.env import read_int -from pyrannic.contracts.config.respository import ConfigRepositoryInterface +from pyrannic.contracts.config.repository import ConfigRepositoryInterface from pyrannic.support.facades.facade import Facade diff --git a/tests/unit/config/provider/test_provider.py b/tests/unit/config/provider/test_provider.py index 3ef5ad8..21e18dc 100644 --- a/tests/unit/config/provider/test_provider.py +++ b/tests/unit/config/provider/test_provider.py @@ -2,7 +2,7 @@ from pyrannic.config.provider import ConfigRepositoryProvider from pyrannic.contracts.application import ApplicationInterface -from pyrannic.contracts.config.respository import ConfigRepositoryInterface +from pyrannic.contracts.config.repository import ConfigRepositoryInterface @pytest.mark.asyncio diff --git a/tests/unit/config/repository/conftest.py b/tests/unit/config/repository/conftest.py index f958e98..e462ffe 100644 --- a/tests/unit/config/repository/conftest.py +++ b/tests/unit/config/repository/conftest.py @@ -1,7 +1,7 @@ import pytest from pyrannic.config.repository import ConfigRepository -from pyrannic.contracts.config.respository import ConfigRepositoryInterface +from pyrannic.contracts.config.repository import ConfigRepositoryInterface @pytest.fixture() diff --git a/tests/unit/config/repository/test_boolean.py b/tests/unit/config/repository/test_boolean.py index dcbbf38..f2e3bb8 100644 --- a/tests/unit/config/repository/test_boolean.py +++ b/tests/unit/config/repository/test_boolean.py @@ -1,4 +1,4 @@ -from pyrannic.contracts.config.respository import ConfigRepositoryInterface +from pyrannic.contracts.config.repository import ConfigRepositoryInterface def test_optional_boolean(repository: ConfigRepositoryInterface): diff --git a/tests/unit/config/repository/test_config_repository_all_method.py b/tests/unit/config/repository/test_config_repository_all_method.py index 73a45a0..8ab87b5 100644 --- a/tests/unit/config/repository/test_config_repository_all_method.py +++ b/tests/unit/config/repository/test_config_repository_all_method.py @@ -1,4 +1,4 @@ -from pyrannic.contracts.config.respository import ConfigRepositoryInterface +from pyrannic.contracts.config.repository import ConfigRepositoryInterface def test_all_method(repository: ConfigRepositoryInterface): diff --git a/tests/unit/config/repository/test_config_repository_get_method.py b/tests/unit/config/repository/test_config_repository_get_method.py index 50e1991..25784de 100644 --- a/tests/unit/config/repository/test_config_repository_get_method.py +++ b/tests/unit/config/repository/test_config_repository_get_method.py @@ -1,6 +1,6 @@ import pytest -from pyrannic.contracts.config.respository import ConfigRepositoryInterface +from pyrannic.contracts.config.repository import ConfigRepositoryInterface def test_get_method(repository: ConfigRepositoryInterface): diff --git a/tests/unit/config/repository/test_config_repository_has_method.py b/tests/unit/config/repository/test_config_repository_has_method.py index 9643869..f8cf8e0 100644 --- a/tests/unit/config/repository/test_config_repository_has_method.py +++ b/tests/unit/config/repository/test_config_repository_has_method.py @@ -1,4 +1,4 @@ -from pyrannic.contracts.config.respository import ConfigRepositoryInterface +from pyrannic.contracts.config.repository import ConfigRepositoryInterface def test_has_method(repository: ConfigRepositoryInterface): diff --git a/tests/unit/config/repository/test_config_repository_set_method.py b/tests/unit/config/repository/test_config_repository_set_method.py index 5007507..69536a8 100644 --- a/tests/unit/config/repository/test_config_repository_set_method.py +++ b/tests/unit/config/repository/test_config_repository_set_method.py @@ -1,4 +1,4 @@ -from pyrannic.contracts.config.respository import ConfigRepositoryInterface +from pyrannic.contracts.config.repository import ConfigRepositoryInterface def test_set_a_new_key(repository: ConfigRepositoryInterface): diff --git a/tests/unit/config/repository/test_float.py b/tests/unit/config/repository/test_float.py index f0472a7..b388d12 100644 --- a/tests/unit/config/repository/test_float.py +++ b/tests/unit/config/repository/test_float.py @@ -1,4 +1,4 @@ -from pyrannic.contracts.config.respository import ConfigRepositoryInterface +from pyrannic.contracts.config.repository import ConfigRepositoryInterface def test_optional_float(repository: ConfigRepositoryInterface): diff --git a/tests/unit/config/repository/test_integer.py b/tests/unit/config/repository/test_integer.py index d9780dd..88ea2ce 100644 --- a/tests/unit/config/repository/test_integer.py +++ b/tests/unit/config/repository/test_integer.py @@ -1,4 +1,4 @@ -from pyrannic.contracts.config.respository import ConfigRepositoryInterface +from pyrannic.contracts.config.repository import ConfigRepositoryInterface def test_optional_integer(repository: ConfigRepositoryInterface): diff --git a/tests/unit/container/conftest.py b/tests/unit/container/conftest.py index 4b8286c..70777bf 100644 --- a/tests/unit/container/conftest.py +++ b/tests/unit/container/conftest.py @@ -1,10 +1,11 @@ -from abc import ABC, abstractmethod import asyncio -from typing import Generator +from abc import ABC, abstractmethod +from typing import Generator, Generic -from fastapi import Request import pytest +from fastapi import Request +from pyrannic.container.container import T from pyrannic.container.decorators import scoped, singleton from pyrannic.contracts.application import ApplicationInterface from pyrannic.contracts.container.container import ContainerInterface @@ -37,6 +38,23 @@ def __init__(self, foo_service: Resolve[FooInterface]): self.foo_service = foo_service +class BazServiceWithParams: + def __init__( + self, value1: str, value2: str, foo_service: Resolve[FooImplementation] + ): + self.value1 = value1 + self.value2 = value2 + self.foo_service = foo_service + + +class FooModel: + pass + + +class FooGeneric(Generic[T]): + pass + + @scoped class ScopedClass(FooInterface): def foo_method(self) -> str: @@ -64,10 +82,7 @@ async def async_callable_with_dependencies(foo: Resolve[FooImplementation]) -> N assert foo.foo_method() == "FooImplementation" -async def resolve_foo_interface( - app: ApplicationInterface, - request: Request, -) -> FooInterface: +async def resolve_foo_interface(_: ApplicationInterface, __: Request) -> FooInterface: await asyncio.sleep(0.1) # Simulate async work return FooImplementation() diff --git a/tests/unit/container/container/test_make.py b/tests/unit/container/container/test_make.py new file mode 100644 index 0000000..f4391e3 --- /dev/null +++ b/tests/unit/container/container/test_make.py @@ -0,0 +1,85 @@ +import pytest +from fastapi.exceptions import RequestValidationError + +from pyrannic.contracts.container.container import ContainerInterface +from pyrannic.support.reflection import get_generic_type +from tests.unit.container.conftest import ( + BazServiceWithParams, + FooGeneric, + FooImplementation, + FooInterface, + FooModel, + FooSecondaryImplementation, +) + + +@pytest.mark.asyncio +async def test_make_concrete_class(container: ContainerInterface): + instance = await container.make(FooSecondaryImplementation) + + assert not container.is_bound(FooSecondaryImplementation) + assert isinstance(instance, FooSecondaryImplementation) + + +@pytest.mark.asyncio +async def test_make_generic_class(container: ContainerInterface): + instance = await container.make(FooGeneric[FooModel]) + + assert not container.is_bound(FooGeneric[FooModel]) + assert isinstance(instance, FooGeneric) + assert get_generic_type(instance) == FooModel + + +@pytest.mark.asyncio +async def test_make_with_positional_parameters(container: ContainerInterface): + instance = await container.make(BazServiceWithParams, "value1", "value2") + + assert not container.is_bound(BazServiceWithParams) + assert isinstance(instance, BazServiceWithParams) + assert instance.value1 == "value1" + assert instance.value2 == "value2" + assert isinstance(instance.foo_service, FooImplementation) + + +@pytest.mark.asyncio +async def test_make_with_named_parameters(container: ContainerInterface): + instance = await container.make( + BazServiceWithParams, value1="value1", value2="value2" + ) + + assert not container.is_bound(BazServiceWithParams) + assert isinstance(instance, BazServiceWithParams) + assert instance.value1 == "value1" + assert instance.value2 == "value2" + assert isinstance(instance.foo_service, FooImplementation) + + +@pytest.mark.asyncio +async def test_make_with_mixed_parameters(container: ContainerInterface): + instance = await container.make(BazServiceWithParams, "value1", value2="value2") + + assert not container.is_bound(BazServiceWithParams) + assert isinstance(instance, BazServiceWithParams) + assert instance.value1 == "value1" + assert instance.value2 == "value2" + assert isinstance(instance.foo_service, FooImplementation) + + +@pytest.mark.asyncio +async def test_make_with_interface_not_bound(container: ContainerInterface): + with pytest.raises(RequestValidationError) as exc_info: + await container.make(FooInterface) + + error = str(exc_info.value) + print(error) + assert "No binding found for interface FooInterface" in error + + +@pytest.mark.asyncio +async def test_make_with_key_not_bound(container: ContainerInterface): + with pytest.raises(RequestValidationError) as exc_info: + await container.make("FooInterface") + + error = str(exc_info.value) + print(error) + assert "No binding found for key FooInterface" in error diff --git a/tests/unit/container/container/test_resolve.py b/tests/unit/container/container/test_resolve.py index e55afea..e69117c 100644 --- a/tests/unit/container/container/test_resolve.py +++ b/tests/unit/container/container/test_resolve.py @@ -1,8 +1,16 @@ -from fastapi.exceptions import RequestValidationError import pytest +from fastapi.exceptions import RequestValidationError from pyrannic.contracts.container.container import ContainerInterface -from tests.unit.container.conftest import FooInterface, FooSecondaryImplementation +from pyrannic.support.reflection import get_generic_type +from tests.unit.container.conftest import ( + BazServiceWithParams, + FooGeneric, + FooImplementation, + FooInterface, + FooModel, + FooSecondaryImplementation, +) @pytest.mark.asyncio @@ -13,6 +21,50 @@ async def test_resolve_concrete_class(container: ContainerInterface): assert isinstance(instance, FooSecondaryImplementation) +@pytest.mark.asyncio +async def test_resolve_generic_class(container: ContainerInterface): + instance = await container.resolve(FooGeneric[FooModel]) + + assert not container.is_bound(FooGeneric[FooModel]) + assert isinstance(instance, FooGeneric) + assert get_generic_type(instance) == FooModel + + +@pytest.mark.asyncio +async def test_resolve_with_positional_parameters(container: ContainerInterface): + instance = await container.resolve(BazServiceWithParams, "value1", "value2") + + assert not container.is_bound(BazServiceWithParams) + assert isinstance(instance, BazServiceWithParams) + assert instance.value1 == "value1" + assert instance.value2 == "value2" + assert isinstance(instance.foo_service, FooImplementation) + + +@pytest.mark.asyncio +async def test_resolve_with_named_parameters(container: ContainerInterface): + instance = await container.resolve( + BazServiceWithParams, value1="value1", value2="value2" + ) + + assert not container.is_bound(BazServiceWithParams) + assert isinstance(instance, BazServiceWithParams) + assert instance.value1 == "value1" + assert instance.value2 == "value2" + assert isinstance(instance.foo_service, FooImplementation) + + +@pytest.mark.asyncio +async def test_resolve_with_mixed_parameters(container: ContainerInterface): + instance = await container.resolve(BazServiceWithParams, "value1", value2="value2") + + assert not container.is_bound(BazServiceWithParams) + assert isinstance(instance, BazServiceWithParams) + assert instance.value1 == "value1" + assert instance.value2 == "value2" + assert isinstance(instance.foo_service, FooImplementation) + + @pytest.mark.asyncio async def test_resolve_with_interface_not_bound(container: ContainerInterface): with pytest.raises(RequestValidationError) as exc_info: diff --git a/tests/unit/database/test_database_provider.py b/tests/unit/database/test_database_provider.py index f9ce7af..d98943f 100644 --- a/tests/unit/database/test_database_provider.py +++ b/tests/unit/database/test_database_provider.py @@ -8,7 +8,7 @@ from pyrannic.contracts.database.manager import DatabaseManagerInterface from pyrannic.database.manager import DatabaseManager from pyrannic.database.provider import DatabaseServiceProvider -from pyrannic.orm.sqlalchemy.connector import SqlAlchemyConnector +from pyrannic.orm.sqlalchemy.connector import Connector from tests.unit.database.utils import MockDatabaseServiceProvider @@ -23,7 +23,7 @@ async def test_register_singletons(): connector_1 = await application.container.resolve(ConnectorInterface) manager_1 = await application.container.resolve(DatabaseManagerInterface) - assert isinstance(connector_1, SqlAlchemyConnector) + assert isinstance(connector_1, Connector) assert isinstance(manager_1, DatabaseManager) connector_2 = await application.container.resolve(ConnectorInterface) diff --git a/tests/unit/http/resources/test_collection.py b/tests/unit/http/resources/test_collection.py index 994a6bb..f04c47a 100644 --- a/tests/unit/http/resources/test_collection.py +++ b/tests/unit/http/resources/test_collection.py @@ -6,8 +6,8 @@ from tests.unit.http.resources.utils import ( BarCollection, FooCollection, - FooResource, FooModel, + FooResource, OptionalMetaCollection, RequiredMetaCollection, ) @@ -110,11 +110,20 @@ def test_collection__exception_required_meta(): RequiredMetaCollection([]) error = str(exc_info.value) + + assert ( + "The 'meta' attribute is defined as required in your RequiredMetaCollection class." + in error + ) + + assert "RequiredMetaCollection(your_repository.paginate())" in error assert ( - "The 'meta' attribute is defined as required in your ResourceCollection subclass" + "Make the 'meta' attribute optional in your RequiredMetaCollection class." in error ) + assert "Remove the 'meta' attribute from your RequiredMetaCollection class" in error + def test_collection__serialize_empty_data(): collection = FooCollection([]) diff --git a/tests/unit/orm/abstract_model/test_tablename_abstract_model.py b/tests/unit/orm/abstract_model/test_tablename_abstract_model.py new file mode 100644 index 0000000..8116ca3 --- /dev/null +++ b/tests/unit/orm/abstract_model/test_tablename_abstract_model.py @@ -0,0 +1,43 @@ +from pyrannic.orm.abstract_model import AbstractModel + + +def test_tablename_with_all_supported_suffixes(): + class UserModel(AbstractModel): + pass + + class ProductEntity(AbstractModel): + pass + + class OrderSchema(AbstractModel): + pass + + class CustomerTable(AbstractModel): + pass + + assert UserModel.tablename() == "users" + assert ProductEntity.tablename() == "products" + assert OrderSchema.tablename() == "orders" + assert CustomerTable.tablename() == "customers" + + +def test_tablename_with_no_suffix(): + class Category(AbstractModel): + pass + + assert Category.tablename() == "categories" + + +def test_tablename_with_mixed_case(): + class MixedCaseModel(AbstractModel): + pass + + assert MixedCaseModel.tablename() == "mixed_cases" + + +def test_tablename_for_intermediate_table(): + class UserRole(AbstractModel): + @classmethod + def is_intermediate_table(cls) -> bool: + return True + + assert UserRole.tablename() == "users_roles" diff --git a/tests/unit/orm/sqlalchemy/conftest.py b/tests/unit/orm/sqlalchemy/conftest.py new file mode 100644 index 0000000..b16124b --- /dev/null +++ b/tests/unit/orm/sqlalchemy/conftest.py @@ -0,0 +1,24 @@ +import pytest_asyncio + +from pyrannic.contracts import ( + ApplicationInterface, + ConnectorInterface, + DatabaseManagerInterface, + RepositoryInterface, +) +from pyrannic.database.manager import DatabaseManager +from pyrannic.facades import Config +from pyrannic.orm.sqlalchemy import Connector, Repository +from tests.unit.orm.sqlalchemy.utils import BarModel + + +@pytest_asyncio.fixture() +async def repository( + application: ApplicationInterface, +) -> RepositoryInterface[BarModel]: + Config.set("database.connections.sqlite.database", ":memory:") + + application.container.singleton(ConnectorInterface, Connector) + application.container.singleton(DatabaseManagerInterface, DatabaseManager) + + return await application.container.make(Repository[BarModel]) diff --git a/tests/unit/orm/sqlalchemy/test_sa_model.py b/tests/unit/orm/sqlalchemy/test_sa_model.py new file mode 100644 index 0000000..5615f4d --- /dev/null +++ b/tests/unit/orm/sqlalchemy/test_sa_model.py @@ -0,0 +1,53 @@ +import pytest + +from pyrannic.contracts import ApplicationInterface +from pyrannic.contracts.database.manager import DatabaseManagerInterface +from pyrannic.contracts.orm import RepositoryInterface +from tests.unit.orm.sqlalchemy.utils import BarModel, BarsTable + + +def test_tablename_property(): + assert BarModel.tablename() == "bars" + assert BarModel.__tablename__ == "bars" + + +def test_primary_key_column(): + assert BarModel.primary_key_column().name == "id" # pyright: ignore[reportFunctionMemberAccess, reportUnknownMemberType, reportAttributeAccessIssue] + + +def test_primary_key_value(): + bar = BarModel(id=1) + assert bar.primary_key_value == 1 + + +@pytest.mark.asyncio +async def test_is_dirty_and_is_clean( + application: ApplicationInterface, + repository: RepositoryInterface[BarModel], +): + manager = await application.container.resolve(DatabaseManagerInterface) + await manager.migrate([BarsTable]) + + bar = BarModel(id=1, name="Bar") + + # Initially, the model is not saved, so it should be considered dirty + assert bar.is_dirty() + assert not bar.is_clean() + + # Save the model to the repository + bar = repository.create(bar) + assert not bar.is_dirty() + assert bar.is_clean() + + # Simulate a change to the model + bar.id = 2 + assert bar.is_dirty() + assert not bar.is_clean() + + # Check specific attribute + assert bar.is_dirty("id") + assert not bar.is_clean("id") + + # Check with an attribute that hasn't changed + assert not bar.is_dirty("name") + assert not bar.is_dirty("non_existent_attr") diff --git a/tests/unit/orm/sqlalchemy/test_query_builder.py b/tests/unit/orm/sqlalchemy/test_sa_query_builder.py similarity index 88% rename from tests/unit/orm/sqlalchemy/test_query_builder.py rename to tests/unit/orm/sqlalchemy/test_sa_query_builder.py index e69cbc4..cf98976 100644 --- a/tests/unit/orm/sqlalchemy/test_query_builder.py +++ b/tests/unit/orm/sqlalchemy/test_sa_query_builder.py @@ -1,18 +1,17 @@ import pytest -from sqlalchemy.orm import Session from sqlalchemy.ext.asyncio import AsyncSession +from sqlalchemy.orm import Session -from pyrannic import Config from pyrannic.contracts.application import ApplicationInterface from pyrannic.contracts.database.connector import ConnectorInterface from pyrannic.contracts.database.manager import DatabaseManagerInterface from pyrannic.database.manager import DatabaseManager +from pyrannic.facades import Config from pyrannic.orm.sqlalchemy import ( - Repository, + AsyncConnector, AsyncRepository, - SqlAlchemyAsyncConnector, + Repository, ) - from tests.unit.orm.sqlalchemy.utils import BarModel @@ -27,7 +26,7 @@ async def test_query_builder(application: ApplicationInterface): async def test_async_query_builder(application: ApplicationInterface): Config.set("database.connections.sqlite.driver", "sqlite+aiosqlite") - application.container.singleton(ConnectorInterface, SqlAlchemyAsyncConnector) + application.container.singleton(ConnectorInterface, AsyncConnector) application.container.singleton(DatabaseManagerInterface, DatabaseManager) repository = await application.container.resolve(AsyncRepository[BarModel]) diff --git a/tests/unit/orm/sqlalchemy/test_sa_schema.py b/tests/unit/orm/sqlalchemy/test_sa_schema.py new file mode 100644 index 0000000..7e2e096 --- /dev/null +++ b/tests/unit/orm/sqlalchemy/test_sa_schema.py @@ -0,0 +1,148 @@ +from logging import Logger +from typing import cast + +import pytest +import sqlalchemy + +from pyrannic.contracts import ( + ApplicationInterface, + ConnectorInterface, + DatabaseManagerInterface, +) +from pyrannic.database.manager import DatabaseManager +from pyrannic.orm.sqlalchemy import AsyncConnector, Connector, Schema +from pyrannic.support.facades.config import Config +from tests.unit.orm.sqlalchemy.utils import BarModel + + +@pytest.mark.asyncio +async def test_create(application: ApplicationInterface): + Config.set("database.connections.sqlite.driver", "sqlite") + Config.set("database.connections.sqlite.database", ":memory:") + + application.container.singleton(ConnectorInterface, Connector) + application.container.singleton(DatabaseManagerInterface, DatabaseManager) + + connector = cast( + Connector, + await application.container.resolve(ConnectorInterface), + ) + + logger = await application.container.resolve(Logger) + + assert not sqlalchemy.inspect(connector.engine).has_table("bars") + + schema = Schema(connector.engine, logger) + await schema.create(BarModel) + + assert sqlalchemy.inspect(connector.engine).has_table("bars") + + +@pytest.mark.asyncio +async def test_async_create(application: ApplicationInterface): + Config.set("database.connections.sqlite.driver", "sqlite+aiosqlite") + Config.set("database.connections.sqlite.database", ":memory:") + + application.container.singleton(ConnectorInterface, AsyncConnector) + application.container.singleton(DatabaseManagerInterface, DatabaseManager) + + connector = cast( + AsyncConnector, + await application.container.resolve(ConnectorInterface), + ) + + logger = await application.container.resolve(Logger) + + async with connector.engine.connect() as conn: + assert not await conn.run_sync( + lambda sync_conn: sqlalchemy.inspect(sync_conn).has_table("bars") + ) + + schema = Schema(connector.engine, logger) + await schema.create(BarModel) + + async with connector.engine.connect() as conn: + assert await conn.run_sync( + lambda sync_conn: sqlalchemy.inspect(sync_conn).has_table("bars") + ) + + +@pytest.mark.asyncio +async def test_drop(application: ApplicationInterface): + Config.set("database.connections.sqlite.driver", "sqlite") + Config.set("database.connections.sqlite.database", ":memory:") + + application.container.singleton(ConnectorInterface, Connector) + application.container.singleton(DatabaseManagerInterface, DatabaseManager) + + connector = cast( + Connector, + await application.container.resolve(ConnectorInterface), + ) + + logger = await application.container.resolve(Logger) + + schema = Schema(connector.engine, logger) + await schema.create(BarModel) + + assert sqlalchemy.inspect(connector.engine).has_table("bars") + + await schema.drop(BarModel) + + assert not sqlalchemy.inspect(connector.engine).has_table("bars") + + +@pytest.mark.asyncio +async def test_async_drop(application: ApplicationInterface): + Config.set("database.connections.sqlite.driver", "sqlite+aiosqlite") + Config.set("database.connections.sqlite.database", ":memory:") + + application.container.singleton(ConnectorInterface, AsyncConnector) + application.container.singleton(DatabaseManagerInterface, DatabaseManager) + + connector = cast( + AsyncConnector, + await application.container.resolve(ConnectorInterface), + ) + + logger = await application.container.resolve(Logger) + + schema = Schema(connector.engine, logger) + await schema.create(BarModel) + + async with connector.engine.connect() as conn: + assert await conn.run_sync( + lambda sync_conn: sqlalchemy.inspect(sync_conn).has_table("bars") + ) + + await schema.drop(BarModel) + + async with connector.engine.connect() as conn: + assert not await conn.run_sync( + lambda sync_conn: sqlalchemy.inspect(sync_conn).has_table("bars") + ) + + +@pytest.mark.asyncio +async def test_log_on_exception( + caplog: pytest.LogCaptureFixture, application: ApplicationInterface +): + Config.set("database.connections.sqlite.database", ":memory:") + + application.container.singleton(ConnectorInterface, Connector) + application.container.singleton(DatabaseManagerInterface, DatabaseManager) + + connector = cast( + Connector, + await application.container.resolve(ConnectorInterface), + ) + + logger = await application.container.resolve(Logger) + + schema = Schema(connector.engine, logger) + + await schema.create(BarModel) + assert "Failed to create bars table:" in caplog.text + + await schema.drop(BarModel) + assert "Failed to drop bars table:" in caplog.text diff --git a/tests/unit/orm/sqlalchemy/utils.py b/tests/unit/orm/sqlalchemy/utils.py index d8e1199..d488509 100644 --- a/tests/unit/orm/sqlalchemy/utils.py +++ b/tests/unit/orm/sqlalchemy/utils.py @@ -1,13 +1,18 @@ +from sqlalchemy import Integer, String from sqlalchemy.orm import Mapped, mapped_column -from sqlalchemy import Integer -from pyrannic.orm.sqlalchemy import Model +from pyrannic.database.migration import Migration +from pyrannic.orm.sqlalchemy.model import Model class BarModel(Model): - id: Mapped[int] = mapped_column( - Integer, - primary_key=True, - autoincrement=True, - comment="Unique identifier for the hero; serves as the primary key.", - ) + id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True) + name: Mapped[str] = mapped_column(String(255)) + + +class BarsTable(Migration): + async def up(self) -> None: + await self.schema.create(BarModel) + + async def down(self) -> None: + await self.schema.drop(BarModel) diff --git a/tests/unit/orm/test_abstract_model.py b/tests/unit/orm/test_abstract_model.py deleted file mode 100644 index 01b5c95..0000000 --- a/tests/unit/orm/test_abstract_model.py +++ /dev/null @@ -1,8 +0,0 @@ -from tests.unit.orm.utils import FooModel, FooEntity, FooSchema, FooTable - - -def test_tablename(): - assert FooModel.tablename() == "foos" - assert FooEntity.tablename() == "foos" - assert FooSchema.tablename() == "foos" - assert FooTable.tablename() == "foos" diff --git a/tests/unit/orm/utils.py b/tests/unit/orm/utils.py deleted file mode 100644 index e549c72..0000000 --- a/tests/unit/orm/utils.py +++ /dev/null @@ -1,17 +0,0 @@ -from pyrannic.orm.abstract_model import AbstractModel - - -class FooModel(AbstractModel): - pass - - -class FooEntity(AbstractModel): - pass - - -class FooSchema(AbstractModel): - pass - - -class FooTable(AbstractModel): - pass diff --git a/tests/unit/support/facades/test_config_facade.py b/tests/unit/support/facades/test_config_facade.py index 10bf15c..1309270 100644 --- a/tests/unit/support/facades/test_config_facade.py +++ b/tests/unit/support/facades/test_config_facade.py @@ -1,5 +1,5 @@ from pyrannic.contracts.application import ApplicationInterface -from pyrannic.contracts.config.respository import ConfigRepositoryInterface +from pyrannic.contracts.config.repository import ConfigRepositoryInterface from pyrannic.facades import Config