Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 9 additions & 2 deletions app/domain/containers/repository.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,14 @@
from typing import Protocol

from app.domain.containers.entities import Container
from app.domain.protocols import AsyncRepositoryProtocol
from cleanstack import EntityId


class ContainerRepositoryProtocol(AsyncRepositoryProtocol[Container], Protocol): ...
class ContainerRepositoryProtocol(Protocol):
async def get_by_id(self, entity_id: EntityId, /) -> Container | None: ...

async def save(self, entity: Container, /) -> None: ...

async def update(self, entity: Container, /) -> None: ...

async def remove(self, entity: Container, /) -> None: ...
19 changes: 17 additions & 2 deletions app/domain/items/repository.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,22 @@
from typing import Protocol

from app.domain.items.entities import Item
from app.domain.protocols import AsyncRepositoryProtocol
from cleanstack import EntityId, FilterEntity, PaginatedResponse, Pagination, SortEntity


class ItemRepositoryProtocol(AsyncRepositoryProtocol[Item], Protocol): ...
class ItemRepositoryProtocol(Protocol):
async def get_all(
self,
search: str | None = None,
filters: list[FilterEntity] | None = None,
sort: list[SortEntity] | None = None,
pagination: Pagination | None = None,
) -> PaginatedResponse[Item]: ...

async def get_by_id(self, entity_id: EntityId, /) -> Item | None: ...

async def save(self, entity: Item, /) -> None: ...

async def update(self, entity: Item, /) -> None: ...

async def remove(self, entity: Item, /) -> None: ...
1 change: 0 additions & 1 deletion app/domain/items/use_cases.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,7 +38,6 @@ async def get_item(context: ContextProtocol, /, item_id: EntityId) -> Item:


async def create_item(context: ContextProtocol, /, data: ItemCreate) -> Item:
# Explicitly write all fields for clarity
item = Item(
id=uuid.uuid7(),
uuid_field=uuid.uuid7(),
Expand Down
46 changes: 0 additions & 46 deletions app/domain/protocols.py

This file was deleted.

33 changes: 31 additions & 2 deletions app/infrastructure/mongo/containers.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,36 @@
from pymongo.asynchronous.client_session import AsyncClientSession
from pymongo.asynchronous.database import AsyncDatabase

from app.domain.containers.entities import Container
from cleanstack.mongo import AsyncMongoRepository
from app.domain.containers.repository import ContainerRepositoryProtocol
from cleanstack import EntityId
from cleanstack.mongo import AsyncMongoRepository, MongoDocument


class ContainerMongoRepository(AsyncMongoRepository[Container]):
class ContainerMongoRepository(ContainerRepositoryProtocol):
domain_entity_type = Container
collection_name = "containers"
searchable_fields = ()

def __init__(
self,
database: AsyncDatabase[MongoDocument],
session: AsyncClientSession | None = None,
) -> None:
self.repository = AsyncMongoRepository[Container].from_spec(
binding=self,
database=database,
session=session,
)

async def get_by_id(self, entity_id: EntityId, /) -> Container | None:
return await self.repository.get_by_id(entity_id)

async def save(self, entity: Container, /) -> None:
await self.repository.save(entity)

async def update(self, entity: Container, /) -> None:
await self.repository.update(entity)

async def remove(self, entity: Container, /) -> None:
await self.repository.remove(entity)
46 changes: 44 additions & 2 deletions app/infrastructure/mongo/items.py
Original file line number Diff line number Diff line change
@@ -1,8 +1,50 @@
from pymongo.asynchronous.client_session import AsyncClientSession
from pymongo.asynchronous.database import AsyncDatabase

from app.domain.items.entities import Item
from cleanstack.mongo import AsyncMongoRepository
from app.domain.items.repository import ItemRepositoryProtocol
from cleanstack import EntityId, FilterEntity, PaginatedResponse, Pagination, SortEntity
from cleanstack.mongo import AsyncMongoRepository, MongoDocument


class ItemMongoRepository(AsyncMongoRepository[Item]):
class ItemMongoRepository(ItemRepositoryProtocol):
domain_entity_type = Item
collection_name = "items"
searchable_fields = ("string_field",)

def __init__(
self,
database: AsyncDatabase[MongoDocument],
session: AsyncClientSession | None = None,
) -> None:
self.repository = AsyncMongoRepository[Item].from_spec(
binding=self,
database=database,
session=session,
)

async def get_all(
self,
search: str | None = None,
filters: list[FilterEntity] | None = None,
sort: list[SortEntity] | None = None,
pagination: Pagination | None = None,
) -> PaginatedResponse[Item]:
return await self.repository.get_all(
search=search,
filters=filters,
sort=sort,
pagination=pagination,
)

async def get_by_id(self, entity_id: EntityId, /) -> Item | None:
return await self.repository.get_by_id(entity_id)

async def save(self, entity: Item, /) -> None:
await self.repository.save(entity)

async def update(self, entity: Item, /) -> None:
await self.repository.update(entity)

async def remove(self, entity: Item, /) -> None:
await self.repository.remove(entity)
33 changes: 29 additions & 4 deletions app/infrastructure/sql/containers.py
Original file line number Diff line number Diff line change
@@ -1,15 +1,15 @@
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import selectinload
from sqlalchemy.sql.base import ExecutableOption

from app.domain.containers.entities import Container
from app.domain.containers.repository import ContainerRepositoryProtocol
from app.infrastructure.sql.tables import OrmContainer, OrmNode
from cleanstack import EntityId
from cleanstack.sql import AsyncSQLRepository


class ContainerSQLRepository(AsyncSQLRepository[Container, OrmContainer]):
domain_entity_type = Container
orm_model_type = OrmContainer

class ContainerSQLAdapter(AsyncSQLRepository[Container, OrmContainer]):
def to_database_entity(self, entity: Container) -> OrmContainer:
return OrmContainer(
id=entity.id,
Expand All @@ -21,3 +21,28 @@ def to_database_entity(self, entity: Container) -> OrmContainer:
def load_options(self) -> list[ExecutableOption]:
# SELECT * FROM node WHERE container_id IN (...);
return [selectinload(OrmContainer.nodes)]


class ContainerSQLRepository(ContainerRepositoryProtocol):
domain_entity_type = Container
orm_model_type = OrmContainer
searchable_fields = ()

def __init__(self, session: AsyncSession) -> None:
self.session = session
self.repository = ContainerSQLAdapter.from_spec(
binding=self,
session=session,
)

async def get_by_id(self, entity_id: EntityId, /) -> Container | None:
return await self.repository.get_by_id(entity_id)

async def save(self, entity: Container, /) -> None:
await self.repository.save(entity)

async def update(self, entity: Container, /) -> None:
await self.repository.update(entity)

async def remove(self, entity: Container, /) -> None:
await self.repository.remove(entity)
39 changes: 38 additions & 1 deletion app/infrastructure/sql/items.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,46 @@
from sqlalchemy.ext.asyncio import AsyncSession

from app.domain.items.entities import Item
from app.domain.items.repository import ItemRepositoryProtocol
from app.infrastructure.sql.tables import OrmItem
from cleanstack import EntityId, FilterEntity, PaginatedResponse, Pagination, SortEntity
from cleanstack.sql.asynchronous.repository import AsyncSQLRepository


class ItemSQLRepository(AsyncSQLRepository[Item, OrmItem]):
class ItemSQLRepository(ItemRepositoryProtocol):
domain_entity_type = Item
orm_model_type = OrmItem
searchable_fields = ("string_field",)

def __init__(self, session: AsyncSession) -> None:
self.session = session
self.repository = AsyncSQLRepository[Item, OrmItem].from_spec(
binding=self,
session=session,
)

async def get_all(
self,
search: str | None = None,
filters: list[FilterEntity] | None = None,
sort: list[SortEntity] | None = None,
pagination: Pagination | None = None,
) -> PaginatedResponse[Item]:
return await self.repository.get_all(
search=search,
filters=filters,
sort=sort,
pagination=pagination,
)

async def get_by_id(self, entity_id: EntityId, /) -> Item | None:
return await self.repository.get_by_id(entity_id)

async def save(self, entity: Item, /) -> None:
await self.repository.save(entity)

async def update(self, entity: Item, /) -> None:
await self.repository.update(entity)

async def remove(self, entity: Item, /) -> None:
await self.repository.remove(entity)
7 changes: 5 additions & 2 deletions cleanstack/factories/asynchronous.py
Original file line number Diff line number Diff line change
@@ -1,12 +1,15 @@
from abc import ABC, abstractmethod
from collections.abc import AsyncIterator
from contextlib import asynccontextmanager
from typing import Any
from typing import Any, Protocol

from app.domain.protocols import AsyncRepositoryProtocol
from cleanstack.entities.base import BaseEntity


class AsyncRepositoryProtocol[T: BaseEntity](Protocol):
async def save(self, entity: T, /) -> None: ...


class BaseFactory[T: BaseEntity](ABC):
async def create_one(self, **kwargs: Any) -> T: # noqa: ANN401
entity = self.build(**kwargs)
Expand Down
7 changes: 5 additions & 2 deletions cleanstack/factories/synchronous.py
Original file line number Diff line number Diff line change
@@ -1,12 +1,15 @@
from abc import ABC, abstractmethod
from collections.abc import Iterator
from contextlib import contextmanager
from typing import Any
from typing import Any, Protocol

from app.domain.protocols import SyncRepositoryProtocol
from cleanstack.entities.base import BaseEntity


class SyncRepositoryProtocol[T: BaseEntity](Protocol):
def save(self, entity: T, /) -> None: ...


class BaseFactory[T: BaseEntity](ABC):
def create_one(self, **kwargs: Any) -> T: # noqa: ANN401
entity = self.build(**kwargs)
Expand Down
23 changes: 22 additions & 1 deletion cleanstack/mongo/asynchronous/repository.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,19 +9,40 @@
Pagination,
SortEntity,
)
from cleanstack.mongo.mixin import MongoMixin
from cleanstack.mongo.mixin import MongoBinding, MongoMixin
from cleanstack.mongo.types import MongoDocument


class AsyncMongoRepository[T: BaseEntity](MongoMixin[T]):
def __init__(
self,
domain_entity_type: type[T],
collection_name: str,
searchable_fields: tuple[str, ...],
database: AsyncDatabase[MongoDocument],
session: AsyncClientSession | None = None,
) -> None:
self.domain_entity_type = domain_entity_type
self.collection_name = collection_name
self.searchable_fields = searchable_fields
self.collection = database[self.collection_name]
self.session = session

@classmethod
def from_spec(
cls,
binding: MongoBinding[T],
database: AsyncDatabase[MongoDocument],
session: AsyncClientSession | None = None,
) -> AsyncMongoRepository[T]:
return cls(
domain_entity_type=binding.domain_entity_type,
collection_name=binding.collection_name,
searchable_fields=binding.searchable_fields,
database=database,
session=session,
)

async def get_all(
self,
search: str | None = None,
Expand Down
15 changes: 12 additions & 3 deletions cleanstack/mongo/mixin.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
from typing import cast
from typing import Protocol

from pydantic import BaseModel
from pydantic.fields import ComputedFieldInfo, FieldInfo
Expand All @@ -16,6 +16,15 @@
from cleanstack.utils import convert_filter_value_generic


class MongoBinding[T: BaseEntity](Protocol):
@property
def domain_entity_type(self) -> type[T]: ...
@property
def collection_name(self) -> str: ...
@property
def searchable_fields(self) -> tuple[str, ...]: ...


class Pipeline(BaseModel):
data: list[MongoDocument]
count: list[MongoDocument]
Expand Down Expand Up @@ -108,10 +117,10 @@ def build_pipeline(
def _get_field(self, field: str, /) -> FieldInfo | ComputedFieldInfo:
field_info = self.domain_entity_type.model_fields.get(field)
if field_info:
return cast(FieldInfo, field_info)
return field_info # ty: ignore[unsound-return-statement]

computed_field_info = self.domain_entity_type.model_computed_fields.get(field)
if computed_field_info:
return cast(ComputedFieldInfo, computed_field_info)
return computed_field_info # ty: ignore[unsound-return-statement]

raise InvalidFieldError("Invalid field")
Loading