fix: stream long-running website operations

This commit is contained in:
SmartUp Developer
2026-07-12 12:07:47 +08:00
parent c0b3b8276d
commit 40e75c0b51
6 changed files with 1248 additions and 273 deletions
+340 -42
View File
@@ -2,10 +2,15 @@ from __future__ import annotations
import json import json
import logging import logging
import secrets
import queue as _queue
import threading as _threading
import time
from datetime import datetime, timezone from datetime import datetime, timezone
from typing import List from typing import Any, List
from fastapi import APIRouter, Depends, HTTPException, Query, status from fastapi import APIRouter, Depends, HTTPException, Query, status
from fastapi.responses import StreamingResponse
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from app.database import get_db 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"} 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: def _mask(cfg: dict) -> dict:
masked = {} masked = {}
@@ -1480,6 +1611,7 @@ def _sync_upstream_models_generator(wid: int, db: Session):
return return
yield "start", {"total_accounts": len(candidates)} yield "start", {"total_accounts": len(candidates)}
last_event_at = time.monotonic()
upstream_ids = {cand["upstream_id"] for cand in candidates.values()} 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()} 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) items.append(item)
yield "item", item.model_dump() yield "item", item.model_dump()
if time.monotonic() - last_event_at >= 10:
yield "heartbeat", {}
last_event_at = time.monotonic()
continue continue
acc_name = remote_acc.get("name") acc_name = remote_acc.get("name")
@@ -1518,9 +1653,13 @@ def _sync_upstream_models_generator(wid: int, db: Session):
) )
items.append(item) items.append(item)
yield "item", item.model_dump() yield "item", item.model_dump()
if time.monotonic() - last_event_at >= 10:
yield "heartbeat", {}
last_event_at = time.monotonic()
continue continue
upstream = upstreams_map.get(cand["upstream_id"]) upstream = upstreams_map.get(cand["upstream_id"])
upstream_base_url = upstream.base_url if upstream else None upstream_base_url = upstream.base_url if upstream else None
if not upstream_base_url or not upstream_base_url.strip(): 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) items.append(item)
yield "item", item.model_dump() yield "item", item.model_dump()
if time.monotonic() - last_event_at >= 10:
yield "heartbeat", {}
last_event_at = time.monotonic()
continue continue
# 1. 先安全修复 base_url (不依赖同步模型结果) # 1. 先安全修复 base_url (不依赖同步模型结果)
@@ -1563,6 +1705,9 @@ def _sync_upstream_models_generator(wid: int, db: Session):
) )
items.append(item) items.append(item)
yield "item", item.model_dump() yield "item", item.model_dump()
if time.monotonic() - last_event_at >= 10:
yield "heartbeat", {}
last_event_at = time.monotonic()
continue continue
current_creds["api_key"] = local_key_value current_creds["api_key"] = local_key_value
@@ -1579,6 +1724,9 @@ def _sync_upstream_models_generator(wid: int, db: Session):
) )
items.append(item) items.append(item)
yield "item", item.model_dump() yield "item", item.model_dump()
if time.monotonic() - last_event_at >= 10:
yield "heartbeat", {}
last_event_at = time.monotonic()
continue continue
# 2. 调用 sub2api 同步模型 # 2. 调用 sub2api 同步模型
@@ -1598,6 +1746,9 @@ def _sync_upstream_models_generator(wid: int, db: Session):
) )
items.append(item) items.append(item)
yield "item", item.model_dump() yield "item", item.model_dump()
if time.monotonic() - last_event_at >= 10:
yield "heartbeat", {}
last_event_at = time.monotonic()
continue continue
# 3. 模型同步成功后,在 current_creds 基础上替换 model_mapping 并写回 # 3. 模型同步成功后,在 current_creds 基础上替换 model_mapping 并写回
@@ -1615,6 +1766,9 @@ def _sync_upstream_models_generator(wid: int, db: Session):
) )
items.append(item) items.append(item)
yield "item", item.model_dump() yield "item", item.model_dump()
if time.monotonic() - last_event_at >= 10:
yield "heartbeat", {}
last_event_at = time.monotonic()
except Exception as e: except Exception as e:
item = SyncUpstreamModelsItem( item = SyncUpstreamModelsItem(
account_id=aid, account_id=aid,
@@ -1626,6 +1780,9 @@ def _sync_upstream_models_generator(wid: int, db: Session):
) )
items.append(item) items.append(item)
yield "item", item.model_dump() 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") success_count = sum(1 for item in items if item.status == "success")
failed_count = sum(1 for item in items if item.status == "failed") failed_count = sum(1 for item in items if item.status == "failed")
@@ -1655,11 +1812,18 @@ def sync_website_accounts_upstream_models(
_=Depends(get_current_user), _=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 = [] items = []
success = False success = False
message = "" message = ""
error_occurred = False error_occurred = False
try:
for event, data in _sync_upstream_models_generator(wid, db): for event, data in _sync_upstream_models_generator(wid, db):
if event == "item": if event == "item":
items.append(SyncUpstreamModelsItem(**data)) items.append(SyncUpstreamModelsItem(**data))
@@ -1669,18 +1833,24 @@ def sync_website_accounts_upstream_models(
elif event == "error": elif event == "error":
error_occurred = True error_occurred = True
message = data["message"] message = data["message"]
finally:
_finish_task(wid, "sync_models", _sync_token)
if error_occurred: if error_occurred:
if message == "website not found": if message == "website not found":
raise HTTPException(404, "website not found") raise HTTPException(404, "website not found")
if "only sub2api" in message: if "only sub2api" in message:
raise HTTPException(400, 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( return SyncUpstreamModelsResponse(
success=False, success=False,
message=message, message=message,
items=[] 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( return SyncUpstreamModelsResponse(
success=success, success=success,
message=message, message=message,
@@ -1695,27 +1865,43 @@ def sync_website_accounts_upstream_models_stream(
_=Depends(get_current_user), _=Depends(get_current_user),
): ):
"""流式一键同步上游模型""" """流式一键同步上游模型"""
from fastapi.responses import StreamingResponse
def event_generator(): 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" 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_generator(wid: int, db: Session):
def organize_website_groups( """一键整理分组生成器,逐条产生事件供流式或批量消费。"""
wid: int,
db: Session = Depends(get_db),
_=Depends(get_current_user),
):
"""一键整理分组:按现有绑定关系导入/自愈账号"""
website = db.query(Website).filter(Website.id == wid).first() website = db.query(Website).filter(Website.id == wid).first()
if not website: if not website:
raise HTTPException(404, "website not found") yield "error", {"message": "website not found"}
return
if website.site_type != "sub2api": if website.site_type != "sub2api":
raise HTTPException(400, "目前只支持 sub2api") yield "error", {"message": "目前只支持 sub2api"}
return
# 1. 读取当前网站的所有启用绑定关系 # 1. 读取当前网站的所有启用绑定关系
bindings = ( bindings = (
@@ -1742,6 +1928,23 @@ def organize_website_groups(
items: list[OrganizeGroupsItem] = [] 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 的状态最新 # 2. 收集所有的 upstream_id 并执行对账,以保证本地 Key 的状态最新
upstream_ids = set() upstream_ids = set()
for b in bindings: for b in bindings:
@@ -1854,8 +2057,7 @@ def organize_website_groups(
) )
if not keys: if not keys:
items.append( _item = OrganizeGroupsItem(
OrganizeGroupsItem(
target_group_id=str(target_group_id), target_group_id=str(target_group_id),
target_group_name=target_group_name, target_group_name=target_group_name,
upstream_name=upstream_name, upstream_name=upstream_name,
@@ -1867,7 +2069,11 @@ def organize_website_groups(
status="missing_key", status="missing_key",
message="请先生成上游 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 continue
# 获取来源上游的 Base URL 用于账号创建 # 获取来源上游的 Base URL 用于账号创建
@@ -1876,8 +2082,7 @@ def organize_website_groups(
for row in keys: for row in keys:
# 跳过状态为 failed 的 Key # 跳过状态为 failed 的 Key
if row.status == "failed": if row.status == "failed":
items.append( _item = OrganizeGroupsItem(
OrganizeGroupsItem(
target_group_id=str(target_group_id), target_group_id=str(target_group_id),
target_group_name=target_group_name, target_group_name=target_group_name,
upstream_name=upstream_name, upstream_name=upstream_name,
@@ -1889,7 +2094,11 @@ def organize_website_groups(
status="failed", status="failed",
message="上游 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 continue
# 检测平台类型 # 检测平台类型
@@ -1906,8 +2115,7 @@ def organize_website_groups(
if remote_account_ids is None: if remote_account_ids is None:
# 账号列表获取失败,无法校验状态,保守跳过 # 账号列表获取失败,无法校验状态,保守跳过
items.append( _item = OrganizeGroupsItem(
OrganizeGroupsItem(
target_group_id=str(target_group_id), target_group_id=str(target_group_id),
target_group_name=target_group_name, target_group_name=target_group_name,
upstream_name=upstream_name, upstream_name=upstream_name,
@@ -1918,7 +2126,11 @@ def organize_website_groups(
status="failed", status="failed",
message="无法校验目标账号状态,已保守跳过", message="无法校验目标账号状态,已保守跳过",
) )
) items.append(_item)
yield "item", _item.model_dump()
if time.monotonic() - last_event_at >= 10:
yield "heartbeat", {}
last_event_at = time.monotonic()
continue continue
remote_acc = remote_account_map.get(str(old_account_id)) remote_acc = remote_account_map.get(str(old_account_id))
@@ -1989,8 +2201,7 @@ def organize_website_groups(
row.imported_target_group_name = target_group_names.get(str(target_group_id)) row.imported_target_group_name = target_group_names.get(str(target_group_id))
db.commit() db.commit()
items.append( items.append(_item := OrganizeGroupsItem(
OrganizeGroupsItem(
target_group_id=str(target_group_id), target_group_id=str(target_group_id),
target_group_name=target_group_name, target_group_name=target_group_name,
upstream_name=upstream_name, upstream_name=upstream_name,
@@ -2001,12 +2212,14 @@ def organize_website_groups(
account_name=str(remote_acc.get("name") or ""), account_name=str(remote_acc.get("name") or ""),
status="exists", status="exists",
message=msg, message=msg,
) ))
) yield "item", _item.model_dump()
if time.monotonic() - last_event_at >= 10:
yield "heartbeat", {}
last_event_at = time.monotonic()
continue continue
except Exception as exc: except Exception as exc:
items.append( _item = OrganizeGroupsItem(
OrganizeGroupsItem(
target_group_id=str(target_group_id), target_group_id=str(target_group_id),
target_group_name=target_group_name, target_group_name=target_group_name,
upstream_name=upstream_name, upstream_name=upstream_name,
@@ -2017,7 +2230,11 @@ def organize_website_groups(
status="failed", status="failed",
message=f"更新绑定失败: {exc}", 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 continue
else: else:
# 远端已删除,清理标记后进行重建 # 远端已删除,清理标记后进行重建
@@ -2032,8 +2249,7 @@ def organize_website_groups(
# 2. 检查是否有明文 Key # 2. 检查是否有明文 Key
if not row.key_value: if not row.key_value:
items.append( _item = OrganizeGroupsItem(
OrganizeGroupsItem(
target_group_id=str(target_group_id), target_group_id=str(target_group_id),
target_group_name=target_group_name, target_group_name=target_group_name,
upstream_name=upstream_name, upstream_name=upstream_name,
@@ -2043,7 +2259,11 @@ def organize_website_groups(
status="failed", status="failed",
message="该 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 continue
# 3. 创建账号并绑定当前目标分组 # 3. 创建账号并绑定当前目标分组
@@ -2091,8 +2311,7 @@ def organize_website_groups(
if account_id: if account_id:
remote_account_map[str(account_id)] = created remote_account_map[str(account_id)] = created
items.append( _item = OrganizeGroupsItem(
OrganizeGroupsItem(
target_group_id=str(target_group_id), target_group_id=str(target_group_id),
target_group_name=target_group_name, target_group_name=target_group_name,
upstream_name=upstream_name, upstream_name=upstream_name,
@@ -2104,14 +2323,17 @@ def organize_website_groups(
status="recreated" if is_recreated else "created", status="recreated" if is_recreated else "created",
message="清理后重建账号" if is_recreated else "已创建账号", 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: except Exception as exc:
logger.exception("organize create account failed website=%s key=%s", wid, row.id) logger.exception("organize create account failed website=%s key=%s", wid, row.id)
row.status = "import_failed" row.status = "import_failed"
row.error = str(exc) row.error = str(exc)
db.commit() db.commit()
items.append( _item = OrganizeGroupsItem(
OrganizeGroupsItem(
target_group_id=str(target_group_id), target_group_id=str(target_group_id),
target_group_name=target_group_name, target_group_name=target_group_name,
upstream_name=upstream_name, upstream_name=upstream_name,
@@ -2121,7 +2343,11 @@ def organize_website_groups(
status="failed", status="failed",
message=str(exc), 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") created_count = sum(1 for item in items if item.status == "created")
recreated_count = sum(1 for item in items if item.status == "recreated") 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) logger.warning("failed to sync priorities after organize for website %s: %s", wid, exc)
message = "整理完成:" + " / ".join(parts) if parts else "整理完成:无任何绑定或数据" message = "整理完成:" + " / ".join(parts) if parts else "整理完成:无任何绑定或数据"
return OrganizeGroupsResponse( yield "complete", {
success=failed_count == 0, "success": failed_count == 0,
message=message, "message": message,
items=items, "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",
},
) )
+159
View File
@@ -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)
+339 -1
View File
@@ -327,9 +327,347 @@ def test_organize_groups_list_accounts_none_conservatively_skips(db_session, mon
assert k3.status == "imported" 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): def test_organize_groups_force_alignment(db_session, monkeypatch):
"""测试强对齐场景: """测试强对齐场景:
1. 已存在账号在旧 SmartUp 分组,整理后移除旧分组并加入新分组。
2. 已存在账号同时有非 SmartUp 分组,整理后保留非 SmartUp 分组。 2. 已存在账号同时有非 SmartUp 分组,整理后保留非 SmartUp 分组。
3. 已存在账号已经完全一致,整理不重复调用更新接口。 3. 已存在账号已经完全一致,整理不重复调用更新接口。
4. 多目标分组、多上游来源时,每个账号按自己的当前绑定关系对齐。 4. 多目标分组、多上游来源时,每个账号按自己的当前绑定关系对齐。
+76
View File
@@ -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}" assert len(update_calls) == 0, f"不应调用 update_account,实际调用:{update_calls}"
finally: finally:
app.dependency_overrides.clear() 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()
+111
View File
@@ -67,6 +67,117 @@ export const authApi = {
me: () => api.get<{ email: string }>('/api/auth/me'), me: () => api.get<{ email: string }>('/api/auth/me'),
} }
// ——— NDJSON 流式读取工具 ———
export type NdjsonEvent = {
event: string
data: Record<string, unknown>
}
export type NdjsonStreamHandlers = {
onStart?: (data: Record<string, unknown>) => void
onItem?: (data: Record<string, unknown>) => void
onComplete?: (data: Record<string, unknown>) => void
onError?: (data: Record<string, unknown>) => 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<string, string> = {
'Content-Type': 'application/json',
...((fetchOpts.headers as Record<string, string>) || {}),
}
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 ——— // ——— Upstreams ———
export type FinanceCostMode = 'usage_stats' | 'balance_delta' export type FinanceCostMode = 'usage_stats' | 'balance_delta'
+93 -100
View File
@@ -751,11 +751,35 @@
</template> </template>
</el-dialog> </el-dialog>
<el-dialog v-model="organizeDialog" title="一键整理分组结果" width="850px" destroy-on-close> <el-dialog
v-model="organizeDialog"
title="一键整理分组结果"
width="850px"
destroy-on-close
:show-close="!organizeLoading"
:close-on-click-modal="!organizeLoading"
:close-on-press-escape="!organizeLoading"
:before-close="handleOrganizeDialogBeforeClose"
>
<div>
<div style="margin-bottom: 15px; display: flex; align-items: center; justify-content: space-between; font-size: 14px;">
<div>
<span style="margin-right: 15px;">已处理: <strong style="color: var(--el-color-primary);">{{ organizeProcessedCount }}{{ organizeTotal ? '/' + organizeTotal : '' }}</strong></span>
<span style="margin-right: 15px;">成功: <strong style="color: var(--el-color-success);">{{ organizeSuccessCount }}</strong></span>
<span style="margin-right: 15px;">跳过: <strong style="color: var(--el-color-warning);">{{ organizeSkippedCount }}</strong></span>
<span style="margin-right: 15px;">失败: <strong style="color: var(--el-color-danger);">{{ organizeFailedCount }}</strong></span>
</div>
<div v-if="organizeLoading" style="color: var(--el-color-info); display: flex; align-items: center;">
<span class="is-loading" style="margin-right: 5px; display: inline-flex;"><el-icon><Refresh /></el-icon></span>
正在整理分组...
</div>
</div>
<div v-if="organizeMessage" style="margin-bottom:12px; font-weight:bold; color:var(--el-color-primary)"> <div v-if="organizeMessage" style="margin-bottom:12px; font-weight:bold; color:var(--el-color-primary)">
{{ organizeMessage }} {{ organizeMessage }}
</div> </div>
<el-table :data="organizeResults" size="small" border style="width:100%">
<el-table :data="organizeResults" size="small" border style="width:100%; max-height: 400px; overflow-y: auto;">
<el-table-column label="目标分组" min-width="150"> <el-table-column label="目标分组" min-width="150">
<template #default="{ row }"> <template #default="{ row }">
<div>{{ row.target_group_name }}</div> <div>{{ row.target_group_name }}</div>
@@ -798,8 +822,9 @@
</template> </template>
</el-table-column> </el-table-column>
</el-table> </el-table>
</div>
<template #footer> <template #footer>
<el-button @click="organizeDialog = false">关闭</el-button> <el-button @click="organizeDialog = false" :disabled="organizeLoading">关闭</el-button>
</template> </template>
</el-dialog> </el-dialog>
@@ -1125,6 +1150,7 @@ import {
type CleanupInvalidAccountsItem, type CleanupInvalidAccountsItem,
type SetConcurrencyItem, type SetConcurrencyItem,
type SyncUpstreamModelsItem, type SyncUpstreamModelsItem,
readNdjsonStream,
} from '@/api' } from '@/api'
const websites = ref<(WebsiteData & { _testing?: boolean })[]>([]) const websites = ref<(WebsiteData & { _testing?: boolean })[]>([])
@@ -1276,6 +1302,18 @@ const organizeDialog = ref(false)
const organizeLoading = ref(false) const organizeLoading = ref(false)
const organizeResults = ref<OrganizeGroupsItem[]>([]) const organizeResults = ref<OrganizeGroupsItem[]>([])
const organizeMessage = 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 cleanupDialog = ref(false)
const cleanupLoading = ref(false) const cleanupLoading = ref(false)
@@ -1971,19 +2009,42 @@ async function organizeWebsiteGroups() {
return return
} }
organizeResults.value = []
organizeMessage.value = ''
organizeTotal.value = 0
organizeLoading.value = true 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()]) 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() { async function openCleanupDialog() {
@@ -2127,105 +2188,37 @@ async function triggerSyncUpstreamModels() {
syncModelsDialog.value = true syncModelsDialog.value = true
let hasProcessedStart = false let hasProcessedStart = false
let gotComplete = false
try {
const authStore = useAuthStore() const authStore = useAuthStore()
const headers: Record<string, string> = { const { ok, message } = await readNdjsonStream(
'Content-Type': 'application/json', `/api/websites/${selectedWebsite.value.id}/accounts/sync-upstream-models/stream`,
} { method: 'POST', token: authStore.token },
if (authStore.token) { {
headers['Authorization'] = `Bearer ${authStore.token}` onStart(data) {
} syncModelsTotal.value = (data.total_accounts as number) || 0
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 hasProcessedStart = true
} else if (event === 'item') { },
syncModelsResults.value.push(data) onItem(data: any) { syncModelsResults.value.push(data) },
} else if (event === 'complete') { onComplete(data) {
syncModelsMessage.value = data.message syncModelsMessage.value = (data.message as string) || ''
gotComplete = true if (data.success) ElMessage.success('同步完成')
if (data.success) { else ElMessage.warning((data.message as string) || '部分账号同步模型失败')
ElMessage.success('同步完成') },
} else { onError() { /* handled via !ok */ },
ElMessage.warning(data.message || '部分账号同步模型失败') },
} )
} else if (event === 'error') {
throw new Error(data.message || '流式同步出错')
}
}
}
if (buffer.trim()) { if (!ok) {
const eventObj = JSON.parse(buffer) if (message.includes('正在执行中')) {
const { event, data } = eventObj ElMessage.warning(message)
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 { } else {
ElMessage.warning(data.message || '部分账号同步模型失败') ElMessage.error(message || '同步上游模型失败')
} }
} else if (event === 'error') {
throw new Error(data.message || '流式同步出错')
}
}
if (!gotComplete) {
throw new Error('同步流异常中断,未获取到完整执行报告')
}
} catch (e: any) {
ElMessage.error(e.message || '同步上游模型失败')
if (!hasProcessedStart || syncModelsResults.value.length === 0) { if (!hasProcessedStart || syncModelsResults.value.length === 0) {
syncModelsDialog.value = false syncModelsDialog.value = false
} }
} finally {
syncModelsExecuting.value = false
} }
syncModelsExecuting.value = false
} }
async function toggleBinding(row: GroupBindingData) { async function toggleBinding(row: GroupBindingData) {