chore: minor fixes and formatting
This commit is contained in:
@@ -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