feat: sync upstream models with real-time streaming updates
This commit is contained in:
+103
-40
@@ -1288,35 +1288,25 @@ def set_website_accounts_concurrency(
|
||||
)
|
||||
|
||||
|
||||
@router.post("/api/websites/{wid}/accounts/sync-upstream-models", response_model=SyncUpstreamModelsResponse)
|
||||
def sync_website_accounts_upstream_models(
|
||||
wid: int,
|
||||
db: Session = Depends(get_db),
|
||||
_=Depends(get_current_user),
|
||||
):
|
||||
"""一键同步上游模型"""
|
||||
def _sync_upstream_models_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, "only sub2api site supports account model syncing")
|
||||
yield "error", {"message": "only sub2api site supports account model syncing"}
|
||||
return
|
||||
|
||||
with _client(website) as c:
|
||||
try:
|
||||
remote_accounts = c.list_accounts()
|
||||
except Exception as e:
|
||||
return SyncUpstreamModelsResponse(
|
||||
success=False,
|
||||
message=f"拉取远端账号列表失败: {e}",
|
||||
items=[]
|
||||
)
|
||||
yield "error", {"message": f"拉取远端账号列表失败: {e}"}
|
||||
return
|
||||
|
||||
if remote_accounts is None:
|
||||
return SyncUpstreamModelsResponse(
|
||||
success=False,
|
||||
message="拉取远端账号列表失败,无法同步上游模型",
|
||||
items=[]
|
||||
)
|
||||
yield "error", {"message": "拉取远端账号列表失败,无法同步上游模型"}
|
||||
return
|
||||
|
||||
remote_map = {}
|
||||
for acc in remote_accounts:
|
||||
@@ -1341,11 +1331,15 @@ def sync_website_accounts_upstream_models(
|
||||
}
|
||||
|
||||
if not candidates:
|
||||
return SyncUpstreamModelsResponse(
|
||||
success=True,
|
||||
message="没有找到 SmartUp 导入的有效账号",
|
||||
items=[]
|
||||
)
|
||||
yield "start", {"total_accounts": 0}
|
||||
yield "complete", {
|
||||
"success": True,
|
||||
"message": "没有找到 SmartUp 导入的有效账号",
|
||||
"items": []
|
||||
}
|
||||
return
|
||||
|
||||
yield "start", {"total_accounts": len(candidates)}
|
||||
|
||||
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()}
|
||||
@@ -1356,14 +1350,16 @@ def sync_website_accounts_upstream_models(
|
||||
# 校验是否存在于远端
|
||||
remote_acc = remote_map.get(aid)
|
||||
if not remote_acc:
|
||||
items.append(SyncUpstreamModelsItem(
|
||||
item = SyncUpstreamModelsItem(
|
||||
account_id=aid,
|
||||
account_name=None,
|
||||
model_count=0,
|
||||
models=[],
|
||||
status="skipped",
|
||||
message="账号在远端已被删除或不存在"
|
||||
))
|
||||
)
|
||||
items.append(item)
|
||||
yield "item", item.model_dump()
|
||||
continue
|
||||
|
||||
acc_name = remote_acc.get("name")
|
||||
@@ -1372,28 +1368,32 @@ def sync_website_accounts_upstream_models(
|
||||
try:
|
||||
int(aid)
|
||||
except ValueError:
|
||||
items.append(SyncUpstreamModelsItem(
|
||||
item = SyncUpstreamModelsItem(
|
||||
account_id=aid,
|
||||
account_name=acc_name,
|
||||
model_count=0,
|
||||
models=[],
|
||||
status="skipped",
|
||||
message="账号 ID 非数字,跳过同步"
|
||||
))
|
||||
)
|
||||
items.append(item)
|
||||
yield "item", item.model_dump()
|
||||
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():
|
||||
items.append(SyncUpstreamModelsItem(
|
||||
item = SyncUpstreamModelsItem(
|
||||
account_id=aid,
|
||||
account_name=acc_name,
|
||||
model_count=0,
|
||||
models=[],
|
||||
status="failed",
|
||||
message="上游 base_url 为空,跳过处理避免继续写坏账号"
|
||||
))
|
||||
)
|
||||
items.append(item)
|
||||
yield "item", item.model_dump()
|
||||
continue
|
||||
|
||||
# 1. 先安全修复 base_url (不依赖同步模型结果)
|
||||
@@ -1407,14 +1407,16 @@ def sync_website_accounts_upstream_models(
|
||||
try:
|
||||
c.update_account(aid, {"credentials": current_creds})
|
||||
except Exception as e:
|
||||
items.append(SyncUpstreamModelsItem(
|
||||
item = SyncUpstreamModelsItem(
|
||||
account_id=aid,
|
||||
account_name=acc_name,
|
||||
model_count=0,
|
||||
models=[],
|
||||
status="failed",
|
||||
message=f"修复 Base URL 失败: {e}"
|
||||
))
|
||||
)
|
||||
items.append(item)
|
||||
yield "item", item.model_dump()
|
||||
continue
|
||||
|
||||
# 2. 调用 sub2api 同步模型
|
||||
@@ -1424,14 +1426,16 @@ def sync_website_accounts_upstream_models(
|
||||
valid_models = sorted(list(set(m.strip() for m in raw_models if m and m.strip())))
|
||||
|
||||
if not valid_models:
|
||||
items.append(SyncUpstreamModelsItem(
|
||||
item = SyncUpstreamModelsItem(
|
||||
account_id=aid,
|
||||
account_name=acc_name,
|
||||
model_count=0,
|
||||
models=[],
|
||||
status="failed",
|
||||
message="已修复 Base URL,但上游同步返回模型列表为空"
|
||||
))
|
||||
)
|
||||
items.append(item)
|
||||
yield "item", item.model_dump()
|
||||
continue
|
||||
|
||||
# 3. 模型同步成功后,在 current_creds 基础上替换 model_mapping 并写回
|
||||
@@ -1439,23 +1443,27 @@ def sync_website_accounts_upstream_models(
|
||||
updated_creds["model_mapping"] = {m: m for m in valid_models}
|
||||
c.update_account(aid, {"credentials": updated_creds})
|
||||
|
||||
items.append(SyncUpstreamModelsItem(
|
||||
item = SyncUpstreamModelsItem(
|
||||
account_id=aid,
|
||||
account_name=acc_name,
|
||||
model_count=len(valid_models),
|
||||
models=valid_models,
|
||||
status="success",
|
||||
message=f"已修复 Base URL 并同步 {len(valid_models)} 个模型"
|
||||
))
|
||||
)
|
||||
items.append(item)
|
||||
yield "item", item.model_dump()
|
||||
except Exception as e:
|
||||
items.append(SyncUpstreamModelsItem(
|
||||
item = SyncUpstreamModelsItem(
|
||||
account_id=aid,
|
||||
account_name=acc_name,
|
||||
model_count=0,
|
||||
models=[],
|
||||
status="failed",
|
||||
message=f"已修复 Base URL,但模型同步失败: {e}"
|
||||
))
|
||||
)
|
||||
items.append(item)
|
||||
yield "item", item.model_dump()
|
||||
|
||||
success_count = sum(1 for item in items if item.status == "success")
|
||||
failed_count = sum(1 for item in items if item.status == "failed")
|
||||
@@ -1471,13 +1479,68 @@ def sync_website_accounts_upstream_models(
|
||||
msg_parts.append(f"跳过 {skip_count} 个")
|
||||
|
||||
message = "同步上游模型执行完毕:" + ",".join(msg_parts)
|
||||
yield "complete", {
|
||||
"success": success,
|
||||
"message": message,
|
||||
"items": [item.model_dump() for item in items]
|
||||
}
|
||||
|
||||
|
||||
@router.post("/api/websites/{wid}/accounts/sync-upstream-models", response_model=SyncUpstreamModelsResponse)
|
||||
def sync_website_accounts_upstream_models(
|
||||
wid: int,
|
||||
db: Session = Depends(get_db),
|
||||
_=Depends(get_current_user),
|
||||
):
|
||||
"""一键同步上游模型"""
|
||||
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"]
|
||||
|
||||
if error_occurred:
|
||||
if message == "website not found":
|
||||
raise HTTPException(404, "website not found")
|
||||
if "only sub2api" in message:
|
||||
raise HTTPException(400, message)
|
||||
return SyncUpstreamModelsResponse(
|
||||
success=success,
|
||||
success=False,
|
||||
message=message,
|
||||
items=items
|
||||
items=[]
|
||||
)
|
||||
|
||||
return SyncUpstreamModelsResponse(
|
||||
success=success,
|
||||
message=message,
|
||||
items=items
|
||||
)
|
||||
|
||||
|
||||
@router.post("/api/websites/{wid}/accounts/sync-upstream-models/stream")
|
||||
def sync_website_accounts_upstream_models_stream(
|
||||
wid: int,
|
||||
db: Session = Depends(get_db),
|
||||
_=Depends(get_current_user),
|
||||
):
|
||||
"""流式一键同步上游模型"""
|
||||
from fastapi.responses import StreamingResponse
|
||||
|
||||
def event_generator():
|
||||
for event, data in _sync_upstream_models_generator(wid, db):
|
||||
yield json.dumps({"event": event, "data": data}, ensure_ascii=False) + "\n"
|
||||
|
||||
return StreamingResponse(event_generator(), media_type="application/x-ndjson")
|
||||
|
||||
|
||||
@router.post("/api/websites/{wid}/groups/organize", response_model=OrganizeGroupsResponse)
|
||||
def organize_website_groups(
|
||||
|
||||
Reference in New Issue
Block a user