diff --git a/routers/record.py b/routers/record.py index de3e071..0d2d791 100644 --- a/routers/record.py +++ b/routers/record.py @@ -14,18 +14,14 @@ def create(record_in: RecordCreate, db: Session = Depends(get_db)): return record_service.create_record(db, record_in) -@router.get("", response_model=List[RecordResponse], summary="获取所有记录") -def list_all(skip: int = 0, limit: int = 20, db: Session = Depends(get_db)): - return record_service.get_all_records(db, skip, limit) - - @router.get("/page", response_model=RecordPageResponse, summary="分页获取所有记录") def list_all_by_page( + user_id: int, page: int = Query(1, ge=1, description="页码,从1开始"), size: int = Query(20, ge=1, le=100, description="每页条数,最大100"), db: Session = Depends(get_db), ): - return record_service.get_all_records_by_page(db, page, size) + return record_service.get_all_records_by_page(db, user_id, page, size) @router.get("/{record_id}", response_model=RecordResponse, summary="获取记录详情") diff --git a/schemas/record.py b/schemas/record.py index bea0070..b9ccb18 100644 --- a/schemas/record.py +++ b/schemas/record.py @@ -27,6 +27,7 @@ class RecordResponse(BaseModel): image_list: List[str] video_list: List[str] like_count: int + liked: bool = False comment_count: int comments: List[CommentResponse] = [] create_time: datetime diff --git a/service/record_service.py b/service/record_service.py index f4efec6..b9144c7 100644 --- a/service/record_service.py +++ b/service/record_service.py @@ -1,29 +1,18 @@ -from sqlalchemy import select, desc, func +from sqlalchemy import select, func from sqlalchemy.orm import Session, selectinload -from models.family import Records, Comments +from models.family import Records, Comments, Likes from schemas.record import RecordCreate, RecordUpdate, RecordResponse, RecordPageResponse -def get_all_records(db: Session, skip: int = 0, limit: int = 20) -> list[RecordResponse]: - result = db.execute( +def get_all_records_by_page(db: Session, user_id: int, page: int = 1, size: int = 20) -> RecordPageResponse: + base_query = ( select(Records) .options( selectinload(Records.user), selectinload(Records.comments).selectinload(Comments.user), ) - .order_by(desc(Records.create_time)) - .offset(skip) - .limit(limit) - ) - records = result.scalars().all() - return [RecordResponse.model_validate(r) for r in records] - - -def get_all_records_by_page(db: Session, page: int = 1, size: int = 20) -> RecordPageResponse: - base_query = select(Records).options( - selectinload(Records.user), - selectinload(Records.comments).selectinload(Comments.user), + .order_by(Records.create_time.desc()) ) # 总数 @@ -32,15 +21,34 @@ def get_all_records_by_page(db: Session, page: int = 1, size: int = 20) -> Recor ) # 分页数据 - records = db.execute( - base_query - .order_by(Records.create_time.desc()) - .offset((page - 1) * size) - .limit(size) - ).scalars().all() + records = ( + db.execute(base_query.offset((page - 1) * size).limit(size)) + .scalars() + .all() + ) + + # 批量判断当前用户点赞状态 + liked_record_ids: set[int] = set() + + if user_id and records: + liked_record_ids = set( + db.scalars( + select(Likes.record_id).where( + Likes.record_id.in_([r.id for r in records]), + Likes.user_id == user_id, + ) + ).all() + ) + + # 组装响应 + response_records = [] + for r in records: + resp = RecordResponse.model_validate(r, from_attributes=True) + resp.liked = r.id in liked_record_ids + response_records.append(resp) return RecordPageResponse( - records=[RecordResponse.model_validate(r) for r in records], + records=response_records, total=total, page=page, size=size,