diff --git a/backend/app/routers/websites.py b/backend/app/routers/websites.py index 12892fd..43883f4 100644 --- a/backend/app/routers/websites.py +++ b/backend/app/routers/websites.py @@ -2,10 +2,15 @@ from __future__ import annotations import json import logging +import secrets +import queue as _queue +import threading as _threading +import time from datetime import datetime, timezone -from typing import List +from typing import Any, List from fastapi import APIRouter, Depends, HTTPException, Query, status +from fastapi.responses import StreamingResponse from sqlalchemy.orm import Session from app.database import get_db @@ -75,6 +80,132 @@ SENSITIVE_CREDENTIAL_KEYS = { } ALGORITHMS = {"max_plus_percent", "average_plus_percent", "min_plus_percent", "priority_weighted_plus_percent"} +# ——— 服务端任务锁(防止同一网站同类任务并发执行) ——— +_running_tasks: dict[str, dict[str, Any]] = {} +_tasks_lock = _threading.Lock() + + +def _try_start_task(wid: int, task_type: str, timeout_minutes: int = 240) -> tuple[bool, str | None]: + """尝试获取任务锁,成功返回 (True, token),失败返回 (False, None)。 + + token 用于 _finish_task 校验锁所有权:只有当前锁记录的 token 与传入 token + 一致时才释放锁。防止旧 worker(强制过期后仍存活的线程)误删新任务的锁。 + """ + key = f"task:{task_type}:{wid}" + now = time.monotonic() + token = secrets.token_hex(16) + with _tasks_lock: + entry = _running_tasks.get(key) + if entry and (now - entry["started_at"]) < timeout_minutes * 60: + return False, None + if entry: + logger.warning("task lock force-expired wid=%s task_type=%s elapsed=%.1fmin", wid, task_type, (now - entry["started_at"]) / 60) + _running_tasks[key] = {"started_at": now, "token": token} + return True, token + + +def _finish_task(wid: int, task_type: str, token: str | None): + """释放任务锁,但仅当 token 与当前锁所有者一致时执行。 + + 如果锁已被其他任务强制过期并替换,token 不匹配,则仅记 warning, + 不删除新任务的锁。 + """ + key = f"task:{task_type}:{wid}" + with _tasks_lock: + entry = _running_tasks.get(key) + if entry is None: + return + if entry.get("token") != token: + logger.warning("task lock not released: token mismatch (stale worker) wid=%s task_type=%s", wid, task_type) + return + _running_tasks.pop(key, None) + + +# ——— 心跳包装器:在同步生成器执行期间确保定期有事件输出 ——— +def _with_background_heartbeat(gen_fn, interval: float = 10, timeout_factor: int = 3, *, task_wid: int | None = None, task_type: str | None = None, task_token: str | None = None, **gen_kwargs): + """在后台线程运行 gen_fn(**gen_kwargs),主线程通过队列消费事件。 + + 独立心跳线程每 *interval* 秒检查一次,若主线程消费落后(队列为空)则输出 + heartbeat 事件。防止在单个远端请求阻塞期间无事件产生。 + + 任务锁生命周期(解决客户端断开后锁提前释放的问题): + - 任务锁的获取(_try_start_task)在路由处理函数中完成。 + - 任务锁的释放(_finish_task)在 _run() 的 finally 中,由业务线程 + 自身执行。因此锁的释放与业务线程的结束严格绑定。 + - 客户端断开只触发 cancel 信号和队列停止,不会调用 _finish_task。 + - 即使业务线程阻塞在远端 HTTP 请求中,锁仍然保持持有状态, + 直到请求返回(超时或完成)且线程退出。 + - 取消信号拦截在下一次 Python 代码执行边界(队列写/循环头), + 无法中断阻塞中的系统调用,但锁不会提前释放。 + + 队列满保护: + - q.put(timeout=5) + 循环重试,避免队列满后后台线程永久阻塞。 + 重试时检查 cancel,允许线程在被阻塞时响应取消。 + """ + q = _queue.Queue(maxsize=100) + stop = _threading.Event() + cancel = _threading.Event() + + def _run(): + try: + for event, data in gen_fn(**gen_kwargs): + if cancel.is_set(): + return + # 使用超时 put 防止队列满后永久阻塞 + while True: + try: + q.put((event, data), timeout=5) + break + except _queue.Full: + if cancel.is_set(): + return + except Exception as exc: + try: + q.put(("error", {"message": str(exc)}), timeout=2) + except _queue.Full: + pass + finally: + stop.set() + # 业务线程结束后才释放任务锁,并通过 token 校验所有权, + # 防止过期后仍存活的旧 worker 误删新任务锁。 + if task_wid is not None and task_type is not None: + _finish_task(task_wid, task_type, task_token) + + def _heartbeat(): + while not stop.wait(timeout=interval): + if cancel.is_set(): + return + try: + q.put(("heartbeat", {}), timeout=2) + except _queue.Full: + pass + + t_work = _threading.Thread(target=_run, daemon=True) + t_hb = _threading.Thread(target=_heartbeat, daemon=True) + t_work.start() + t_hb.start() + + try: + while True: + try: + event, data = q.get(timeout=interval * timeout_factor) + yield event, data + if event in ("complete", "error"): + break + except _queue.Empty: + cancel.set() + yield "error", {"message": "内部处理线程长时间无响应,任务异常中断"} + break + except GeneratorExit: + # 客户端断开连接或消费方异常退出 → 通知后台线程停止 + cancel.set() + raise + finally: + cancel.set() + stop.set() + # 注意:不在此处释放任务锁。_finish_task 由 _run() 的 finally 负责, + # 确保锁释放与业务线程终止严格同步。 + def _mask(cfg: dict) -> dict: masked = {} @@ -1480,6 +1611,7 @@ def _sync_upstream_models_generator(wid: int, db: Session): return yield "start", {"total_accounts": len(candidates)} + last_event_at = time.monotonic() upstream_ids = {cand["upstream_id"] for cand in candidates.values()} upstreams_map = {up.id: up for up in db.query(Upstream).filter(Upstream.id.in_(upstream_ids)).all()} @@ -1500,6 +1632,9 @@ def _sync_upstream_models_generator(wid: int, db: Session): ) items.append(item) yield "item", item.model_dump() + if time.monotonic() - last_event_at >= 10: + yield "heartbeat", {} + last_event_at = time.monotonic() continue acc_name = remote_acc.get("name") @@ -1518,9 +1653,13 @@ def _sync_upstream_models_generator(wid: int, db: Session): ) items.append(item) yield "item", item.model_dump() + if time.monotonic() - last_event_at >= 10: + yield "heartbeat", {} + last_event_at = time.monotonic() continue upstream = upstreams_map.get(cand["upstream_id"]) + upstream_base_url = upstream.base_url if upstream else None if not upstream_base_url or not upstream_base_url.strip(): @@ -1534,6 +1673,9 @@ def _sync_upstream_models_generator(wid: int, db: Session): ) items.append(item) yield "item", item.model_dump() + if time.monotonic() - last_event_at >= 10: + yield "heartbeat", {} + last_event_at = time.monotonic() continue # 1. 先安全修复 base_url (不依赖同步模型结果) @@ -1563,6 +1705,9 @@ def _sync_upstream_models_generator(wid: int, db: Session): ) items.append(item) yield "item", item.model_dump() + if time.monotonic() - last_event_at >= 10: + yield "heartbeat", {} + last_event_at = time.monotonic() continue current_creds["api_key"] = local_key_value @@ -1579,6 +1724,9 @@ def _sync_upstream_models_generator(wid: int, db: Session): ) items.append(item) yield "item", item.model_dump() + if time.monotonic() - last_event_at >= 10: + yield "heartbeat", {} + last_event_at = time.monotonic() continue # 2. 调用 sub2api 同步模型 @@ -1598,6 +1746,9 @@ def _sync_upstream_models_generator(wid: int, db: Session): ) items.append(item) yield "item", item.model_dump() + if time.monotonic() - last_event_at >= 10: + yield "heartbeat", {} + last_event_at = time.monotonic() continue # 3. 模型同步成功后,在 current_creds 基础上替换 model_mapping 并写回 @@ -1615,6 +1766,9 @@ def _sync_upstream_models_generator(wid: int, db: Session): ) items.append(item) yield "item", item.model_dump() + if time.monotonic() - last_event_at >= 10: + yield "heartbeat", {} + last_event_at = time.monotonic() except Exception as e: item = SyncUpstreamModelsItem( account_id=aid, @@ -1626,6 +1780,9 @@ def _sync_upstream_models_generator(wid: int, db: Session): ) items.append(item) yield "item", item.model_dump() + if time.monotonic() - last_event_at >= 10: + yield "heartbeat", {} + last_event_at = time.monotonic() success_count = sum(1 for item in items if item.status == "success") failed_count = sum(1 for item in items if item.status == "failed") @@ -1655,32 +1812,45 @@ def sync_website_accounts_upstream_models( _=Depends(get_current_user), ): """一键同步上游模型""" + ok, _sync_token = _try_start_task(wid, "sync_models") + if not ok: + raise HTTPException(409, "该网站的同步上游模型正在执行中") + + logger.info("sync_models start wid=%s", wid) + t0 = time.monotonic() items = [] success = False message = "" error_occurred = False - for event, data in _sync_upstream_models_generator(wid, db): - if event == "item": - items.append(SyncUpstreamModelsItem(**data)) - elif event == "complete": - success = data["success"] - message = data["message"] - elif event == "error": - error_occurred = True - message = data["message"] + try: + for event, data in _sync_upstream_models_generator(wid, db): + if event == "item": + items.append(SyncUpstreamModelsItem(**data)) + elif event == "complete": + success = data["success"] + message = data["message"] + elif event == "error": + error_occurred = True + message = data["message"] + finally: + _finish_task(wid, "sync_models", _sync_token) if error_occurred: if message == "website not found": raise HTTPException(404, "website not found") if "only sub2api" in message: raise HTTPException(400, message) + elapsed = time.monotonic() - t0 + logger.info("sync_models complete wid=%s error=%s elapsed=%.1fs", wid, message, elapsed) return SyncUpstreamModelsResponse( success=False, message=message, items=[] ) + elapsed = time.monotonic() - t0 + logger.info("sync_models complete wid=%s items=%d success=%s elapsed=%.1fs", wid, len(items), success, elapsed) return SyncUpstreamModelsResponse( success=success, message=message, @@ -1695,27 +1865,43 @@ def sync_website_accounts_upstream_models_stream( _=Depends(get_current_user), ): """流式一键同步上游模型""" - from fastapi.responses import StreamingResponse def event_generator(): - for event, data in _sync_upstream_models_generator(wid, db): + ok, _sync_token = _try_start_task(wid, "sync_models") + if not ok: + yield json.dumps({"event": "error", "data": {"message": "该网站的同步上游模型正在执行中"}}, ensure_ascii=False) + "\n" + return + + logger.info("sync_models/stream start wid=%s", wid) + t0 = time.monotonic() + for event, data in _with_background_heartbeat( + _sync_upstream_models_generator, wid=wid, db=db, + task_wid=wid, task_type="sync_models", task_token=_sync_token, + ): yield json.dumps({"event": event, "data": data}, ensure_ascii=False) + "\n" - return StreamingResponse(event_generator(), media_type="application/x-ndjson") + elapsed = time.monotonic() - t0 + logger.info("sync_models/stream complete wid=%s elapsed=%.1fs", wid, elapsed) + + return StreamingResponse( + event_generator(), + media_type="application/x-ndjson", + headers={ + "Cache-Control": "no-cache", + "X-Accel-Buffering": "no", + }, + ) -@router.post("/api/websites/{wid}/groups/organize", response_model=OrganizeGroupsResponse) -def organize_website_groups( - wid: int, - db: Session = Depends(get_db), - _=Depends(get_current_user), -): - """一键整理分组:按现有绑定关系导入/自愈账号""" +def _organize_website_groups_generator(wid: int, db: Session): + """一键整理分组生成器,逐条产生事件供流式或批量消费。""" website = db.query(Website).filter(Website.id == wid).first() if not website: - raise HTTPException(404, "website not found") + yield "error", {"message": "website not found"} + return if website.site_type != "sub2api": - raise HTTPException(400, "目前只支持 sub2api") + yield "error", {"message": "目前只支持 sub2api"} + return # 1. 读取当前网站的所有启用绑定关系 bindings = ( @@ -1742,6 +1928,23 @@ def organize_website_groups( items: list[OrganizeGroupsItem] = [] + # 预计算总处理项数(用于流式进度);missing_key 也算 1 项 + total_items = 0 + for b in bindings: + for src in binding_sources(b): + uid = src.get("upstream_id") + gid = src.get("group_id") + if uid and gid: + cnt = db.query(UpstreamGeneratedKey).filter( + UpstreamGeneratedKey.upstream_id == uid, + UpstreamGeneratedKey.group_id == gid, + UpstreamGeneratedKey.status != "orphaned", + ).count() + total_items += max(1, cnt) + + yield "start", {"total_items": total_items, "total_bindings": len(bindings)} + last_event_at = time.monotonic() + # 2. 收集所有的 upstream_id 并执行对账,以保证本地 Key 的状态最新 upstream_ids = set() for b in bindings: @@ -1854,20 +2057,23 @@ def organize_website_groups( ) if not keys: - items.append( - OrganizeGroupsItem( - target_group_id=str(target_group_id), - target_group_name=target_group_name, - upstream_name=upstream_name, - source_group_id=str(gid), - source_group_name=str(source_group_name), - key_name="", - account_id=None, - account_name=None, - status="missing_key", - message="请先生成上游 Key", - ) + _item = OrganizeGroupsItem( + target_group_id=str(target_group_id), + target_group_name=target_group_name, + upstream_name=upstream_name, + source_group_id=str(gid), + source_group_name=str(source_group_name), + key_name="", + account_id=None, + account_name=None, + status="missing_key", + message="请先生成上游 Key", ) + items.append(_item) + yield "item", _item.model_dump() + if time.monotonic() - last_event_at >= 10: + yield "heartbeat", {} + last_event_at = time.monotonic() continue # 获取来源上游的 Base URL 用于账号创建 @@ -1876,20 +2082,23 @@ def organize_website_groups( for row in keys: # 跳过状态为 failed 的 Key if row.status == "failed": - items.append( - OrganizeGroupsItem( - target_group_id=str(target_group_id), - target_group_name=target_group_name, - upstream_name=upstream_name, - source_group_id=str(gid), - source_group_name=str(source_group_name), - key_name=row.key_name, - account_id=None, - account_name=None, - status="failed", - message="上游 Key 生成状态为失败", - ) + _item = OrganizeGroupsItem( + target_group_id=str(target_group_id), + target_group_name=target_group_name, + upstream_name=upstream_name, + source_group_id=str(gid), + source_group_name=str(source_group_name), + key_name=row.key_name, + account_id=None, + account_name=None, + status="failed", + message="上游 Key 生成状态为失败", ) + items.append(_item) + yield "item", _item.model_dump() + if time.monotonic() - last_event_at >= 10: + yield "heartbeat", {} + last_event_at = time.monotonic() continue # 检测平台类型 @@ -1906,19 +2115,22 @@ def organize_website_groups( if remote_account_ids is None: # 账号列表获取失败,无法校验状态,保守跳过 - items.append( - OrganizeGroupsItem( - target_group_id=str(target_group_id), - target_group_name=target_group_name, - upstream_name=upstream_name, - source_group_id=str(gid), - source_group_name=str(source_group_name), - key_name=row.key_name, - account_id=old_account_id, - status="failed", - message="无法校验目标账号状态,已保守跳过", - ) + _item = OrganizeGroupsItem( + target_group_id=str(target_group_id), + target_group_name=target_group_name, + upstream_name=upstream_name, + source_group_id=str(gid), + source_group_name=str(source_group_name), + key_name=row.key_name, + account_id=old_account_id, + status="failed", + message="无法校验目标账号状态,已保守跳过", ) + items.append(_item) + yield "item", _item.model_dump() + if time.monotonic() - last_event_at >= 10: + yield "heartbeat", {} + last_event_at = time.monotonic() continue remote_acc = remote_account_map.get(str(old_account_id)) @@ -1989,35 +2201,40 @@ def organize_website_groups( row.imported_target_group_name = target_group_names.get(str(target_group_id)) db.commit() - items.append( - OrganizeGroupsItem( - target_group_id=str(target_group_id), - target_group_name=target_group_name, - upstream_name=upstream_name, - source_group_id=str(gid), - source_group_name=str(source_group_name), - key_name=row.key_name, - account_id=old_account_id, - account_name=str(remote_acc.get("name") or ""), - status="exists", - message=msg, - ) - ) + items.append(_item := OrganizeGroupsItem( + target_group_id=str(target_group_id), + target_group_name=target_group_name, + upstream_name=upstream_name, + source_group_id=str(gid), + source_group_name=str(source_group_name), + key_name=row.key_name, + account_id=old_account_id, + account_name=str(remote_acc.get("name") or ""), + status="exists", + message=msg, + )) + yield "item", _item.model_dump() + if time.monotonic() - last_event_at >= 10: + yield "heartbeat", {} + last_event_at = time.monotonic() continue except Exception as exc: - items.append( - OrganizeGroupsItem( - target_group_id=str(target_group_id), - target_group_name=target_group_name, - upstream_name=upstream_name, - source_group_id=str(gid), - source_group_name=str(source_group_name), - key_name=row.key_name, - account_id=old_account_id, - status="failed", - message=f"更新绑定失败: {exc}", - ) + _item = OrganizeGroupsItem( + target_group_id=str(target_group_id), + target_group_name=target_group_name, + upstream_name=upstream_name, + source_group_id=str(gid), + source_group_name=str(source_group_name), + key_name=row.key_name, + account_id=old_account_id, + status="failed", + message=f"更新绑定失败: {exc}", ) + items.append(_item) + yield "item", _item.model_dump() + if time.monotonic() - last_event_at >= 10: + yield "heartbeat", {} + last_event_at = time.monotonic() continue else: # 远端已删除,清理标记后进行重建 @@ -2032,18 +2249,21 @@ def organize_website_groups( # 2. 检查是否有明文 Key if not row.key_value: - items.append( - OrganizeGroupsItem( - target_group_id=str(target_group_id), - target_group_name=target_group_name, - upstream_name=upstream_name, - source_group_id=str(gid), - source_group_name=str(source_group_name), - key_name=row.key_name, - status="failed", - message="该 Key 无明文值,无法导入", - ) + _item = OrganizeGroupsItem( + target_group_id=str(target_group_id), + target_group_name=target_group_name, + upstream_name=upstream_name, + source_group_id=str(gid), + source_group_name=str(source_group_name), + key_name=row.key_name, + status="failed", + message="该 Key 无明文值,无法导入", ) + items.append(_item) + yield "item", _item.model_dump() + if time.monotonic() - last_event_at >= 10: + yield "heartbeat", {} + last_event_at = time.monotonic() continue # 3. 创建账号并绑定当前目标分组 @@ -2091,37 +2311,43 @@ def organize_website_groups( if account_id: remote_account_map[str(account_id)] = created - items.append( - OrganizeGroupsItem( - target_group_id=str(target_group_id), - target_group_name=target_group_name, - upstream_name=upstream_name, - source_group_id=str(gid), - source_group_name=str(source_group_name), - key_name=row.key_name, - account_id=account_id or None, - account_name=str(created.get("name") or account_name), - status="recreated" if is_recreated else "created", - message="清理后重建账号" if is_recreated else "已创建账号", - ) + _item = OrganizeGroupsItem( + target_group_id=str(target_group_id), + target_group_name=target_group_name, + upstream_name=upstream_name, + source_group_id=str(gid), + source_group_name=str(source_group_name), + key_name=row.key_name, + account_id=account_id or None, + account_name=str(created.get("name") or account_name), + status="recreated" if is_recreated else "created", + message="清理后重建账号" if is_recreated else "已创建账号", ) + items.append(_item) + yield "item", _item.model_dump() + if time.monotonic() - last_event_at >= 10: + yield "heartbeat", {} + last_event_at = time.monotonic() except Exception as exc: logger.exception("organize create account failed website=%s key=%s", wid, row.id) row.status = "import_failed" row.error = str(exc) db.commit() - items.append( - OrganizeGroupsItem( - target_group_id=str(target_group_id), - target_group_name=target_group_name, - upstream_name=upstream_name, - source_group_id=str(gid), - source_group_name=str(source_group_name), - key_name=row.key_name, - status="failed", - message=str(exc), - ) + _item = OrganizeGroupsItem( + target_group_id=str(target_group_id), + target_group_name=target_group_name, + upstream_name=upstream_name, + source_group_id=str(gid), + source_group_name=str(source_group_name), + key_name=row.key_name, + status="failed", + message=str(exc), ) + items.append(_item) + yield "item", _item.model_dump() + if time.monotonic() - last_event_at >= 10: + yield "heartbeat", {} + last_event_at = time.monotonic() created_count = sum(1 for item in items if item.status == "created") recreated_count = sum(1 for item in items if item.status == "recreated") @@ -2151,10 +2377,82 @@ def organize_website_groups( logger.warning("failed to sync priorities after organize for website %s: %s", wid, exc) message = "整理完成:" + " / ".join(parts) if parts else "整理完成:无任何绑定或数据" - return OrganizeGroupsResponse( - success=failed_count == 0, - message=message, - items=items, + yield "complete", { + "success": failed_count == 0, + "message": message, + "items": [item.model_dump() for item in items], + } + + +@router.post("/api/websites/{wid}/groups/organize", response_model=OrganizeGroupsResponse) +def organize_website_groups( + wid: int, + db: Session = Depends(get_db), + _=Depends(get_current_user), +): + """一键整理分组:按现有绑定关系导入/自愈账号(JSON 接口,兼容前端和测试)""" + ok, _organize_token = _try_start_task(wid, "organize") + if not ok: + raise HTTPException(409, "该网站的一键整理正在执行中") + + logger.info("organize start wid=%s", wid) + t0 = time.monotonic() + items: list[OrganizeGroupsItem] = [] + message = "" + success = False + try: + for event, data in _organize_website_groups_generator(wid, db): + if event == "item": + items.append(OrganizeGroupsItem(**data)) + elif event == "complete": + success = data["success"] + message = data["message"] + elif event == "error": + if data["message"] == "website not found": + raise HTTPException(404, "website not found") + if "only sub2api" in data["message"] or "目前只支持" in data["message"]: + raise HTTPException(400, data["message"]) + return OrganizeGroupsResponse(success=False, message=data["message"], items=[]) + finally: + _finish_task(wid, "organize", _organize_token) + + elapsed = time.monotonic() - t0 + logger.info("organize complete wid=%s items=%d success=%s elapsed=%.1fs", wid, len(items), success, elapsed) + return OrganizeGroupsResponse(success=success, message=message, items=items) + + +@router.post("/api/websites/{wid}/groups/organize/stream") +def organize_website_groups_stream( + wid: int, + db: Session = Depends(get_db), + _=Depends(get_current_user), +): + """流式一键整理分组""" + + def event_generator(): + ok, _organize_token = _try_start_task(wid, "organize") + if not ok: + yield json.dumps({"event": "error", "data": {"message": "该网站的一键整理正在执行中"}}, ensure_ascii=False) + "\n" + return + + logger.info("organize/stream start wid=%s", wid) + t0 = time.monotonic() + for event, data in _with_background_heartbeat( + _organize_website_groups_generator, wid=wid, db=db, + task_wid=wid, task_type="organize", task_token=_organize_token, + ): + yield json.dumps({"event": event, "data": data}, ensure_ascii=False) + "\n" + + elapsed = time.monotonic() - t0 + logger.info("organize/stream complete wid=%s elapsed=%.1fs", wid, elapsed) + + return StreamingResponse( + event_generator(), + media_type="application/x-ndjson", + headers={ + "Cache-Control": "no-cache", + "X-Accel-Buffering": "no", + }, ) diff --git a/backend/test_background_heartbeat.py b/backend/test_background_heartbeat.py new file mode 100644 index 0000000..f007849 --- /dev/null +++ b/backend/test_background_heartbeat.py @@ -0,0 +1,159 @@ +import time +import threading +from app.routers.websites import _with_background_heartbeat, _try_start_task, _finish_task +import app.routers.websites as _websites_module + + +def test_heartbeat_emitted_during_blocking_operation(): + """生成器阻塞超过心跳间隔时,仍然能收到 heartbeat 事件。""" + # 模拟一个每 6 秒产生一条 item 的生成器 + def slow_gen(): + yield "item", {"n": 1} + time.sleep(7) # 阻塞时间 > interval=4 + yield "item", {"n": 2} + yield "complete", {"success": True, "message": "done"} + + events = list(_with_background_heartbeat(slow_gen, interval=4, timeout_factor=5)) + + heartbeat_count = sum(1 for e, _ in events if e == "heartbeat") + assert heartbeat_count >= 1, f"阻塞 7s 应产生至少 1 个 heartbeat,实际 {heartbeat_count}" + items = [(e, d) for e, d in events if e == "item"] + assert len(items) == 2 + assert items[0][1]["n"] == 1 + assert items[1][1]["n"] == 2 + + +def test_heartbeat_multiple_during_long_block(): + """长时间阻塞(>2 个心跳周期)时产生多个 heartbeat。""" + def long_block_gen(): + yield "item", {"n": 1} + time.sleep(11) # > 2 * interval=5 + yield "complete", {"success": True, "message": "done"} + + events = [] + for event, data in _with_background_heartbeat(long_block_gen, interval=5, timeout_factor=5): + events.append(event) + if event == "complete": + break + + heartbeat_count = sum(1 for e in events if e == "heartbeat") + assert heartbeat_count >= 2, f"11s 阻塞 (interval=5) 应产生至少 2 个 heartbeat,实际 {heartbeat_count}" + + +def test_queue_full_does_not_deadlock(): + """生产速度超过消费速度且队列满时,后台线程不会永久阻塞。""" + n_items = 300 + + def fast_gen(): + for i in range(n_items): + yield "item", {"n": i} + yield "complete", {"success": True, "message": "done"} + + # 使用大 timeout_factor 确保主线程有充足时间消费 + events = list(_with_background_heartbeat(fast_gen, interval=10, timeout_factor=20)) + + item_count = sum(1 for e, _ in events if e == "item") + assert item_count == n_items, f"应有 {n_items} 个 item,实际 {item_count}" + assert any(e == "complete" for e, _ in events), "应收到 complete 事件" + + +def test_generator_close_stops_background_thread_immediately(): + """生成器关闭后后台线程停止,不阻塞 join 超时。""" + thread_done = threading.Event() + + def infinite_gen(): + try: + i = 0 + while True: + yield "item", {"n": i} + i += 1 + finally: + thread_done.set() + + gen = _with_background_heartbeat(infinite_gen, interval=10, timeout_factor=10) + + # 消费 3 个事件 + count = 0 + for event, data in gen: + count += 1 + if count >= 3: + break + + # 关闭生成器模拟客户端断开 + gen.close() + + # 后台线程应快速响应 cancel 并退出 + assert thread_done.wait(timeout=10), "后台线程未在 10s 内停止" + + +def test_worker_exception_caught_as_error_event(): + """生成器内部抛出异常时,产生 error 事件而非崩溃。""" + def crashing_gen(): + yield "item", {"n": 1} + raise RuntimeError("模拟意外错误") + + events = list(_with_background_heartbeat(crashing_gen, interval=10, timeout_factor=5)) + + error_events = [(e, d) for e, d in events if e == "error"] + assert len(error_events) == 1 + assert "模拟意外错误" in error_events[0][1]["message"] + # 不应收到 complete + assert not any(e == "complete" for e, _ in events) + + +def test_worker_timeout_emits_error_and_cancels(): + """worker 长时间不产生事件时,主线程超时并产生 error。""" + def stuck_gen(): + yield "start", {} + time.sleep(15) # 远远超过 queue.get 超时时间 + yield "complete", {"success": True, "message": "done"} + + # interval=5, timeout_factor=0.5 → queue.get 超时 = 2.5s + # 心跳间隔(5s) > queue.get 超时(2.5s),所以不会因心跳重置超时计数器 + events = list(_with_background_heartbeat(stuck_gen, interval=5, timeout_factor=0.5)) + + error_events = [(e, d) for e, d in events if e == "error"] + assert len(error_events) == 1 + assert "长时间无响应" in error_events[0][1]["message"] + + +def test_lock_held_until_worker_finishes_after_disconnect(): + """生成器关闭后任务锁保持持有,直到业务线程真正退出。""" + wid = 9999 + + ok, _task_token = _try_start_task(wid, "test_lock") + assert ok, "应成功获取锁" + + worker_done = threading.Event() + + def blocking_gen(): + try: + yield "item", {"n": 1} + time.sleep(5) # 模拟长时间阻塞(HTTP 请求等) + yield "complete", {"success": True, "message": "done"} + finally: + worker_done.set() + + gen = _with_background_heartbeat( + blocking_gen, interval=10, timeout_factor=10, + task_wid=wid, task_type="test_lock", task_token=_task_token, + ) + + # 消费一个 item 后关闭生成器(模拟客户端断开) + next(gen) + gen.close() + + # 关闭后锁应当仍然被持有(worker 还在 sleep(5) 中) + ok, _ = _try_start_task(wid, "test_lock") + assert not ok, "断开后锁应继续保持持有" + + # 等待 worker 真正结束 + assert worker_done.wait(timeout=10), "worker 未在预期时间内结束" + + # worker 结束后,_finish_task 已在 _run() 的 finally 中执行, + # 新请求应能获取锁 + ok, _ = _try_start_task(wid, "test_lock") + assert ok, "worker 结束后应能重新获取锁" + + # 清理 + _finish_task(wid, "test_lock", _task_token) diff --git a/backend/test_organize_groups.py b/backend/test_organize_groups.py index ac11de4..1034832 100644 --- a/backend/test_organize_groups.py +++ b/backend/test_organize_groups.py @@ -327,9 +327,347 @@ def test_organize_groups_list_accounts_none_conservatively_skips(db_session, mon assert k3.status == "imported" +def test_organize_groups_streaming_basic_event_sequence(db_session, monkeypatch): + """流式一键整理:验证基本 NDJSON 事件序列、响应头、total_items 准确性。""" + w = Website( + name="W1", + site_type="sub2api", + base_url="http://w1", + enabled=True, + auth_config_json="{}", + timeout_seconds=30 + ) + u1 = Upstream(name="U1", base_url="http://u1") + db_session.add_all([w, u1]) + db_session.commit() + db_session.refresh(w) + db_session.refresh(u1) + + b1 = WebsiteGroupBinding( + website_id=w.id, + target_group_id="TG1", + target_group_name="TG1-Group", + source_groups_json=json.dumps([ + {"upstream_id": u1.id, "group_id": "G1"}, + ]), + enabled=True + ) + db_session.add(b1) + + k1 = UpstreamGeneratedKey( + upstream_id=u1.id, + group_id="G1", + group_name="G1-Name", + key_name="Key-G1", + key_value="sk-g1-secret", + status="created", + ) + db_session.add(k1) + db_session.commit() + + created_accounts = [] + + class MockClient: + def __init__(self, **kwargs): pass + def __enter__(self): return self + def __exit__(self, *a): pass + def get_groups(self, *a, **kw): + return [{"id": "TG1", "name": "TG1-Group"}] + def list_accounts(self): return [] + def extract_id(self, val): return val.get("id") if isinstance(val, dict) else str(val) + def create_account(self, body): + acc_id = f"NEW-{body['name']}" + new_acc = {"id": acc_id, "name": body["name"], "group_ids": body["group_ids"]} + created_accounts.append(new_acc) + return new_acc + + monkeypatch.setattr("app.routers.websites.Sub2ApiWebsiteClient", MockClient) + monkeypatch.setattr("app.routers.websites.sync_account_priorities_for_website", lambda db, wid: []) + + from fastapi.testclient import TestClient + from app.main import app + from app.database import get_db + from app.utils.auth import get_current_user + + app.dependency_overrides[get_db] = lambda: db_session + app.dependency_overrides[get_current_user] = lambda: None + + try: + client = TestClient(app) + resp = client.post(f"/api/websites/{w.id}/groups/organize/stream") + assert resp.status_code == 200 + assert resp.headers["content-type"] == "application/x-ndjson" + assert resp.headers.get("cache-control") == "no-cache" + assert resp.headers.get("x-accel-buffering") == "no" + + lines = [line if isinstance(line, str) else line.decode("utf-8") for line in resp.iter_lines() if line] + parsed_events = [json.loads(line) for line in lines] + + # Event sequence: start → item → complete + assert len(parsed_events) >= 2 + assert parsed_events[0]["event"] == "start" + assert parsed_events[0]["data"]["total_items"] == 1 + + item_events = [e for e in parsed_events if e["event"] == "item"] + assert len(item_events) == 1 + assert item_events[0]["data"]["source_group_id"] == "G1" + assert item_events[0]["data"]["status"] == "created" + + complete_event = next(e for e in parsed_events if e["event"] == "complete") + assert complete_event["data"]["success"] is True + finally: + app.dependency_overrides.clear() + + +def test_organize_groups_streaming_409_concurrent_lock(db_session, monkeypatch): + """流式一键整理:同一网站并发调用返回 409 错误。""" + w = Website( + name="W1", + site_type="sub2api", + base_url="http://w1", + enabled=True, + auth_config_json="{}", + timeout_seconds=30 + ) + db_session.add(w) + db_session.commit() + db_session.refresh(w) + + # Mock to make the first call run long enough to trigger concurrency + import time + original_generator = None + + class MockClient: + def __init__(self, **kwargs): pass + def __enter__(self): return self + def __exit__(self, *a): pass + def get_groups(self, *a, **kw): return [{"id": "TG1", "name": "TG1-Group"}] + def list_accounts(self): return [] + def extract_id(self, val): return val.get("id") if isinstance(val, dict) else str(val) + + monkeypatch.setattr("app.routers.websites.Sub2ApiWebsiteClient", MockClient) + monkeypatch.setattr("app.routers.websites.sync_account_priorities_for_website", lambda db, wid: []) + + from fastapi.testclient import TestClient + from app.main import app + from app.database import get_db + from app.utils.auth import get_current_user + + app.dependency_overrides[get_db] = lambda: db_session + app.dependency_overrides[get_current_user] = lambda: None + + try: + client = TestClient(app) + # First request starts the task lock + resp1 = client.post(f"/api/websites/{w.id}/groups/organize/stream") + assert resp1.status_code == 200 + + # Second request while first is still "running" (lock not yet released) + # But the first request already completed, so we need to simulate + # by calling the task lock manually. + from app.routers.websites import _try_start_task, _finish_task + ok, _task_token = _try_start_task(w.id, "organize") + assert ok + try: + resp2 = client.post(f"/api/websites/{w.id}/groups/organize/stream") + assert resp2.status_code == 200 + lines = [line if isinstance(line, str) else line.decode("utf-8") for line in resp2.iter_lines() if line] + parsed_events = [json.loads(line) for line in lines] + assert len(parsed_events) == 1 + assert parsed_events[0]["event"] == "error" + assert "正在执行中" in parsed_events[0]["data"]["message"] + finally: + _finish_task(w.id, "organize", _task_token) + + # Task lock released: third request should succeed + resp3 = client.post(f"/api/websites/{w.id}/groups/organize/stream") + assert resp3.status_code == 200 + finally: + app.dependency_overrides.clear() + + +def test_organize_groups_streaming_non_sub2api_returns_error_event(db_session, monkeypatch): + """流式一键整理:非 sub2api 网站返回 error 事件而非 HTTP 错误。""" + w = Website( + name="W1", + site_type="other", + base_url="http://w1", + enabled=True, + auth_config_json="{}", + timeout_seconds=30 + ) + db_session.add(w) + db_session.commit() + db_session.refresh(w) + + from fastapi.testclient import TestClient + from app.main import app + from app.database import get_db + from app.utils.auth import get_current_user + + app.dependency_overrides[get_db] = lambda: db_session + app.dependency_overrides[get_current_user] = lambda: None + + try: + client = TestClient(app) + resp = client.post(f"/api/websites/{w.id}/groups/organize/stream") + assert resp.status_code == 200 + lines = [line if isinstance(line, str) else line.decode("utf-8") for line in resp.iter_lines() if line] + parsed_events = [json.loads(line) for line in lines] + assert len(parsed_events) == 1 + assert parsed_events[0]["event"] == "error" + finally: + app.dependency_overrides.clear() + + +def test_organize_groups_streaming_total_items_matches_items(db_session, monkeypatch): + """流式一键整理:total_items 应当等于实际产生的 item 数量。""" + w = Website( + name="W1", + site_type="sub2api", + base_url="http://w1", + enabled=True, + auth_config_json="{}", + timeout_seconds=30 + ) + u1 = Upstream(name="U1", base_url="http://u1") + db_session.add_all([w, u1]) + db_session.commit() + db_session.refresh(w) + db_session.refresh(u1) + + # Two bindings, each with one key + b1 = WebsiteGroupBinding( + website_id=w.id, + target_group_id="TG1", + target_group_name="TG1-Group", + source_groups_json=json.dumps([{"upstream_id": u1.id, "group_id": "G1"}]), + enabled=True + ) + b2 = WebsiteGroupBinding( + website_id=w.id, + target_group_id="TG2", + target_group_name="TG2-Group", + source_groups_json=json.dumps([{"upstream_id": u1.id, "group_id": "G2"}]), + enabled=True + ) + db_session.add_all([b1, b2]) + + k1 = UpstreamGeneratedKey( + upstream_id=u1.id, group_id="G1", group_name="G1-Name", + key_name="Key-G1", key_value="sk-g1", status="created", + ) + k2 = UpstreamGeneratedKey( + upstream_id=u1.id, group_id="G2", group_name="G2-Name", + key_name="Key-G2", key_value="sk-g2", status="created", + ) + db_session.add_all([k1, k2]) + db_session.commit() + + class MockClient: + def __init__(self, **kwargs): pass + def __enter__(self): return self + def __exit__(self, *a): pass + def get_groups(self, *a, **kw): + return [{"id": "TG1", "name": "TG1-Group"}, {"id": "TG2", "name": "TG2-Group"}] + def list_accounts(self): return [] + def extract_id(self, val): return val.get("id") if isinstance(val, dict) else str(val) + def create_account(self, body): + return {"id": "NEW-" + body["name"], "name": body["name"], "group_ids": body["group_ids"]} + + monkeypatch.setattr("app.routers.websites.Sub2ApiWebsiteClient", MockClient) + monkeypatch.setattr("app.routers.websites.sync_account_priorities_for_website", lambda db, wid: []) + + from fastapi.testclient import TestClient + from app.main import app + from app.database import get_db + from app.utils.auth import get_current_user + + app.dependency_overrides[get_db] = lambda: db_session + app.dependency_overrides[get_current_user] = lambda: None + + try: + client = TestClient(app) + resp = client.post(f"/api/websites/{w.id}/groups/organize/stream") + assert resp.status_code == 200 + lines = [line if isinstance(line, str) else line.decode("utf-8") for line in resp.iter_lines() if line] + parsed_events = [json.loads(line) for line in lines] + + start_event = next(e for e in parsed_events if e["event"] == "start") + item_events = [e for e in parsed_events if e["event"] == "item"] + assert start_event["data"]["total_items"] == len(item_events) + finally: + app.dependency_overrides.clear() + + +def test_organize_groups_original_json_endpoint_still_works(db_session, monkeypatch): + """一键整理 JSON 接口的回归测试:HTTP 状态码和错误处理仍然正确。""" + # Non-sub2api → HTTP 400 + w = Website( + name="W1", site_type="other", base_url="http://w1", + enabled=True, auth_config_json="{}", timeout_seconds=30 + ) + db_session.add(w) + db_session.commit() + db_session.refresh(w) + + from fastapi.testclient import TestClient + from app.main import app + from app.database import get_db + from app.utils.auth import get_current_user + + app.dependency_overrides[get_db] = lambda: db_session + app.dependency_overrides[get_current_user] = lambda: None + + try: + client = TestClient(app) + resp = client.post(f"/api/websites/{w.id}/groups/organize") + assert resp.status_code == 400 + assert "只支持" in resp.json()["detail"] or "only sub2api" in resp.json()["detail"] + finally: + app.dependency_overrides.clear() + + +def test_organize_groups_json_409_concurrent_lock(db_session, monkeypatch): + """一键整理 JSON 接口:并发调用返回 HTTP 409。""" + w = Website( + name="W1", site_type="sub2api", base_url="http://w1", + enabled=True, auth_config_json="{}", timeout_seconds=30, + ) + db_session.add(w) + db_session.commit() + db_session.refresh(w) + + from fastapi.testclient import TestClient + from app.main import app + from app.database import get_db + from app.utils.auth import get_current_user + + app.dependency_overrides[get_db] = lambda: db_session + app.dependency_overrides[get_current_user] = lambda: None + + try: + client = TestClient(app) + # Hold lock manually to simulate concurrent request + from app.routers.websites import _try_start_task, _finish_task + ok, _task_token = _try_start_task(w.id, "organize") + assert ok + try: + resp = client.post(f"/api/websites/{w.id}/groups/organize") + assert resp.status_code == 409 + assert "正在执行中" in resp.json()["detail"] + finally: + _finish_task(w.id, "organize", _task_token) + + # Lock released: second request should succeed + resp2 = client.post(f"/api/websites/{w.id}/groups/organize") + assert resp2.status_code == 200 + finally: + app.dependency_overrides.clear() + + def test_organize_groups_force_alignment(db_session, monkeypatch): """测试强对齐场景: - 1. 已存在账号在旧 SmartUp 分组,整理后移除旧分组并加入新分组。 2. 已存在账号同时有非 SmartUp 分组,整理后保留非 SmartUp 分组。 3. 已存在账号已经完全一致,整理不重复调用更新接口。 4. 多目标分组、多上游来源时,每个账号按自己的当前绑定关系对齐。 diff --git a/backend/test_sync_upstream_models.py b/backend/test_sync_upstream_models.py index b7d242e..ba9a11e 100644 --- a/backend/test_sync_upstream_models.py +++ b/backend/test_sync_upstream_models.py @@ -408,3 +408,79 @@ def test_sync_upstream_models_empty_key_value_skips_account(db_session, monkeypa assert len(update_calls) == 0, f"不应调用 update_account,实际调用:{update_calls}" finally: app.dependency_overrides.clear() + + +def test_sync_upstream_models_concurrent_lock_json_endpoint(db_session, monkeypatch): + """同步上游模型 JSON 接口:并发调用返回 HTTP 409。""" + w = Website( + name="W1", site_type="sub2api", base_url="http://w1", + enabled=True, auth_config_json="{}", timeout_seconds=30, + ) + db_session.add(w) + db_session.commit() + db_session.refresh(w) + + from fastapi.testclient import TestClient + from app.main import app + from app.database import get_db + from app.utils.auth import get_current_user + + app.dependency_overrides[get_db] = lambda: db_session + app.dependency_overrides[get_current_user] = lambda: None + + try: + client = TestClient(app) + from app.routers.websites import _try_start_task, _finish_task + ok, _sync_token = _try_start_task(w.id, "sync_models") + assert ok + try: + resp = client.post(f"/api/websites/{w.id}/accounts/sync-upstream-models") + assert resp.status_code == 409 + assert "正在执行中" in resp.json()["detail"] + finally: + _finish_task(w.id, "sync_models", _sync_token) + + resp2 = client.post(f"/api/websites/{w.id}/accounts/sync-upstream-models") + assert resp2.status_code == 200 + finally: + app.dependency_overrides.clear() + + +def test_sync_upstream_models_concurrent_lock_stream_endpoint(db_session, monkeypatch): + """同步上游模型流式接口:并发调用返回 error 事件。""" + w = Website( + name="W1", site_type="sub2api", base_url="http://w1", + enabled=True, auth_config_json="{}", timeout_seconds=30, + ) + db_session.add(w) + db_session.commit() + db_session.refresh(w) + + from fastapi.testclient import TestClient + from app.main import app + from app.database import get_db + from app.utils.auth import get_current_user + + app.dependency_overrides[get_db] = lambda: db_session + app.dependency_overrides[get_current_user] = lambda: None + + try: + client = TestClient(app) + from app.routers.websites import _try_start_task, _finish_task + ok, _sync_token = _try_start_task(w.id, "sync_models") + assert ok + try: + resp = client.post(f"/api/websites/{w.id}/accounts/sync-upstream-models/stream") + assert resp.status_code == 200 + lines = [line if isinstance(line, str) else line.decode("utf-8") for line in resp.iter_lines() if line] + parsed_events = [json.loads(line) for line in lines] + assert len(parsed_events) == 1 + assert parsed_events[0]["event"] == "error" + assert "正在执行中" in parsed_events[0]["data"]["message"] + finally: + _finish_task(w.id, "sync_models", _sync_token) + + resp2 = client.post(f"/api/websites/{w.id}/accounts/sync-upstream-models/stream") + assert resp2.status_code == 200 + finally: + app.dependency_overrides.clear() diff --git a/frontend/src/api/index.ts b/frontend/src/api/index.ts index 48e7f9f..c5769b1 100644 --- a/frontend/src/api/index.ts +++ b/frontend/src/api/index.ts @@ -67,6 +67,117 @@ export const authApi = { me: () => api.get<{ email: string }>('/api/auth/me'), } +// ——— NDJSON 流式读取工具 ——— +export type NdjsonEvent = { + event: string + data: Record +} + +export type NdjsonStreamHandlers = { + onStart?: (data: Record) => void + onItem?: (data: Record) => void + onComplete?: (data: Record) => void + onError?: (data: Record) => void + onHeartbeat?: () => void +} + +/** + * 读取 NDJSON 流,按 event 类型分派处理器。 + * 返回 { ok, message },ok=false 表示流未正常完成(含 HTTP 错误、解析错误、断流)。 + */ +export async function readNdjsonStream( + url: string, + options: RequestInit & { token?: string }, + handlers: NdjsonStreamHandlers, +): Promise<{ ok: boolean; message: string }> { + const { token, ...fetchOpts } = options + const headers: Record = { + 'Content-Type': 'application/json', + ...((fetchOpts.headers as Record) || {}), + } + if (token) { + headers['Authorization'] = `Bearer ${token}` + } + + let response: Response + try { + response = await fetch(url, { ...fetchOpts, headers }) + } catch (e: any) { + return { ok: false, message: e.message || '网络请求失败' } + } + + if (!response.ok) { + let errMsg = `请求失败 (${response.status})` + try { + const errJson = await response.json() + if (errJson?.detail) errMsg = errJson.detail + } catch {} + return { ok: false, message: errMsg } + } + + const reader = response.body?.getReader() + if (!reader) { + return { ok: false, message: '无法读取响应流' } + } + + const decoder = new TextDecoder() + let buffer = '' + let gotComplete = false + + try { + while (true) { + const { done, value } = await reader.read() + if (done) break + + buffer += decoder.decode(value, { stream: true }) + const lines = buffer.split('\n') + buffer = lines.pop() || '' + + for (const line of lines) { + if (!line.trim()) continue + let parsed: NdjsonEvent + try { + parsed = JSON.parse(line) + } catch { + continue + } + const { event, data } = parsed + if (event === 'start') handlers.onStart?.(data) + else if (event === 'item') handlers.onItem?.(data) + else if (event === 'complete') { + handlers.onComplete?.(data) + gotComplete = true + } else if (event === 'error') handlers.onError?.(data) + else if (event === 'heartbeat') handlers.onHeartbeat?.() + } + } + + // 处理尾部缓冲 + if (buffer.trim()) { + try { + const parsed = JSON.parse(buffer) + const { event, data } = parsed + if (event === 'complete') { + handlers.onComplete?.(data) + gotComplete = true + } else if (event === 'error') { + handlers.onError?.(data) + } else if (event === 'item') { + handlers.onItem?.(data) + } + } catch {} + } + + if (!gotComplete) { + return { ok: false, message: '流异常中断,未获取到完整执行报告' } + } + } catch (e: any) { + return { ok: false, message: e.message || '流读取异常' } + } + + return { ok: true, message: '' } +} + // ——— Upstreams ——— export type FinanceCostMode = 'usage_stats' | 'balance_delta' diff --git a/frontend/src/views/Websites.vue b/frontend/src/views/Websites.vue index cfe136f..3f60971 100644 --- a/frontend/src/views/Websites.vue +++ b/frontend/src/views/Websites.vue @@ -751,55 +751,80 @@ - -
- {{ organizeMessage }} + +
+
+
+ 已处理: {{ organizeProcessedCount }}{{ organizeTotal ? '/' + organizeTotal : '' }} + 成功: {{ organizeSuccessCount }} + 跳过: {{ organizeSkippedCount }} + 失败: {{ organizeFailedCount }} +
+
+ + 正在整理分组... +
+
+ +
+ {{ organizeMessage }} +
+ + + + + + + + + + + + + + + + + + + + +
- - - - - - - - - - - - - - - - - - - -
@@ -1125,6 +1150,7 @@ import { type CleanupInvalidAccountsItem, type SetConcurrencyItem, type SyncUpstreamModelsItem, + readNdjsonStream, } from '@/api' const websites = ref<(WebsiteData & { _testing?: boolean })[]>([]) @@ -1276,6 +1302,18 @@ const organizeDialog = ref(false) const organizeLoading = ref(false) const organizeResults = ref([]) const organizeMessage = ref('') +const organizeTotal = ref(0) +const organizeSuccessCount = computed(() => organizeResults.value.filter(r => r.status === 'created' || r.status === 'recreated' || r.status === 'exists').length) +const organizeFailedCount = computed(() => organizeResults.value.filter(r => r.status === 'failed').length) +const organizeSkippedCount = computed(() => organizeResults.value.filter(r => r.status === 'missing_key').length) +const organizeProcessedCount = computed(() => organizeResults.value.length) + +function handleOrganizeDialogBeforeClose(done: () => void) { + if (organizeLoading.value) { + return + } + done() +} const cleanupDialog = ref(false) const cleanupLoading = ref(false) @@ -1971,19 +2009,42 @@ async function organizeWebsiteGroups() { return } + organizeResults.value = [] + organizeMessage.value = '' + organizeTotal.value = 0 organizeLoading.value = true - try { - const res = await websitesApi.organizeGroups(selectedWebsite.value.id) - organizeMessage.value = res.data.message - organizeResults.value = res.data.items - organizeDialog.value = true + organizeDialog.value = true + + const authStore = useAuthStore() + const { ok, message } = await readNdjsonStream( + `/api/websites/${selectedWebsite.value.id}/groups/organize/stream`, + { method: 'POST', token: authStore.token }, + { + onStart(data) { organizeTotal.value = (data.total_items as number) || 0 }, + onItem(data: any) { organizeResults.value.push(data) }, + onComplete(data) { + organizeMessage.value = (data.message as string) || '' + if (data.success) ElMessage.success('整理完成') + else ElMessage.warning((data.message as string) || '部分账号整理失败') + }, + onError(data) { throw new Error((data.message as string) || '流式整理出错') }, + }, + ) + + if (!ok) { + if (message.includes('正在执行中')) { + ElMessage.warning(message) + } else { + ElMessage.error(message || '一键整理分组失败') + } + if (organizeResults.value.length === 0) { + organizeDialog.value = false + } + } else { // 刷新相关数据 await Promise.all([loadLogs(), loadWebsiteGroups(), loadBindings()]) - } catch (e: any) { - ElMessage.error(e.response?.data?.detail || '一键整理分组失败') - } finally { - organizeLoading.value = false } + organizeLoading.value = false } async function openCleanupDialog() { @@ -2127,105 +2188,37 @@ async function triggerSyncUpstreamModels() { syncModelsDialog.value = true let hasProcessedStart = false - let gotComplete = false - try { - const authStore = useAuthStore() - const headers: Record = { - 'Content-Type': 'application/json', - } - if (authStore.token) { - headers['Authorization'] = `Bearer ${authStore.token}` - } - - const response = await fetch(`/api/websites/${selectedWebsite.value.id}/accounts/sync-upstream-models/stream`, { - method: 'POST', - headers, - }) - - if (!response.ok) { - let errMsg = `请求失败 (${response.status})` - try { - const errJson = await response.json() - if (errJson?.detail) errMsg = errJson.detail - } catch {} - throw new Error(errMsg) - } - - const reader = response.body?.getReader() - if (!reader) { - throw new Error('无法读取响应流') - } - - const decoder = new TextDecoder() - let buffer = '' - - while (true) { - const { done, value } = await reader.read() - if (done) { - break - } - - buffer += decoder.decode(value, { stream: true }) - const lines = buffer.split('\n') - buffer = lines.pop() || '' - - for (const line of lines) { - if (!line.trim()) continue - const eventObj = JSON.parse(line) - const { event, data } = eventObj - - if (event === 'start') { - syncModelsTotal.value = data.total_accounts || 0 - hasProcessedStart = true - } else if (event === 'item') { - syncModelsResults.value.push(data) - } else if (event === 'complete') { - syncModelsMessage.value = data.message - gotComplete = true - if (data.success) { - ElMessage.success('同步完成') - } else { - ElMessage.warning(data.message || '部分账号同步模型失败') - } - } else if (event === 'error') { - throw new Error(data.message || '流式同步出错') - } - } - } - - if (buffer.trim()) { - const eventObj = JSON.parse(buffer) - const { event, data } = eventObj - if (event === 'start') { - syncModelsTotal.value = data.total_accounts || 0 + const authStore = useAuthStore() + const { ok, message } = await readNdjsonStream( + `/api/websites/${selectedWebsite.value.id}/accounts/sync-upstream-models/stream`, + { method: 'POST', token: authStore.token }, + { + onStart(data) { + syncModelsTotal.value = (data.total_accounts as number) || 0 hasProcessedStart = true - } else if (event === 'item') { - syncModelsResults.value.push(data) - } else if (event === 'complete') { - syncModelsMessage.value = data.message - gotComplete = true - if (data.success) { - ElMessage.success('同步完成') - } else { - ElMessage.warning(data.message || '部分账号同步模型失败') - } - } else if (event === 'error') { - throw new Error(data.message || '流式同步出错') - } - } + }, + onItem(data: any) { syncModelsResults.value.push(data) }, + onComplete(data) { + syncModelsMessage.value = (data.message as string) || '' + if (data.success) ElMessage.success('同步完成') + else ElMessage.warning((data.message as string) || '部分账号同步模型失败') + }, + onError() { /* handled via !ok */ }, + }, + ) - if (!gotComplete) { - throw new Error('同步流异常中断,未获取到完整执行报告') + if (!ok) { + if (message.includes('正在执行中')) { + ElMessage.warning(message) + } else { + ElMessage.error(message || '同步上游模型失败') } - } catch (e: any) { - ElMessage.error(e.message || '同步上游模型失败') if (!hasProcessedStart || syncModelsResults.value.length === 0) { syncModelsDialog.value = false } - } finally { - syncModelsExecuting.value = false } + syncModelsExecuting.value = false } async function toggleBinding(row: GroupBindingData) {