fix: stream long-running website operations
This commit is contained in:
+340
-42
@@ -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,11 +1812,18 @@ 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
|
||||
|
||||
try:
|
||||
for event, data in _sync_upstream_models_generator(wid, db):
|
||||
if event == "item":
|
||||
items.append(SyncUpstreamModelsItem(**data))
|
||||
@@ -1669,18 +1833,24 @@ def sync_website_accounts_upstream_models(
|
||||
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,8 +2057,7 @@ def organize_website_groups(
|
||||
)
|
||||
|
||||
if not keys:
|
||||
items.append(
|
||||
OrganizeGroupsItem(
|
||||
_item = OrganizeGroupsItem(
|
||||
target_group_id=str(target_group_id),
|
||||
target_group_name=target_group_name,
|
||||
upstream_name=upstream_name,
|
||||
@@ -1867,7 +2069,11 @@ def organize_website_groups(
|
||||
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,8 +2082,7 @@ def organize_website_groups(
|
||||
for row in keys:
|
||||
# 跳过状态为 failed 的 Key
|
||||
if row.status == "failed":
|
||||
items.append(
|
||||
OrganizeGroupsItem(
|
||||
_item = OrganizeGroupsItem(
|
||||
target_group_id=str(target_group_id),
|
||||
target_group_name=target_group_name,
|
||||
upstream_name=upstream_name,
|
||||
@@ -1889,7 +2094,11 @@ def organize_website_groups(
|
||||
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,8 +2115,7 @@ def organize_website_groups(
|
||||
|
||||
if remote_account_ids is None:
|
||||
# 账号列表获取失败,无法校验状态,保守跳过
|
||||
items.append(
|
||||
OrganizeGroupsItem(
|
||||
_item = OrganizeGroupsItem(
|
||||
target_group_id=str(target_group_id),
|
||||
target_group_name=target_group_name,
|
||||
upstream_name=upstream_name,
|
||||
@@ -1918,7 +2126,11 @@ def organize_website_groups(
|
||||
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,8 +2201,7 @@ def organize_website_groups(
|
||||
row.imported_target_group_name = target_group_names.get(str(target_group_id))
|
||||
db.commit()
|
||||
|
||||
items.append(
|
||||
OrganizeGroupsItem(
|
||||
items.append(_item := OrganizeGroupsItem(
|
||||
target_group_id=str(target_group_id),
|
||||
target_group_name=target_group_name,
|
||||
upstream_name=upstream_name,
|
||||
@@ -2001,12 +2212,14 @@ def organize_website_groups(
|
||||
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(
|
||||
_item = OrganizeGroupsItem(
|
||||
target_group_id=str(target_group_id),
|
||||
target_group_name=target_group_name,
|
||||
upstream_name=upstream_name,
|
||||
@@ -2017,7 +2230,11 @@ def organize_website_groups(
|
||||
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,8 +2249,7 @@ def organize_website_groups(
|
||||
|
||||
# 2. 检查是否有明文 Key
|
||||
if not row.key_value:
|
||||
items.append(
|
||||
OrganizeGroupsItem(
|
||||
_item = OrganizeGroupsItem(
|
||||
target_group_id=str(target_group_id),
|
||||
target_group_name=target_group_name,
|
||||
upstream_name=upstream_name,
|
||||
@@ -2043,7 +2259,11 @@ def organize_website_groups(
|
||||
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,8 +2311,7 @@ def organize_website_groups(
|
||||
if account_id:
|
||||
remote_account_map[str(account_id)] = created
|
||||
|
||||
items.append(
|
||||
OrganizeGroupsItem(
|
||||
_item = OrganizeGroupsItem(
|
||||
target_group_id=str(target_group_id),
|
||||
target_group_name=target_group_name,
|
||||
upstream_name=upstream_name,
|
||||
@@ -2104,14 +2323,17 @@ def organize_website_groups(
|
||||
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(
|
||||
_item = OrganizeGroupsItem(
|
||||
target_group_id=str(target_group_id),
|
||||
target_group_name=target_group_name,
|
||||
upstream_name=upstream_name,
|
||||
@@ -2121,7 +2343,11 @@ def organize_website_groups(
|
||||
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",
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
@@ -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. 多目标分组、多上游来源时,每个账号按自己的当前绑定关系对齐。
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -67,6 +67,117 @@ export const authApi = {
|
||||
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 ———
|
||||
export type FinanceCostMode = 'usage_stats' | 'balance_delta'
|
||||
|
||||
|
||||
+93
-100
@@ -751,11 +751,35 @@
|
||||
</template>
|
||||
</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)">
|
||||
{{ organizeMessage }}
|
||||
</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">
|
||||
<template #default="{ row }">
|
||||
<div>{{ row.target_group_name }}</div>
|
||||
@@ -798,8 +822,9 @@
|
||||
</template>
|
||||
</el-table-column>
|
||||
</el-table>
|
||||
</div>
|
||||
<template #footer>
|
||||
<el-button @click="organizeDialog = false">关闭</el-button>
|
||||
<el-button @click="organizeDialog = false" :disabled="organizeLoading">关闭</el-button>
|
||||
</template>
|
||||
</el-dialog>
|
||||
|
||||
@@ -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<OrganizeGroupsItem[]>([])
|
||||
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
|
||||
|
||||
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<string, string> = {
|
||||
'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
|
||||
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 (buffer.trim()) {
|
||||
const eventObj = JSON.parse(buffer)
|
||||
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('同步完成')
|
||||
if (!ok) {
|
||||
if (message.includes('正在执行中')) {
|
||||
ElMessage.warning(message)
|
||||
} 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) {
|
||||
syncModelsDialog.value = false
|
||||
}
|
||||
} finally {
|
||||
syncModelsExecuting.value = false
|
||||
}
|
||||
syncModelsExecuting.value = false
|
||||
}
|
||||
|
||||
async function toggleBinding(row: GroupBindingData) {
|
||||
|
||||
Reference in New Issue
Block a user