Files
SmartUp/backend/test_upstream_key_account_import.py
T

413 lines
14 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
import json
import sys
from pathlib import Path
import pytest
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from sqlalchemy.pool import StaticPool
sys.path.insert(0, str(Path(__file__).resolve().parent))
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 import websites as websites_router
from app.schemas.website import ImportAccountsRequest
@pytest.fixture()
def db_session():
engine = create_engine(
"sqlite://",
connect_args={"check_same_thread": False},
poolclass=StaticPool,
)
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 seed_account_import_rows(db_session):
website = Website(
name="My Sub2API",
site_type="sub2api",
base_url="http://sub2api.local",
api_prefix="/api/v1/admin",
auth_type="api_key",
auth_config_json=json.dumps({"key": "admin-key", "header": "x-api-key"}),
groups_endpoint="/groups",
group_update_endpoint="/groups/{id}",
)
upstream = Upstream(
name="Packy",
base_url="http://packy.local",
api_prefix="/api/v1",
auth_type="login_password",
auth_config_json="{}",
)
db_session.add_all([website, upstream])
db_session.commit()
db_session.refresh(website)
db_session.refresh(upstream)
generated = UpstreamGeneratedKey(
upstream_id=upstream.id,
group_id="vip",
group_name="VIP",
key_id="up-key-id",
key_name="SmartUp-VIP",
key_value="sk-upstream-generated",
masked_key="sk-u...ated",
raw_json="{}",
status="created",
)
db_session.add(generated)
db_session.commit()
db_session.refresh(generated)
return website, generated
def test_import_upstream_key_creates_sub2api_account_management_apikey(monkeypatch, db_session):
website, generated = seed_account_import_rows(db_session)
account_bodies = []
class FakeWebsiteClient:
def __init__(self, **kwargs):
self.kwargs = kwargs
def __enter__(self):
return self
def __exit__(self, exc_type, exc, tb):
return False
def create_account(self, body, endpoint="/accounts"):
account_bodies.append((endpoint, body))
return {"id": 101, "name": body["name"]}
def account_exists(self, account_id):
return True
@staticmethod
def extract_id(data):
return str(data.get("id"))
monkeypatch.setattr(websites_router, "Sub2ApiWebsiteClient", FakeWebsiteClient)
response = websites_router.import_upstream_keys_as_accounts(
website.id,
ImportAccountsRequest(
upstream_key_ids=[generated.id],
target_group_map={"vip": "7"},
account_name_prefix="SmartUp",
default_platform="anthropic",
),
db_session,
object(),
)
assert "新建 1" in response.message
assert len(account_bodies) == 1
endpoint, body = account_bodies[0]
assert endpoint == "/accounts"
assert body["type"] == "apikey"
assert body["platform"] == "anthropic"
assert body["credentials"]["api_key"] == "sk-upstream-generated"
assert body["credentials"]["base_url"] == "http://packy.local"
assert body["group_ids"] == [7]
assert body["concurrency"] == 10
assert body["priority"] == 1
assert body["credentials"]["base_url"] == "http://packy.local"
db_session.refresh(generated)
assert generated.status == "imported"
assert generated.imported_website_id == website.id
assert generated.imported_account_id == "101"
def test_import_upstream_key_idempotent_skips_already_imported(monkeypatch, db_session):
"""已导入过的 Key 再次调用不会重复创建。"""
website, generated = seed_account_import_rows(db_session)
generated.imported_website_id = website.id
generated.imported_account_id = "101"
db_session.commit()
create_called = False
class FakeWebsiteClient:
def __init__(self, **kwargs):
self.kwargs = kwargs
def __enter__(self):
return self
def __exit__(self, exc_type, exc, tb):
return False
def create_account(self, body, endpoint="/accounts"):
nonlocal create_called
create_called = True
return {"id": 999, "name": body["name"]}
def account_exists(self, account_id):
return True # 模拟远端账号仍存在
@staticmethod
def extract_id(data):
return str(data.get("id"))
monkeypatch.setattr(websites_router, "Sub2ApiWebsiteClient", FakeWebsiteClient)
response = websites_router.import_upstream_keys_as_accounts(
website.id,
ImportAccountsRequest(
upstream_key_ids=[generated.id],
target_group_map={"vip": "7"},
account_name_prefix="SmartUp",
default_platform="anthropic",
),
db_session,
object(),
)
# create_account 不应被调用
assert not create_called, "create_account should NOT be called for already imported key"
# 返回 exists
assert len(response.items) == 1
assert response.items[0].status == "exists"
assert response.items[0].account_id == "101"
assert response.items[0].message == "已导入过,已跳过"
# success 应为 true(没有 failed
assert response.success
def test_sync_imported_upstream_keys_efficient_and_correct(monkeypatch, db_session):
"""验证 sync_imported_upstream_keys 优化:只拉一次账号列表,处理 exists 和 stale_cleared。"""
from app.routers import websites as websites_router
from app.schemas.website import SyncImportStatusRequest
website, generated1 = seed_account_import_rows(db_session)
generated1.imported_website_id = website.id
generated1.imported_account_id = "101"
# 增加另一个已导入行以测试单次请求 & 多个状态
generated2 = UpstreamGeneratedKey(
upstream_id=generated1.upstream_id,
group_id="normal",
group_name="normal",
key_name="normal-key",
key_value="sk-normal",
imported_website_id=website.id,
imported_account_id="102",
status="imported",
)
db_session.add(generated2)
db_session.commit()
list_accounts_calls = 0
class FakeClient:
def __init__(self, **kwargs):
pass
def __enter__(self):
return self
def __exit__(self, *a):
return False
def list_accounts(self):
nonlocal list_accounts_calls
list_accounts_calls += 1
# 远端只有 A1 (id: 101),没有 A2 (id: 102)
return [
{"id": "101", "name": "SmartUp-A1"}
]
@staticmethod
def extract_id(data):
return str(data.get("id"))
monkeypatch.setattr(websites_router, "Sub2ApiWebsiteClient", lambda **kw: FakeClient(**kw))
response = websites_router.sync_imported_upstream_keys(
website.id, SyncImportStatusRequest(upstream_id=generated1.upstream_id),
db_session, object(),
)
# 1. 验证 list_accounts() 只被调用了 1 次
assert list_accounts_calls == 1
# 2. 检查结果
assert len(response.items) == 2
item_101 = next(item for item in response.items if item.account_id == "101")
item_102 = next(item for item in response.items if item.account_id == "102")
# 101 仍存在
assert item_101.status == "exists"
# 102 已被清除
assert item_102.status == "stale_cleared"
# 3. 校验 DB 反映
db_session.refresh(generated1)
db_session.refresh(generated2)
assert generated1.imported_account_id == "101"
assert generated2.imported_account_id is None
def test_sync_imported_upstream_keys_check_failed_on_exception(monkeypatch, db_session):
"""验证 list_accounts 报错或返回 None 时返回 check_failed,且不清空本地导入标记。"""
from app.routers import websites as websites_router
from app.schemas.website import SyncImportStatusRequest
website, generated = seed_account_import_rows(db_session)
generated.imported_website_id = website.id
generated.imported_account_id = "101"
db_session.commit()
class FakeClient:
def __init__(self, **kwargs):
pass
def __enter__(self):
return self
def __exit__(self, *a):
return False
def list_accounts(self):
raise Exception("List accounts failure")
@staticmethod
def extract_id(data):
return str(data.get("id"))
monkeypatch.setattr(websites_router, "Sub2ApiWebsiteClient", lambda **kw: FakeClient(**kw))
response = websites_router.sync_imported_upstream_keys(
website.id, SyncImportStatusRequest(upstream_id=generated.upstream_id),
db_session, object(),
)
assert len(response.items) == 1
assert response.items[0].status == "check_failed"
db_session.refresh(generated)
assert generated.imported_account_id == "101"
def test_import_rebuilds_when_remote_deleted(monkeypatch, db_session):
"""导入接口:远端账号已删除时自动清标记并重新创建。"""
from app.routers import websites as websites_router
from app.schemas.website import ImportAccountsRequest
website, generated = seed_account_import_rows(db_session)
generated.imported_website_id = website.id
generated.imported_account_id = "101"
db_session.commit()
class FakeClient:
def __init__(self, **kwargs):
self.kwargs = kwargs
def __enter__(self):
return self
def __exit__(self, *a):
return False
def account_exists(self, account_id):
return False
def create_account(self, body, endpoint="/accounts"):
return {"id": 202, "name": body["name"]}
@staticmethod
def extract_id(data):
return str(data.get("id"))
monkeypatch.setattr(websites_router, "Sub2ApiWebsiteClient", lambda **kw: FakeClient(**kw))
response = websites_router.import_upstream_keys_as_accounts(
website.id,
ImportAccountsRequest(upstream_key_ids=[generated.id], target_group_map={"vip": "7"},
account_name_prefix="SmartUp", default_platform="openai"),
db_session, object(),
)
assert len(response.items) == 1
assert response.items[0].status == "created", f"expected created, got {response.items[0].status}"
db_session.refresh(generated)
assert generated.imported_account_id == "202"
def test_import_skips_on_check_failed(monkeypatch, db_session):
"""导入接口:校验失败时保守跳过,不创建也不清除。"""
from app.routers import websites as websites_router
from app.schemas.website import ImportAccountsRequest
website, generated = seed_account_import_rows(db_session)
generated.imported_website_id = website.id
generated.imported_account_id = "101"
db_session.commit()
class FakeClient:
def __init__(self, **kwargs):
self.kwargs = kwargs
def __enter__(self):
return self
def __exit__(self, *a):
return False
def account_exists(self, account_id):
return None
def create_account(self, body, endpoint="/accounts"):
raise RuntimeError("should not be called")
@staticmethod
def extract_id(data):
return str(data.get("id"))
monkeypatch.setattr(websites_router, "Sub2ApiWebsiteClient", lambda **kw: FakeClient(**kw))
response = websites_router.import_upstream_keys_as_accounts(
website.id,
ImportAccountsRequest(upstream_key_ids=[generated.id], target_group_map={"vip": "7"},
account_name_prefix="SmartUp", default_platform="openai"),
db_session, object(),
)
assert len(response.items) == 1
assert response.items[0].status == "check_failed"
db_session.refresh(generated)
assert generated.imported_account_id == "101" # 未被清除
assert not response.success
def test_import_upstream_key_with_custom_concurrency_and_priority(monkeypatch, db_session):
website, generated = seed_account_import_rows(db_session)
account_bodies = []
class FakeWebsiteClient:
def __init__(self, **kwargs):
self.kwargs = kwargs
def __enter__(self):
return self
def __exit__(self, exc_type, exc, tb):
return False
def create_account(self, body, endpoint="/accounts"):
account_bodies.append((endpoint, body))
return {"id": 202, "name": body["name"]}
def account_exists(self, account_id):
return True
@staticmethod
def extract_id(data):
return str(data.get("id"))
monkeypatch.setattr(websites_router, "Sub2ApiWebsiteClient", FakeWebsiteClient)
websites_router.import_upstream_keys_as_accounts(
website.id,
ImportAccountsRequest(
upstream_key_ids=[generated.id],
target_group_map={"vip": "7"},
account_name_prefix="SmartUp",
default_platform="openai",
concurrency=20,
priority=5,
),
db_session,
object(),
)
assert len(account_bodies) == 1
_, body = account_bodies[0]
assert body["concurrency"] == 20
assert body["priority"] == 5
assert body["credentials"]["base_url"] == "http://packy.local"