219 lines
9.2 KiB
Python
219 lines
9.2 KiB
Python
from uuid import uuid4
|
|
|
|
from fastapi import APIRouter, HTTPException, status
|
|
from sqlmodel import Session, select
|
|
|
|
from app.access import product_barcode_accessible, user_can_access_dish
|
|
from app.deps import CurrentUserDep, SessionDep
|
|
from app.models import Dish, DishIngredient, DishShare, Product, User
|
|
from app.schemas import DishCreate, DishIngredientRead, DishRead, DishShareCreate, DishShareRead
|
|
|
|
router = APIRouter(prefix="/api/dishes", tags=["dishes"])
|
|
|
|
NUTRIENT_FIELDS = ["calories", "carbs", "protein", "fat", "sugar", "fiber", "saturated_fat", "salt"]
|
|
|
|
|
|
def _compute_nutrition_per_100g(session: Session, ingredients: list, user_id: int) -> tuple[dict, float]:
|
|
total_weight = sum(ingredient.amount_g for ingredient in ingredients)
|
|
if total_weight <= 0:
|
|
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Zutaten ergeben insgesamt 0g")
|
|
|
|
totals = dict.fromkeys(NUTRIENT_FIELDS, 0.0)
|
|
for ingredient in ingredients:
|
|
product = session.exec(select(Product).where(Product.barcode == ingredient.barcode)).first()
|
|
if product is None or not product_barcode_accessible(session, ingredient.barcode, user_id):
|
|
raise HTTPException(
|
|
status_code=status.HTTP_404_NOT_FOUND,
|
|
detail=f"Zutat mit Barcode {ingredient.barcode} nicht gefunden",
|
|
)
|
|
factor = ingredient.amount_g / 100.0
|
|
for field in NUTRIENT_FIELDS:
|
|
totals[field] += getattr(product, field) * factor
|
|
|
|
per_100g = {field: totals[field] / total_weight * 100 for field in NUTRIENT_FIELDS}
|
|
return per_100g, total_weight
|
|
|
|
|
|
def _get_owned_dish(session: Session, dish_id: int, user_id: int) -> Dish:
|
|
dish = session.get(Dish, dish_id)
|
|
if dish is None or dish.user_id != user_id:
|
|
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Gericht nicht gefunden")
|
|
return dish
|
|
|
|
|
|
def _get_accessible_dish(session: Session, dish_id: int, user_id: int) -> Dish:
|
|
dish = session.get(Dish, dish_id)
|
|
if dish is None or not user_can_access_dish(session, dish, user_id):
|
|
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Gericht nicht gefunden")
|
|
return dish
|
|
|
|
|
|
def _to_dish_read(session: Session, dish: Dish, current_user_id: int) -> DishRead:
|
|
product = session.exec(select(Product).where(Product.barcode == dish.product_barcode)).first()
|
|
ingredient_rows = session.exec(select(DishIngredient).where(DishIngredient.dish_id == dish.id)).all()
|
|
owner = session.get(User, dish.user_id)
|
|
|
|
ingredients = []
|
|
for row in ingredient_rows:
|
|
ingredient_product = session.exec(select(Product).where(Product.barcode == row.product_barcode)).first()
|
|
ingredients.append(
|
|
DishIngredientRead(
|
|
barcode=row.product_barcode,
|
|
name=ingredient_product.name if ingredient_product else "Unbekannt",
|
|
amount_g=row.amount_g,
|
|
**{
|
|
field: (getattr(ingredient_product, field) if ingredient_product else 0.0)
|
|
for field in NUTRIENT_FIELDS
|
|
},
|
|
)
|
|
)
|
|
|
|
return DishRead(
|
|
id=dish.id,
|
|
name=dish.name,
|
|
instructions=dish.instructions,
|
|
barcode=dish.product_barcode,
|
|
total_weight_g=dish.total_weight_g,
|
|
ingredients=ingredients,
|
|
is_owner=dish.user_id == current_user_id,
|
|
owner_username=owner.username if owner else "?",
|
|
**{field: getattr(product, field) for field in NUTRIENT_FIELDS},
|
|
)
|
|
|
|
|
|
def _replace_ingredients(session: Session, dish: Dish, ingredients_in: list) -> None:
|
|
existing = session.exec(select(DishIngredient).where(DishIngredient.dish_id == dish.id)).all()
|
|
for row in existing:
|
|
session.delete(row)
|
|
session.flush()
|
|
for ingredient in ingredients_in:
|
|
session.add(DishIngredient(dish_id=dish.id, product_barcode=ingredient.barcode, amount_g=ingredient.amount_g))
|
|
|
|
|
|
@router.post("", response_model=DishRead, status_code=status.HTTP_201_CREATED)
|
|
def create_dish(dish_in: DishCreate, current_user: CurrentUserDep, session: SessionDep):
|
|
per_100g, total_weight = _compute_nutrition_per_100g(session, dish_in.ingredients, current_user.id)
|
|
|
|
barcode = f"dish-{uuid4().hex[:12]}"
|
|
product = Product(barcode=barcode, name=dish_in.name, **per_100g)
|
|
session.add(product)
|
|
|
|
dish = Dish(
|
|
user_id=current_user.id,
|
|
product_barcode=barcode,
|
|
name=dish_in.name,
|
|
instructions=dish_in.instructions,
|
|
total_weight_g=total_weight,
|
|
)
|
|
session.add(dish)
|
|
session.flush()
|
|
|
|
_replace_ingredients(session, dish, dish_in.ingredients)
|
|
|
|
session.commit()
|
|
session.refresh(dish)
|
|
return _to_dish_read(session, dish, current_user.id)
|
|
|
|
|
|
@router.get("", response_model=list[DishRead])
|
|
def list_dishes(current_user: CurrentUserDep, session: SessionDep):
|
|
owned = session.exec(select(Dish).where(Dish.user_id == current_user.id)).all()
|
|
shared_dish_ids = session.exec(
|
|
select(DishShare.dish_id).where(DishShare.shared_with_user_id == current_user.id)
|
|
).all()
|
|
shared = session.exec(select(Dish).where(Dish.id.in_(shared_dish_ids))).all() if shared_dish_ids else []
|
|
dishes = sorted(owned + shared, key=lambda dish: dish.name.lower())
|
|
return [_to_dish_read(session, dish, current_user.id) for dish in dishes]
|
|
|
|
|
|
@router.get("/{dish_id}", response_model=DishRead)
|
|
def get_dish(dish_id: int, current_user: CurrentUserDep, session: SessionDep):
|
|
dish = _get_accessible_dish(session, dish_id, current_user.id)
|
|
return _to_dish_read(session, dish, current_user.id)
|
|
|
|
|
|
@router.put("/{dish_id}", response_model=DishRead)
|
|
def update_dish(dish_id: int, dish_in: DishCreate, current_user: CurrentUserDep, session: SessionDep):
|
|
dish = _get_owned_dish(session, dish_id, current_user.id)
|
|
per_100g, total_weight = _compute_nutrition_per_100g(session, dish_in.ingredients, current_user.id)
|
|
|
|
product = session.exec(select(Product).where(Product.barcode == dish.product_barcode)).first()
|
|
product.name = dish_in.name
|
|
for field in NUTRIENT_FIELDS:
|
|
setattr(product, field, per_100g[field])
|
|
session.add(product)
|
|
|
|
dish.name = dish_in.name
|
|
dish.instructions = dish_in.instructions
|
|
dish.total_weight_g = total_weight
|
|
session.add(dish)
|
|
|
|
_replace_ingredients(session, dish, dish_in.ingredients)
|
|
|
|
session.commit()
|
|
session.refresh(dish)
|
|
return _to_dish_read(session, dish, current_user.id)
|
|
|
|
|
|
@router.delete("/{dish_id}", status_code=status.HTTP_204_NO_CONTENT)
|
|
def delete_dish(dish_id: int, current_user: CurrentUserDep, session: SessionDep):
|
|
dish = _get_owned_dish(session, dish_id, current_user.id)
|
|
|
|
ingredient_rows = session.exec(select(DishIngredient).where(DishIngredient.dish_id == dish.id)).all()
|
|
for row in ingredient_rows:
|
|
session.delete(row)
|
|
|
|
share_rows = session.exec(select(DishShare).where(DishShare.dish_id == dish.id)).all()
|
|
for row in share_rows:
|
|
session.delete(row)
|
|
|
|
product = session.exec(select(Product).where(Product.barcode == dish.product_barcode)).first()
|
|
session.delete(dish)
|
|
if product:
|
|
session.delete(product)
|
|
session.commit()
|
|
|
|
|
|
@router.get("/{dish_id}/shares", response_model=list[DishShareRead])
|
|
def list_shares(dish_id: int, current_user: CurrentUserDep, session: SessionDep):
|
|
dish = _get_owned_dish(session, dish_id, current_user.id)
|
|
shares = session.exec(select(DishShare).where(DishShare.dish_id == dish.id)).all()
|
|
users = {user.id: user for user in session.exec(select(User)).all()}
|
|
return [DishShareRead(username=users[s.shared_with_user_id].username) for s in shares if s.shared_with_user_id in users]
|
|
|
|
|
|
@router.post("/{dish_id}/shares", response_model=list[DishShareRead], status_code=status.HTTP_201_CREATED)
|
|
def add_share(dish_id: int, share_in: DishShareCreate, current_user: CurrentUserDep, session: SessionDep):
|
|
dish = _get_owned_dish(session, dish_id, current_user.id)
|
|
|
|
target = session.exec(select(User).where(User.username == share_in.username)).first()
|
|
if target is None:
|
|
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Nutzer nicht gefunden")
|
|
if target.id == current_user.id:
|
|
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Kann nicht mit sich selbst geteilt werden")
|
|
|
|
existing = session.exec(
|
|
select(DishShare).where(DishShare.dish_id == dish.id, DishShare.shared_with_user_id == target.id)
|
|
).first()
|
|
if existing is None:
|
|
session.add(DishShare(dish_id=dish.id, shared_with_user_id=target.id))
|
|
session.commit()
|
|
|
|
return list_shares(dish_id, current_user, session)
|
|
|
|
|
|
@router.delete("/{dish_id}/shares/{username}", status_code=status.HTTP_204_NO_CONTENT)
|
|
def remove_share(dish_id: int, username: str, current_user: CurrentUserDep, session: SessionDep):
|
|
dish = _get_owned_dish(session, dish_id, current_user.id)
|
|
|
|
target = session.exec(select(User).where(User.username == username)).first()
|
|
if target is None:
|
|
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Nutzer nicht gefunden")
|
|
|
|
row = session.exec(
|
|
select(DishShare).where(DishShare.dish_id == dish.id, DishShare.shared_with_user_id == target.id)
|
|
).first()
|
|
if row:
|
|
session.delete(row)
|
|
session.commit()
|