import json import random from fastapi import APIRouter, Depends, HTTPException, Query from sqlalchemy import func from sqlalchemy.orm import Session from app.database import get_db from app.models.joke import Joke from app.models.category import JokeCrowd, JokeType from app.schemas.joke import JokeResponse, PaginatedJokeResponse router = APIRouter(prefix="/jokes", tags=["笑话"]) def _parse_ids(raw) -> list[int]: if not raw: return [] if isinstance(raw, list): return raw try: return json.loads(raw) except Exception: return [] def _get_type_names(db: Session, ids: list[int]) -> list[str]: if not ids: return [] rows = db.query(JokeType).filter(JokeType.id.in_(ids)).all() return [r.name for r in rows] def _get_crowd_names(db: Session, ids: list[int]) -> list[str]: if not ids: return [] rows = db.query(JokeCrowd).filter(JokeCrowd.id.in_(ids)).all() return [r.name for r in rows] def joke_to_response(joke: Joke, db: Session = None) -> JokeResponse: """Convert Joke model to JokeResponse schema.""" ids = _parse_ids(joke.type_ids) crowd_ids = _parse_ids(joke.crowd_ids) type_names = _get_type_names(db, ids) if db else [] crowd_names = _get_crowd_names(db, crowd_ids) if db else [] return JokeResponse( id=joke.id, title=joke.title, content=joke.content, type_ids=ids, crowd_ids=crowd_ids, status=joke.status, view_count=joke.view_count, like_count=joke.like_count, created_at=joke.created_at, updated_at=joke.updated_at, type_names=type_names, crowd_names=crowd_names, ) @router.get("/", response_model=PaginatedJokeResponse) def list_jokes( page: int = 1, page_size: int = 20, type_ids: str | None = Query(None, description="逗号分隔的类型 ID,如 1,3,5"), crowd_ids: str | None = Query(None, description="逗号分隔的人群 ID,如 2,4"), db: Session = Depends(get_db), ): """获取笑话列表(仅返回已审核通过的笑话)""" query = db.query(Joke).filter(Joke.status == "approved") # 支持多选过滤:逗号分隔的 ID if type_ids: filter_set = {int(x.strip()) for x in type_ids.split(",") if x.strip().isdigit()} if filter_set: # 兼容旧单值字段 type_id old_filter = Joke.type_id.in_(filter_set) # 兼容新数组字段 type_ids(JSON 包含任一) new_filter = [ Joke.type_ids.contains(str(tid)) for tid in filter_set ] combined = old_filter for nf in new_filter: combined = combined | nf query = query.filter(combined) if crowd_ids: filter_set = {int(x.strip()) for x in crowd_ids.split(",") if x.strip().isdigit()} if filter_set: old_filter = Joke.crowd_id.in_(filter_set) new_filter = [ Joke.crowd_ids.contains(str(cid)) for cid in filter_set ] combined = old_filter for nf in new_filter: combined = combined | nf query = query.filter(combined) total = query.count() offset = (page - 1) * page_size jokes = query.order_by(Joke.created_at.desc()).offset(offset).limit(page_size).all() return PaginatedJokeResponse( items=[joke_to_response(j, db) for j in jokes], total=total, page=page, page_size=page_size, ) @router.get("/{joke_id}", response_model=JokeResponse) def get_joke(joke_id: int, db: Session = Depends(get_db)): """获取单条笑话详情""" # 只返回已审核通过的笑话 joke = db.query(Joke).filter( Joke.id == joke_id, Joke.status == "approved" ).first() if not joke: raise HTTPException(status_code=404, detail="笑话不存在") # 使用原子更新避免并发竞态 db.query(Joke).filter(Joke.id == joke_id).update({Joke.view_count: Joke.view_count + 1}) db.commit() # 重新查询获取更新后的数据 db.refresh(joke) return joke_to_response(joke, db) @router.get("/random", response_model=JokeResponse) def get_random_joke(db: Session = Depends(get_db)): """随机获取一条已审核通过的笑话""" # 使用效率更高的方式:随机 ID 取模 max_id = db.query(func.max(Joke.id)).filter(Joke.status == "approved").scalar() if not max_id: raise HTTPException(status_code=404, detail="暂无笑话") # 尝试最多 10 次找到有效笑话 for _ in range(10): random_id = random.randint(1, max_id) joke = db.query(Joke).filter( Joke.id >= random_id, Joke.status == "approved" ).first() if joke: return joke_to_response(joke, db) # 兜底:全表随机 joke = db.query(Joke).filter(Joke.status == "approved").order_by(func.random()).first() if not joke: raise HTTPException(status_code=404, detail="暂无笑话") return joke_to_response(joke, db) @router.post("/{joke_id}/like") def like_joke(joke_id: int, db: Session = Depends(get_db)): """为笑话点赞(仅允许已审核通过的笑话)""" # 先检查笑话是否存在且已审核 joke = db.query(Joke).filter( Joke.id == joke_id, Joke.status == "approved" ).first() if not joke: raise HTTPException(status_code=404, detail="笑话不存在") # 使用原子更新避免并发竞态 db.query(Joke).filter(Joke.id == joke_id).update({Joke.like_count: Joke.like_count + 1}) db.commit() # 获取更新后的值 db.refresh(joke) return {"message": "点赞成功", "like_count": joke.like_count}