- 后端: 新增 dislike_count 字段(模型/Schema/数据库迁移)
- 后端: 新增 POST /api/jokes/{id}/dislike 端点
- 后端: 管理后台统计新增 total_dislikes
- 前端: 新增 useVote 组合函数(localStorage 持久化防重复)
- 前端: JokeCard 列表卡片新增 👍/👎 可点击投票按钮
- 前端: 详情页 ❤️ 改为 👍 赞一下 / 👎 踩一脚 双按钮
236 lines
8.2 KiB
Python
236 lines
8.2 KiB
Python
import json
|
||
import random
|
||
|
||
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"),
|
||
db: Session = Depends(get_db),
|
||
):
|
||
"""获取笑话列表(仅返回已审核通过的笑话)"""
|
||
query = db.query(Joke).filter(Joke.status == "approved")
|
||
|
||
# 支持多选过滤:逗号分隔的 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("/{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)
|
||
|
||
|
||
@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)
|
||
|
||
|
||
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) |