360 lines
10 KiB
Python
360 lines
10 KiB
Python
import pytest
|
|
from sqlalchemy import create_engine
|
|
from sqlalchemy.orm import sessionmaker
|
|
from sqlalchemy.pool import StaticPool
|
|
|
|
from app.database import Base
|
|
from app.models.upstream import Upstream
|
|
from app.models.upstream_key import UpstreamGeneratedKey
|
|
from app.models.website import Website
|
|
from app.routers.websites import set_website_accounts_concurrency
|
|
from app.schemas.website import SetConcurrencyRequest
|
|
from app.services.website_client import Sub2ApiWebsiteClient
|
|
|
|
@pytest.fixture()
|
|
def db_session():
|
|
engine = create_engine(
|
|
"sqlite://",
|
|
connect_args={"check_same_thread": False},
|
|
poolclass=StaticPool,
|
|
)
|
|
from app.models import admin_user, upstream, snapshot, webhook_config, notification_log, custom_page, website, revoked_token, upstream_key
|
|
Base.metadata.create_all(bind=engine)
|
|
TestingSessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine)
|
|
db = TestingSessionLocal()
|
|
try:
|
|
yield db
|
|
finally:
|
|
db.close()
|
|
Base.metadata.drop_all(bind=engine)
|
|
|
|
|
|
def test_set_website_accounts_concurrency_workflow(db_session, monkeypatch):
|
|
# 1. 创建测试用的网站与上游
|
|
w = Website(
|
|
id=1,
|
|
name="Sub2Api site",
|
|
site_type="sub2api",
|
|
base_url="http://sub2api",
|
|
api_prefix="api/v1",
|
|
auth_type="bearer",
|
|
auth_config_json='{"token": "tok1"}'
|
|
)
|
|
db_session.add(w)
|
|
|
|
up = Upstream(id=10, name="Upstream A", base_url="http://upstream-a")
|
|
db_session.add(up)
|
|
db_session.commit()
|
|
|
|
# 2. 插入本地导入的 Key 记录
|
|
# k1: 属于本站,远端正常存在的账号
|
|
k1 = UpstreamGeneratedKey(
|
|
id=101,
|
|
upstream_id=up.id,
|
|
group_id="g1",
|
|
key_name="key-101",
|
|
key_value="val-101",
|
|
status="active",
|
|
imported_website_id=w.id,
|
|
imported_account_id="1001",
|
|
)
|
|
# k2: 属于本站,但账号 ID 无法转为数字 (例如字符串 'abc')
|
|
k2 = UpstreamGeneratedKey(
|
|
id=102,
|
|
upstream_id=up.id,
|
|
group_id="g1",
|
|
key_name="key-102",
|
|
key_value="val-102",
|
|
status="active",
|
|
imported_website_id=w.id,
|
|
imported_account_id="abc",
|
|
)
|
|
# k3: 属于其他站的账号(不应受本次批量设置影响)
|
|
k3 = UpstreamGeneratedKey(
|
|
id=103,
|
|
upstream_id=up.id,
|
|
group_id="g1",
|
|
key_name="key-103",
|
|
key_value="val-103",
|
|
status="active",
|
|
imported_website_id=999, # 属于其他站
|
|
imported_account_id="1003",
|
|
)
|
|
# k4: 属于本站,但账号在远端已被删除
|
|
k4 = UpstreamGeneratedKey(
|
|
id=104,
|
|
upstream_id=up.id,
|
|
group_id="g1",
|
|
key_name="key-104",
|
|
key_value="val-104",
|
|
status="active",
|
|
imported_website_id=w.id,
|
|
imported_account_id="1004",
|
|
)
|
|
db_session.add_all([k1, k2, k3, k4])
|
|
db_session.commit()
|
|
|
|
# Mock Sub2Api 客户端
|
|
bulk_calls = []
|
|
update_calls = []
|
|
list_calls_count = 0
|
|
closed_count = 0
|
|
|
|
class MockClient:
|
|
def __enter__(self):
|
|
return self
|
|
|
|
def __exit__(self, exc_type, exc, tb):
|
|
nonlocal closed_count
|
|
closed_count += 1
|
|
return False
|
|
|
|
def list_accounts(self):
|
|
nonlocal list_calls_count
|
|
list_calls_count += 1
|
|
return [
|
|
{"id": 1001, "name": "acc-1001", "concurrency": 10},
|
|
{"id": "abc", "name": "acc-abc", "concurrency": 10},
|
|
]
|
|
|
|
def extract_id(self, val):
|
|
return str(val.get("id"))
|
|
|
|
def bulk_update_accounts(self, account_ids, body):
|
|
bulk_calls.append((account_ids, body))
|
|
|
|
def update_account(self, account_id, body):
|
|
update_calls.append((account_id, body))
|
|
|
|
monkeypatch.setattr("app.routers.websites._client", lambda row: MockClient())
|
|
|
|
# ── 1. 执行成功设置并发数,期待调用 bulk_update_accounts ──
|
|
req = SetConcurrencyRequest(concurrency=200)
|
|
res = set_website_accounts_concurrency(wid=w.id, body=req, db=db_session)
|
|
|
|
assert res.success is True
|
|
assert "成功 1 个" in res.message
|
|
assert "跳过 2 个" in res.message
|
|
|
|
items = res.items
|
|
assert len(items) == 3
|
|
|
|
item_1001 = next(i for i in items if i.account_id == "1001")
|
|
assert item_1001.status == "success"
|
|
assert item_1001.old_concurrency == 10
|
|
assert item_1001.target_concurrency == 200
|
|
|
|
item_abc = next(i for i in items if i.account_id == "abc")
|
|
assert item_abc.status == "skipped"
|
|
assert "无法转换为数字" in item_abc.message
|
|
|
|
item_1004 = next(i for i in items if i.account_id == "1004")
|
|
assert item_1004.status == "skipped"
|
|
assert "不存在" in item_1004.message
|
|
|
|
assert len(bulk_calls) == 1
|
|
assert bulk_calls[0] == (["1001"], {"concurrency": 200})
|
|
assert len(update_calls) == 0
|
|
# 验证是否调用了 __exit__ 关闭连接
|
|
assert closed_count == 1
|
|
|
|
|
|
def test_set_website_accounts_concurrency_fallback(db_session, monkeypatch):
|
|
w = Website(
|
|
id=1,
|
|
name="Sub2Api site",
|
|
site_type="sub2api",
|
|
base_url="http://sub2api",
|
|
api_prefix="api/v1",
|
|
auth_type="bearer",
|
|
auth_config_json='{"token": "tok1"}'
|
|
)
|
|
db_session.add(w)
|
|
|
|
up = Upstream(id=10, name="Upstream A", base_url="http://upstream-a")
|
|
db_session.add(up)
|
|
db_session.commit()
|
|
|
|
k1 = UpstreamGeneratedKey(
|
|
id=101,
|
|
upstream_id=up.id,
|
|
group_id="g1",
|
|
key_name="key-101",
|
|
key_value="val-101",
|
|
status="active",
|
|
imported_website_id=w.id,
|
|
imported_account_id="1001",
|
|
)
|
|
db_session.add(k1)
|
|
db_session.commit()
|
|
|
|
bulk_calls = []
|
|
update_calls = []
|
|
closed_count = 0
|
|
|
|
class MockClientError:
|
|
def __enter__(self):
|
|
return self
|
|
|
|
def __exit__(self, exc_type, exc, tb):
|
|
nonlocal closed_count
|
|
closed_count += 1
|
|
return False
|
|
|
|
def list_accounts(self):
|
|
return [{"id": 1001, "name": "acc-1001", "concurrency": 10}]
|
|
|
|
def extract_id(self, val):
|
|
return str(val.get("id"))
|
|
|
|
def bulk_update_accounts(self, account_ids, body):
|
|
bulk_calls.append((account_ids, body))
|
|
raise Exception("Bulk update is not supported by target server (404)")
|
|
|
|
def update_account(self, account_id, body):
|
|
update_calls.append((account_id, body))
|
|
|
|
monkeypatch.setattr("app.routers.websites._client", lambda row: MockClientError())
|
|
|
|
req = SetConcurrencyRequest(concurrency=150)
|
|
res = set_website_accounts_concurrency(wid=w.id, body=req, db=db_session)
|
|
|
|
assert res.success is True
|
|
assert "成功 1 个" in res.message
|
|
|
|
assert len(bulk_calls) == 1
|
|
assert len(update_calls) == 1
|
|
assert update_calls[0] == ("1001", {"concurrency": 150})
|
|
|
|
items = res.items
|
|
assert len(items) == 1
|
|
assert items[0].status == "success"
|
|
assert "逐个更新" in items[0].message
|
|
assert closed_count == 1
|
|
|
|
|
|
def test_set_website_accounts_concurrency_pool_mode_merges_credentials(db_session, monkeypatch):
|
|
w = Website(
|
|
id=1,
|
|
name="Sub2Api site",
|
|
site_type="sub2api",
|
|
base_url="http://sub2api",
|
|
api_prefix="api/v1",
|
|
auth_type="bearer",
|
|
auth_config_json='{"token": "tok1"}'
|
|
)
|
|
db_session.add(w)
|
|
|
|
up = Upstream(id=10, name="Upstream A", base_url="http://upstream-a")
|
|
db_session.add(up)
|
|
db_session.commit()
|
|
|
|
k1 = UpstreamGeneratedKey(
|
|
id=101,
|
|
upstream_id=up.id,
|
|
group_id="g1",
|
|
key_name="key-101",
|
|
key_value="val-101",
|
|
status="active",
|
|
imported_website_id=w.id,
|
|
imported_account_id="1001",
|
|
)
|
|
db_session.add(k1)
|
|
db_session.commit()
|
|
|
|
bulk_calls = []
|
|
update_calls = []
|
|
|
|
class MockClient:
|
|
def __enter__(self):
|
|
return self
|
|
|
|
def __exit__(self, exc_type, exc, tb):
|
|
return False
|
|
|
|
def list_accounts(self):
|
|
return [{
|
|
"id": 1001,
|
|
"name": "acc-1001",
|
|
"concurrency": 10,
|
|
"credentials": {
|
|
"base_url": "https://relay.example.com",
|
|
"api_key": "sk-old",
|
|
"pool_mode": False,
|
|
},
|
|
}]
|
|
|
|
def extract_id(self, val):
|
|
return str(val.get("id"))
|
|
|
|
def bulk_update_accounts(self, account_ids, body):
|
|
bulk_calls.append((account_ids, body))
|
|
|
|
def update_account(self, account_id, body):
|
|
update_calls.append((account_id, body))
|
|
|
|
monkeypatch.setattr("app.routers.websites._client", lambda row: MockClient())
|
|
|
|
req = SetConcurrencyRequest(
|
|
concurrency=80,
|
|
configure_pool_mode=True,
|
|
pool_mode=True,
|
|
pool_mode_retry_count=1,
|
|
pool_mode_retry_status_codes=[529, 429, 999, 502, 429],
|
|
)
|
|
res = set_website_accounts_concurrency(wid=w.id, body=req, db=db_session)
|
|
|
|
assert res.success is True
|
|
assert "设置账号参数执行完毕" in res.message
|
|
assert bulk_calls == []
|
|
assert update_calls == [("1001", {
|
|
"concurrency": 80,
|
|
"credentials": {
|
|
"base_url": "https://relay.example.com",
|
|
"api_key": "sk-old",
|
|
"pool_mode": True,
|
|
"pool_mode_retry_count": 1,
|
|
"pool_mode_retry_status_codes": [429, 502, 529],
|
|
},
|
|
})]
|
|
assert res.items[0].old_pool_mode is False
|
|
assert res.items[0].target_pool_mode is True
|
|
|
|
|
|
def test_set_website_accounts_concurrency_list_accounts_none(db_session, monkeypatch):
|
|
w = Website(
|
|
id=1,
|
|
name="Sub2Api site",
|
|
site_type="sub2api",
|
|
base_url="http://sub2api",
|
|
api_prefix="api/v1",
|
|
auth_type="bearer",
|
|
auth_config_json='{"token": "tok1"}'
|
|
)
|
|
db_session.add(w)
|
|
db_session.commit()
|
|
|
|
closed_count = 0
|
|
|
|
class MockClientNone:
|
|
def __enter__(self):
|
|
return self
|
|
|
|
def __exit__(self, exc_type, exc, tb):
|
|
nonlocal closed_count
|
|
closed_count += 1
|
|
return False
|
|
|
|
def list_accounts(self):
|
|
return None
|
|
|
|
monkeypatch.setattr("app.routers.websites._client", lambda row: MockClientNone())
|
|
|
|
req = SetConcurrencyRequest(concurrency=150)
|
|
res = set_website_accounts_concurrency(wid=w.id, body=req, db=db_session)
|
|
|
|
assert res.success is False
|
|
assert "拉取远端账号列表失败" in res.message
|
|
assert len(res.items) == 0
|
|
assert closed_count == 1
|