优化任务列表性能和本地登录代理

This commit is contained in:
yml2213
2026-08-05 10:03:18 +08:00
parent dff0d7db9e
commit 2fda47c631
13 changed files with 252 additions and 29 deletions
+7
View File
@@ -6,6 +6,7 @@ import uvicorn
from contextlib import asynccontextmanager
from fastapi import FastAPI, Request
from fastapi.middleware.cors import CORSMiddleware
from fastapi.middleware.gzip import GZipMiddleware
from fastapi.staticfiles import StaticFiles
from fastapi.responses import FileResponse
from starlette.middleware.base import BaseHTTPMiddleware
@@ -66,6 +67,9 @@ app.add_middleware(
allow_headers=["*"],
)
# 压缩 JSON 和静态资源响应,避免任务列表轮询反复传输大体积文本。
app.add_middleware(GZipMiddleware, minimum_size=1024)
# 安全响应头中间件
class SecurityHeadersMiddleware(BaseHTTPMiddleware):
@@ -74,6 +78,9 @@ class SecurityHeadersMiddleware(BaseHTTPMiddleware):
response.headers["X-Content-Type-Options"] = "nosniff"
response.headers["X-Frame-Options"] = "DENY"
response.headers["Referrer-Policy"] = "strict-origin-when-cross-origin"
if request.url.path.startswith("/assets/"):
# Vite 产物文件名带 hash,可以长期缓存。
response.headers["Cache-Control"] = "public, max-age=31536000, immutable"
return response
@@ -0,0 +1,47 @@
"""补充任务状态清理索引
Revision ID: 20260805_0015
Revises: 20260728_0014
Create Date: 2026-08-05
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
revision: str = "20260805_0015"
down_revision: Union[str, None] = "20260728_0014"
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
INDEXES = [
("ix_douyu_tasks_status_id", "douyu_tasks", ["status", "id"]),
("ix_huya_tasks_status_id", "huya_tasks", ["status", "id"]),
]
def _has_table(bind, table_name: str) -> bool:
return sa.inspect(bind).has_table(table_name)
def _indexes(bind, table_name: str) -> set[str]:
if not _has_table(bind, table_name):
return set()
return {index["name"] for index in sa.inspect(bind).get_indexes(table_name)}
def upgrade() -> None:
bind = op.get_bind()
for name, table_name, columns in INDEXES:
if name not in _indexes(bind, table_name):
op.create_index(name, table_name, columns)
def downgrade() -> None:
bind = op.get_bind()
for name, table_name, _ in reversed(INDEXES):
if name in _indexes(bind, table_name):
op.drop_index(name, table_name=table_name)
+75 -4
View File
@@ -118,7 +118,73 @@ def _account_out(account: Account) -> DouyuTaskAccountOut:
)
def _task_out(task: DouyuTask) -> DouyuTaskOut:
def _slim_goods(goods: object) -> object:
"""保留兑换图片和列表展示需要的商品小字段。"""
if not isinstance(goods, dict):
return None
return {
key: value
for key, value in goods.items()
if key in {"commodityId", "commodity_id", "commodityName", "name", "webPic", "pic", "score", "status"}
}
def _slim_limited_goods(goods: object) -> list[dict[str, object]]:
"""限兑列表只返回前几个商品名,避免任务列表携带完整原始数组。"""
if not isinstance(goods, list):
return []
result: list[dict[str, object]] = []
for item in goods[:5]:
if not isinstance(item, dict):
continue
result.append({
key: value
for key, value in item.items()
if key in {"commodityId", "commodity_id", "commodityName", "name"}
})
return result
def _sanitize_task_result(result: dict | None, task_type: str, *, include_detail: bool = False) -> dict | None:
"""列表接口剥离原始快照/大数组;详情接口保留完整 result。"""
if not isinstance(result, dict):
return result
if include_detail:
return result
data = dict(result)
for key in (
"raw",
"bind_info",
"before_bind_info",
"bind_candidates",
"cooldown_bind_info",
"after_bind_info",
"activity_bind_snapshot",
"points_query",
"exchange_balance_query",
"query_act_aliases",
"records",
):
data.pop(key, None)
goods = data.get("goods")
if isinstance(goods, list):
data.pop("goods", None)
elif isinstance(goods, dict):
data["goods"] = _slim_goods(goods)
if "limited_goods" in data:
data["limited_goods"] = _slim_limited_goods(data.get("limited_goods"))
if task_type not in {"get_bind_qr", "prepare_esports_bind", "get_esports_bind_qr"}:
data.pop("url", None)
if task_type not in {"create_elite_qr", "create_esports_qr", "create_gold_qr"}:
data.pop("pay_url", None)
return data
def _task_out(task: DouyuTask, *, include_detail: bool = False) -> DouyuTaskOut:
account = task.account
return DouyuTaskOut(
id=task.id,
@@ -130,7 +196,11 @@ def _task_out(task: DouyuTask) -> DouyuTaskOut:
task_type=task.task_type,
status=task.status or "",
message=task.message or "",
result=task.result if isinstance(task.result, dict) else None,
result=_sanitize_task_result(
task.result if isinstance(task.result, dict) else None,
task.task_type or "",
include_detail=include_detail,
),
created_by=task.created_by,
created_at=task.created_at,
finished_at=task.finished_at,
@@ -299,6 +369,7 @@ async def create_task_batch(
@router.get("/tasks", response_model=list[DouyuTaskOut])
def list_tasks(
batch_id: str | None = None,
include_detail: bool = Query(False, description="是否返回完整任务结果(默认否,轮询请保持 false)"),
db: Session = Depends(get_db),
current: User = Depends(require_permission("douyu:task")),
):
@@ -313,7 +384,7 @@ def list_tasks(
if batch_id:
query = query.filter(DouyuTask.batch_id == batch_id)
rows = query.order_by(DouyuTask.id.desc()).limit(300).all()
return [_task_out(task) for task in rows]
return [_task_out(task, include_detail=include_detail) for task in rows]
@router.get("/tasks/{task_id}", response_model=DouyuTaskOut)
@@ -326,7 +397,7 @@ def get_task(
task = _visible_tasks_query(db, current).filter(DouyuTask.id == task_id).first()
if not task:
raise HTTPException(status_code=404, detail="任务不存在")
return _task_out(task)
return _task_out(task, include_detail=True)
@router.post("/stop/{batch_id}")
+27
View File
@@ -208,6 +208,24 @@ def _visible_huya_tasks_query(db: Session, current: User):
raise HTTPException(status_code=403, detail="无权查看虎牙任务")
def _huya_task_summary(query):
"""按状态汇总虎牙任务,避免概览页拉完整任务列表。"""
rows = (
query.enable_eagerloads(False)
.order_by(None)
.with_entities(HuyaTask.status, func.count(HuyaTask.id))
.group_by(HuyaTask.status)
.all()
)
status_counts = {status or "": count for status, count in rows}
return {
"total": sum(status_counts.values()),
"success": status_counts.get("success", 0),
"failed": sum(status_counts.get(status, 0) for status in ("failed", "error")),
"status_counts": status_counts,
}
def _require_huya_task_account_access(db: Session, current: User, account_ids: list[int]) -> None:
"""确保任务只会提交到当前用户可操作的虎牙账号。"""
requested_ids = set(account_ids)
@@ -1379,6 +1397,15 @@ def list_tasks(
return [_task_out(task, include_images=include_images) for task in tasks]
@router.get("/tasks/summary")
def tasks_summary(
db: Session = Depends(get_db),
current: User = Depends(require_permission("huya:task")),
):
"""虎牙任务统计。"""
return _huya_task_summary(_visible_huya_tasks_query(db, current))
@router.get("/tasks/{task_id}", response_model=HuyaTaskOut)
def get_task(
task_id: int,
+33 -3
View File
@@ -4,7 +4,8 @@ import asyncio
import threading
from datetime import datetime
from fastapi import APIRouter, Depends, HTTPException, WebSocket, WebSocketDisconnect
from sqlalchemy.orm import Session
from sqlalchemy import func
from sqlalchemy.orm import Session, defer
from ..database import get_db, SessionLocal
from ..models import User, Account, LoginTask, ProxyConfig as ProxyConfigModel
@@ -16,6 +17,23 @@ from ..services.login_service import LoginBatchRunner, batch_registry
router = APIRouter(prefix="/api/login", tags=["登录任务"])
def _login_task_summary(query):
"""按状态汇总登录任务,避免概览页拉完整任务列表。"""
rows = (
query.order_by(None)
.with_entities(LoginTask.status, func.count(LoginTask.id))
.group_by(LoginTask.status)
.all()
)
status_counts = {status or "": count for status, count in rows}
return {
"total": sum(status_counts.values()),
"success": status_counts.get("success", 0),
"failed": sum(status_counts.get(status, 0) for status in ("failed", "error")),
"status_counts": status_counts,
}
@router.post("/batch")
async def create_batch(
req: LoginBatchRequest,
@@ -84,7 +102,7 @@ def list_tasks(
current: User = Depends(get_current_user),
):
"""查看登录任务列表。"""
query = db.query(LoginTask).join(Account, LoginTask.account_id == Account.id)
query = db.query(LoginTask).options(defer(LoginTask.cookie)).join(Account, LoginTask.account_id == Account.id)
# 客服只能看自己账号的任务
if not user_has_permission(current, "login:view_all"):
@@ -106,12 +124,24 @@ def list_tasks(
result.append(LoginTaskOut(
id=t.id, batch_id=t.batch_id, account_id=t.account_id,
account_username=accounts_map.get(t.account_id, ""),
status=t.status, cookie=t.cookie or "", message=t.message or "",
status=t.status, cookie="", message=t.message or "",
created_by=t.created_by, created_at=t.created_at, finished_at=t.finished_at,
))
return result
@router.get("/tasks/summary")
def tasks_summary(
db: Session = Depends(get_db),
current: User = Depends(get_current_user),
):
"""登录任务统计。"""
query = db.query(LoginTask).join(Account, LoginTask.account_id == Account.id)
if not user_has_permission(current, "login:view_all"):
query = query.filter(Account.assigned_to == current.id)
return _login_task_summary(query)
@router.delete("/tasks/{task_id}")
def delete_task(
task_id: int,