"""登录任务路由 + WebSocket 实时日志""" import asyncio import threading from fastapi import APIRouter, Depends, HTTPException, WebSocket, WebSocketDisconnect from sqlalchemy import func from sqlalchemy.orm import Session, defer, load_only from ..database import SessionLocal, get_db from ..deps import authenticate_websocket, get_current_user, require_permission from ..models import Account, LoginTask, User from ..models import ProxyConfig as ProxyConfigModel from ..permissions import get_user_permissions, user_has_permission from ..schemas import LoginBatchRequest, LoginTaskOut 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, db: Session = Depends(get_db), current: User = Depends(require_permission("login:batch")), ): """创建批量登录或账号检测任务。""" if not req.account_ids: raise HTTPException(status_code=400, detail="请选择账号") # 读取代理配置 proxy = db.query(ProxyConfigModel).first() # 权限过滤账号 valid_ids = [] seen_ids = set() for aid in req.account_ids: if aid in seen_ids: continue seen_ids.add(aid) acc = db.query(Account).filter(Account.id == aid).first() if not acc: continue if ( not user_has_permission(current, "login:view_all") and acc.assigned_to != current.id ): continue valid_ids.append(aid) if not valid_ids: action_name = "检测" if req.mode == "check" else "登录" raise HTTPException(status_code=403, detail=f"没有可{action_name}的账号") # 在主事件循环中创建 log_queue,传给后台线程 log_queue = asyncio.Queue() loop = asyncio.get_running_loop() # 创建执行器(用独立的 DB 会话,因为在线程中运行) thread_db = SessionLocal() runner = LoginBatchRunner( db=thread_db, account_ids=valid_ids, created_by=current.id, creator_permissions=get_user_permissions(current), max_login_retries=req.max_login_retries, max_total_time=req.max_total_time, proxy_config=proxy, log_queue=log_queue, loop=loop, concurrency=req.concurrency, api_strategy=req.api_strategy, mode=req.mode, ) batch_id = runner.batch_id # 先注册到批次注册表,再启动线程,确保 WebSocket 连接时能找到 batch_registry.register(batch_id, log_queue, loop, runner) # 启动线程 thread = threading.Thread(target=runner.run, daemon=True) thread.start() return {"batch_id": batch_id, "count": len(valid_ids), "success": True} @router.get("/tasks", response_model=list[LoginTaskOut]) def list_tasks( batch_id: str | None = None, db: Session = Depends(get_db), current: User = Depends(get_current_user), ): """查看登录任务列表。""" query = ( db.query(LoginTask) .options(defer(LoginTask.cookie)) .join(Account, LoginTask.account_id == Account.id) ) # 客服只能看自己账号的任务 if not user_has_permission(current, "login:view_all"): query = query.filter(Account.assigned_to == current.id) if batch_id: query = query.filter(LoginTask.batch_id == batch_id) rows = query.order_by(LoginTask.id.desc()).limit(200).all() # 批量收集 account_id,一次性查出 username account_ids = [t.account_id for t in rows] accounts_map = {} if account_ids: accs = ( db.query(Account) .options(load_only(Account.id, Account.username)) .filter(Account.id.in_(account_ids)) .all() ) accounts_map = {a.id: a.username for a in accs} result = [] for t in rows: 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="", 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, db: Session = Depends(get_db), current: User = Depends(require_permission("login:batch")), ): """删除单个登录任务。""" task = db.query(LoginTask).filter(LoginTask.id == task_id).first() if not task: raise HTTPException(status_code=404, detail="任务不存在") db.delete(task) db.commit() return {"message": "已删除", "success": True} @router.delete("/tasks") def delete_tasks( task_ids: str = "", db: Session = Depends(get_db), current: User = Depends(require_permission("login:batch")), ): """批量删除登录任务。""" if not task_ids: raise HTTPException(status_code=400, detail="请指定任务ID") ids = [int(x) for x in task_ids.split(",") if x.strip().isdigit()] if not ids: raise HTTPException(status_code=400, detail="无效的任务ID") deleted = ( db.query(LoginTask) .filter(LoginTask.id.in_(ids)) .delete(synchronize_session=False) ) db.commit() return {"message": f"已删除 {deleted} 个任务", "deleted": deleted, "success": True} @router.post("/stop/{batch_id}") def stop_batch( batch_id: str, current: User = Depends(require_permission("login:batch")), ): batch = batch_registry.get(batch_id) if batch: if batch.get("finished"): # 已结束但 WS 尚未清理(或从未有 WS 订阅):回收注册表条目。 batch_registry.pop(batch_id) raise HTTPException(status_code=404, detail="批次已结束") batch["runner"].stop() return {"message": "已发送停止信号", "success": True} raise HTTPException(status_code=404, detail="批次不存在或已结束") @router.websocket("/ws/login/{batch_id}") async def ws_login_logs(websocket: WebSocket, batch_id: str): """WebSocket 推送登录实时日志(需认证)。""" # 认证:从 cookie 或 token query param 验证用户身份 user = authenticate_websocket(websocket) if not user: await websocket.close(code=1008, reason="未授权") return # 归属校验:非 view_all 只能订阅自己创建的批次(与虎牙批次一致)。 if not user_has_permission(user, "login:view_all"): check_db = SessionLocal() try: owned = ( check_db.query(LoginTask.id) .filter(LoginTask.batch_id == batch_id, LoginTask.created_by == user.id) .first() ) finally: check_db.close() if not owned: await websocket.close(code=1008, reason="无权访问该任务批次") return await websocket.accept() # 从已注册的批次中获取 log_queue(由 create_batch 创建) batch = batch_registry.get(batch_id) if not batch: await websocket.send_json({"level": "error", "message": "批次不存在或已结束"}) await websocket.close() return log_queue: asyncio.Queue = batch["log_queue"] try: while True: try: msg = await asyncio.wait_for(log_queue.get(), timeout=30) await websocket.send_json(msg) # 收到 result 表示任务结束 if msg.get("level") == "result": await asyncio.sleep(0.1) break except TimeoutError: await websocket.send_json({"level": "heartbeat", "message": ""}) except WebSocketDisconnect: pass finally: # 只在批次真正结束后清理注册表,客户端断线/刷新页面后重连仍可继续订阅。 latest = batch_registry.get(batch_id) if latest and latest.get("finished"): batch_registry.pop(batch_id)