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 || '—'