chore: minor fixes and formatting
This commit is contained in:
@@ -1,11 +1,11 @@
|
||||
from .orders import OrderRepository
|
||||
from .categories import CategoriesRepository
|
||||
from .products import ProductRepository
|
||||
from .invoices import InvoiceRepository
|
||||
from .orders import OrderRepository
|
||||
from .products import ProductRepository
|
||||
|
||||
__all__ = [
|
||||
"OrderRepository",
|
||||
"CategoriesRepository",
|
||||
"ProductRepository",
|
||||
"InvoiceRepository",
|
||||
"OrderRepository",
|
||||
"ProductRepository",
|
||||
]
|
||||
|
||||
@@ -1,20 +1,17 @@
|
||||
from typing import Optional
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.orm import aliased
|
||||
|
||||
from db.models import Category
|
||||
|
||||
|
||||
class CategoriesRepository:
|
||||
async def get_category_by_id(
|
||||
self, session: AsyncSession, category_id: int
|
||||
) -> Optional[Category]:
|
||||
async def get_category_by_id(self, session: AsyncSession, category_id: int) -> Category | None:
|
||||
stmt = select(Category).where(Category.id == category_id)
|
||||
return await session.scalar(stmt)
|
||||
|
||||
async def get_categories_by_parent_id(
|
||||
self, session: AsyncSession, parent_id: Optional[int] = None
|
||||
self, session: AsyncSession, parent_id: int | None = None
|
||||
) -> list[Category]:
|
||||
stmt = select(Category).where(Category.parent_id == parent_id)
|
||||
result = await session.scalars(stmt)
|
||||
@@ -31,9 +28,7 @@ class CategoriesRepository:
|
||||
parent = aliased(Category)
|
||||
|
||||
cte = cte.union_all(
|
||||
select(parent.id, parent.parent_id, parent.name).join(
|
||||
cte, cte.c.parent_id == parent.id
|
||||
)
|
||||
select(parent.id, parent.parent_id, parent.name).join(cte, cte.c.parent_id == parent.id)
|
||||
)
|
||||
|
||||
stmt = select(cte)
|
||||
@@ -45,7 +40,7 @@ class CategoriesRepository:
|
||||
return list(reversed(rows))
|
||||
|
||||
async def add_category(
|
||||
self, session: AsyncSession, *, name: str, parent_id: Optional[int]
|
||||
self, session: AsyncSession, *, name: str, parent_id: int | None
|
||||
) -> Category:
|
||||
category = Category(name=name, parent_id=parent_id)
|
||||
|
||||
@@ -54,9 +49,7 @@ class CategoriesRepository:
|
||||
|
||||
return category
|
||||
|
||||
async def update_category_name(
|
||||
self, session: AsyncSession, category_id: int, value: str
|
||||
):
|
||||
async def update_category_name(self, session: AsyncSession, category_id: int, value: str):
|
||||
category = await self.get_category_by_id(session, category_id)
|
||||
category.name = value
|
||||
|
||||
|
||||
@@ -1,9 +1,7 @@
|
||||
from typing import Optional
|
||||
|
||||
from sqlalchemy import select
|
||||
from db.models import Invoice
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from db.models import Invoice
|
||||
from db.models.invoices import InvoiceStatus
|
||||
|
||||
|
||||
@@ -14,7 +12,7 @@ class InvoiceRepository:
|
||||
*,
|
||||
amount: int,
|
||||
creator_id: int,
|
||||
inline_message_id: Optional[int] = None,
|
||||
inline_message_id: int | None = None,
|
||||
status: InvoiceStatus = InvoiceStatus.PENDING,
|
||||
) -> Invoice:
|
||||
invoice = Invoice(
|
||||
@@ -28,9 +26,7 @@ class InvoiceRepository:
|
||||
await session.commit()
|
||||
return invoice
|
||||
|
||||
async def get_invoice_by_id(
|
||||
self, session: AsyncSession, invoice_id: int
|
||||
) -> Optional[Invoice]:
|
||||
async def get_invoice_by_id(self, session: AsyncSession, invoice_id: int) -> Invoice | None:
|
||||
stmt = select(Invoice).where(Invoice.id == invoice_id)
|
||||
return await session.scalar(stmt)
|
||||
|
||||
|
||||
@@ -1,26 +1,18 @@
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy import delete, func, select
|
||||
from typing import Optional
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from db.models.orders import Order, OrderItem, OrderStatus
|
||||
|
||||
|
||||
class OrderItemRepository:
|
||||
async def get_item_by_id(
|
||||
self, session: AsyncSession, item_id: int
|
||||
) -> Optional[OrderItem]:
|
||||
async def get_item_by_id(self, session: AsyncSession, item_id: int) -> OrderItem | None:
|
||||
stmt = select(OrderItem).where(OrderItem.id == item_id)
|
||||
return await session.scalar(stmt)
|
||||
|
||||
async def get_items_by_order(
|
||||
self, session: AsyncSession, *, order_id: int, limit: int, offset: int = 0
|
||||
) -> list[OrderItem]:
|
||||
stmt = (
|
||||
select(OrderItem)
|
||||
.where(OrderItem.order_id == order_id)
|
||||
.limit(limit)
|
||||
.offset(offset)
|
||||
)
|
||||
stmt = select(OrderItem).where(OrderItem.order_id == order_id).limit(limit).offset(offset)
|
||||
result = await session.scalars(stmt)
|
||||
|
||||
return list(result)
|
||||
@@ -29,9 +21,7 @@ class OrderItemRepository:
|
||||
stmt = select(func.count(OrderItem.id)).where(OrderItem.order_id == order_id)
|
||||
return await session.scalar(stmt) or 0
|
||||
|
||||
async def get_items_count_by_customer(
|
||||
self, session: AsyncSession, customer: int
|
||||
) -> int:
|
||||
async def get_items_count_by_customer(self, session: AsyncSession, customer: int) -> int:
|
||||
count = await session.scalar(
|
||||
select(func.count(OrderItem.id))
|
||||
.join(Order)
|
||||
@@ -44,7 +34,7 @@ class OrderItemRepository:
|
||||
|
||||
async def get_item_by_order_and_product(
|
||||
self, session: AsyncSession, *, order_id: int, product_id: int
|
||||
) -> Optional[OrderItem]:
|
||||
) -> OrderItem | None:
|
||||
stmt = (
|
||||
select(OrderItem)
|
||||
.where(OrderItem.order_id == order_id)
|
||||
@@ -68,9 +58,7 @@ class OrderItemRepository:
|
||||
product_id: int,
|
||||
quantity: int = 1,
|
||||
) -> OrderItem:
|
||||
order_item = OrderItem(
|
||||
order_id=order_id, product_id=product_id, quantity=quantity
|
||||
)
|
||||
order_item = OrderItem(order_id=order_id, product_id=product_id, quantity=quantity)
|
||||
|
||||
session.add(order_item)
|
||||
await session.commit()
|
||||
|
||||
@@ -1,30 +1,22 @@
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy import delete, select
|
||||
from typing import Optional
|
||||
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.orm import selectinload
|
||||
|
||||
from db.models import Order, OrderStatus, OrderItem
|
||||
from db.models import Order, OrderItem, OrderStatus
|
||||
|
||||
|
||||
class OrderRepository:
|
||||
async def get_order_by_id(
|
||||
self, session: AsyncSession, order_id: int
|
||||
) -> Optional[Order]:
|
||||
async def get_order_by_id(self, session: AsyncSession, order_id: int) -> Order | None:
|
||||
stmt = select(Order).where(Order.id == order_id)
|
||||
return await session.scalar(stmt)
|
||||
|
||||
async def get_orders_by_user(
|
||||
self, session: AsyncSession, customer: int
|
||||
) -> list[Order]:
|
||||
async def get_orders_by_user(self, session: AsyncSession, customer: int) -> list[Order]:
|
||||
stmt = select(Order).where(Order.customer == customer)
|
||||
result = await session.scalars(stmt)
|
||||
|
||||
return list(result)
|
||||
|
||||
async def get_draft_order_by_user(
|
||||
self, session: AsyncSession, customer: int
|
||||
) -> Optional[Order]:
|
||||
async def get_draft_order_by_user(self, session: AsyncSession, customer: int) -> Order | None:
|
||||
stmt = (
|
||||
select(Order)
|
||||
.where(Order.customer == customer)
|
||||
@@ -59,9 +51,7 @@ class OrderRepository:
|
||||
await session.execute(stmt)
|
||||
await session.commit()
|
||||
|
||||
async def update_order_status(
|
||||
self, session: AsyncSession, order_id: int, status: OrderStatus
|
||||
):
|
||||
async def update_order_status(self, session: AsyncSession, order_id: int, status: OrderStatus):
|
||||
order = await self.get_order_by_id(session, order_id)
|
||||
order.status = status
|
||||
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.sql.expression import func
|
||||
from typing import Optional
|
||||
|
||||
from config import PAGE_SIZE
|
||||
from db.models import Product
|
||||
@@ -9,21 +8,28 @@ from db.models import Product
|
||||
|
||||
class ProductRepository:
|
||||
async def get_product_by_id(
|
||||
self, session: AsyncSession, product_id: int
|
||||
) -> Optional[Product]:
|
||||
self, session: AsyncSession, product_id: int, *, include_hidden: bool = True
|
||||
) -> Product | None:
|
||||
stmt = select(Product).where(Product.id == product_id)
|
||||
if not include_hidden:
|
||||
stmt = stmt.where(Product.is_hidden.is_(False))
|
||||
return await session.scalar(stmt)
|
||||
|
||||
async def get_product_by_category(
|
||||
self, session: AsyncSession, category_id: int, *, limit: int, offset: int
|
||||
self,
|
||||
session: AsyncSession,
|
||||
category_id: int,
|
||||
*,
|
||||
limit: int,
|
||||
offset: int,
|
||||
include_hidden: bool = False,
|
||||
) -> list[Product]:
|
||||
stmt = (
|
||||
select(Product)
|
||||
.where(Product.category_id == category_id)
|
||||
.order_by(Product.id)
|
||||
.offset(offset)
|
||||
.limit(limit)
|
||||
)
|
||||
stmt = select(Product).where(Product.category_id == category_id)
|
||||
|
||||
if not include_hidden:
|
||||
stmt = stmt.where(Product.is_hidden.is_(False))
|
||||
|
||||
stmt = stmt.order_by(Product.id).offset(offset).limit(limit)
|
||||
|
||||
result = await session.scalars(stmt)
|
||||
|
||||
@@ -44,13 +50,17 @@ class ProductRepository:
|
||||
*,
|
||||
limit: int = PAGE_SIZE + 1,
|
||||
offset: int = 0,
|
||||
) -> list[Optional[Product]]:
|
||||
include_hidden: bool = False,
|
||||
) -> list[Product | None]:
|
||||
ts_query = func.plainto_tsquery("simple", query)
|
||||
|
||||
stmt = select(Product).where(Product.search_vector.op("@@")(ts_query))
|
||||
|
||||
if not include_hidden:
|
||||
stmt = stmt.where(Product.is_hidden.is_(False))
|
||||
|
||||
stmt = (
|
||||
select(Product)
|
||||
.where(Product.search_vector.op("@@")(ts_query))
|
||||
.order_by(func.ts_rank(Product.search_vector, ts_query).desc())
|
||||
stmt.order_by(func.ts_rank(Product.search_vector, ts_query).desc())
|
||||
.limit(limit)
|
||||
.offset(offset)
|
||||
)
|
||||
@@ -59,9 +69,7 @@ class ProductRepository:
|
||||
|
||||
return list(res)
|
||||
|
||||
async def update_product_name_by_id(
|
||||
self, session: AsyncSession, product_id: int, name: str
|
||||
):
|
||||
async def update_product_name_by_id(self, session: AsyncSession, product_id: int, name: str):
|
||||
product = await self.get_product_by_id(session, product_id=product_id)
|
||||
product.name = name
|
||||
|
||||
@@ -77,11 +85,21 @@ class ProductRepository:
|
||||
await session.commit()
|
||||
return product
|
||||
|
||||
async def update_product_price_by_id(
|
||||
self, session: AsyncSession, product_id: int, price: str
|
||||
):
|
||||
async def update_product_price_by_id(self, session: AsyncSession, product_id: int, price: str):
|
||||
product = await self.get_product_by_id(session, product_id=product_id)
|
||||
product.price = price
|
||||
|
||||
await session.commit()
|
||||
return product
|
||||
|
||||
async def toggle_product_hidden_by_id(
|
||||
self, session: AsyncSession, product_id: int
|
||||
) -> Product | None:
|
||||
product = await self.get_product_by_id(session, product_id=product_id)
|
||||
if not product:
|
||||
return None
|
||||
|
||||
product.is_hidden = not product.is_hidden
|
||||
|
||||
await session.commit()
|
||||
return product
|
||||
|
||||
Reference in New Issue
Block a user