fix: 一键同步模型接口同步修复并携带上游 base_url,防止 sub2api 账号 base_url 被清空或置为默认值
This commit is contained in:
@@ -1330,6 +1330,7 @@ def sync_website_accounts_upstream_models(
|
|||||||
if aid not in candidates:
|
if aid not in candidates:
|
||||||
candidates[aid] = {
|
candidates[aid] = {
|
||||||
"account_id": aid,
|
"account_id": aid,
|
||||||
|
"upstream_id": key.upstream_id,
|
||||||
"db_key": key
|
"db_key": key
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1340,6 +1341,9 @@ def sync_website_accounts_upstream_models(
|
|||||||
items=[]
|
items=[]
|
||||||
)
|
)
|
||||||
|
|
||||||
|
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()}
|
||||||
|
|
||||||
items = []
|
items = []
|
||||||
|
|
||||||
for aid, cand in candidates.items():
|
for aid, cand in candidates.items():
|
||||||
@@ -1372,29 +1376,56 @@ def sync_website_accounts_upstream_models(
|
|||||||
))
|
))
|
||||||
continue
|
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(
|
||||||
|
account_id=aid,
|
||||||
|
account_name=acc_name,
|
||||||
|
model_count=0,
|
||||||
|
models=[],
|
||||||
|
status="failed",
|
||||||
|
message="上游 base_url 为空,跳过处理避免继续写坏账号"
|
||||||
|
))
|
||||||
|
continue
|
||||||
|
|
||||||
|
# 1. 先安全修复 base_url (不依赖同步模型结果)
|
||||||
|
try:
|
||||||
|
c.update_account(aid, {"credentials": {"base_url": upstream_base_url}})
|
||||||
|
except Exception as e:
|
||||||
|
items.append(SyncUpstreamModelsItem(
|
||||||
|
account_id=aid,
|
||||||
|
account_name=acc_name,
|
||||||
|
model_count=0,
|
||||||
|
models=[],
|
||||||
|
status="failed",
|
||||||
|
message=f"修复 Base URL 失败: {e}"
|
||||||
|
))
|
||||||
|
continue
|
||||||
|
|
||||||
|
# 2. 调用 sub2api 同步模型
|
||||||
try:
|
try:
|
||||||
raw_models = c.sync_account_upstream_models(aid)
|
raw_models = c.sync_account_upstream_models(aid)
|
||||||
# 过滤空、去重、排序
|
# 过滤空、去重、排序
|
||||||
valid_models = sorted(list(set(m.strip() for m in raw_models if m and m.strip())))
|
valid_models = sorted(list(set(m.strip() for m in raw_models if m and m.strip())))
|
||||||
|
|
||||||
if not valid_models:
|
if not valid_models:
|
||||||
# 如果返回空模型,保守处理为失败,不清空已有的模型白名单配置
|
|
||||||
items.append(SyncUpstreamModelsItem(
|
items.append(SyncUpstreamModelsItem(
|
||||||
account_id=aid,
|
account_id=aid,
|
||||||
account_name=acc_name,
|
account_name=acc_name,
|
||||||
model_count=0,
|
model_count=0,
|
||||||
models=[],
|
models=[],
|
||||||
status="failed",
|
status="failed",
|
||||||
message="上游同步返回模型列表为空"
|
message="已修复 Base URL,但上游同步返回模型列表为空"
|
||||||
))
|
))
|
||||||
continue
|
continue
|
||||||
|
|
||||||
# 构造 model_mapping
|
# 3. 模型同步成功后,写回 base_url 和 model_mapping
|
||||||
model_mapping = {m: m for m in valid_models}
|
c.update_account(aid, {"credentials": {
|
||||||
|
"base_url": upstream_base_url,
|
||||||
# 直接发送增量字段进行更新,依赖远端的 JSONB merge 逻辑,避免覆盖/丢失敏感 credentials
|
"model_mapping": {m: m for m in valid_models}
|
||||||
c.update_account(aid, {"credentials": {"model_mapping": model_mapping}})
|
}})
|
||||||
|
|
||||||
items.append(SyncUpstreamModelsItem(
|
items.append(SyncUpstreamModelsItem(
|
||||||
account_id=aid,
|
account_id=aid,
|
||||||
@@ -1402,7 +1433,7 @@ def sync_website_accounts_upstream_models(
|
|||||||
model_count=len(valid_models),
|
model_count=len(valid_models),
|
||||||
models=valid_models,
|
models=valid_models,
|
||||||
status="success",
|
status="success",
|
||||||
message=f"成功同步 {len(valid_models)} 个模型"
|
message=f"已修复 Base URL 并同步 {len(valid_models)} 个模型"
|
||||||
))
|
))
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
items.append(SyncUpstreamModelsItem(
|
items.append(SyncUpstreamModelsItem(
|
||||||
@@ -1411,7 +1442,7 @@ def sync_website_accounts_upstream_models(
|
|||||||
model_count=0,
|
model_count=0,
|
||||||
models=[],
|
models=[],
|
||||||
status="failed",
|
status="failed",
|
||||||
message=f"同步/保存上游模型失败: {e}"
|
message=f"已修复 Base URL,但模型同步失败: {e}"
|
||||||
))
|
))
|
||||||
|
|
||||||
success_count = sum(1 for item in items if item.status == "success")
|
success_count = sum(1 for item in items if item.status == "success")
|
||||||
|
|||||||
@@ -43,6 +43,9 @@ def test_sync_upstream_models_workflow(db_session, monkeypatch):
|
|||||||
|
|
||||||
up = Upstream(id=10, name="Upstream A", base_url="http://upstream-a")
|
up = Upstream(id=10, name="Upstream A", base_url="http://upstream-a")
|
||||||
db_session.add(up)
|
db_session.add(up)
|
||||||
|
|
||||||
|
up_empty = Upstream(id=11, name="Upstream Empty", base_url="")
|
||||||
|
db_session.add(up_empty)
|
||||||
db_session.commit()
|
db_session.commit()
|
||||||
|
|
||||||
# 2. 插入本地导入的 Key 记录
|
# 2. 插入本地导入的 Key 记录
|
||||||
@@ -90,7 +93,7 @@ def test_sync_upstream_models_workflow(db_session, monkeypatch):
|
|||||||
imported_website_id=w.id,
|
imported_website_id=w.id,
|
||||||
imported_account_id="1004",
|
imported_account_id="1004",
|
||||||
)
|
)
|
||||||
# k5: 属于本站,但同步接口返回空模型(应当标记为 failed,且不覆盖)
|
# k5: 属于本站,但同步接口返回空模型(应当标记为 failed,但应先安全修复 base_url)
|
||||||
k5 = UpstreamGeneratedKey(
|
k5 = UpstreamGeneratedKey(
|
||||||
id=105,
|
id=105,
|
||||||
upstream_id=up.id,
|
upstream_id=up.id,
|
||||||
@@ -101,7 +104,7 @@ def test_sync_upstream_models_workflow(db_session, monkeypatch):
|
|||||||
imported_website_id=w.id,
|
imported_website_id=w.id,
|
||||||
imported_account_id="1005",
|
imported_account_id="1005",
|
||||||
)
|
)
|
||||||
# k6: 属于本站,但同步接口抛错(应当标记为 failed,且不覆盖)
|
# k6: 属于本站,但同步接口抛错(应当标记为 failed,但应先安全修复 base_url)
|
||||||
k6 = UpstreamGeneratedKey(
|
k6 = UpstreamGeneratedKey(
|
||||||
id=106,
|
id=106,
|
||||||
upstream_id=up.id,
|
upstream_id=up.id,
|
||||||
@@ -112,7 +115,18 @@ def test_sync_upstream_models_workflow(db_session, monkeypatch):
|
|||||||
imported_website_id=w.id,
|
imported_website_id=w.id,
|
||||||
imported_account_id="1006",
|
imported_account_id="1006",
|
||||||
)
|
)
|
||||||
db_session.add_all([k1, k2, k3, k4, k5, k6])
|
# k7: 属于本站,但对应上游 base_url 为空(应当直接标记为 failed 且不调用同步和写回)
|
||||||
|
k7 = UpstreamGeneratedKey(
|
||||||
|
id=107,
|
||||||
|
upstream_id=up_empty.id,
|
||||||
|
group_id="g1",
|
||||||
|
key_name="key-107",
|
||||||
|
key_value="val-107",
|
||||||
|
status="active",
|
||||||
|
imported_website_id=w.id,
|
||||||
|
imported_account_id="1007",
|
||||||
|
)
|
||||||
|
db_session.add_all([k1, k2, k3, k4, k5, k6, k7])
|
||||||
db_session.commit()
|
db_session.commit()
|
||||||
|
|
||||||
# Mock Sub2Api 客户端
|
# Mock Sub2Api 客户端
|
||||||
@@ -131,10 +145,11 @@ def test_sync_upstream_models_workflow(db_session, monkeypatch):
|
|||||||
|
|
||||||
def list_accounts(self):
|
def list_accounts(self):
|
||||||
return [
|
return [
|
||||||
{"id": 1001, "name": "acc-1001", "credentials": {"api_key": "k1", "model_mapping": {"old": "old"}}},
|
{"id": 1001, "name": "acc-1001", "credentials": {"api_key": "k1", "model_mapping": {"old": "old"}, "base_url": "http://default-base-url"}},
|
||||||
{"id": "abc", "name": "acc-abc", "credentials": {}},
|
{"id": "abc", "name": "acc-abc", "credentials": {}},
|
||||||
{"id": 1005, "name": "acc-1005", "credentials": {"api_key": "k5"}},
|
{"id": 1005, "name": "acc-1005", "credentials": {"api_key": "k5"}},
|
||||||
{"id": 1006, "name": "acc-1006", "credentials": {"api_key": "k6"}},
|
{"id": 1006, "name": "acc-1006", "credentials": {"api_key": "k6"}},
|
||||||
|
{"id": 1007, "name": "acc-1007", "credentials": {"api_key": "k7"}},
|
||||||
]
|
]
|
||||||
|
|
||||||
def extract_id(self, val):
|
def extract_id(self, val):
|
||||||
@@ -160,25 +175,31 @@ def test_sync_upstream_models_workflow(db_session, monkeypatch):
|
|||||||
# 验证总体结果
|
# 验证总体结果
|
||||||
assert res.success is False # 存在 failed 账号,success 应该为 False
|
assert res.success is False # 存在 failed 账号,success 应该为 False
|
||||||
assert "成功 1 个" in res.message
|
assert "成功 1 个" in res.message
|
||||||
assert "失败 2 个" in res.message
|
assert "失败 3 个" in res.message
|
||||||
assert "跳过 2 个" in res.message
|
assert "跳过 2 个" in res.message
|
||||||
|
|
||||||
items = res.items
|
items = res.items
|
||||||
assert len(items) == 5
|
assert len(items) == 6
|
||||||
|
|
||||||
# 1001 成功
|
# 1001 成功 (修复 Base URL 并同步)
|
||||||
item_1001 = next(i for i in items if i.account_id == "1001")
|
item_1001 = next(i for i in items if i.account_id == "1001")
|
||||||
assert item_1001.status == "success"
|
assert item_1001.status == "success"
|
||||||
assert item_1001.model_count == 2
|
assert item_1001.model_count == 2
|
||||||
assert item_1001.models == ["gpt-3.5-turbo", "gpt-4"]
|
assert item_1001.models == ["gpt-3.5-turbo", "gpt-4"]
|
||||||
|
assert "已修复 Base URL" in item_1001.message
|
||||||
|
|
||||||
# 1001 的 update 应当只传递增量的 model_mapping,防止敏感字段丢失
|
# 1001 的两次 update 调用检测
|
||||||
assert len(update_calls) == 1
|
calls_1001 = [call for call in update_calls if call[0] == "1001"]
|
||||||
assert update_calls[0][0] == "1001"
|
assert len(calls_1001) == 2
|
||||||
assert "api_key" not in update_calls[0][1]["credentials"]
|
# 第一次安全修复:只传递 base_url
|
||||||
assert update_calls[0][1]["credentials"]["model_mapping"] == {
|
assert calls_1001[0][1]["credentials"] == {"base_url": "http://upstream-a"}
|
||||||
"gpt-3.5-turbo": "gpt-3.5-turbo",
|
# 第二次同步成功写入:传递 base_url + model_mapping
|
||||||
"gpt-4": "gpt-4"
|
assert calls_1001[1][1]["credentials"] == {
|
||||||
|
"base_url": "http://upstream-a",
|
||||||
|
"model_mapping": {
|
||||||
|
"gpt-3.5-turbo": "gpt-3.5-turbo",
|
||||||
|
"gpt-4": "gpt-4"
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
# abc 跳过
|
# abc 跳过
|
||||||
@@ -191,15 +212,31 @@ def test_sync_upstream_models_workflow(db_session, monkeypatch):
|
|||||||
assert item_1004.status == "skipped"
|
assert item_1004.status == "skipped"
|
||||||
assert "不存在" in item_1004.message
|
assert "不存在" in item_1004.message
|
||||||
|
|
||||||
# 1005 失败 (空模型)
|
# 1005 失败 (空模型,但也应当执行了第一次安全修复的 update_account)
|
||||||
item_1005 = next(i for i in items if i.account_id == "1005")
|
item_1005 = next(i for i in items if i.account_id == "1005")
|
||||||
assert item_1005.status == "failed"
|
assert item_1005.status == "failed"
|
||||||
|
assert "已修复 Base URL" in item_1005.message
|
||||||
assert "为空" in item_1005.message
|
assert "为空" in item_1005.message
|
||||||
|
calls_1005 = [call for call in update_calls if call[0] == "1005"]
|
||||||
|
assert len(calls_1005) == 1
|
||||||
|
assert calls_1005[0][1]["credentials"] == {"base_url": "http://upstream-a"}
|
||||||
|
|
||||||
# 1006 失败 (抛错)
|
# 1006 失败 (抛错,但应当执行了第一次安全修复的 update_account)
|
||||||
item_1006 = next(i for i in items if i.account_id == "1006")
|
item_1006 = next(i for i in items if i.account_id == "1006")
|
||||||
assert item_1006.status == "failed"
|
assert item_1006.status == "failed"
|
||||||
|
assert "已修复 Base URL" in item_1006.message
|
||||||
assert "Network Error" in item_1006.message
|
assert "Network Error" in item_1006.message
|
||||||
|
calls_1006 = [call for call in update_calls if call[0] == "1006"]
|
||||||
|
assert len(calls_1006) == 1
|
||||||
|
assert calls_1006[0][1]["credentials"] == {"base_url": "http://upstream-a"}
|
||||||
|
|
||||||
|
# 1007 失败 (上游 base_url 为空,不应当调用 update_account 且不调用同步模型)
|
||||||
|
item_1007 = next(i for i in items if i.account_id == "1007")
|
||||||
|
assert item_1007.status == "failed"
|
||||||
|
assert "base_url 为空" in item_1007.message
|
||||||
|
calls_1007 = [call for call in update_calls if call[0] == "1007"]
|
||||||
|
assert len(calls_1007) == 0
|
||||||
|
assert "1007" not in sync_calls
|
||||||
|
|
||||||
assert closed_count == 1
|
assert closed_count == 1
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user