Files
joke/api/app/routers/jokes.py
T
bwstudio cab4d6a9bd refactor(web): 前台代码全面改进
- 搜索逻辑改为后端关键词搜索(修复前端100条限制)
- 合并重复的 handleSelectType/handleDrawerSelectType
- 提取 getLevelText 到公共工具 utils/joke.js
- 删除废弃组件 Header.vue
- 删除重复 API searchJokes
- 修复 AppHeader 滚动事件监听未清理
- 统一 Rightbar URL 参数为 type_ids/crowd_ids
- 清除调试 console.log 代码
- 清理未使用导入 watchEffect
- onMounted 请求并行化(Promise.allSettled)
2026-06-12 08:58:23 +08:00

285 lines
10 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 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)