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
+423 -125
View File
@@ -2,10 +2,15 @@ from __future__ import annotations
import json
import logging
import secrets
import queue as _queue
import threading as _threading
import time
from datetime import datetime, timezone
from typing import List
from typing import Any, List
from fastapi import APIRouter, Depends, HTTPException, Query, status
from fastapi.responses import StreamingResponse
from sqlalchemy.orm import Session
from app.database import get_db
@@ -75,6 +80,132 @@ SENSITIVE_CREDENTIAL_KEYS = {
}
ALGORITHMS = {"max_plus_percent", "average_plus_percent", "min_plus_percent", "priority_weighted_plus_percent"}
# ——— 服务端任务锁(防止同一网站同类任务并发执行) ———
_running_tasks: dict[str, dict[str, Any]] = {}
_tasks_lock = _threading.Lock()
def _try_start_task(wid: int, task_type: str, timeout_minutes: int = 240) -> tuple[bool, str | None]:
"""尝试获取任务锁,成功返回 (True, token),失败返回 (False, None)。
token 用于 _finish_task 校验锁所有权:只有当前锁记录的 token 与传入 token
一致时才释放锁。防止旧 worker(强制过期后仍存活的线程)误删新任务的锁。
"""
key = f"task:{task_type}:{wid}"
now = time.monotonic()
token = secrets.token_hex(16)
with _tasks_lock:
entry = _running_tasks.get(key)
if entry and (now - entry["started_at"]) < timeout_minutes * 60:
return False, None
if entry:
logger.warning("task lock force-expired wid=%s task_type=%s elapsed=%.1fmin", wid, task_type, (now - entry["started_at"]) / 60)
_running_tasks[key] = {"started_at": now, "token": token}
return True, token
def _finish_task(wid: int, task_type: str, token: str | None):
"""释放任务锁,但仅当 token 与当前锁所有者一致时执行。
如果锁已被其他任务强制过期并替换,token 不匹配,则仅记 warning
不删除新任务的锁。
"""
key = f"task:{task_type}:{wid}"
with _tasks_lock:
entry = _running_tasks.get(key)
if entry is None:
return
if entry.get("token") != token:
logger.warning("task lock not released: token mismatch (stale worker) wid=%s task_type=%s", wid, task_type)
return
_running_tasks.pop(key, None)
# ——— 心跳包装器:在同步生成器执行期间确保定期有事件输出 ———
def _with_background_heartbeat(gen_fn, interval: float = 10, timeout_factor: int = 3, *, task_wid: int | None = None, task_type: str | None = None, task_token: str | None = None, **gen_kwargs):
"""在后台线程运行 gen_fn(**gen_kwargs),主线程通过队列消费事件。
独立心跳线程每 *interval* 秒检查一次,若主线程消费落后(队列为空)则输出
heartbeat 事件。防止在单个远端请求阻塞期间无事件产生。
任务锁生命周期(解决客户端断开后锁提前释放的问题):
- 任务锁的获取(_try_start_task)在路由处理函数中完成。
- 任务锁的释放(_finish_task)在 _run() 的 finally 中,由业务线程
自身执行。因此锁的释放与业务线程的结束严格绑定。
- 客户端断开只触发 cancel 信号和队列停止,不会调用 _finish_task。
- 即使业务线程阻塞在远端 HTTP 请求中,锁仍然保持持有状态,
直到请求返回(超时或完成)且线程退出。
- 取消信号拦截在下一次 Python 代码执行边界(队列写/循环头),
无法中断阻塞中的系统调用,但锁不会提前释放。
队列满保护:
- q.put(timeout=5) + 循环重试,避免队列满后后台线程永久阻塞。
重试时检查 cancel,允许线程在被阻塞时响应取消。
"""
q = _queue.Queue(maxsize=100)
stop = _threading.Event()
cancel = _threading.Event()
def _run():
try:
for event, data in gen_fn(**gen_kwargs):
if cancel.is_set():
return
# 使用超时 put 防止队列满后永久阻塞
while True:
try:
q.put((event, data), timeout=5)
break
except _queue.Full:
if cancel.is_set():
return
except Exception as exc:
try:
q.put(("error", {"message": str(exc)}), timeout=2)
except _queue.Full:
pass
finally:
stop.set()
# 业务线程结束后才释放任务锁,并通过 token 校验所有权,
# 防止过期后仍存活的旧 worker 误删新任务锁。
if task_wid is not None and task_type is not None:
_finish_task(task_wid, task_type, task_token)
def _heartbeat():
while not stop.wait(timeout=interval):
if cancel.is_set():
return
try:
q.put(("heartbeat", {}), timeout=2)
except _queue.Full:
pass
t_work = _threading.Thread(target=_run, daemon=True)
t_hb = _threading.Thread(target=_heartbeat, daemon=True)
t_work.start()
t_hb.start()
try:
while True:
try:
event, data = q.get(timeout=interval * timeout_factor)
yield event, data
if event in ("complete", "error"):
break
except _queue.Empty:
cancel.set()
yield "error", {"message": "内部处理线程长时间无响应,任务异常中断"}
break
except GeneratorExit:
# 客户端断开连接或消费方异常退出 → 通知后台线程停止
cancel.set()
raise
finally:
cancel.set()
stop.set()
# 注意:不在此处释放任务锁。_finish_task 由 _run() 的 finally 负责,
# 确保锁释放与业务线程终止严格同步。
def _mask(cfg: dict) -> dict:
masked = {}
@@ -1480,6 +1611,7 @@ def _sync_upstream_models_generator(wid: int, db: Session):
return
yield "start", {"total_accounts": len(candidates)}
last_event_at = time.monotonic()
upstream_ids = {cand["upstream_id"] for cand in candidates.values()}
upstreams_map = {up.id: up for up in db.query(Upstream).filter(Upstream.id.in_(upstream_ids)).all()}
@@ -1500,6 +1632,9 @@ def _sync_upstream_models_generator(wid: int, db: Session):
)
items.append(item)
yield "item", item.model_dump()
if time.monotonic() - last_event_at >= 10:
yield "heartbeat", {}
last_event_at = time.monotonic()
continue
acc_name = remote_acc.get("name")
@@ -1518,9 +1653,13 @@ def _sync_upstream_models_generator(wid: int, db: Session):
)
items.append(item)
yield "item", item.model_dump()
if time.monotonic() - last_event_at >= 10:
yield "heartbeat", {}
last_event_at = time.monotonic()
continue
upstream = upstreams_map.get(cand["upstream_id"])
upstream_base_url = upstream.base_url if upstream else None
if not upstream_base_url or not upstream_base_url.strip():
@@ -1534,6 +1673,9 @@ def _sync_upstream_models_generator(wid: int, db: Session):
)
items.append(item)
yield "item", item.model_dump()
if time.monotonic() - last_event_at >= 10:
yield "heartbeat", {}
last_event_at = time.monotonic()
continue
# 1. 先安全修复 base_url (不依赖同步模型结果)
@@ -1563,6 +1705,9 @@ def _sync_upstream_models_generator(wid: int, db: Session):
)
items.append(item)
yield "item", item.model_dump()
if time.monotonic() - last_event_at >= 10:
yield "heartbeat", {}
last_event_at = time.monotonic()
continue
current_creds["api_key"] = local_key_value
@@ -1579,6 +1724,9 @@ def _sync_upstream_models_generator(wid: int, db: Session):
)
items.append(item)
yield "item", item.model_dump()
if time.monotonic() - last_event_at >= 10:
yield "heartbeat", {}
last_event_at = time.monotonic()
continue
# 2. 调用 sub2api 同步模型
@@ -1598,6 +1746,9 @@ def _sync_upstream_models_generator(wid: int, db: Session):
)
items.append(item)
yield "item", item.model_dump()
if time.monotonic() - last_event_at >= 10:
yield "heartbeat", {}
last_event_at = time.monotonic()
continue
# 3. 模型同步成功后,在 current_creds 基础上替换 model_mapping 并写回
@@ -1615,6 +1766,9 @@ def _sync_upstream_models_generator(wid: int, db: Session):
)
items.append(item)
yield "item", item.model_dump()
if time.monotonic() - last_event_at >= 10:
yield "heartbeat", {}
last_event_at = time.monotonic()
except Exception as e:
item = SyncUpstreamModelsItem(
account_id=aid,
@@ -1626,6 +1780,9 @@ def _sync_upstream_models_generator(wid: int, db: Session):
)
items.append(item)
yield "item", item.model_dump()
if time.monotonic() - last_event_at >= 10:
yield "heartbeat", {}
last_event_at = time.monotonic()
success_count = sum(1 for item in items if item.status == "success")
failed_count = sum(1 for item in items if item.status == "failed")
@@ -1655,32 +1812,45 @@ def sync_website_accounts_upstream_models(
_=Depends(get_current_user),
):
"""一键同步上游模型"""
ok, _sync_token = _try_start_task(wid, "sync_models")
if not ok:
raise HTTPException(409, "该网站的同步上游模型正在执行中")
logger.info("sync_models start wid=%s", wid)
t0 = time.monotonic()
items = []
success = False
message = ""
error_occurred = False
for event, data in _sync_upstream_models_generator(wid, db):
if event == "item":
items.append(SyncUpstreamModelsItem(**data))
elif event == "complete":
success = data["success"]
message = data["message"]
elif event == "error":
error_occurred = True
message = data["message"]
try:
for event, data in _sync_upstream_models_generator(wid, db):
if event == "item":
items.append(SyncUpstreamModelsItem(**data))
elif event == "complete":
success = data["success"]
message = data["message"]
elif event == "error":
error_occurred = True
message = data["message"]
finally:
_finish_task(wid, "sync_models", _sync_token)
if error_occurred:
if message == "website not found":
raise HTTPException(404, "website not found")
if "only sub2api" in message:
raise HTTPException(400, message)
elapsed = time.monotonic() - t0
logger.info("sync_models complete wid=%s error=%s elapsed=%.1fs", wid, message, elapsed)
return SyncUpstreamModelsResponse(
success=False,
message=message,
items=[]
)
elapsed = time.monotonic() - t0
logger.info("sync_models complete wid=%s items=%d success=%s elapsed=%.1fs", wid, len(items), success, elapsed)
return SyncUpstreamModelsResponse(
success=success,
message=message,
@@ -1695,27 +1865,43 @@ def sync_website_accounts_upstream_models_stream(
_=Depends(get_current_user),
):
"""流式一键同步上游模型"""
from fastapi.responses import StreamingResponse
def event_generator():
for event, data in _sync_upstream_models_generator(wid, db):
ok, _sync_token = _try_start_task(wid, "sync_models")
if not ok:
yield json.dumps({"event": "error", "data": {"message": "该网站的同步上游模型正在执行中"}}, ensure_ascii=False) + "\n"
return
logger.info("sync_models/stream start wid=%s", wid)
t0 = time.monotonic()
for event, data in _with_background_heartbeat(
_sync_upstream_models_generator, wid=wid, db=db,
task_wid=wid, task_type="sync_models", task_token=_sync_token,
):
yield json.dumps({"event": event, "data": data}, ensure_ascii=False) + "\n"
return StreamingResponse(event_generator(), media_type="application/x-ndjson")
elapsed = time.monotonic() - t0
logger.info("sync_models/stream complete wid=%s elapsed=%.1fs", wid, elapsed)
return StreamingResponse(
event_generator(),
media_type="application/x-ndjson",
headers={
"Cache-Control": "no-cache",
"X-Accel-Buffering": "no",
},
)
@router.post("/api/websites/{wid}/groups/organize", response_model=OrganizeGroupsResponse)
def organize_website_groups(
wid: int,
db: Session = Depends(get_db),
_=Depends(get_current_user),
):
"""一键整理分组:按现有绑定关系导入/自愈账号"""
def _organize_website_groups_generator(wid: int, db: Session):
"""一键整理分组生成器,逐条产生事件供流式或批量消费。"""
website = db.query(Website).filter(Website.id == wid).first()
if not website:
raise HTTPException(404, "website not found")
yield "error", {"message": "website not found"}
return
if website.site_type != "sub2api":
raise HTTPException(400, "目前只支持 sub2api")
yield "error", {"message": "目前只支持 sub2api"}
return
# 1. 读取当前网站的所有启用绑定关系
bindings = (
@@ -1742,6 +1928,23 @@ def organize_website_groups(
items: list[OrganizeGroupsItem] = []
# 预计算总处理项数(用于流式进度);missing_key 也算 1 项
total_items = 0
for b in bindings:
for src in binding_sources(b):
uid = src.get("upstream_id")
gid = src.get("group_id")
if uid and gid:
cnt = db.query(UpstreamGeneratedKey).filter(
UpstreamGeneratedKey.upstream_id == uid,
UpstreamGeneratedKey.group_id == gid,
UpstreamGeneratedKey.status != "orphaned",
).count()
total_items += max(1, cnt)
yield "start", {"total_items": total_items, "total_bindings": len(bindings)}
last_event_at = time.monotonic()
# 2. 收集所有的 upstream_id 并执行对账,以保证本地 Key 的状态最新
upstream_ids = set()
for b in bindings:
@@ -1854,20 +2057,23 @@ def organize_website_groups(
)
if not keys:
items.append(
OrganizeGroupsItem(
target_group_id=str(target_group_id),
target_group_name=target_group_name,
upstream_name=upstream_name,
source_group_id=str(gid),
source_group_name=str(source_group_name),
key_name="",
account_id=None,
account_name=None,
status="missing_key",
message="请先生成上游 Key",
)
_item = OrganizeGroupsItem(
target_group_id=str(target_group_id),
target_group_name=target_group_name,
upstream_name=upstream_name,
source_group_id=str(gid),
source_group_name=str(source_group_name),
key_name="",
account_id=None,
account_name=None,
status="missing_key",
message="请先生成上游 Key",
)
items.append(_item)
yield "item", _item.model_dump()
if time.monotonic() - last_event_at >= 10:
yield "heartbeat", {}
last_event_at = time.monotonic()
continue
# 获取来源上游的 Base URL 用于账号创建
@@ -1876,20 +2082,23 @@ def organize_website_groups(
for row in keys:
# 跳过状态为 failed 的 Key
if row.status == "failed":
items.append(
OrganizeGroupsItem(
target_group_id=str(target_group_id),
target_group_name=target_group_name,
upstream_name=upstream_name,
source_group_id=str(gid),
source_group_name=str(source_group_name),
key_name=row.key_name,
account_id=None,
account_name=None,
status="failed",
message="上游 Key 生成状态为失败",
)
_item = OrganizeGroupsItem(
target_group_id=str(target_group_id),
target_group_name=target_group_name,
upstream_name=upstream_name,
source_group_id=str(gid),
source_group_name=str(source_group_name),
key_name=row.key_name,
account_id=None,
account_name=None,
status="failed",
message="上游 Key 生成状态为失败",
)
items.append(_item)
yield "item", _item.model_dump()
if time.monotonic() - last_event_at >= 10:
yield "heartbeat", {}
last_event_at = time.monotonic()
continue
# 检测平台类型
@@ -1906,19 +2115,22 @@ def organize_website_groups(
if remote_account_ids is None:
# 账号列表获取失败,无法校验状态,保守跳过
items.append(
OrganizeGroupsItem(
target_group_id=str(target_group_id),
target_group_name=target_group_name,
upstream_name=upstream_name,
source_group_id=str(gid),
source_group_name=str(source_group_name),
key_name=row.key_name,
account_id=old_account_id,
status="failed",
message="无法校验目标账号状态,已保守跳过",
)
_item = OrganizeGroupsItem(
target_group_id=str(target_group_id),
target_group_name=target_group_name,
upstream_name=upstream_name,
source_group_id=str(gid),
source_group_name=str(source_group_name),
key_name=row.key_name,
account_id=old_account_id,
status="failed",
message="无法校验目标账号状态,已保守跳过",
)
items.append(_item)
yield "item", _item.model_dump()
if time.monotonic() - last_event_at >= 10:
yield "heartbeat", {}
last_event_at = time.monotonic()
continue
remote_acc = remote_account_map.get(str(old_account_id))
@@ -1989,35 +2201,40 @@ def organize_website_groups(
row.imported_target_group_name = target_group_names.get(str(target_group_id))
db.commit()
items.append(
OrganizeGroupsItem(
target_group_id=str(target_group_id),
target_group_name=target_group_name,
upstream_name=upstream_name,
source_group_id=str(gid),
source_group_name=str(source_group_name),
key_name=row.key_name,
account_id=old_account_id,
account_name=str(remote_acc.get("name") or ""),
status="exists",
message=msg,
)
)
items.append(_item := OrganizeGroupsItem(
target_group_id=str(target_group_id),
target_group_name=target_group_name,
upstream_name=upstream_name,
source_group_id=str(gid),
source_group_name=str(source_group_name),
key_name=row.key_name,
account_id=old_account_id,
account_name=str(remote_acc.get("name") or ""),
status="exists",
message=msg,
))
yield "item", _item.model_dump()
if time.monotonic() - last_event_at >= 10:
yield "heartbeat", {}
last_event_at = time.monotonic()
continue
except Exception as exc:
items.append(
OrganizeGroupsItem(
target_group_id=str(target_group_id),
target_group_name=target_group_name,
upstream_name=upstream_name,
source_group_id=str(gid),
source_group_name=str(source_group_name),
key_name=row.key_name,
account_id=old_account_id,
status="failed",
message=f"更新绑定失败: {exc}",
)
_item = OrganizeGroupsItem(
target_group_id=str(target_group_id),
target_group_name=target_group_name,
upstream_name=upstream_name,
source_group_id=str(gid),
source_group_name=str(source_group_name),
key_name=row.key_name,
account_id=old_account_id,
status="failed",
message=f"更新绑定失败: {exc}",
)
items.append(_item)
yield "item", _item.model_dump()
if time.monotonic() - last_event_at >= 10:
yield "heartbeat", {}
last_event_at = time.monotonic()
continue
else:
# 远端已删除,清理标记后进行重建
@@ -2032,18 +2249,21 @@ def organize_website_groups(
# 2. 检查是否有明文 Key
if not row.key_value:
items.append(
OrganizeGroupsItem(
target_group_id=str(target_group_id),
target_group_name=target_group_name,
upstream_name=upstream_name,
source_group_id=str(gid),
source_group_name=str(source_group_name),
key_name=row.key_name,
status="failed",
message="该 Key 无明文值,无法导入",
)
_item = OrganizeGroupsItem(
target_group_id=str(target_group_id),
target_group_name=target_group_name,
upstream_name=upstream_name,
source_group_id=str(gid),
source_group_name=str(source_group_name),
key_name=row.key_name,
status="failed",
message="该 Key 无明文值,无法导入",
)
items.append(_item)
yield "item", _item.model_dump()
if time.monotonic() - last_event_at >= 10:
yield "heartbeat", {}
last_event_at = time.monotonic()
continue
# 3. 创建账号并绑定当前目标分组
@@ -2091,37 +2311,43 @@ def organize_website_groups(
if account_id:
remote_account_map[str(account_id)] = created
items.append(
OrganizeGroupsItem(
target_group_id=str(target_group_id),
target_group_name=target_group_name,
upstream_name=upstream_name,
source_group_id=str(gid),
source_group_name=str(source_group_name),
key_name=row.key_name,
account_id=account_id or None,
account_name=str(created.get("name") or account_name),
status="recreated" if is_recreated else "created",
message="清理后重建账号" if is_recreated else "已创建账号",
)
_item = OrganizeGroupsItem(
target_group_id=str(target_group_id),
target_group_name=target_group_name,
upstream_name=upstream_name,
source_group_id=str(gid),
source_group_name=str(source_group_name),
key_name=row.key_name,
account_id=account_id or None,
account_name=str(created.get("name") or account_name),
status="recreated" if is_recreated else "created",
message="清理后重建账号" if is_recreated else "已创建账号",
)
items.append(_item)
yield "item", _item.model_dump()
if time.monotonic() - last_event_at >= 10:
yield "heartbeat", {}
last_event_at = time.monotonic()
except Exception as exc:
logger.exception("organize create account failed website=%s key=%s", wid, row.id)
row.status = "import_failed"
row.error = str(exc)
db.commit()
items.append(
OrganizeGroupsItem(
target_group_id=str(target_group_id),
target_group_name=target_group_name,
upstream_name=upstream_name,
source_group_id=str(gid),
source_group_name=str(source_group_name),
key_name=row.key_name,
status="failed",
message=str(exc),
)
_item = OrganizeGroupsItem(
target_group_id=str(target_group_id),
target_group_name=target_group_name,
upstream_name=upstream_name,
source_group_id=str(gid),
source_group_name=str(source_group_name),
key_name=row.key_name,
status="failed",
message=str(exc),
)
items.append(_item)
yield "item", _item.model_dump()
if time.monotonic() - last_event_at >= 10:
yield "heartbeat", {}
last_event_at = time.monotonic()
created_count = sum(1 for item in items if item.status == "created")
recreated_count = sum(1 for item in items if item.status == "recreated")
@@ -2151,10 +2377,82 @@ def organize_website_groups(
logger.warning("failed to sync priorities after organize for website %s: %s", wid, exc)
message = "整理完成:" + " / ".join(parts) if parts else "整理完成:无任何绑定或数据"
return OrganizeGroupsResponse(
success=failed_count == 0,
message=message,
items=items,
yield "complete", {
"success": failed_count == 0,
"message": message,
"items": [item.model_dump() for item in items],
}
@router.post("/api/websites/{wid}/groups/organize", response_model=OrganizeGroupsResponse)
def organize_website_groups(
wid: int,
db: Session = Depends(get_db),
_=Depends(get_current_user),
):
"""一键整理分组:按现有绑定关系导入/自愈账号(JSON 接口,兼容前端和测试)"""
ok, _organize_token = _try_start_task(wid, "organize")
if not ok:
raise HTTPException(409, "该网站的一键整理正在执行中")
logger.info("organize start wid=%s", wid)
t0 = time.monotonic()
items: list[OrganizeGroupsItem] = []
message = ""
success = False
try:
for event, data in _organize_website_groups_generator(wid, db):
if event == "item":
items.append(OrganizeGroupsItem(**data))
elif event == "complete":
success = data["success"]
message = data["message"]
elif event == "error":
if data["message"] == "website not found":
raise HTTPException(404, "website not found")
if "only sub2api" in data["message"] or "目前只支持" in data["message"]:
raise HTTPException(400, data["message"])
return OrganizeGroupsResponse(success=False, message=data["message"], items=[])
finally:
_finish_task(wid, "organize", _organize_token)
elapsed = time.monotonic() - t0
logger.info("organize complete wid=%s items=%d success=%s elapsed=%.1fs", wid, len(items), success, elapsed)
return OrganizeGroupsResponse(success=success, message=message, items=items)
@router.post("/api/websites/{wid}/groups/organize/stream")
def organize_website_groups_stream(
wid: int,
db: Session = Depends(get_db),
_=Depends(get_current_user),
):
"""流式一键整理分组"""
def event_generator():
ok, _organize_token = _try_start_task(wid, "organize")
if not ok:
yield json.dumps({"event": "error", "data": {"message": "该网站的一键整理正在执行中"}}, ensure_ascii=False) + "\n"
return
logger.info("organize/stream start wid=%s", wid)
t0 = time.monotonic()
for event, data in _with_background_heartbeat(
_organize_website_groups_generator, wid=wid, db=db,
task_wid=wid, task_type="organize", task_token=_organize_token,
):
yield json.dumps({"event": event, "data": data}, ensure_ascii=False) + "\n"
elapsed = time.monotonic() - t0
logger.info("organize/stream complete wid=%s elapsed=%.1fs", wid, elapsed)
return StreamingResponse(
event_generator(),
media_type="application/x-ndjson",
headers={
"Cache-Control": "no-cache",
"X-Accel-Buffering": "no",
},
)