from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.sql.expression import func from config import PAGE_SIZE from db.models import Product class ProductRepository: async def delete_product(self, session: AsyncSession, product_id: int) -> Product | None: product = await self.get_product_by_id(session, product_id) if not product: return None await session.delete(product) await session.commit() return product async def add_product( self, session: AsyncSession, *, category_id: int | None, name: str, description: str, price: int, ) -> Product: product = Product( category_id=category_id, name=name, description=description, price=price, ) session.add(product) await session.commit() await session.refresh(product) return product async def get_product_by_id( 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, include_hidden: bool = False, ) -> list[Product]: 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) return list(result) async def add_product_file_id( self, session: AsyncSession, product_id: int, *, file_id: str ) -> None: product = await self.get_product_by_id(session, product_id) product.file_id = file_id await session.commit() async def clear_product_photo(self, session: AsyncSession, product_id: int) -> Product | None: product = await self.get_product_by_id(session, product_id) if not product: return None product.file_id = None product.img_path = None await session.commit() return product async def search( self, session: AsyncSession, query: str, *, limit: int = PAGE_SIZE + 1, offset: int = 0, 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 = ( stmt.order_by(func.ts_rank(Product.search_vector, ts_query).desc()) .limit(limit) .offset(offset) ) res = await session.scalars(stmt) return list(res) 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 await session.commit() return product async def update_product_description_by_id( self, session: AsyncSession, product_id: int, description: str ): product = await self.get_product_by_id(session, product_id=product_id) product.description = description await session.commit() return product 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