Add upstream_key_account_links table mapping generated keys to remote accounts per website/platform, surface imported_accounts on key responses, and update website sync/routers to manage the links.
297 lines
11 KiB
Python
297 lines
11 KiB
Python
"""Grok / XAI 平台识别测试。"""
|
|
import json
|
|
|
|
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.snapshot import UpstreamRateSnapshot
|
|
from app.models.website import Website, WebsiteGroupBinding
|
|
from app.routers.websites import (
|
|
_normalize_platform,
|
|
organize_website_groups,
|
|
import_upstream_keys_as_accounts,
|
|
)
|
|
from app.schemas.website import ImportAccountsRequest
|
|
|
|
|
|
# ── 单元测试:_normalize_platform ─────────────────────
|
|
|
|
def test_normalize_platform_xai_to_grok():
|
|
assert _normalize_platform("xai") == "grok"
|
|
|
|
|
|
def test_normalize_platform_xai_mixed_case():
|
|
assert _normalize_platform("XAI") == "grok"
|
|
|
|
|
|
def test_normalize_platform_grok_preserved():
|
|
assert _normalize_platform("grok") == "grok"
|
|
|
|
|
|
def test_normalize_platform_openai_unchanged():
|
|
assert _normalize_platform("openai") == "openai"
|
|
|
|
|
|
def test_normalize_platform_gemini_unchanged():
|
|
assert _normalize_platform("GEMINI") == "gemini"
|
|
|
|
|
|
def test_normalize_platform_unknown_preserved():
|
|
assert _normalize_platform("some_unknown") == "some_unknown"
|
|
|
|
|
|
# ── 单元测试:_detect_platform ───────────────────────
|
|
|
|
# ── 单元测试:_resolve_platform ───────────────────────
|
|
|
|
# ── 集成测试夹具 ──────────────────────────────────────
|
|
|
|
@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)
|
|
|
|
|
|
# ── 一键整理:Grok 平台检测与创建 ─────────────────────
|
|
|
|
def test_organize_groups_creates_account_with_grok_platform(db_session, monkeypatch):
|
|
"""一键整理时,快照平台为 grok → 新建账号的 platform = grok。"""
|
|
w = Website(
|
|
name="W1", site_type="sub2api", base_url="http://w1",
|
|
enabled=True, auth_config_json="{}", timeout_seconds=30,
|
|
)
|
|
u1 = Upstream(name="U1", base_url="http://u1")
|
|
db_session.add_all([w, u1])
|
|
db_session.commit()
|
|
db_session.refresh(w)
|
|
db_session.refresh(u1)
|
|
|
|
b1 = WebsiteGroupBinding(
|
|
website_id=w.id, target_group_id="TG1", target_group_name="TG1-Group",
|
|
source_groups_json=json.dumps([{"upstream_id": u1.id, "group_id": "G1"}]),
|
|
enabled=True,
|
|
platform="grok",
|
|
)
|
|
db_session.add(b1)
|
|
|
|
k1 = UpstreamGeneratedKey(
|
|
upstream_id=u1.id, group_id="G1", group_name="G1-Name",
|
|
key_name="Key-G1", key_value="sk-g1", status="created",
|
|
)
|
|
db_session.add(k1)
|
|
|
|
# 快照明确指定 platform=grok
|
|
snapshot = UpstreamRateSnapshot(
|
|
upstream_id=u1.id,
|
|
snapshot_json=json.dumps({
|
|
"groups": {"G1": {"group_name": "G1-Name", "rate": 0.1, "platform": "grok"}}
|
|
}),
|
|
)
|
|
db_session.add(snapshot)
|
|
db_session.commit()
|
|
|
|
created_bodies = []
|
|
|
|
class MockClient:
|
|
def __init__(self, **kwargs): pass
|
|
def __enter__(self): return self
|
|
def __exit__(self, *a): pass
|
|
def get_groups(self, *a, **kw): return [{"id": "TG1", "name": "TG1-Group", "platform": "grok"}]
|
|
def list_accounts(self): return []
|
|
def extract_id(self, val): return val.get("id") if isinstance(val, dict) else str(val)
|
|
def create_account(self, body):
|
|
created_bodies.append(body)
|
|
return {"id": "NEW-ACC", "name": body["name"], "group_ids": body["group_ids"]}
|
|
|
|
monkeypatch.setattr("app.routers.websites.Sub2ApiWebsiteClient", MockClient)
|
|
monkeypatch.setattr("app.routers.websites.sync_account_priorities_for_website", lambda db, wid: [])
|
|
|
|
response = organize_website_groups(wid=w.id, db=db_session)
|
|
assert response.success is True
|
|
assert len(created_bodies) == 1
|
|
assert created_bodies[0]["platform"] == "grok"
|
|
|
|
|
|
def test_organize_groups_detects_grok_from_group_name(db_session, monkeypatch):
|
|
"""一键整理时,分组名包含 Grok 且无快照 → 新建账号的 platform = grok。"""
|
|
w = Website(
|
|
name="W1", site_type="sub2api", base_url="http://w1",
|
|
enabled=True, auth_config_json="{}", timeout_seconds=30,
|
|
)
|
|
u1 = Upstream(name="U1", base_url="http://u1")
|
|
db_session.add_all([w, u1])
|
|
db_session.commit()
|
|
db_session.refresh(w)
|
|
db_session.refresh(u1)
|
|
|
|
b1 = WebsiteGroupBinding(
|
|
website_id=w.id, target_group_id="TG1", target_group_name="TG1-Group",
|
|
source_groups_json=json.dumps([{"upstream_id": u1.id, "group_id": "Grok01"}]),
|
|
enabled=True,
|
|
platform="grok",
|
|
)
|
|
db_session.add(b1)
|
|
|
|
k1 = UpstreamGeneratedKey(
|
|
upstream_id=u1.id, group_id="Grok01", group_name="Grok01-Group",
|
|
key_name="Key-G1", key_value="sk-g1", status="created",
|
|
)
|
|
db_session.add(k1)
|
|
db_session.commit()
|
|
|
|
created_bodies = []
|
|
|
|
class MockClient:
|
|
def __init__(self, **kwargs): pass
|
|
def __enter__(self): return self
|
|
def __exit__(self, *a): pass
|
|
def get_groups(self, *a, **kw): return [{"id": "TG1", "name": "TG1-Group", "platform": "grok"}]
|
|
def list_accounts(self): return []
|
|
def extract_id(self, val): return val.get("id") if isinstance(val, dict) else str(val)
|
|
def create_account(self, body):
|
|
created_bodies.append(body)
|
|
return {"id": "NEW-ACC", "name": body["name"], "group_ids": body["group_ids"]}
|
|
|
|
monkeypatch.setattr("app.routers.websites.Sub2ApiWebsiteClient", MockClient)
|
|
monkeypatch.setattr("app.routers.websites.sync_account_priorities_for_website", lambda db, wid: [])
|
|
|
|
response = organize_website_groups(wid=w.id, db=db_session)
|
|
assert response.success is True
|
|
assert len(created_bodies) == 1
|
|
assert created_bodies[0]["platform"] == "grok"
|
|
|
|
|
|
def test_organize_groups_grok_platform_creates_independent_account(db_session, monkeypatch):
|
|
"""binding platform=grok 但有旧 openai 账号时 → 创建独立 grok 账号,不修改旧账号。"""
|
|
w = Website(
|
|
name="W1", site_type="sub2api", base_url="http://w1",
|
|
enabled=True, auth_config_json="{}", timeout_seconds=30,
|
|
)
|
|
u1 = Upstream(name="U1", base_url="http://u1")
|
|
db_session.add_all([w, u1])
|
|
db_session.commit()
|
|
db_session.refresh(w)
|
|
db_session.refresh(u1)
|
|
|
|
b1 = WebsiteGroupBinding(
|
|
website_id=w.id, target_group_id="TG1", target_group_name="TG1-Group",
|
|
source_groups_json=json.dumps([{"upstream_id": u1.id, "group_id": "G1"}]),
|
|
enabled=True,
|
|
platform="grok",
|
|
)
|
|
db_session.add(b1)
|
|
|
|
k1 = UpstreamGeneratedKey(
|
|
upstream_id=u1.id, group_id="G1", group_name="G1-Group",
|
|
key_name="Key-G1", key_value="sk-g1", status="created",
|
|
)
|
|
db_session.add(k1)
|
|
db_session.commit()
|
|
|
|
created_bodies = []
|
|
updated_accounts = []
|
|
|
|
class MockClient:
|
|
def __init__(self, **kwargs): pass
|
|
def __enter__(self): return self
|
|
def __exit__(self, *a): pass
|
|
def get_groups(self, *a, **kw): return [{"id": "TG1", "name": "TG1-Group", "platform": "grok"}]
|
|
def list_accounts(self):
|
|
return [{"id": "ACC-OLD", "name": "Old-OpenAI", "group_ids": ["TG1"], "platform": "openai"}]
|
|
def extract_id(self, val): return val.get("id") if isinstance(val, dict) else str(val)
|
|
def create_account(self, body):
|
|
created_bodies.append(body)
|
|
acc_id = f"NEW-{body['platform']}-{len(created_bodies)}"
|
|
return {"id": acc_id, "name": body["name"], "group_ids": body["group_ids"], "platform": body["platform"]}
|
|
def update_account(self, account_id, body):
|
|
updated_accounts.append((account_id, body))
|
|
return {"id": account_id}
|
|
|
|
monkeypatch.setattr("app.routers.websites.Sub2ApiWebsiteClient", MockClient)
|
|
monkeypatch.setattr("app.routers.websites.sync_account_priorities_for_website", lambda db, wid: [])
|
|
|
|
response = organize_website_groups(wid=w.id, db=db_session)
|
|
assert response.success is True
|
|
assert len(updated_accounts) == 0
|
|
assert len(created_bodies) == 1
|
|
assert created_bodies[0]["platform"] == "grok"
|
|
assert response.items[0].status == "created"
|
|
|
|
|
|
# ── 手动导入:Grok default_platform ───────────────────
|
|
|
|
def test_import_upstream_key_uses_target_group_platform(monkeypatch, db_session):
|
|
"""手动导入只使用目标网站分组平台,忽略请求中的平台字段。"""
|
|
website = Website(
|
|
name="My Sub2API", site_type="sub2api", base_url="http://sub2api.local",
|
|
api_prefix="/api/v1", 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="Up1", base_url="http://up1.local")
|
|
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", key_name="SmartUp-VIP", key_value="sk-upstream-generated",
|
|
masked_key="sk-u...", raw_json="{}", status="created",
|
|
)
|
|
db_session.add(generated)
|
|
db_session.commit()
|
|
db_session.refresh(generated)
|
|
|
|
account_bodies = []
|
|
|
|
class FakeClient:
|
|
def __init__(self, **kwargs): pass
|
|
def __enter__(self): return self
|
|
def __exit__(self, *a): pass
|
|
def create_account(self, body, endpoint="/accounts"):
|
|
account_bodies.append(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"))
|
|
def get_groups(self, **kw): return [{"id": "7", "name": "Grok", "platform": "grok"}]
|
|
def list_accounts(self): return []
|
|
|
|
monkeypatch.setattr("app.routers.websites.Sub2ApiWebsiteClient", FakeClient)
|
|
monkeypatch.setattr("app.routers.websites.reconcile_upstream_keys_full", lambda db, uid: None)
|
|
monkeypatch.setattr("app.routers.websites.latest_rate_map", lambda db, uid: {})
|
|
monkeypatch.setattr("app.routers.websites.build_target_group_priority_map", lambda db, src: {})
|
|
|
|
response = import_upstream_keys_as_accounts(
|
|
website.id,
|
|
ImportAccountsRequest(
|
|
upstream_key_ids=[generated.id],
|
|
target_group_map={"vip": "7"},
|
|
default_platform="openai",
|
|
platform_mode="manual",
|
|
),
|
|
db_session,
|
|
object(),
|
|
)
|
|
|
|
assert "新建 1" in response.message
|
|
assert len(account_bodies) == 1
|
|
assert account_bodies[0]["platform"] == "grok"
|