chore: minor fixes and formatting

This commit is contained in:
2026-08-12 19:52:30 +07:00
parent a011b2724b
commit ca276cdd61
48 changed files with 394 additions and 354 deletions

View File

@@ -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",
]

View File

@@ -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

View File

@@ -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)

View File

@@ -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()

View File

@@ -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

View File

@@ -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