Files
joke/api/app/routers/jokes.py
T
bwstudio fde86b90d2 feat: 新增点踩功能,列表卡片和详情页支持 👍/👎 投票
- 后端: 新增 dislike_count 字段(模型/Schema/数据库迁移)
- 后端: 新增 POST /api/jokes/{id}/dislike 端点
- 后端: 管理后台统计新增 total_dislikes
- 前端: 新增 useVote 组合函数(localStorage 持久化防重复)
- 前端: JokeCard 列表卡片新增 👍/👎 可点击投票按钮
- 前端: 详情页 ❤️ 改为 👍 赞一下 / 👎 踩一脚 双按钮
2026-06-11 19:17:12 +08:00

236 lines
8.2 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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)