Spaces:
Sleeping
Sleeping
| from typing import Any, Optional | |
| from fastapi import APIRouter, Depends, HTTPException, status, Header | |
| from sqlalchemy.orm import Session, joinedload | |
| from datetime import datetime, timezone | |
| from app.db.base import get_db | |
| from app.models.user import User | |
| from app.models.product import Product | |
| from app.models.cart import Cart, CartItem, Coupon | |
| from app.schemas.cart import ( | |
| CartResponse, CartItemCreate, CartItemUpdate, | |
| CartItemResponse, CouponApply, CouponResponse | |
| ) | |
| from app.schemas.product import ProductListItem | |
| from app.core.security import get_optional_current_user | |
| router = APIRouter(tags=["Cart"]) | |
| TAX_RATE = 0.0 # Tax removed (0%) | |
| SHIPPING_FLAT_RATE = 20.0 # Flat shipping rate | |
| FREE_SHIPPING_THRESHOLD = 500.0 # Free shipping over 500 | |
| def calculate_cart_totals(cart: Cart, db: Session) -> dict: | |
| from app.models.settings import StoreSettings | |
| settings = db.query(StoreSettings).first() | |
| global_multiplier = (1 - (settings.global_discount / 100.0)) if settings else 1.0 | |
| subtotal = sum(item.quantity * item.unit_price * global_multiplier for item in cart.items) | |
| # Calculate discount | |
| discount_amount = 0.0 | |
| applied_coupon = None | |
| if cart.coupon and cart.coupon.is_active: | |
| if cart.coupon.expires_at is None or cart.coupon.expires_at > datetime.now(timezone.utc): | |
| discount_amount = subtotal * (cart.coupon.discount_percent / 100.0) | |
| applied_coupon = CouponResponse.model_validate(cart.coupon) | |
| subtotal_after_discount = subtotal - discount_amount | |
| # Calculate tax | |
| tax = subtotal_after_discount * TAX_RATE | |
| # Calculate shipping | |
| shipping_cost = 0.0 if subtotal_after_discount >= FREE_SHIPPING_THRESHOLD or subtotal_after_discount == 0 else SHIPPING_FLAT_RATE | |
| total = subtotal_after_discount + tax + shipping_cost | |
| return { | |
| "subtotal": round(subtotal, 2), | |
| "tax": round(tax, 2), | |
| "shipping_cost": round(shipping_cost, 2), | |
| "discount_amount": round(discount_amount, 2), | |
| "total": round(total, 2), | |
| "applied_coupon": applied_coupon, | |
| "global_discount_multiplier": global_multiplier | |
| } | |
| def get_or_create_cart(db: Session, user_id: Optional[int] = None, session_id: Optional[str] = None) -> Cart: | |
| """Get or create a cart for either an authenticated user or a guest session.""" | |
| if user_id: | |
| cart = db.query(Cart).filter(Cart.user_id == user_id).first() | |
| if not cart: | |
| cart = Cart(user_id=user_id) | |
| db.add(cart) | |
| db.commit() | |
| db.refresh(cart) | |
| return cart | |
| elif session_id: | |
| cart = db.query(Cart).filter(Cart.session_id == session_id).first() | |
| if not cart: | |
| cart = Cart(session_id=session_id) | |
| db.add(cart) | |
| db.commit() | |
| db.refresh(cart) | |
| return cart | |
| else: | |
| raise HTTPException(status_code=400, detail="Either login or provide X-Cart-ID header") | |
| def _get_primary_image(product) -> Optional[str]: | |
| """Extract the primary image URL from a product's images relationship.""" | |
| if not product: | |
| return None | |
| if product.images: | |
| sorted_imgs = sorted(product.images, key=lambda x: x.sort_order) | |
| return sorted_imgs[0].image_url if sorted_imgs else None | |
| return None | |
| def get_cart( | |
| current_user_id: Optional[int] = Depends(get_optional_current_user), | |
| x_cart_id: Optional[str] = Header(None), | |
| db: Session = Depends(get_db) | |
| ) -> Any: | |
| """Get the current user's (or guest session's) shopping cart.""" | |
| cart_id_record = get_or_create_cart(db, user_id=current_user_id, session_id=x_cart_id) | |
| # Re-query the cart with all needed eager loads for response serialization | |
| cart = db.query(Cart).options( | |
| joinedload(Cart.items) | |
| .joinedload(CartItem.product) | |
| .joinedload(Product.images), | |
| joinedload(Cart.items) | |
| .joinedload(CartItem.product) | |
| .joinedload(Product.category), | |
| joinedload(Cart.coupon) | |
| ).filter(Cart.id == cart_id_record.id).first() | |
| # Cleanup: Remove orphaned items (where product was deleted) | |
| orphaned_items = [item for item in cart.items if not item.product] | |
| if orphaned_items: | |
| for item in orphaned_items: | |
| db.delete(item) | |
| db.commit() | |
| db.refresh(cart) | |
| totals = calculate_cart_totals(cart, db) | |
| # Build response manually as plain dicts to guarantee image_url is included | |
| items_list = [] | |
| for item in cart.items: | |
| # We checked if item.product exists above, but let's be safe | |
| if not item.product: | |
| continue | |
| primary_image = _get_primary_image(item.product) | |
| multiplier = totals.get("global_discount_multiplier", 1.0) | |
| product_original_price = item.unit_price # Use the stored unit price (handles variants) | |
| product_original_compare = item.product.compare_price if item.product.compare_price else item.unit_price | |
| product_dict = { | |
| "id": item.product.id, | |
| "name_ar": item.product.name_ar, | |
| "name_en": item.product.name_en, | |
| "price": round(product_original_price * multiplier, 2), | |
| "compare_price": round(product_original_compare, 2) if (multiplier < 1.0 or item.product.compare_price) else None, | |
| "stock": item.product.stock, | |
| "category_id": item.product.category_id, | |
| "category": None, | |
| "rating": item.product.rating, | |
| "rating_count": item.product.rating_count, | |
| "is_featured": item.product.is_featured, | |
| "image_url": primary_image, | |
| "created_at": item.product.created_at.isoformat() if item.product.created_at else None, | |
| } | |
| if item.product.category: | |
| product_dict["category"] = { | |
| "id": item.product.category.id, | |
| "name_ar": item.product.category.name_ar, | |
| "name_en": item.product.category.name_en, | |
| "icon": item.product.category.icon, | |
| "sort_order": item.product.category.sort_order, | |
| } | |
| items_list.append({ | |
| "id": item.id, | |
| "cart_id": item.cart_id, | |
| "product_id": item.product_id, | |
| "variant_label": item.variant_label, | |
| "variant_id": item.variant_id, | |
| "quantity": item.quantity, | |
| "unit_price": round(item.unit_price * multiplier, 2), | |
| "total_price": round(item.quantity * item.unit_price * multiplier, 2), | |
| "product": product_dict, | |
| }) | |
| response_value = { | |
| "id": cart.id, | |
| "user_id": cart.user_id, | |
| "items": items_list, | |
| **totals | |
| } | |
| return { | |
| "isSuccess": True, | |
| "value": response_value, | |
| "statusCode": 200 | |
| } | |
| def add_item_to_cart( | |
| item_in: CartItemCreate, | |
| current_user_id: Optional[int] = Depends(get_optional_current_user), | |
| x_cart_id: Optional[str] = Header(None), | |
| db: Session = Depends(get_db) | |
| ) -> Any: | |
| """Add a product to the cart.""" | |
| cart = get_or_create_cart(db, user_id=current_user_id, session_id=x_cart_id) | |
| # Check product availability | |
| product = db.query(Product).filter(Product.id == item_in.product_id).first() | |
| if not product: | |
| raise HTTPException(status_code=404, detail="Product not found") | |
| if not product.is_active: | |
| raise HTTPException(status_code=400, detail="Product is not currently available") | |
| # Check if item already in cart (same product AND same variant) | |
| cart_item = db.query(CartItem).filter( | |
| CartItem.cart_id == cart.id, | |
| CartItem.product_id == item_in.product_id, | |
| CartItem.variant_id == item_in.variant_id | |
| ).first() | |
| # Calculate variant price if applicable | |
| final_price = product.price | |
| variant_label = item_in.variant_label | |
| if item_in.variant_id and product.specs: | |
| variants = product.specs.get("variants") or product.specs.get("options", [{}])[0].get("values") | |
| if variants and isinstance(variants, list): | |
| matched_variant = None | |
| # 1. Try matching by ID string | |
| for v in variants: | |
| if str(v.get("id")) == str(item_in.variant_id): | |
| matched_variant = v | |
| break | |
| # 2. Try matching by index if variant_id is numeric and no ID match found | |
| if not matched_variant and str(item_in.variant_id).isdigit(): | |
| idx = int(item_in.variant_id) | |
| if 0 <= idx < len(variants): | |
| matched_variant = variants[idx] | |
| if matched_variant: | |
| # Update price | |
| if matched_variant.get("price_modifier") is not None: | |
| final_price += float(matched_variant.get("price_modifier", 0)) | |
| elif matched_variant.get("price") is not None: | |
| final_price = float(matched_variant.get("price")) | |
| # Auto-assign label if not provided | |
| if not variant_label: | |
| variant_label = matched_variant.get("name_ar") or matched_variant.get("name_en") or matched_variant.get("label") | |
| if cart_item: | |
| if cart_item.quantity + item_in.quantity > product.stock: | |
| raise HTTPException(status_code=400, detail=f"Not enough stock. Only {product.stock} available.") | |
| cart_item.quantity += item_in.quantity | |
| cart_item.unit_price = final_price | |
| cart_item.variant_label = variant_label | |
| else: | |
| if item_in.quantity > product.stock: | |
| raise HTTPException(status_code=400, detail=f"Not enough stock. Only {product.stock} available.") | |
| cart_item = CartItem( | |
| cart_id=cart.id, | |
| product_id=item_in.product_id, | |
| quantity=item_in.quantity, | |
| unit_price=final_price, | |
| variant_id=item_in.variant_id, | |
| variant_label=variant_label | |
| ) | |
| db.add(cart_item) | |
| db.commit() | |
| return {"isSuccess": True, "value": {"message": "Item added to cart"}, "statusCode": 200} | |
| def update_cart_item( | |
| item_id: int, | |
| item_in: CartItemUpdate, | |
| current_user_id: Optional[int] = Depends(get_optional_current_user), | |
| x_cart_id: Optional[str] = Header(None), | |
| db: Session = Depends(get_db) | |
| ) -> Any: | |
| """Update quantity of an item in the cart.""" | |
| cart = get_or_create_cart(db, user_id=current_user_id, session_id=x_cart_id) | |
| cart_item = db.query(CartItem).filter( | |
| CartItem.id == item_id, | |
| CartItem.cart_id == cart.id | |
| ).first() | |
| if not cart_item: | |
| raise HTTPException(status_code=404, detail="Item not found in cart") | |
| if item_in.quantity > cart_item.product.stock: | |
| raise HTTPException(status_code=400, detail=f"Not enough stock. Only {cart_item.product.stock} available.") | |
| cart_item.quantity = item_in.quantity | |
| db.commit() | |
| return {"isSuccess": True, "value": {"message": "Cart updated"}, "statusCode": 200} | |
| def remove_cart_item( | |
| item_id: int, | |
| current_user_id: Optional[int] = Depends(get_optional_current_user), | |
| x_cart_id: Optional[str] = Header(None), | |
| db: Session = Depends(get_db) | |
| ) -> Any: | |
| """Remove an item from the cart.""" | |
| cart = get_or_create_cart(db, user_id=current_user_id, session_id=x_cart_id) | |
| cart_item = db.query(CartItem).filter( | |
| CartItem.id == item_id, | |
| CartItem.cart_id == cart.id | |
| ).first() | |
| if not cart_item: | |
| raise HTTPException(status_code=404, detail="Item not found in cart") | |
| db.delete(cart_item) | |
| db.commit() | |
| return {"isSuccess": True, "value": {"message": "Item removed from cart"}, "statusCode": 200} | |
| def apply_coupon( | |
| coupon_in: CouponApply, | |
| current_user_id: Optional[int] = Depends(get_optional_current_user), | |
| x_cart_id: Optional[str] = Header(None), | |
| db: Session = Depends(get_db) | |
| ) -> Any: | |
| """Apply a discount coupon to the cart.""" | |
| cart = get_or_create_cart(db, user_id=current_user_id, session_id=x_cart_id) | |
| coupon = db.query(Coupon).filter(Coupon.code == coupon_in.code.upper()).first() | |
| if not coupon or not coupon.is_active: | |
| raise HTTPException(status_code=400, detail="Invalid coupon code") | |
| if coupon.expires_at and coupon.expires_at < datetime.now(timezone.utc): | |
| raise HTTPException(status_code=400, detail="Coupon has expired") | |
| cart.coupon_id = coupon.id | |
| db.commit() | |
| return {"isSuccess": True, "value": {"message": f"Coupon {coupon.code} applied successfully"}, "statusCode": 200} | |
| def clear_cart( | |
| current_user_id: Optional[int] = Depends(get_optional_current_user), | |
| x_cart_id: Optional[str] = Header(None), | |
| db: Session = Depends(get_db) | |
| ) -> Any: | |
| """Remove all items from the cart and detach coupons.""" | |
| cart = get_or_create_cart(db, user_id=current_user_id, session_id=x_cart_id) | |
| db.query(CartItem).filter(CartItem.cart_id == cart.id).delete() | |
| cart.coupon_id = None | |
| db.commit() | |
| return {"isSuccess": True, "value": {"message": "Cart cleared"}, "statusCode": 200} | |