import json import random from datetime import datetime from fastapi import APIRouter, Depends, HTTPException, Query from sqlalchemy import func, text 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 _batch_load_categories(db: Session, jokes: list[Joke]) -> tuple[dict, dict]: """批量加载类型和人群名称,避免 N+1 查询""" # 收集所有需要的 id all_type_ids = set() all_crowd_ids = set() for joke in jokes: all_type_ids.update(_parse_ids(joke.type_ids)) all_crowd_ids.update(_parse_ids(joke.crowd_ids)) # 批量查询 type_map = {} if all_type_ids: types = db.query(JokeType).filter(JokeType.id.in_(all_type_ids)).all() type_map = {t.id: t.name for t in types} crowd_map = {} if all_crowd_ids: crowds = db.query(JokeCrowd).filter(JokeCrowd.id.in_(all_crowd_ids)).all() crowd_map = {c.id: c.name for c in crowds} return type_map, crowd_map def joke_to_response(joke: Joke, type_map: dict = None, crowd_map: dict = None) -> JokeResponse: """Convert Joke model to JokeResponse schema. Args: type_map: 预加载的类型 id→name 映射 crowd_map: 预加载的人群 id→name 映射 """ ids = _parse_ids(joke.type_ids) crowd_ids = _parse_ids(joke.crowd_ids) if type_map is not None: type_names = [type_map[i] for i in ids if i in type_map] else: type_names = [] if crowd_map is not None: crowd_names = [crowd_map[i] for i in crowd_ids if i in crowd_map] else: crowd_names = [] 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, dislike_count=joke.dislike_count, created_at=joke.created_at, updated_at=joke.updated_at, type_names=type_names, crowd_names=crowd_names, polished_content=joke.polished_content, ai_score=joke.ai_score, ai_level=joke.ai_level, ) def _build_json_filter(column, target_ids: set[int]) -> list: """构建精确的 JSON 数组匹配过滤器(SQLite)""" filters = [] for tid in target_ids: # 匹配 [n,...] 或 [...,n,...] 或 [...,n] # 使用 JSON_EACH 函数检查 json_str = json.dumps(tid) filters.append( func.json_extract(column, '$').cast(type(target_ids)).contains(tid) ) return filters @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"), keyword: str | None = Query(None, description="关键词搜索(标题+内容)"), db: Session = Depends(get_db), ): """获取笑话列表(仅返回已审核通过的笑话)""" query = db.query(Joke).filter(Joke.status == "approved") # 关键词搜索:标题和内容模糊匹配 if keyword: kw = f"%{keyword}%" query = query.filter( (Joke.title.like(kw)) | (Joke.content.like(kw)) ) # 支持多选过滤:逗号分隔的 ID # 使用自定义 JSON 匹配函数避免 LIKE %N% 的错误匹配 if type_ids: filter_set = {int(x.strip()) for x in type_ids.split(",") if x.strip().isdigit()} if filter_set: # 使用文本匹配,确保精确匹配 JSON 数组中的 id # 匹配模式:[1,2,3] 或 [1] 或带尾部逗号的 type_filters = [] for tid in filter_set: # 构建精确匹配的正则:在逗号、引号、方括号包围中的数字 import re pattern = rf'["\s\[,]{tid}[",\s\]]' type_filters.append(Joke.type_ids.regexp_match(pattern)) # 只要匹配任一类型即可 type_filter = type_filters[0] for f in type_filters[1:]: type_filter = type_filter | f # 还需要兼容旧的单值 type_id 字段 old_filter = Joke.type_id.in_(filter_set) query = query.filter(old_filter | type_filter) if crowd_ids: filter_set = {int(x.strip()) for x in crowd_ids.split(",") if x.strip().isdigit()} if filter_set: import re crowd_filters = [] for cid in filter_set: pattern = rf'["\s\[,]{cid}[",\s\]]' crowd_filters.append(Joke.crowd_ids.regexp_match(pattern)) crowd_filter = crowd_filters[0] for f in crowd_filters[1:]: crowd_filter = crowd_filter | f old_filter = Joke.crowd_id.in_(filter_set) query = query.filter(old_filter | crowd_filter) total = query.count() offset = (page - 1) * page_size jokes = query.order_by(Joke.created_at.desc()).offset(offset).limit(page_size).all() # 批量预加载类型和人群名称 type_map, crowd_map = _batch_load_categories(db, jokes) return PaginatedJokeResponse( items=[joke_to_response(j, type_map, crowd_map) for j in jokes], total=total, page=page, page_size=page_size, ) @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: type_map, crowd_map = _batch_load_categories(db, [joke]) return joke_to_response(joke, type_map, crowd_map) # 兜底:全表随机 joke = db.query(Joke).filter(Joke.status == "approved").order_by(func.random()).first() if not joke: raise HTTPException(status_code=404, detail="暂无笑话") type_map, crowd_map = _batch_load_categories(db, [joke]) return joke_to_response(joke, type_map, crowd_map) @router.get("/hot-monthly") def hot_monthly_jokes(db: Session = Depends(get_db)): """当月热门笑话:本月浏览量最高的前 10 条""" now = datetime.now() start_of_month = now.replace(day=1, hour=0, minute=0, second=0, microsecond=0) jokes = ( db.query(Joke) .filter( Joke.status == "approved", Joke.created_at >= start_of_month, ) .order_by(Joke.view_count.desc()) .limit(10) .all() ) type_map, crowd_map = _batch_load_categories(db, jokes) return [joke_to_response(j, type_map, crowd_map) for j in jokes] @router.get("/stats") def public_stats(db: Session = Depends(get_db)): """公开统计数据:笑话总数、审核/待审/拒绝数量、总浏览/赞/踩""" total = db.query(Joke).count() approved = db.query(Joke).filter(Joke.status == "approved").count() pending = db.query(Joke).filter(Joke.status == "pending").count() rejected = db.query(Joke).filter(Joke.status == "rejected").count() total_views = db.query(Joke).with_entities(func.sum(Joke.view_count)).scalar() or 0 total_likes = db.query(Joke).with_entities(func.sum(Joke.like_count)).scalar() or 0 total_dislikes = db.query(Joke).with_entities(func.sum(Joke.dislike_count)).scalar() or 0 return { "total_jokes": total, "approved_jokes": approved, "pending_jokes": pending, "rejected_jokes": rejected, "total_views": total_views, "total_likes": total_likes, "total_dislikes": total_dislikes, } @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) # 批量加载类型和人群名称 type_map, crowd_map = _batch_load_categories(db, [joke]) return joke_to_response(joke, type_map, crowd_map) def _vote_joke(joke_id: int, field: str, db: Session): """通用投票逻辑:点赞或点踩(仅允许已审核通过的笑话)""" joke = db.query(Joke).filter( Joke.id == joke_id, Joke.status == "approved" ).first() if not joke: raise HTTPException(status_code=404, detail="笑话不存在") # 使用原子更新避免并发竞态 update_data = {getattr(Joke, field): getattr(Joke, field) + 1} db.query(Joke).filter(Joke.id == joke_id).update(update_data) db.commit() db.refresh(joke) return {"message": "操作成功", field: getattr(joke, field)} @router.post("/{joke_id}/like") def like_joke(joke_id: int, db: Session = Depends(get_db)): """为笑话点赞(仅允许已审核通过的笑话)""" return _vote_joke(joke_id, "like_count", db) @router.post("/{joke_id}/dislike") def dislike_joke(joke_id: int, db: Session = Depends(get_db)): """为笑话点踩(仅允许已审核通过的笑话)""" return _vote_joke(joke_id, "dislike_count", db)