refactor: Antigravity/Codex 服务重构为插件化适配器架构

- 将 Antigravity 和 Codex 从独立模块迁移至 src/services/provider/adapters/ 插件体系
- 新增 provider_types 和 oauth_token 模块,移除 maintenance_scheduler 中的 OAuth 定时刷新
- 增强 admin API:扩展 keys 和 provider_query 端点,新增 dashboard 路由
- 大幅增强 ProviderDetailDrawer 组件,新增 AntigravityQuotaDialog
- 改进 handler 基类(chat/cli)和错误分类器
- 优化 fetch_scheduler 和 upstream_fetcher
- 前端 UI 组件清理和优化
- 更新测试以匹配新模块结构
This commit is contained in:
fawney19
2026-02-06 16:37:06 +08:00
parent e01dfee41d
commit b88fb6273b
112 changed files with 5193 additions and 1992 deletions

View File

@@ -21,14 +21,16 @@ from src.core.crypto import crypto_service
from src.core.exceptions import InvalidRequestException, NotFoundException
from src.core.key_capabilities import get_capability
from src.core.logger import logger
from src.core.provider_types import ProviderType
from src.database import get_db
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint, User
from src.models.endpoint_models import (
EndpointAPIKeyCreate,
EndpointAPIKeyResponse,
EndpointAPIKeyUpdate,
)
from src.services.cache.provider_cache import ProviderCacheService
from src.utils.auth_utils import require_admin
router = APIRouter(tags=["Provider Keys"])
pipeline = ApiRequestPipeline()
@@ -146,6 +148,42 @@ async def delete_endpoint_key(
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.post("/keys/{key_id}/clear-oauth-invalid")
async def clear_oauth_invalid(
key_id: str,
request: Request,
db: Session = Depends(get_db),
_: User = Depends(require_admin),
) -> dict:
"""
清除 Key 的 OAuth 失效标记
手动清除指定 Key 的 oauth_invalid_at / oauth_invalid_reason 状态,
通常在管理员确认账号已完成验证后使用。
**路径参数**:
- `key_id`: Key ID
**返回字段**:
- `message`: 操作结果消息
"""
key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == key_id).first()
if not key:
raise NotFoundException(f"Key {key_id} 不存在")
if not key.oauth_invalid_at:
return {"message": "该 Key 当前无失效标记,无需清除"}
old_reason = key.oauth_invalid_reason
key.oauth_invalid_at = None
key.oauth_invalid_reason = None
db.commit()
logger.info("[OK] 手动清除 Key {}... 的 OAuth 失效标记 (原因: {})", key_id[:8], old_reason)
return {"message": "已清除 OAuth 失效标记"}
# ========== Provider Keys API ==========
@@ -640,14 +678,12 @@ def _build_key_response(
oauth_expires_at = auth_config.get("expires_at")
oauth_email = auth_config.get("email")
oauth_plan_type = auth_config.get("plan_type") # Codex: plus/free/team/enterprise
# Antigravity 使用 "tier" 字段(如 "PAID"/"FREE"),做小写化 fallback
if not oauth_plan_type:
ag_tier = auth_config.get("tier")
if ag_tier and isinstance(ag_tier, str):
oauth_plan_type = ag_tier.lower()
oauth_account_id = auth_config.get("account_id") # Codex: chatgpt_account_id
logger.debug(
"OAuth key {} auth_config: email={} plan_type={} account_id={}",
key.id,
oauth_email,
oauth_plan_type,
oauth_account_id,
)
except Exception as e:
logger.error("Failed to decrypt auth_config for key {}: {}", key.id, e)
@@ -924,9 +960,9 @@ class AdminRefreshProviderQuotaAdapter(AdminApiAdapter):
if not provider:
raise NotFoundException(f"Provider {self.provider_id} 不存在")
# 检查是否是 Codex 类型
if provider.provider_type != "codex":
raise InvalidRequestException("仅支持 Codex 类型的 Provider 刷新限额")
provider_type = str(getattr(provider, "provider_type", "") or "").strip().lower()
if provider_type not in {ProviderType.CODEX, ProviderType.ANTIGRAVITY}:
raise InvalidRequestException("仅支持 Codex / Antigravity 类型的 Provider 刷新限额")
# 获取所有活跃的 Keys
keys = (
@@ -947,15 +983,24 @@ class AdminRefreshProviderQuotaAdapter(AdminApiAdapter):
"message": "没有活跃的 Key",
}
# 获取 openai:cli 端点
# 获取端点
# - Codex: openai:cli
# - Antigravity: gemini:cli用于触发 oauth 刷新 + 提供 auth_config.project_id
endpoint = None
for ep in provider.endpoints:
if ep.api_format == "openai:cli" and ep.is_active:
endpoint = ep
break
if not endpoint:
raise InvalidRequestException("找不到有效的 openai:cli 端点")
if provider_type == ProviderType.CODEX:
for ep in provider.endpoints:
if ep.api_format == "openai:cli" and ep.is_active:
endpoint = ep
break
if not endpoint:
raise InvalidRequestException("找不到有效的 openai:cli 端点")
else:
for ep in provider.endpoints:
if ep.api_format == "gemini:cli" and ep.is_active:
endpoint = ep
break
if not endpoint:
raise InvalidRequestException("找不到有效的 gemini:cli 端点")
results: list[dict] = []
success_count = 0
@@ -967,56 +1012,57 @@ class AdminRefreshProviderQuotaAdapter(AdminApiAdapter):
# 单个 Key 刷新函数
async def refresh_single_key(key: ProviderAPIKey) -> dict:
try:
# 获取认证信息
auth_info = await get_provider_auth(endpoint, key)
if provider_type == ProviderType.CODEX:
# 获取认证信息
auth_info = await get_provider_auth(endpoint, key)
# 构建请求 URL
url = build_provider_url(endpoint, key=key)
# 构建请求 URL
url = build_provider_url(endpoint, key=key)
# 构建请求头
headers = {
"Content-Type": "application/json",
"Accept": "application/json",
}
if auth_info:
headers[auth_info.auth_header] = auth_info.auth_value
else:
# 标准 API Key
decrypted_key = crypto_service.decrypt(key.api_key)
headers["Authorization"] = f"Bearer {decrypted_key}"
# 发送最小的测试请求,使用 Codex Responses API 格式
test_body = {
"model": CODEX_QUOTA_REFRESH_MODEL,
"input": [
{
"type": "message",
"role": "user",
"content": [{"type": "input_text", "text": "hi"}],
}
],
"instructions": "",
"stream": True,
"store": False,
}
async with httpx.AsyncClient(timeout=30.0, verify=get_ssl_context()) as client:
response = await client.post(url, json=test_body, headers=headers)
# 解析响应头中的限额信息
response_headers = dict(response.headers)
metadata = MetadataCollectorRegistry.collect("codex", response_headers)
if metadata:
# 收集元数据,稍后统一更新数据库
metadata_updates[key.id] = metadata
return {
"key_id": key.id,
"key_name": key.name,
"status": "success",
"metadata": metadata,
# 构建请求头
headers = {
"Content-Type": "application/json",
"Accept": "application/json",
}
else:
if auth_info:
headers[auth_info.auth_header] = auth_info.auth_value
else:
# 标准 API Key
decrypted_key = crypto_service.decrypt(key.api_key)
headers["Authorization"] = f"Bearer {decrypted_key}"
# 发送最小的测试请求,使用 Codex Responses API 格式
test_body = {
"model": CODEX_QUOTA_REFRESH_MODEL,
"input": [
{
"type": "message",
"role": "user",
"content": [{"type": "input_text", "text": "hi"}],
}
],
"instructions": "",
"stream": True,
"store": False,
}
async with httpx.AsyncClient(timeout=30.0, verify=get_ssl_context()) as client:
response = await client.post(url, json=test_body, headers=headers)
# 解析响应头中的限额信息
response_headers = dict(response.headers)
metadata = MetadataCollectorRegistry.collect("codex", response_headers)
if metadata:
# 收集元数据,稍后统一更新数据库
metadata_updates[key.id] = metadata
return {
"key_id": key.id,
"key_name": key.name,
"status": "success",
"metadata": metadata,
}
# 响应成功但没有限额头
return {
"key_id": key.id,
@@ -1026,6 +1072,82 @@ class AdminRefreshProviderQuotaAdapter(AdminApiAdapter):
"status_code": response.status_code,
}
elif provider_type == ProviderType.ANTIGRAVITY:
# 直接调用 /v1internal:fetchAvailableModels 获取 quotaInfo无需发送真实对话请求
auth_info = await get_provider_auth(endpoint, key)
if not auth_info:
return {
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "缺少 OAuth 认证信息,请先授权/刷新 Token",
}
access_token = str(auth_info.auth_value).removeprefix("Bearer ").strip()
from src.services.model.upstream_fetcher import (
UpstreamModelsFetchContext,
fetch_models_for_key,
)
fetch_ctx = UpstreamModelsFetchContext(
provider_type="antigravity",
api_key_value=access_token,
# antigravity fetcher 不依赖 endpoint mapping
format_to_endpoint={},
proxy_config=getattr(provider, "proxy", None),
auth_config=auth_info.decrypted_auth_config,
)
_models, errors, ok, upstream_meta = await fetch_models_for_key(
fetch_ctx, timeout_seconds=10.0
)
if ok and upstream_meta:
metadata_updates[key.id] = upstream_meta
return {
"key_id": key.id,
"key_name": key.name,
"status": "success",
"metadata": upstream_meta,
}
if ok and not upstream_meta:
return {
"key_id": key.id,
"key_name": key.name,
"status": "no_metadata",
"message": "响应中未包含配额信息",
}
error_msg = "; ".join(errors) if errors else "fetchAvailableModels failed"
# 403 "verify your account" → 标记账号异常
if any(
"403" in e and ("verify" in e.lower() or "permission" in e.lower())
for e in errors
):
from datetime import datetime, timezone
from src.services.provider.oauth_token import (
OAUTH_ACCOUNT_BLOCK_PREFIX,
)
key.oauth_invalid_at = datetime.now(timezone.utc)
key.oauth_invalid_reason = (
f"{OAUTH_ACCOUNT_BLOCK_PREFIX}Google 要求验证账号"
)
# 立即持久化异常标记,使负载均衡能尽快跳过此 Key
# 后续 metadata 更新由外层统一 commit
db.commit()
logger.warning("[QUOTA_REFRESH] Key {} 因 403 verify 已标记为异常", key.id)
return {
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": error_msg,
}
except Exception as e:
logger.error("刷新 Key {} 限额失败: {}", key.id, e)
return {
@@ -1054,19 +1176,41 @@ class AdminRefreshProviderQuotaAdapter(AdminApiAdapter):
if metadata_updates:
for key in keys:
if key.id in metadata_updates:
key.upstream_metadata = metadata_updates[key.id]
updates = metadata_updates[key.id]
if isinstance(updates, dict):
# NOTE: upstream_metadata is a plain JSON column (not MutableDict),
# so in-place mutation won't be persisted reliably. Always assign
# a new dict object to mark the column as dirty.
current = key.upstream_metadata
merged: dict = dict(current) if isinstance(current, dict) else {}
merged.update(updates)
key.upstream_metadata = merged
db.add(key)
# 提交数据库更改
db.commit()
logger.info(
"[QUOTA_REFRESH] Provider {}: 成功 {}/{}, 失败 {}",
self.provider_id,
success_count,
len(keys),
failed_count,
)
failed_details = [
f"{r.get('key_name', r.get('key_id', '?'))}: {r.get('message', 'unknown')}"
for r in results
if r["status"] != "success"
]
if failed_details:
logger.info(
"[QUOTA_REFRESH] Provider {}: 成功 {}/{}, 失败 {} [{}]",
self.provider_id,
success_count,
len(keys),
failed_count,
"; ".join(failed_details),
)
else:
logger.info(
"[QUOTA_REFRESH] Provider {}: 成功 {}/{}",
self.provider_id,
success_count,
len(keys),
)
return {
"success": success_count,

View File

@@ -21,6 +21,7 @@ from src.api.base.pipeline import ApiRequestPipeline
from src.core.api_format.signature import parse_signature_key
from src.core.exceptions import InvalidRequestException, NotFoundException
from src.core.logger import logger
from src.core.provider_types import ProviderType
from src.database import get_db
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
from src.models.endpoint_models import (
@@ -283,7 +284,7 @@ class AdminCreateProviderEndpointAdapter(AdminApiAdapter):
# 固定类型 Provider禁止通过该接口新增 Endpoints端点由模板自动创建并锁定
provider_type = (getattr(provider, "provider_type", "custom") or "custom").strip()
if provider_type != "custom":
if provider_type != ProviderType.CUSTOM:
raise InvalidRequestException("固定类型 Provider 不允许手动新增 Endpoint")
if self.endpoint_data.provider_id != self.provider_id:
@@ -424,7 +425,7 @@ class AdminUpdateProviderEndpointAdapter(AdminApiAdapter):
provider = db.query(Provider).filter(Provider.id == endpoint.provider_id).first()
if provider:
provider_type = (getattr(provider, "provider_type", "custom") or "custom").strip()
if provider_type != "custom":
if provider_type != ProviderType.CUSTOM:
if "base_url" in update_data or "custom_path" in update_data:
raise InvalidRequestException(
"固定类型 Provider 的 Endpoint 不允许修改 base_url/custom_path"
@@ -441,8 +442,14 @@ class AdminUpdateProviderEndpointAdapter(AdminApiAdapter):
new_proxy["password"] = old_password
update_data["proxy"] = new_proxy
# proxy 为 None 时保留,用于清除代理配置
# JSON 列需要 flag_modified 以确保 SQLAlchemy 检测到变更
json_fields = {"header_rules", "body_rules", "config", "proxy", "format_acceptance_config"}
for field, value in update_data.items():
setattr(endpoint, field, value)
if field in json_fields:
flag_modified(endpoint, field)
# Phase 3/4: 自动维护新架构字段,确保新增/历史数据都能被调度器按 family/kind 查询
sig = parse_signature_key(endpoint.api_format)

View File

@@ -156,7 +156,7 @@ class ProviderCompleteOAuthResponse(BaseModel):
def _require_fixed_provider(provider: Provider) -> str:
provider_type = (getattr(provider, "provider_type", "custom") or "custom").strip()
if provider_type == "custom":
if provider_type == ProviderType.CUSTOM:
raise InvalidRequestException("该 Provider 不是固定类型,无法使用 provider-oauth")
return provider_type
@@ -415,14 +415,6 @@ async def complete_oauth(
key.auth_config = crypto_service.encrypt(json.dumps(auth_config))
db.commit()
# 触发 OAuth 刷新任务重新调度
try:
from src.services.system import get_maintenance_scheduler
get_maintenance_scheduler().trigger_oauth_refresh_check()
except Exception as e:
logger.debug("trigger_oauth_refresh_check 调用失败: {}", e)
return CompleteOAuthResponse(
provider_type=provider_type,
expires_at=expires_at,
@@ -556,10 +548,24 @@ async def refresh_oauth(
proxy_config=proxy_config,
)
# Antigravityenrich_auth_config 会自动尝试补 project_id
# 即使本次仍未获取到也不阻断刷新token 已成功更新),下次刷新会继续重试
if provider_type == ProviderType.ANTIGRAVITY and not parsed.get("project_id"):
logger.warning(
"[OAUTH_REFRESH] Antigravity key {} 刷新成功但 project_id 仍缺失,"
"下次刷新将继续尝试获取",
key_id,
)
key.auth_config = crypto_service.encrypt(json.dumps(parsed))
# 刷新成功,清除失效标记
key.oauth_invalid_at = None
key.oauth_invalid_reason = None
# 刷新成功,清除 token 级别的失效标记
# 但保留账号级别的失效标记(以 OAUTH_ACCOUNT_BLOCK_PREFIX 开头),
# 这种不是 token 问题,刷新 token 解决不了
from src.services.provider.oauth_token import is_account_level_block
if not is_account_level_block(getattr(key, "oauth_invalid_reason", None)):
key.oauth_invalid_at = None
key.oauth_invalid_reason = None
db.commit()
return CompleteOAuthResponse(
@@ -786,14 +792,6 @@ async def complete_provider_oauth(
db.commit()
db.refresh(new_key)
# 触发 OAuth 刷新任务重新调度
try:
from src.services.system import get_maintenance_scheduler
get_maintenance_scheduler().trigger_oauth_refresh_check()
except Exception as e:
logger.debug("trigger_oauth_refresh_check 调用失败: {}", e)
return ProviderCompleteOAuthResponse(
key_id=str(new_key.id),
provider_type=provider_type,

View File

@@ -6,6 +6,7 @@ Provider Query API 端点
from __future__ import annotations
import asyncio
import json
from typing import Any
import httpx
@@ -17,6 +18,7 @@ from src.config.constants import TimeoutDefaults
from src.core.api_format import get_extra_headers_from_endpoint
from src.core.crypto import crypto_service
from src.core.logger import logger
from src.core.provider_types import ProviderType
from src.database.database import get_db
from src.models.database import Provider, ProviderEndpoint, User
from src.services.model.fetch_scheduler import (
@@ -25,16 +27,89 @@ from src.services.model.fetch_scheduler import (
set_upstream_models_to_cache,
)
from src.services.model.upstream_fetcher import (
_get_adapter_for_format,
build_all_format_configs,
fetch_models_from_endpoints,
UpstreamModelsFetchContext,
fetch_models_for_key,
get_adapter_for_format,
)
from src.services.provider.oauth_token import resolve_oauth_access_token
from src.utils.auth_utils import get_current_user
from src.utils.ssl_utils import get_ssl_context
router = APIRouter(prefix="/api/admin/provider-query", tags=["Provider Query"])
# ---------------------------------------------------------------------------
# Key Auth Resolution (shared by multi-key and single-key paths)
# ---------------------------------------------------------------------------
class _KeyAuthError(Exception):
"""Key 认证解析失败(调用方决定是返回错误还是抛 HTTPException"""
def __init__(self, message: str) -> None:
self.message = message
super().__init__(message)
async def _resolve_key_auth(
api_key: Any,
provider: Any,
) -> tuple[str, dict[str, Any] | None]:
"""统一解析 Key 的 api_key_value 和 auth_config。
Returns:
(api_key_value, auth_config)
Raises:
_KeyAuthError: 解析失败(含可读消息)
"""
auth_type = str(getattr(api_key, "auth_type", "api_key") or "api_key").lower()
provider_type = str(getattr(provider, "provider_type", "") or "").lower()
api_key_value: str | None = None
auth_config: dict[str, Any] | None = None
if auth_type == "oauth":
endpoint_api_format = "gemini:cli" if provider_type == ProviderType.ANTIGRAVITY else None
try:
resolved = await resolve_oauth_access_token(
key_id=str(api_key.id),
encrypted_api_key=str(api_key.api_key or ""),
encrypted_auth_config=(
str(api_key.auth_config)
if getattr(api_key, "auth_config", None) is not None
else None
),
provider_proxy_config=getattr(provider, "proxy", None),
endpoint_api_format=endpoint_api_format,
)
api_key_value = resolved.access_token
auth_config = resolved.decrypted_auth_config
except Exception as e:
logger.error("[provider-query] OAuth auth failed for key {}: {}", api_key.id, e)
raise _KeyAuthError("oauth auth failed") from e
if not api_key_value:
raise _KeyAuthError("oauth token missing")
else:
try:
api_key_value = crypto_service.decrypt(api_key.api_key)
except Exception as e:
logger.error("Failed to decrypt API key {}: {}", api_key.id, e)
raise _KeyAuthError("decrypt failed") from e
# Best-effort: 解密 auth_config 元数据(如 Antigravity project_id
if getattr(api_key, "auth_config", None):
try:
decrypted = crypto_service.decrypt(api_key.auth_config)
parsed = json.loads(decrypted)
auth_config = parsed if isinstance(parsed, dict) else None
except Exception:
auth_config = None
return api_key_value, auth_config
# ============ Request/Response Models ============
@@ -130,22 +205,28 @@ async def query_available_models(
# 缓存未命中或强制刷新,实时获取
try:
api_key_value = crypto_service.decrypt(api_key.api_key)
except Exception as e:
logger.error(f"Failed to decrypt API key {api_key.id}: {e}")
return [], f"Key {api_key.name or api_key.id}: decrypt failed", False
api_key_value, auth_config = await _resolve_key_auth(api_key, provider)
except _KeyAuthError as e:
return [], f"Key {api_key.name or api_key.id}: {e.message}", False
endpoint_configs = build_all_format_configs(api_key_value, format_to_endpoint)
models, errors, has_success = await fetch_models_from_endpoints(
endpoint_configs, timeout=MODEL_FETCH_HTTP_TIMEOUT
fetch_ctx = UpstreamModelsFetchContext(
provider_type=str(getattr(provider, "provider_type", "") or ""),
api_key_value=str(api_key_value or ""),
format_to_endpoint=format_to_endpoint,
proxy_config=getattr(provider, "proxy", None),
auth_config=auth_config,
)
models, errors, has_success, _meta = await fetch_models_for_key(
fetch_ctx, timeout_seconds=MODEL_FETCH_HTTP_TIMEOUT
)
# 写入缓存
if models:
await set_upstream_models_to_cache(request.provider_id, api_key.id, models)
# 写入缓存(按 model id 聚合,保证返回 api_formats 数组,避免前端 schema 不一致)
unique_models = _aggregate_models_by_id([m for m in models if isinstance(m, dict)])
if unique_models:
await set_upstream_models_to_cache(request.provider_id, api_key.id, unique_models)
error = f"Key {api_key.name or api_key.id}: {'; '.join(errors)}" if errors else None
return models, error, False # models, error, from_cache
return unique_models, error, False # models, error, from_cache
# 并发执行所有 Key 的获取
results = await asyncio.gather(*[fetch_for_key(key) for key in active_keys])
@@ -260,9 +341,16 @@ async def _fetch_models_for_single_key(
if not force_refresh:
cached_models = await get_upstream_models_from_cache(provider.id, api_key_id)
if cached_models is not None:
safe_models = [m for m in cached_models if isinstance(m, dict)]
unique_cached = _aggregate_models_by_id(safe_models)
# 修复遗留缓存格式(以前可能缓存了未聚合的 api_format 版本)
if unique_cached and (
not safe_models or "api_formats" not in safe_models[0] # type: ignore[operator]
):
await set_upstream_models_to_cache(provider.id, api_key_id, unique_cached)
return {
"success": True,
"data": {"models": cached_models, "error": None, "from_cache": True},
"data": {"models": unique_cached, "error": None, "from_cache": True},
"provider": {
"id": provider.id,
"name": provider.name,
@@ -271,14 +359,19 @@ async def _fetch_models_for_single_key(
# 缓存未命中或强制刷新,实时获取
try:
api_key_value = crypto_service.decrypt(api_key.api_key)
except Exception as e:
logger.error(f"Failed to decrypt API key: {e}")
raise HTTPException(status_code=500, detail="Failed to decrypt API key")
api_key_value, auth_config = await _resolve_key_auth(api_key, provider)
except _KeyAuthError as e:
raise HTTPException(status_code=500, detail=e.message)
endpoint_configs = build_all_format_configs(api_key_value, format_to_endpoint)
all_models, errors, has_success = await fetch_models_from_endpoints(
endpoint_configs, timeout=MODEL_FETCH_HTTP_TIMEOUT
fetch_ctx = UpstreamModelsFetchContext(
provider_type=str(getattr(provider, "provider_type", "") or ""),
api_key_value=str(api_key_value or ""),
format_to_endpoint=format_to_endpoint,
proxy_config=getattr(provider, "proxy", None),
auth_config=auth_config,
)
all_models, errors, has_success, _meta = await fetch_models_for_key(
fetch_ctx, timeout_seconds=MODEL_FETCH_HTTP_TIMEOUT
)
# 按 model id 聚合,合并所有 api_format
@@ -431,28 +524,44 @@ async def test_model(
if not endpoint or not api_key:
raise HTTPException(status_code=404, detail="No active endpoint or API key found")
auth_type = str(getattr(api_key, "auth_type", "api_key") or "api_key").lower()
try:
api_key_value = crypto_service.decrypt(api_key.api_key)
if auth_type == "oauth":
resolved = await resolve_oauth_access_token(
key_id=str(api_key.id),
encrypted_api_key=str(api_key.api_key or ""),
encrypted_auth_config=(
str(api_key.auth_config) if getattr(api_key, "auth_config", None) else None
),
provider_proxy_config=getattr(provider, "proxy", None),
endpoint_api_format=str(getattr(endpoint, "api_format", "") or ""),
)
api_key_value = resolved.access_token
oauth_meta = resolved.decrypted_auth_config or {}
if not api_key_value:
raise HTTPException(status_code=500, detail="OAuth token missing")
else:
api_key_value = crypto_service.decrypt(api_key.api_key)
oauth_meta = {}
except HTTPException:
raise
except Exception as e:
logger.error(f"[test-model] Failed to decrypt API key: {e}")
raise HTTPException(status_code=500, detail="Failed to decrypt API key")
logger.error(f"[test-model] Failed to resolve API key: {e}")
raise HTTPException(status_code=500, detail="Failed to resolve API key")
# 构建请求配置
extra_headers = get_extra_headers_from_endpoint(endpoint) or {}
# OAuth 认证:从 auth_config 获取 account_id 并添加到请求头
if api_key.auth_type == "oauth" and api_key.auth_config:
# OAuth 认证:Codex 需要 chatgpt-account-id
if auth_type == "oauth":
try:
import json
decrypted_config = crypto_service.decrypt(api_key.auth_config)
auth_config = json.loads(decrypted_config)
account_id = auth_config.get("account_id")
account_id = oauth_meta.get("account_id")
if account_id:
extra_headers["chatgpt-account-id"] = account_id
extra_headers["chatgpt-account-id"] = str(account_id)
logger.debug("[test-model] Added chatgpt-account-id header: {}", account_id)
except Exception as e:
logger.warning("[test-model] Failed to parse OAuth auth_config: {}", e)
logger.warning("[test-model] Failed to apply OAuth extra headers: {}", e)
endpoint_config = {
"api_key": api_key_value,
@@ -465,7 +574,7 @@ async def test_model(
try:
# 获取对应的 Adapter 类
adapter_class = _get_adapter_for_format(endpoint.api_format)
adapter_class = get_adapter_for_format(endpoint.api_format)
if not adapter_class:
return {
"success": False,
@@ -479,6 +588,7 @@ async def test_model(
logger.debug(f"[test-model] 使用 Adapter: {adapter_class.__name__}")
logger.debug(f"[test-model] 端点 API Format: {endpoint.api_format}")
logger.debug(f"[test-model] 使用 Key: {api_key.name or api_key.id} (auth_type={auth_type})")
# 准备测试请求数据
check_request = {
@@ -506,6 +616,9 @@ async def test_model(
) as client:
logger.debug("[test-model] 开始端点测试...")
# Provider 上下文auth_type 用于 OAuth 认证头处理provider_type 用于特殊路由
p_type = str(getattr(provider, "provider_type", "") or "").lower()
response = await adapter_class.check_endpoint(
client,
endpoint_config["base_url"],
@@ -522,6 +635,10 @@ async def test_model(
provider_id=provider.id,
api_key_id=endpoint_config.get("api_key_id"),
model_name=request.model_name,
# Provider 上下文
auth_type=auth_type,
provider_type=p_type if p_type else None,
decrypted_auth_config=oauth_meta if oauth_meta else None,
)
# 记录提供商返回信息
@@ -546,14 +663,69 @@ async def test_model(
error_obj = parsed_body["error"]
# 兼容 error 可能是字典或字符串的情况
if isinstance(error_obj, dict):
logger.debug(f"[test-model] Error Message: {error_obj.get('message')}")
raise HTTPException(status_code=500, detail=error_obj.get("message"))
error_message = error_obj.get("message", "")
logger.debug(f"[test-model] Error Message: {error_message}")
# Antigravity 403 "verify your account" → 标记账号异常
if (
api_key
and auth_type == "oauth"
and error_obj.get("code") == 403
and (
"verify" in error_message.lower()
or "permission" in str(error_obj.get("status", "")).lower()
)
):
from datetime import datetime, timezone
from src.services.provider.oauth_token import (
OAUTH_ACCOUNT_BLOCK_PREFIX,
)
api_key.oauth_invalid_at = datetime.now(timezone.utc)
api_key.oauth_invalid_reason = (
f"{OAUTH_ACCOUNT_BLOCK_PREFIX}Google 要求验证账号"
)
db.commit()
oauth_email = None
if getattr(api_key, "auth_config", None):
try:
decrypted = crypto_service.decrypt(api_key.auth_config)
parsed = json.loads(decrypted)
if isinstance(parsed, dict):
email_val = parsed.get("email")
if isinstance(email_val, str) and email_val.strip():
oauth_email = email_val.strip()
except Exception:
oauth_email = None
if oauth_email:
logger.warning(
"[test-model] Key {} (email={}) 因 403 verify 已标记为异常",
api_key.id,
oauth_email,
)
else:
logger.warning(
"[test-model] Key {} 因 403 verify 已标记为异常", api_key.id
)
raise HTTPException(
status_code=500,
detail=str(error_message)[:500] if error_message else "Provider error",
)
else:
logger.debug(f"[test-model] Error: {error_obj}")
raise HTTPException(status_code=500, detail=error_obj)
# error_obj 可能是字符串,截断以避免泄露过多上游信息
raise HTTPException(
status_code=500,
detail=str(error_obj)[:500] if error_obj else "Provider error",
)
elif "error" in response:
logger.debug(f"[test-model] Error: {response['error']}")
raise HTTPException(status_code=500, detail=response["error"])
raise HTTPException(
status_code=500,
detail=str(response["error"])[:500],
)
else:
# 如果有选择或消息,记录内容预览
if isinstance(response_data, dict):

View File

@@ -33,7 +33,7 @@ class ProviderBillingUpdate(BaseModel):
quota_last_reset_at: str | None = None # 当前周期开始时间
quota_expires_at: str | None = None
rpm_limit: int | None = Field(default=None, ge=0)
provider_priority: int = Field(default=100, ge=0, le=200)
provider_priority: int = Field(default=100, ge=0, le=10000)
@router.put("/providers/{provider_id}/billing")

View File

@@ -19,15 +19,14 @@ from src.core.enums import ProviderBillingType
from src.core.exceptions import InvalidRequestException, NotFoundException
from src.core.logger import logger
from src.core.model_permissions import match_model_with_pattern, parse_allowed_models_to_list
from src.core.provider_templates.fixed_providers import FIXED_PROVIDERS
from src.core.provider_templates.types import ProviderType
from src.database import get_db
from src.models.admin_requests import CreateProviderRequest, UpdateProviderRequest
from src.models.database import GlobalModel, Provider, ProviderAPIKey, ProviderEndpoint
from src.services.cache.model_cache import ModelCacheService
from src.services.cache.provider_cache import ProviderCacheService
from src.core.provider_templates.fixed_providers import FIXED_PROVIDERS
from src.core.provider_templates.types import ProviderType
router = APIRouter(tags=["Provider CRUD"])
pipeline = ApiRequestPipeline()
@@ -291,10 +290,16 @@ class AdminCreateProviderAdapter(AdminApiAdapter):
else ProviderBillingType.PAY_AS_YOU_GO
)
# 有 envelope 包装的 Provider 类型(如 Antigravity、Codex需要格式转换来正确
# 解包上游响应,创建时默认开启 enable_format_conversion。
pt = (validated_data.provider_type or "custom").strip()
envelope_provider_types = {ProviderType.ANTIGRAVITY, ProviderType.CODEX}
default_enable_format_conversion = pt in envelope_provider_types
# 创建 Provider 对象
provider = Provider(
name=validated_data.name,
provider_type=validated_data.provider_type or "custom",
provider_type=pt,
description=validated_data.description,
website=validated_data.website,
billing_type=billing_type,
@@ -311,6 +316,8 @@ class AdminCreateProviderAdapter(AdminApiAdapter):
stream_first_byte_timeout=validated_data.stream_first_byte_timeout,
request_timeout=validated_data.request_timeout,
config=validated_data.config,
# 有 envelope 的反代类型默认开启格式转换
enable_format_conversion=default_enable_format_conversion,
)
db.add(provider)
@@ -318,7 +325,7 @@ class AdminCreateProviderAdapter(AdminApiAdapter):
# 固定类型 Provider自动创建并锁定预置 Endpoints同一事务
provider_type = (provider.provider_type or "custom").strip()
if provider_type != "custom":
if provider_type != ProviderType.CUSTOM:
try:
template = FIXED_PROVIDERS.get(ProviderType(provider_type))
except Exception:

View File

@@ -634,15 +634,6 @@ class AdminSetSystemConfigAdapter(AdminApiAdapter):
except Exception as e:
logger.warning(f"更新用户配额重置任务时间失败: {e}")
# 如果更新的是 OAuth Token 自动刷新开关,触发调度器重新计算
if self.key == "enable_oauth_token_refresh":
try:
from src.services.system.maintenance_scheduler import get_maintenance_scheduler
get_maintenance_scheduler().trigger_oauth_refresh_check()
except Exception as e:
logger.warning("触发 OAuth Token 刷新调度失败: {}", e)
# 返回时不暴露加密后的值
display_value = "********" if self.key in self.ENCRYPTED_KEYS else config.value