From 6906c3b998d70f35a0f9b082550ae086ee1e57f4 Mon Sep 17 00:00:00 2001 From: SmartUp Developer Date: Sun, 12 Jul 2026 13:16:29 +0800 Subject: [PATCH] feat: support Grok platform account imports --- backend/app/routers/websites.py | 24 +- backend/test_grok_platform.py | 383 ++++++++++++++++++++++++++++++++ frontend/src/views/Websites.vue | 3 + 3 files changed, 404 insertions(+), 6 deletions(-) create mode 100644 backend/test_grok_platform.py diff --git a/backend/app/routers/websites.py b/backend/app/routers/websites.py index 43883f4..ec5afc4 100644 --- a/backend/app/routers/websites.py +++ b/backend/app/routers/websites.py @@ -593,14 +593,24 @@ def import_groups_from_upstream( ) -SUPPORTED_PLATFORMS = {"openai", "anthropic", "gemini", "antigravity"} +def _normalize_platform(platform: str) -> str: + """统一平台值:小写、xai → grok。""" + p = platform.lower().strip() + if p == "xai": + return "grok" + return p + + +SUPPORTED_PLATFORMS = {"openai", "anthropic", "gemini", "grok", "antigravity"} def _resolve_platform(text: str, group_snapshot: dict | None, fallback: str = "openai") -> str: if group_snapshot and isinstance(group_snapshot, dict): - platform = group_snapshot.get("platform") - if platform and str(platform).lower() in SUPPORTED_PLATFORMS: - return str(platform).lower() + raw = group_snapshot.get("platform") + if raw: + platform = _normalize_platform(str(raw)) + if platform in SUPPORTED_PLATFORMS: + return platform return _detect_platform(text, fallback) @@ -611,9 +621,11 @@ def _detect_platform(text: str, fallback: str = "openai") -> str: return "anthropic" if "gemini" in lower: return "gemini" + if "grok" in lower or "xai" in lower: + return "grok" if "antigravity" in lower: return "antigravity" - return fallback + return _normalize_platform(fallback) @router.post("/api/websites/{wid}/accounts/sync-imported-upstream-keys", response_model=ImportAccountsResponse) @@ -869,7 +881,7 @@ def import_upstream_keys_as_accounts( body.default_platform, ) else: - platform = body.default_platform + platform = _normalize_platform(body.default_platform) upstream = upstreams_map.get(row.upstream_id) upstream_name = upstream.name if upstream else "Unknown" diff --git a/backend/test_grok_platform.py b/backend/test_grok_platform.py new file mode 100644 index 0000000..3768df2 --- /dev/null +++ b/backend/test_grok_platform.py @@ -0,0 +1,383 @@ +"""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, + _detect_platform, + _resolve_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 ─────────────────────── + +def test_detect_platform_grok_in_name(): + assert _detect_platform("My Grok Key") == "grok" + + +def test_detect_platform_grok_lowercase(): + assert _detect_platform("grok-group") == "grok" + + +def test_detect_platform_xai_in_name(): + assert _detect_platform("Chat xAI Key") == "grok" + + +def test_detect_platform_gemini_still_works(): + assert _detect_platform("Gemini-Pro") == "gemini" + + +def test_detect_platform_claude_still_works(): + assert _detect_platform("Claude-v2") == "anthropic" + + +def test_detect_platform_anthropic_still_works(): + assert _detect_platform("anthropic-key") == "anthropic" + + +def test_detect_platform_antigravity_still_works(): + assert _detect_platform("antigravity-key") == "antigravity" + + +def test_detect_platform_no_match_fallback(): + assert _detect_platform("abc") == "openai" + assert _detect_platform("abc", "anthropic") == "anthropic" + + +def test_detect_platform_fallback_xai_normalized(): + """fallback='xai' 且未匹配任何关键词 → 规范化返回 grok。""" + assert _detect_platform("abc", "xai") == "grok" + + +# ── 单元测试:_resolve_platform ─────────────────────── + +def test_resolve_platform_snapshot_grok(): + """快照 platform=grok → 返回 grok。""" + result = _resolve_platform("some key name", {"platform": "grok"}, "openai") + assert result == "grok" + + +def test_resolve_platform_snapshot_grok_mixed_case(): + """快照 platform=Grok → 统一小写返回 grok。""" + result = _resolve_platform("some key name", {"platform": "Grok"}, "openai") + assert result == "grok" + + +def test_resolve_platform_snapshot_xai_normalized(): + """快照 platform=xai → 规范化为 grok。""" + result = _resolve_platform("some key name", {"platform": "xai"}, "openai") + assert result == "grok" + + +def test_resolve_platform_name_grok_without_snapshot(): + """无快照平台时,名称中包含 grok → 返回 grok。""" + result = _resolve_platform("My-Grok-Key", {}, "openai") + assert result == "grok" + + +def test_resolve_platform_name_xai_without_snapshot(): + """无快照平台时,名称中包含 xai → 返回 grok。""" + result = _resolve_platform("Chat-xAI", {}, "openai") + assert result == "grok" + + +def test_resolve_platform_snapshot_overrides_name(): + """快照有平台时,应优先于名称识别。""" + result = _resolve_platform("claude-key", {"platform": "grok"}, "openai") + assert result == "grok", "快照平台应覆盖名称识别" + + +def test_resolve_platform_gemini_snapshot_still_works(): + """快照 platform=gemini → 返回 gemini(不受 grok 影响)。""" + result = _resolve_platform("xxx", {"platform": "gemini"}, "openai") + assert result == "gemini" + + +# ── 集成测试夹具 ────────────────────────────────────── + +@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, + ) + 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"}] + 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, + ) + 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"}] + 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_corrects_platform_to_grok(db_session, monkeypatch): + """已导入账号平台为 openai,快照为 grok → 修正平台为 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, + ) + 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="imported", + imported_website_id=w.id, imported_account_id="ACC-G1", + imported_target_group_id="TG1", + ) + db_session.add(k1) + + snapshot = UpstreamRateSnapshot( + upstream_id=u1.id, + snapshot_json=json.dumps({ + "groups": {"G1": {"group_name": "G1-Group", "rate": 0.1, "platform": "grok"}} + }), + ) + db_session.add(snapshot) + db_session.commit() + + updated_accounts = [] + mock_platform = "openai" + + 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"}] + def list_accounts(self): + return [{"id": "ACC-G1", "name": "SmartUp-G1", "group_ids": ["TG1"], "platform": mock_platform}] + def extract_id(self, val): return val.get("id") if isinstance(val, dict) else str(val) + def update_account(self, account_id, body): + nonlocal mock_platform + updated_accounts.append((account_id, body)) + if "platform" in body: + mock_platform = body["platform"] + return {"id": account_id, "platform": body.get("platform")} + + 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) == 1 + assert updated_accounts[0][1]["platform"] == "grok" + assert "平台从 openai 修正为 grok" in response.items[0].message + + +# ── 手动导入:Grok default_platform ─────────────────── + +def test_import_upstream_key_with_grok_platform(monkeypatch, db_session): + """手动导入时 default_platform=grok → 创建账号的 platform = grok。""" + 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 [] + 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="grok", + platform_mode="manual", + ), + db_session, + object(), + ) + + assert "新建 1" in response.message + assert len(account_bodies) == 1 + assert account_bodies[0]["platform"] == "grok" diff --git a/frontend/src/views/Websites.vue b/frontend/src/views/Websites.vue index 3f60971..6f7b033 100644 --- a/frontend/src/views/Websites.vue +++ b/frontend/src/views/Websites.vue @@ -617,6 +617,7 @@ + @@ -2339,6 +2340,7 @@ function detectPlatform(item: { group_name?: string; group_id?: string; key_name const text = `${item.group_name || ''} ${item.group_id || ''} ${item.key_name || ''}`.toLowerCase() if (text.includes('claude') || text.includes('anthropic')) return 'Anthropic' if (text.includes('gemini')) return 'Gemini' + if (text.includes('grok') || text.includes('xai')) return 'Grok' if (text.includes('antigravity')) return 'Antigravity' return 'OpenAI 兼容' } @@ -2348,6 +2350,7 @@ function platformLabel(platform: string) { openai: 'OpenAI 兼容', anthropic: 'Anthropic', gemini: 'Gemini', + grok: 'Grok', antigravity: 'Antigravity', } return map[platform] || platform || '—'