mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-03 01:40:21 +08:00
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:
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
# Antigravity:enrich_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,
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -81,13 +81,12 @@ class ApiRequestPipeline:
|
||||
# 高频轮询端点抑制 debug 日志
|
||||
is_quiet = http_request.url.path in QUIET_POLLING_PATHS
|
||||
if not is_quiet:
|
||||
logger.debug("[Pipeline] START | path={}", http_request.url.path)
|
||||
logger.debug(
|
||||
"[Pipeline] Running with mode={}, adapter={}, adapter.mode={}, path={}",
|
||||
mode,
|
||||
adapter.__class__.__name__,
|
||||
adapter.mode,
|
||||
"[Pipeline] {} {} | adapter={}, mode={}",
|
||||
http_request.method,
|
||||
http_request.url.path,
|
||||
adapter.__class__.__name__,
|
||||
mode,
|
||||
)
|
||||
auth_start = PerfRecorder.start(force=perf_sampled)
|
||||
try:
|
||||
@@ -105,12 +104,8 @@ class ApiRequestPipeline:
|
||||
user, management_token = await self._authenticate_management(http_request, db)
|
||||
api_key = None
|
||||
else:
|
||||
if not is_quiet:
|
||||
logger.debug("[Pipeline] 调用 _authenticate_client")
|
||||
user, api_key = self._authenticate_client(http_request, db, adapter, quiet=is_quiet)
|
||||
management_token = None
|
||||
if not is_quiet:
|
||||
logger.debug("[Pipeline] 认证完成 | user={}", user.username if user else None)
|
||||
finally:
|
||||
auth_duration = PerfRecorder.stop(auth_start, "pipeline_auth", labels=perf_labels)
|
||||
_record_perf_metric("auth_ms", auth_duration)
|
||||
@@ -140,11 +135,6 @@ class ApiRequestPipeline:
|
||||
perf_metrics = getattr(http_request.state, "perf_metrics", None)
|
||||
if isinstance(perf_metrics, dict):
|
||||
perf_metrics.setdefault("pipeline", {})["body_bytes"] = int(body_size)
|
||||
if not is_quiet:
|
||||
logger.debug(
|
||||
"[Pipeline] Raw body读取完成 | size={} bytes",
|
||||
len(raw_body) if raw_body is not None else 0,
|
||||
)
|
||||
except TimeoutError:
|
||||
timeout_sec = int(config.request_body_timeout)
|
||||
logger.error(f"读取请求体超时({timeout_sec}s),可能客户端未发送完整请求体")
|
||||
@@ -152,10 +142,6 @@ class ApiRequestPipeline:
|
||||
status_code=408,
|
||||
detail=f"Request timeout: body not received within {timeout_sec} seconds",
|
||||
)
|
||||
else:
|
||||
if not is_quiet:
|
||||
logger.debug("[Pipeline] 非写请求跳过读取Body | method={}", http_request.method)
|
||||
|
||||
context_start = PerfRecorder.start(force=perf_sampled)
|
||||
context = ApiRequestContext.build(
|
||||
request=http_request,
|
||||
@@ -176,23 +162,8 @@ class ApiRequestPipeline:
|
||||
context.management_token = management_token
|
||||
# 存储 quiet 标志到 context,用于审计日志判断
|
||||
context.quiet_logging = is_quiet
|
||||
if not is_quiet:
|
||||
logger.debug(
|
||||
"[Pipeline] Context构建完成 | adapter={} | request_id={}",
|
||||
adapter.name,
|
||||
context.request_id,
|
||||
)
|
||||
|
||||
if mode != ApiMode.ADMIN and user:
|
||||
context.quota_remaining = self._calculate_quota_remaining(user)
|
||||
|
||||
if not is_quiet:
|
||||
logger.debug("[Pipeline] Adapter={} | RequestID={}", adapter.name, context.request_id)
|
||||
logger.debug(
|
||||
"[Pipeline] Calling authorize on {}, user={}",
|
||||
adapter.__class__.__name__,
|
||||
context.user,
|
||||
)
|
||||
# authorize 可能是异步的,需要检查并 await
|
||||
authorize_start = PerfRecorder.start(force=perf_sampled)
|
||||
try:
|
||||
@@ -242,25 +213,14 @@ class ApiRequestPipeline:
|
||||
# --------------------------------------------------------------------- #
|
||||
|
||||
def _authenticate_client(
|
||||
self, request: Request, db: Session, adapter: ApiAdapter, *, quiet: bool = False
|
||||
self, request: Request, db: Session, adapter: ApiAdapter, **_kw: object
|
||||
) -> tuple[User, ApiKey]:
|
||||
if not quiet:
|
||||
logger.debug("[Pipeline._authenticate_client] 开始")
|
||||
# 使用 adapter 的 extract_api_key 方法,支持不同 API 格式的认证头
|
||||
client_api_key = adapter.extract_api_key(request)
|
||||
if not quiet:
|
||||
logger.debug(
|
||||
"[Pipeline._authenticate_client] 提取API密钥完成 | key_prefix={}...",
|
||||
client_api_key[:8] if client_api_key else None,
|
||||
)
|
||||
if not client_api_key:
|
||||
raise HTTPException(status_code=401, detail="请提供API密钥")
|
||||
|
||||
if not quiet:
|
||||
logger.debug("[Pipeline._authenticate_client] 调用 auth_service.authenticate_api_key")
|
||||
auth_result = self.auth_service.authenticate_api_key(db, client_api_key)
|
||||
if not quiet:
|
||||
logger.debug("[Pipeline._authenticate_client] 认证结果 | result={}", bool(auth_result))
|
||||
if not auth_result:
|
||||
raise HTTPException(status_code=401, detail="无效的API密钥")
|
||||
|
||||
|
||||
@@ -1001,6 +1001,37 @@ class DashboardDailyStatsAdapter(DashboardAdapter):
|
||||
}
|
||||
)
|
||||
|
||||
# 补充 unique_models / unique_providers
|
||||
# query_time_series 使用小时粒度数据,不含这些维度统计
|
||||
# 直接从 Usage 表按本地日 UTC 范围查询,避免 StatsDaily 历史数据未回填的问题
|
||||
granularity = (self.time_range.granularity or "day").lower()
|
||||
if formatted and granularity == "day":
|
||||
local_days = self.time_range.get_local_day_hours()
|
||||
enrichment: dict[str, dict] = {}
|
||||
|
||||
for local_date, day_start_utc, day_end_utc in local_days:
|
||||
q = db.query(
|
||||
func.count(func.distinct(Usage.model)).label("um"),
|
||||
func.count(func.distinct(Usage.provider_name)).label("up"),
|
||||
).filter(
|
||||
Usage.created_at >= day_start_utc,
|
||||
Usage.created_at < day_end_utc,
|
||||
)
|
||||
if not is_admin:
|
||||
q = q.filter(Usage.user_id == user.id)
|
||||
row = q.first()
|
||||
if row:
|
||||
enrichment[local_date.isoformat()] = {
|
||||
"unique_models": row.um or 0,
|
||||
"unique_providers": row.up or 0,
|
||||
}
|
||||
|
||||
for item in formatted:
|
||||
date_key = item["date"][:10] # YYYY-MM-DD
|
||||
if date_key in enrichment:
|
||||
item["unique_models"] = enrichment[date_key]["unique_models"]
|
||||
item["unique_providers"] = enrichment[date_key]["unique_providers"]
|
||||
|
||||
# Model summary (use Usage directly for now)
|
||||
start_utc, end_utc = self.time_range.to_utc_datetime_range()
|
||||
model_query = db.query(
|
||||
|
||||
@@ -643,6 +643,10 @@ class ChatAdapterBase(ApiAdapter):
|
||||
provider_id: str | None = None,
|
||||
api_key_id: str | None = None,
|
||||
model_name: str | None = None,
|
||||
# Provider 上下文(Chat 适配器忽略这些参数,仅保持签名兼容)
|
||||
auth_type: str | None = None, # noqa: ARG003
|
||||
provider_type: str | None = None, # noqa: ARG003
|
||||
decrypted_auth_config: dict[str, Any] | None = None, # noqa: ARG003
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
测试模型连接性(非流式)
|
||||
|
||||
@@ -920,13 +920,48 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
try:
|
||||
resp.raise_for_status()
|
||||
except httpx.HTTPStatusError as e:
|
||||
error_body = ""
|
||||
try:
|
||||
error_body = resp.text[:4000] if resp.text else ""
|
||||
except Exception:
|
||||
# OAuth token may be revoked/expired earlier than expires_at indicates.
|
||||
# Best-effort: force refresh once on 401 and retry a single time.
|
||||
if (
|
||||
resp.status_code == 401
|
||||
and str(getattr(key, "auth_type", "") or "").lower() == "oauth"
|
||||
):
|
||||
refreshed_auth = await get_provider_auth(endpoint, key, force_refresh=True)
|
||||
if refreshed_auth:
|
||||
provider_headers[refreshed_auth.auth_header] = refreshed_auth.auth_value
|
||||
ctx.provider_request_headers = provider_headers
|
||||
|
||||
# retry once
|
||||
resp = await http_client.post(
|
||||
url,
|
||||
json=provider_payload,
|
||||
headers=provider_headers,
|
||||
timeout=httpx.Timeout(request_timeout_sync),
|
||||
)
|
||||
ctx.status_code = resp.status_code
|
||||
ctx.response_headers = dict(resp.headers)
|
||||
if envelope:
|
||||
envelope.on_http_status(
|
||||
base_url=ctx.selected_base_url, status_code=ctx.status_code
|
||||
)
|
||||
try:
|
||||
resp.raise_for_status()
|
||||
except httpx.HTTPStatusError as e2:
|
||||
error_body = ""
|
||||
try:
|
||||
error_body = resp.text[:4000] if resp.text else ""
|
||||
except Exception:
|
||||
error_body = ""
|
||||
e2.upstream_response = error_body # type: ignore[attr-defined]
|
||||
raise
|
||||
else:
|
||||
error_body = ""
|
||||
e.upstream_response = error_body # type: ignore[attr-defined]
|
||||
raise
|
||||
try:
|
||||
error_body = resp.text[:4000] if resp.text else ""
|
||||
except Exception:
|
||||
error_body = ""
|
||||
e.upstream_response = error_body # type: ignore[attr-defined]
|
||||
raise
|
||||
|
||||
# Safe JSON parsing.
|
||||
try:
|
||||
@@ -1087,86 +1122,115 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
max_prefetch_lines=config.stream_prefetch_lines,
|
||||
)
|
||||
|
||||
try:
|
||||
# 使用 asyncio.wait_for 包裹整个"建立连接 + 获取首字节"阶段
|
||||
# stream_first_byte_timeout 控制首字节超时,避免上游长时间无响应
|
||||
# 同时检测客户端断连,避免客户端已断开但服务端仍在等待上游响应
|
||||
if is_disconnected is not None:
|
||||
await wait_for_with_disconnect_detection(
|
||||
_connect_and_prefetch(),
|
||||
timeout=request_timeout,
|
||||
is_disconnected=is_disconnected,
|
||||
request_id=self.request_id,
|
||||
)
|
||||
else:
|
||||
await asyncio.wait_for(_connect_and_prefetch(), timeout=request_timeout)
|
||||
|
||||
except ClientDisconnectedException:
|
||||
# 客户端断开连接,清理资源
|
||||
if response_ctx is not None:
|
||||
try:
|
||||
await response_ctx.__aexit__(None, None, None)
|
||||
except Exception:
|
||||
pass
|
||||
await http_client.aclose()
|
||||
logger.warning(f" [{self.request_id}] 客户端在等待首字节时断开连接")
|
||||
ctx.status_code = 499
|
||||
ctx.error_message = "client_disconnected_during_prefetch"
|
||||
raise
|
||||
|
||||
except TimeoutError:
|
||||
# 整体请求超时(建立连接 + 获取首字节)
|
||||
# 清理可能已建立的连接上下文
|
||||
if response_ctx is not None:
|
||||
try:
|
||||
await response_ctx.__aexit__(None, None, None)
|
||||
except Exception:
|
||||
pass
|
||||
await http_client.aclose()
|
||||
logger.warning(
|
||||
f" [{self.request_id}] 请求超时: Provider={provider.name}, timeout={request_timeout}s"
|
||||
)
|
||||
raise ProviderTimeoutException(
|
||||
provider_name=str(provider.name),
|
||||
timeout=int(request_timeout),
|
||||
)
|
||||
|
||||
except (httpx.ConnectError, httpx.ConnectTimeout, httpx.TimeoutException) as e:
|
||||
# 连接/读写超时:清理可能已建立的连接上下文
|
||||
if response_ctx is not None:
|
||||
try:
|
||||
await response_ctx.__aexit__(None, None, None)
|
||||
except Exception:
|
||||
pass
|
||||
await http_client.aclose()
|
||||
if envelope:
|
||||
envelope.on_connection_error(base_url=ctx.selected_base_url, exc=e)
|
||||
if ctx.selected_base_url:
|
||||
logger.warning(
|
||||
f"[{envelope.name}] Connection error: {ctx.selected_base_url} ({e})"
|
||||
)
|
||||
raise
|
||||
|
||||
except httpx.HTTPStatusError as e:
|
||||
error_text = await self._extract_error_text(e)
|
||||
logger.error(f"Provider 返回错误: {e.response.status_code}\n Response: {error_text}")
|
||||
await http_client.aclose()
|
||||
# 将上游错误信息附加到异常,以便故障转移时能够返回给客户端
|
||||
e.upstream_response = error_text # type: ignore[attr-defined]
|
||||
raise
|
||||
|
||||
except EmbeddedErrorException:
|
||||
for attempt in range(2):
|
||||
try:
|
||||
if response_ctx is not None:
|
||||
await response_ctx.__aexit__(None, None, None)
|
||||
except Exception:
|
||||
pass
|
||||
await http_client.aclose()
|
||||
raise
|
||||
# 使用 asyncio.wait_for 包裹整个"建立连接 + 获取首字节"阶段
|
||||
# stream_first_byte_timeout 控制首字节超时,避免上游长时间无响应
|
||||
# 同时检测客户端断连,避免客户端已断开但服务端仍在等待上游响应
|
||||
if is_disconnected is not None:
|
||||
await wait_for_with_disconnect_detection(
|
||||
_connect_and_prefetch(),
|
||||
timeout=request_timeout,
|
||||
is_disconnected=is_disconnected,
|
||||
request_id=self.request_id,
|
||||
)
|
||||
else:
|
||||
await asyncio.wait_for(_connect_and_prefetch(), timeout=request_timeout)
|
||||
break
|
||||
|
||||
except Exception:
|
||||
await http_client.aclose()
|
||||
raise
|
||||
except ClientDisconnectedException:
|
||||
# 客户端断开连接,清理资源
|
||||
if response_ctx is not None:
|
||||
try:
|
||||
await response_ctx.__aexit__(None, None, None)
|
||||
except Exception:
|
||||
pass
|
||||
await http_client.aclose()
|
||||
logger.warning(f" [{self.request_id}] 客户端在等待首字节时断开连接")
|
||||
ctx.status_code = 499
|
||||
ctx.error_message = "client_disconnected_during_prefetch"
|
||||
raise
|
||||
|
||||
except TimeoutError:
|
||||
# 整体请求超时(建立连接 + 获取首字节)
|
||||
# 清理可能已建立的连接上下文
|
||||
if response_ctx is not None:
|
||||
try:
|
||||
await response_ctx.__aexit__(None, None, None)
|
||||
except Exception:
|
||||
pass
|
||||
await http_client.aclose()
|
||||
logger.warning(
|
||||
f" [{self.request_id}] 请求超时: Provider={provider.name}, timeout={request_timeout}s"
|
||||
)
|
||||
raise ProviderTimeoutException(
|
||||
provider_name=str(provider.name),
|
||||
timeout=int(request_timeout),
|
||||
)
|
||||
|
||||
except (httpx.ConnectError, httpx.ConnectTimeout, httpx.TimeoutException) as e:
|
||||
# 连接/读写超时:清理可能已建立的连接上下文
|
||||
if response_ctx is not None:
|
||||
try:
|
||||
await response_ctx.__aexit__(None, None, None)
|
||||
except Exception:
|
||||
pass
|
||||
await http_client.aclose()
|
||||
if envelope:
|
||||
envelope.on_connection_error(base_url=ctx.selected_base_url, exc=e)
|
||||
if ctx.selected_base_url:
|
||||
logger.warning(
|
||||
f"[{envelope.name}] Connection error: {ctx.selected_base_url} ({e})"
|
||||
)
|
||||
raise
|
||||
|
||||
except httpx.HTTPStatusError as e:
|
||||
status = int(getattr(e.response, "status_code", 0) or 0)
|
||||
if (
|
||||
attempt == 0
|
||||
and status == 401
|
||||
and str(getattr(key, "auth_type", "") or "").lower() == "oauth"
|
||||
):
|
||||
# OAuth token may be revoked/expired earlier than expires_at indicates.
|
||||
# Best-effort: force refresh once on 401 and retry a single time.
|
||||
try:
|
||||
if response_ctx is not None:
|
||||
await response_ctx.__aexit__(None, None, None)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
refreshed_auth = await get_provider_auth(endpoint, key, force_refresh=True)
|
||||
if refreshed_auth:
|
||||
provider_headers[refreshed_auth.auth_header] = refreshed_auth.auth_value
|
||||
ctx.provider_request_headers = provider_headers
|
||||
|
||||
# Reset state for the next attempt.
|
||||
byte_iterator = None
|
||||
prefetched_chunks = None
|
||||
response_ctx = None
|
||||
continue
|
||||
|
||||
error_text = await self._extract_error_text(e)
|
||||
logger.error(
|
||||
f"Provider 返回错误: {e.response.status_code}\n Response: {error_text}"
|
||||
)
|
||||
await http_client.aclose()
|
||||
# 将上游错误信息附加到异常,以便故障转移时能够返回给客户端
|
||||
e.upstream_response = error_text # type: ignore[attr-defined]
|
||||
raise
|
||||
|
||||
except EmbeddedErrorException:
|
||||
try:
|
||||
if response_ctx is not None:
|
||||
await response_ctx.__aexit__(None, None, None)
|
||||
except Exception:
|
||||
pass
|
||||
await http_client.aclose()
|
||||
raise
|
||||
|
||||
except Exception:
|
||||
await http_client.aclose()
|
||||
raise
|
||||
|
||||
# 类型断言:成功执行后这些变量不会为 None
|
||||
assert byte_iterator is not None
|
||||
|
||||
@@ -609,6 +609,10 @@ class CliAdapterBase(ApiAdapter):
|
||||
provider_id: str | None = None,
|
||||
api_key_id: str | None = None,
|
||||
model_name: str | None = None,
|
||||
# Provider 上下文(用于 OAuth 认证和 Antigravity 等特殊路由)
|
||||
auth_type: str | None = None,
|
||||
provider_type: str | None = None,
|
||||
decrypted_auth_config: dict[str, Any] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
测试模型连接性(非流式)
|
||||
@@ -634,6 +638,10 @@ class CliAdapterBase(ApiAdapter):
|
||||
provider_id: 提供商ID
|
||||
api_key_id: API密钥ID
|
||||
model_name: 模型名称
|
||||
auth_type: Key 认证类型("api_key"/"oauth"/"vertex_ai"),
|
||||
OAuth 类型自动使用 Authorization: Bearer 替代端点默认认证头
|
||||
provider_type: 提供商类型(用于 Antigravity v1internal 等特殊路由)
|
||||
decrypted_auth_config: 解密后的 OAuth 配置(Antigravity 需要 project_id)
|
||||
|
||||
Returns:
|
||||
测试响应数据
|
||||
@@ -641,37 +649,84 @@ class CliAdapterBase(ApiAdapter):
|
||||
from src.api.handlers.base.endpoint_checker import run_endpoint_check
|
||||
from src.api.handlers.base.request_builder import apply_body_rules
|
||||
from src.core.api_format.headers import HeaderBuilder
|
||||
from src.core.provider_types import ProviderType
|
||||
|
||||
# 构建请求组件
|
||||
url = cls.build_endpoint_url(base_url, request_data, model_name)
|
||||
is_antigravity = provider_type == ProviderType.ANTIGRAVITY
|
||||
is_oauth = auth_type == "oauth"
|
||||
|
||||
# 合并 CLI 额外头部到 extra_headers
|
||||
# ---- URL ----
|
||||
if is_antigravity:
|
||||
# Antigravity 走 v1internal 端点,模型名在请求体 envelope 中,不在 URL 路径里
|
||||
from src.services.provider.adapters.antigravity.constants import (
|
||||
V1INTERNAL_PATH_TEMPLATE,
|
||||
)
|
||||
from src.services.provider.adapters.antigravity.constants import (
|
||||
get_http_user_agent as _get_antigravity_ua,
|
||||
)
|
||||
from src.services.provider.adapters.antigravity.envelope import wrap_v1internal_request
|
||||
from src.services.provider.adapters.antigravity.url_availability import url_availability
|
||||
|
||||
ordered_urls = url_availability.get_ordered_urls(prefer_daily=True)
|
||||
effective_base_url = ordered_urls[0] if ordered_urls else base_url
|
||||
path = V1INTERNAL_PATH_TEMPLATE.format(action="generateContent")
|
||||
url = f"{str(effective_base_url).rstrip('/')}{path}"
|
||||
else:
|
||||
url = cls.build_endpoint_url(base_url, request_data, model_name)
|
||||
|
||||
# ---- Headers ----
|
||||
cli_extra = cls.get_cli_extra_headers(base_url=base_url)
|
||||
merged_extra = dict(extra_headers) if extra_headers else {}
|
||||
merged_extra.update(cli_extra)
|
||||
|
||||
# 使用统一的头部构建函数
|
||||
# Antigravity 需要特定的 User-Agent
|
||||
if is_antigravity:
|
||||
merged_extra["User-Agent"] = _get_antigravity_ua()
|
||||
|
||||
headers = cls.build_headers_with_extra(api_key, merged_extra if merged_extra else None)
|
||||
|
||||
# OAuth 统一处理:替换端点默认认证头为 Authorization: Bearer
|
||||
# (与 get_provider_auth 返回的 ProviderAuthInfo 行为一致)
|
||||
if is_oauth:
|
||||
from src.core.api_format import get_auth_config_for_endpoint
|
||||
|
||||
default_auth_header, _ = get_auth_config_for_endpoint(cls.FORMAT_ID)
|
||||
if default_auth_header.lower() != "authorization":
|
||||
headers.pop(default_auth_header, None)
|
||||
headers["Authorization"] = f"Bearer {api_key}"
|
||||
|
||||
# ---- Body ----
|
||||
body = cls.build_request_body(request_data, base_url=base_url)
|
||||
|
||||
# 应用请求体规则(在格式转换后应用,确保规则效果不被覆盖)
|
||||
if body_rules:
|
||||
body = apply_body_rules(body, body_rules)
|
||||
|
||||
# 应用请求头规则(在请求头构建后应用)
|
||||
if header_rules:
|
||||
# 获取认证头名称,防止被规则覆盖
|
||||
from src.core.api_format import get_auth_config_for_endpoint
|
||||
# Antigravity:用 v1internal envelope 包装请求体
|
||||
if is_antigravity:
|
||||
project_id = (decrypted_auth_config or {}).get("project_id", "")
|
||||
effective_model = model_name or request_data.get("model", "")
|
||||
body = wrap_v1internal_request(
|
||||
body,
|
||||
project_id=project_id,
|
||||
model=effective_model,
|
||||
)
|
||||
|
||||
auth_header, _ = get_auth_config_for_endpoint(cls.FORMAT_ID)
|
||||
protected_keys = {auth_header.lower(), "content-type"}
|
||||
# ---- Header Rules ----
|
||||
if header_rules:
|
||||
from src.core.api_format import get_auth_config_for_endpoint as _get_auth_cfg
|
||||
|
||||
# 保护实际使用的认证头,而非端点默认的
|
||||
if is_oauth:
|
||||
protected_keys = {"authorization", "content-type"}
|
||||
else:
|
||||
ep_auth_header, _ = _get_auth_cfg(cls.FORMAT_ID)
|
||||
protected_keys = {ep_auth_header.lower(), "content-type"}
|
||||
|
||||
header_builder = HeaderBuilder()
|
||||
header_builder.add_many(headers)
|
||||
header_builder.apply_rules(header_rules, protected_keys)
|
||||
headers = header_builder.build()
|
||||
|
||||
# 获取有效的模型名称
|
||||
# ---- Execute ----
|
||||
effective_model_name = model_name or request_data.get("model")
|
||||
|
||||
return await run_endpoint_check(
|
||||
@@ -680,7 +735,6 @@ class CliAdapterBase(ApiAdapter):
|
||||
headers=headers,
|
||||
json_body=body,
|
||||
api_format=cls.FORMAT_ID,
|
||||
# 用量计算参数(现在强制记录)
|
||||
db=db,
|
||||
user=user,
|
||||
provider_name=provider_name,
|
||||
|
||||
@@ -892,13 +892,48 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
try:
|
||||
resp.raise_for_status()
|
||||
except httpx.HTTPStatusError as e:
|
||||
error_body = ""
|
||||
try:
|
||||
error_body = resp.text[:4000] if resp.text else ""
|
||||
except Exception:
|
||||
# OAuth token may be revoked/expired earlier than expires_at indicates.
|
||||
# Best-effort: force refresh once on 401 and retry a single time.
|
||||
if (
|
||||
resp.status_code == 401
|
||||
and str(getattr(key, "auth_type", "") or "").lower() == "oauth"
|
||||
):
|
||||
refreshed_auth = await get_provider_auth(endpoint, key, force_refresh=True)
|
||||
if refreshed_auth:
|
||||
provider_headers[refreshed_auth.auth_header] = refreshed_auth.auth_value
|
||||
ctx.provider_request_headers = provider_headers
|
||||
|
||||
# retry once
|
||||
resp = await http_client.post(
|
||||
url,
|
||||
json=provider_payload,
|
||||
headers=provider_headers,
|
||||
timeout=httpx.Timeout(request_timeout_sync),
|
||||
)
|
||||
ctx.status_code = resp.status_code
|
||||
ctx.response_headers = dict(resp.headers)
|
||||
if envelope:
|
||||
envelope.on_http_status(
|
||||
base_url=ctx.selected_base_url, status_code=ctx.status_code
|
||||
)
|
||||
try:
|
||||
resp.raise_for_status()
|
||||
except httpx.HTTPStatusError as e2:
|
||||
error_body = ""
|
||||
try:
|
||||
error_body = resp.text[:4000] if resp.text else ""
|
||||
except Exception:
|
||||
error_body = ""
|
||||
e2.upstream_response = error_body # type: ignore[attr-defined]
|
||||
raise
|
||||
else:
|
||||
error_body = ""
|
||||
e.upstream_response = error_body # type: ignore[attr-defined]
|
||||
raise
|
||||
try:
|
||||
error_body = resp.text[:4000] if resp.text else ""
|
||||
except Exception:
|
||||
error_body = ""
|
||||
e.upstream_response = error_body # type: ignore[attr-defined]
|
||||
raise
|
||||
|
||||
# Safe JSON parsing.
|
||||
try:
|
||||
@@ -1055,85 +1090,112 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
byte_iterator, provider, endpoint, ctx
|
||||
)
|
||||
|
||||
try:
|
||||
# 使用 asyncio.wait_for 包裹整个"建立连接 + 获取首字节"阶段
|
||||
# stream_first_byte_timeout 控制首字节超时,避免上游长时间无响应
|
||||
# 同时检测客户端断连,避免客户端已断开但服务端仍在等待上游响应
|
||||
if http_request is not None:
|
||||
await wait_for_with_disconnect_detection(
|
||||
_connect_and_prefetch(),
|
||||
timeout=request_timeout,
|
||||
is_disconnected=http_request.is_disconnected,
|
||||
request_id=self.request_id,
|
||||
)
|
||||
else:
|
||||
await asyncio.wait_for(_connect_and_prefetch(), timeout=request_timeout)
|
||||
|
||||
except TimeoutError as e:
|
||||
# 整体请求超时(建立连接 + 获取首字节)
|
||||
# 清理可能已建立的连接上下文
|
||||
if response_ctx is not None:
|
||||
try:
|
||||
await response_ctx.__aexit__(None, None, None)
|
||||
except Exception:
|
||||
pass
|
||||
if envelope:
|
||||
envelope.on_connection_error(base_url=ctx.selected_base_url, exc=e)
|
||||
await http_client.aclose()
|
||||
logger.warning(
|
||||
f" [{self.request_id}] 请求超时: Provider={provider.name}, timeout={request_timeout}s"
|
||||
)
|
||||
raise ProviderTimeoutException(
|
||||
provider_name=str(provider.name),
|
||||
timeout=int(request_timeout),
|
||||
)
|
||||
|
||||
except ClientDisconnectedException:
|
||||
# 客户端断开连接,清理资源
|
||||
if response_ctx is not None:
|
||||
try:
|
||||
await response_ctx.__aexit__(None, None, None)
|
||||
except Exception:
|
||||
pass
|
||||
await http_client.aclose()
|
||||
logger.warning(f" [{self.request_id}] 客户端在等待首字节时断开连接")
|
||||
ctx.status_code = 499
|
||||
ctx.error_message = "client_disconnected_during_prefetch"
|
||||
raise
|
||||
|
||||
except (httpx.ConnectError, httpx.ConnectTimeout, httpx.TimeoutException) as e:
|
||||
if envelope:
|
||||
envelope.on_connection_error(base_url=ctx.selected_base_url, exc=e)
|
||||
if ctx.selected_base_url:
|
||||
logger.warning(
|
||||
f"[{envelope.name}] Connection error: {ctx.selected_base_url} ({e})"
|
||||
)
|
||||
await http_client.aclose()
|
||||
raise
|
||||
|
||||
except httpx.HTTPStatusError as e:
|
||||
error_text = await self._extract_error_text(e)
|
||||
logger.error(
|
||||
f"Provider 返回错误状态: {e.response.status_code}\n Response: {error_text}"
|
||||
)
|
||||
await http_client.aclose()
|
||||
# 将上游错误信息附加到异常,以便故障转移时能够返回给客户端
|
||||
e.upstream_response = error_text # type: ignore[attr-defined]
|
||||
raise
|
||||
|
||||
except EmbeddedErrorException:
|
||||
# 嵌套错误需要触发重试,关闭连接后重新抛出
|
||||
for attempt in range(2):
|
||||
try:
|
||||
if response_ctx is not None:
|
||||
await response_ctx.__aexit__(None, None, None)
|
||||
except Exception:
|
||||
pass
|
||||
await http_client.aclose()
|
||||
raise
|
||||
# 使用 asyncio.wait_for 包裹整个"建立连接 + 获取首字节"阶段
|
||||
# stream_first_byte_timeout 控制首字节超时,避免上游长时间无响应
|
||||
# 同时检测客户端断连,避免客户端已断开但服务端仍在等待上游响应
|
||||
if http_request is not None:
|
||||
await wait_for_with_disconnect_detection(
|
||||
_connect_and_prefetch(),
|
||||
timeout=request_timeout,
|
||||
is_disconnected=http_request.is_disconnected,
|
||||
request_id=self.request_id,
|
||||
)
|
||||
else:
|
||||
await asyncio.wait_for(_connect_and_prefetch(), timeout=request_timeout)
|
||||
break
|
||||
|
||||
except Exception:
|
||||
await http_client.aclose()
|
||||
raise
|
||||
except TimeoutError as e:
|
||||
# 整体请求超时(建立连接 + 获取首字节)
|
||||
# 清理可能已建立的连接上下文
|
||||
if response_ctx is not None:
|
||||
try:
|
||||
await response_ctx.__aexit__(None, None, None)
|
||||
except Exception:
|
||||
pass
|
||||
if envelope:
|
||||
envelope.on_connection_error(base_url=ctx.selected_base_url, exc=e)
|
||||
await http_client.aclose()
|
||||
logger.warning(
|
||||
f" [{self.request_id}] 请求超时: Provider={provider.name}, timeout={request_timeout}s"
|
||||
)
|
||||
raise ProviderTimeoutException(
|
||||
provider_name=str(provider.name),
|
||||
timeout=int(request_timeout),
|
||||
)
|
||||
|
||||
except ClientDisconnectedException:
|
||||
# 客户端断开连接,清理资源
|
||||
if response_ctx is not None:
|
||||
try:
|
||||
await response_ctx.__aexit__(None, None, None)
|
||||
except Exception:
|
||||
pass
|
||||
await http_client.aclose()
|
||||
logger.warning(f" [{self.request_id}] 客户端在等待首字节时断开连接")
|
||||
ctx.status_code = 499
|
||||
ctx.error_message = "client_disconnected_during_prefetch"
|
||||
raise
|
||||
|
||||
except (httpx.ConnectError, httpx.ConnectTimeout, httpx.TimeoutException) as e:
|
||||
if envelope:
|
||||
envelope.on_connection_error(base_url=ctx.selected_base_url, exc=e)
|
||||
if ctx.selected_base_url:
|
||||
logger.warning(
|
||||
f"[{envelope.name}] Connection error: {ctx.selected_base_url} ({e})"
|
||||
)
|
||||
await http_client.aclose()
|
||||
raise
|
||||
|
||||
except httpx.HTTPStatusError as e:
|
||||
status = int(getattr(e.response, "status_code", 0) or 0)
|
||||
if (
|
||||
attempt == 0
|
||||
and status == 401
|
||||
and str(getattr(key, "auth_type", "") or "").lower() == "oauth"
|
||||
):
|
||||
# OAuth token may be revoked/expired earlier than expires_at indicates.
|
||||
# Best-effort: force refresh once on 401 and retry a single time.
|
||||
try:
|
||||
if response_ctx is not None:
|
||||
await response_ctx.__aexit__(None, None, None)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
refreshed_auth = await get_provider_auth(endpoint, key, force_refresh=True)
|
||||
if refreshed_auth:
|
||||
provider_headers[refreshed_auth.auth_header] = refreshed_auth.auth_value
|
||||
ctx.provider_request_headers = provider_headers
|
||||
|
||||
# Reset state for the next attempt.
|
||||
byte_iterator = None
|
||||
prefetched_chunks = None
|
||||
response_ctx = None
|
||||
continue
|
||||
|
||||
error_text = await self._extract_error_text(e)
|
||||
logger.error(
|
||||
f"Provider 返回错误状态: {e.response.status_code}\n Response: {error_text}"
|
||||
)
|
||||
await http_client.aclose()
|
||||
# 将上游错误信息附加到异常,以便故障转移时能够返回给客户端
|
||||
e.upstream_response = error_text # type: ignore[attr-defined]
|
||||
raise
|
||||
|
||||
except EmbeddedErrorException:
|
||||
# 嵌套错误需要触发重试,关闭连接后重新抛出
|
||||
try:
|
||||
if response_ctx is not None:
|
||||
await response_ctx.__aexit__(None, None, None)
|
||||
except Exception:
|
||||
pass
|
||||
await http_client.aclose()
|
||||
raise
|
||||
|
||||
except Exception:
|
||||
await http_client.aclose()
|
||||
raise
|
||||
|
||||
# 类型断言:成功执行后这些变量不会为 None
|
||||
assert byte_iterator is not None
|
||||
|
||||
@@ -599,6 +599,8 @@ def build_passthrough_request(
|
||||
async def get_provider_auth(
|
||||
endpoint: "ProviderEndpoint",
|
||||
key: "ProviderAPIKey",
|
||||
*,
|
||||
force_refresh: bool = False,
|
||||
) -> ProviderAuthInfo | None:
|
||||
"""
|
||||
获取 Provider 的认证信息
|
||||
@@ -638,7 +640,7 @@ async def get_provider_auth(
|
||||
refresh_token = token_meta.get("refresh_token")
|
||||
provider_type = str(token_meta.get("provider_type") or "")
|
||||
|
||||
# 120s skew
|
||||
# 120s skew (or force refresh when upstream returns 401)
|
||||
should_refresh = False
|
||||
try:
|
||||
if expires_at is not None:
|
||||
@@ -646,6 +648,9 @@ async def get_provider_auth(
|
||||
except Exception:
|
||||
should_refresh = False
|
||||
|
||||
if force_refresh:
|
||||
should_refresh = True
|
||||
|
||||
if should_refresh and refresh_token and provider_type:
|
||||
try:
|
||||
from src.core.provider_templates.fixed_providers import FIXED_PROVIDERS
|
||||
|
||||
@@ -578,6 +578,15 @@ class StreamProcessor:
|
||||
if not isinstance(data_obj, dict):
|
||||
return []
|
||||
|
||||
# Provider envelope: unwrap v1internal wrapper before conversion
|
||||
# (e.g. Antigravity {"response": {...}, "traceId": "..."} → inner response)
|
||||
if envelope and isinstance(data_obj, dict):
|
||||
data_obj = envelope.unwrap_response(data_obj)
|
||||
envelope.postprocess_unwrapped_response(
|
||||
model=str(ctx.model or ""),
|
||||
data=data_obj,
|
||||
)
|
||||
|
||||
try:
|
||||
converted_events = registry.convert_stream_chunk(
|
||||
data_obj,
|
||||
|
||||
@@ -160,7 +160,11 @@ class ClaudeChatAdapter(ChatAdapterBase):
|
||||
api_key: str,
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
) -> tuple[list, str | None]:
|
||||
"""查询 Claude API 支持的模型列表"""
|
||||
"""查询 Claude API 支持的模型列表
|
||||
|
||||
Anthropic 的 /v1/models 是分页接口(has_more/first_id/last_id),
|
||||
默认只返回一页。这里做 best-effort 的全量拉取,确保管理端能展示完整模型列表。
|
||||
"""
|
||||
headers = cls.build_headers_with_extra(api_key, extra_headers)
|
||||
|
||||
# 构建 /v1/models URL
|
||||
@@ -171,24 +175,60 @@ class ClaudeChatAdapter(ChatAdapterBase):
|
||||
models_url = f"{base_url}/v1/models"
|
||||
|
||||
try:
|
||||
response = await client.get(models_url, headers=headers)
|
||||
logger.debug(f"Claude models request to {models_url}: status={response.status_code}")
|
||||
if response.status_code == 200:
|
||||
all_models: list[dict] = []
|
||||
seen_ids: set[str] = set()
|
||||
|
||||
after_id: str | None = None
|
||||
limit = 100 # Anthropic 支持 limit,尽量减少分页次数
|
||||
max_pages = 20 # safety guard
|
||||
|
||||
for _ in range(max_pages):
|
||||
params: dict[str, Any] = {"limit": limit}
|
||||
if after_id:
|
||||
params["after_id"] = after_id
|
||||
|
||||
response = await client.get(models_url, headers=headers, params=params)
|
||||
logger.debug(
|
||||
f"Claude models request to {models_url}: status={response.status_code}, after_id={after_id}"
|
||||
)
|
||||
if response.status_code != 200:
|
||||
error_body = response.text[:500] if response.text else "(empty)"
|
||||
error_msg = f"HTTP {response.status_code}: {error_body}"
|
||||
logger.warning(f"Claude models request to {models_url} failed: {error_msg}")
|
||||
return [], error_msg
|
||||
|
||||
data = response.json()
|
||||
models = []
|
||||
if "data" in data:
|
||||
models = data["data"]
|
||||
page_models: list[dict] = []
|
||||
if isinstance(data, dict) and isinstance(data.get("data"), list):
|
||||
page_models = [m for m in data["data"] if isinstance(m, dict)]
|
||||
elif isinstance(data, list):
|
||||
models = data
|
||||
# 为每个模型添加 api_format 字段
|
||||
for m in models:
|
||||
page_models = [m for m in data if isinstance(m, dict)]
|
||||
|
||||
for m in page_models:
|
||||
mid = m.get("id")
|
||||
if isinstance(mid, str) and mid and mid in seen_ids:
|
||||
continue
|
||||
if isinstance(mid, str) and mid:
|
||||
seen_ids.add(mid)
|
||||
m["api_format"] = cls.FORMAT_ID
|
||||
return models, None
|
||||
else:
|
||||
error_body = response.text[:500] if response.text else "(empty)"
|
||||
error_msg = f"HTTP {response.status_code}: {error_body}"
|
||||
logger.warning(f"Claude models request to {models_url} failed: {error_msg}")
|
||||
return [], error_msg
|
||||
all_models.append(m)
|
||||
|
||||
# Pagination (Anthropic list response shape)
|
||||
if not isinstance(data, dict):
|
||||
break
|
||||
|
||||
has_more = bool(data.get("has_more"))
|
||||
last_id = data.get("last_id")
|
||||
if not has_more:
|
||||
break
|
||||
if not isinstance(last_id, str) or not last_id:
|
||||
break
|
||||
if after_id == last_id:
|
||||
# Prevent infinite loops on unexpected upstream behavior.
|
||||
break
|
||||
after_id = last_id
|
||||
|
||||
return all_models, None
|
||||
except Exception as e:
|
||||
error_msg = f"Request error: {str(e)}"
|
||||
logger.warning(f"Failed to fetch Claude models from {models_url}: {e}")
|
||||
|
||||
@@ -137,7 +137,10 @@ class GeminiChatAdapter(ChatAdapterBase):
|
||||
|
||||
contents = getattr(request_obj, "contents", []) or []
|
||||
for content in contents:
|
||||
role = getattr(content, "role", None) or content.get("role", "unknown")
|
||||
if isinstance(content, dict):
|
||||
role = content.get("role", "unknown")
|
||||
else:
|
||||
role = getattr(content, "role", None) or "unknown"
|
||||
role_counts[role] = role_counts.get(role, 0) + 1
|
||||
|
||||
generation_config = getattr(request_obj, "generation_config", None) or {}
|
||||
@@ -262,6 +265,10 @@ class GeminiChatAdapter(ChatAdapterBase):
|
||||
provider_id: str | None = None,
|
||||
api_key_id: str | None = None,
|
||||
model_name: str | None = None,
|
||||
# Provider 上下文(Gemini Chat 适配器忽略,仅保持签名兼容)
|
||||
auth_type: str | None = None, # noqa: ARG003
|
||||
provider_type: str | None = None, # noqa: ARG003
|
||||
decrypted_auth_config: dict[str, Any] | None = None, # noqa: ARG003
|
||||
) -> dict[str, Any]:
|
||||
"""测试 Gemini API 模型连接性(非流式)"""
|
||||
from src.api.handlers.base.endpoint_checker import run_endpoint_check
|
||||
@@ -292,6 +299,7 @@ class GeminiChatAdapter(ChatAdapterBase):
|
||||
if header_rules:
|
||||
# 获取认证头名称,防止被规则覆盖
|
||||
from src.core.api_format import get_auth_config_for_endpoint
|
||||
|
||||
auth_header, _ = get_auth_config_for_endpoint(cls.FORMAT_ID)
|
||||
protected_keys = {auth_header.lower(), "content-type"}
|
||||
|
||||
|
||||
@@ -14,7 +14,7 @@ from fastapi.responses import JSONResponse, Response, StreamingResponse
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.handlers.base.request_builder import get_provider_auth
|
||||
from src.api.handlers.base.request_builder import apply_body_rules, get_provider_auth
|
||||
from src.api.handlers.base.video_handler_base import (
|
||||
VideoHandlerBase,
|
||||
normalize_gemini_operation_id,
|
||||
@@ -145,6 +145,9 @@ class GeminiVeoHandler(VideoHandlerBase):
|
||||
format_conversion_info["provider_format"] = provider_format
|
||||
format_conversion_info["converted"] = needs_conversion
|
||||
|
||||
# 应用端点的请求体规则
|
||||
endpoint_body_rules = getattr(endpoint, "body_rules", None)
|
||||
|
||||
if needs_conversion and provider_format.upper().startswith("OPENAI:"):
|
||||
# Gemini -> OpenAI 格式转换
|
||||
converted_body = format_conversion_registry.convert_video_request(
|
||||
@@ -156,6 +159,9 @@ class GeminiVeoHandler(VideoHandlerBase):
|
||||
if "seconds" in converted_body and converted_body["seconds"] is not None:
|
||||
converted_body["seconds"] = str(converted_body["seconds"])
|
||||
|
||||
if endpoint_body_rules:
|
||||
converted_body = apply_body_rules(converted_body, endpoint_body_rules)
|
||||
|
||||
# 构建 OpenAI 风格的 URL
|
||||
upstream_url = self._build_openai_upstream_url(endpoint.base_url)
|
||||
|
||||
@@ -168,12 +174,18 @@ class GeminiVeoHandler(VideoHandlerBase):
|
||||
return await client.post(upstream_url, headers=headers, json=converted_body)
|
||||
else:
|
||||
# 原始 Gemini 格式
|
||||
request_body = (
|
||||
original_request_body.copy() if endpoint_body_rules else original_request_body
|
||||
)
|
||||
if endpoint_body_rules:
|
||||
request_body = apply_body_rules(request_body, endpoint_body_rules)
|
||||
|
||||
upstream_url = self._build_upstream_url(endpoint.base_url, internal_request.model)
|
||||
headers = self._build_upstream_headers(
|
||||
original_headers, upstream_key, endpoint, auth_info
|
||||
)
|
||||
client = await HTTPClientPool.get_default_client_async()
|
||||
return await client.post(upstream_url, headers=headers, json=original_request_body)
|
||||
return await client.post(upstream_url, headers=headers, json=request_body)
|
||||
|
||||
def _extract_task_id(payload: dict[str, Any]) -> str | None:
|
||||
# 根据响应格式提取 task ID
|
||||
@@ -514,6 +526,7 @@ class GeminiVeoHandler(VideoHandlerBase):
|
||||
endpoint_sig,
|
||||
upstream_key,
|
||||
endpoint_headers=extra_headers,
|
||||
header_rules=getattr(endpoint, "header_rules", None),
|
||||
)
|
||||
if auth_info:
|
||||
# 覆盖为 OAuth2 Bearer(Vertex AI)
|
||||
@@ -557,6 +570,7 @@ class GeminiVeoHandler(VideoHandlerBase):
|
||||
endpoint_sig,
|
||||
upstream_key,
|
||||
endpoint_headers=extra_headers,
|
||||
header_rules=getattr(endpoint, "header_rules", None),
|
||||
)
|
||||
|
||||
def _create_task_record(
|
||||
|
||||
@@ -114,7 +114,9 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
|
||||
Returns:
|
||||
包含 input_tokens, output_tokens, cached_tokens 的字典
|
||||
"""
|
||||
if str(provider_type or "").lower() == "antigravity":
|
||||
from src.core.provider_types import ProviderType
|
||||
|
||||
if str(provider_type or "").lower() == ProviderType.ANTIGRAVITY:
|
||||
return self._extract_antigravity_usage(event)
|
||||
|
||||
from src.api.handlers.gemini.stream_parser import GeminiStreamParser
|
||||
|
||||
@@ -16,7 +16,7 @@ from fastapi.responses import JSONResponse, Response, StreamingResponse
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.handlers.base.request_builder import get_provider_auth
|
||||
from src.api.handlers.base.request_builder import apply_body_rules, get_provider_auth
|
||||
from src.api.handlers.base.video_handler_base import VideoHandlerBase, sanitize_error_message
|
||||
from src.clients.http_client import HTTPClientPool
|
||||
from src.config.settings import config
|
||||
@@ -139,6 +139,9 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
||||
if "seconds" in request_body and request_body["seconds"] is not None:
|
||||
request_body["seconds"] = str(request_body["seconds"])
|
||||
|
||||
# 应用端点的请求体规则
|
||||
endpoint_body_rules = getattr(endpoint, "body_rules", None)
|
||||
|
||||
if needs_conversion and provider_format.upper().startswith("GEMINI:"):
|
||||
# OpenAI -> Gemini 格式转换
|
||||
converted_body = format_conversion_registry.convert_video_request(
|
||||
@@ -150,6 +153,9 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
||||
if "model" not in converted_body:
|
||||
converted_body["model"] = internal_request.model
|
||||
|
||||
if endpoint_body_rules:
|
||||
converted_body = apply_body_rules(converted_body, endpoint_body_rules)
|
||||
|
||||
# 构建 Gemini 风格的 URL
|
||||
upstream_url = self._build_gemini_upstream_url(
|
||||
endpoint.base_url, internal_request.model
|
||||
@@ -165,6 +171,9 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
||||
return await client.post(upstream_url, headers=headers, json=converted_body)
|
||||
else:
|
||||
# 原始 OpenAI 格式
|
||||
if endpoint_body_rules:
|
||||
request_body = apply_body_rules(request_body, endpoint_body_rules)
|
||||
|
||||
upstream_url = self._build_upstream_url(endpoint.base_url)
|
||||
headers = self._build_upstream_headers(original_headers, upstream_key, endpoint)
|
||||
client = await HTTPClientPool.get_default_client_async()
|
||||
@@ -500,6 +509,11 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
||||
if "seconds" in request_body and request_body["seconds"] is not None:
|
||||
request_body["seconds"] = str(request_body["seconds"])
|
||||
|
||||
# 应用端点的请求体规则
|
||||
endpoint_body_rules = getattr(endpoint, "body_rules", None)
|
||||
if endpoint_body_rules:
|
||||
request_body = apply_body_rules(request_body, endpoint_body_rules)
|
||||
|
||||
client = await HTTPClientPool.get_default_client_async()
|
||||
response = await client.post(upstream_url, headers=headers, json=request_body)
|
||||
|
||||
@@ -785,6 +799,7 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
||||
endpoint_sig,
|
||||
upstream_key,
|
||||
endpoint_headers=extra_headers,
|
||||
header_rules=getattr(endpoint, "header_rules", None),
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
@@ -816,6 +831,7 @@ class OpenAIVideoHandler(VideoHandlerBase):
|
||||
endpoint_sig,
|
||||
upstream_key,
|
||||
endpoint_headers=extra_headers,
|
||||
header_rules=getattr(endpoint, "header_rules", None),
|
||||
)
|
||||
if auth_info:
|
||||
# 覆盖为 OAuth2 Bearer(Vertex AI)
|
||||
|
||||
@@ -126,7 +126,7 @@ class OpenAICliAdapter(CliAdapterBase):
|
||||
# 仅 Codex 端点添加特定头部
|
||||
if base_url and is_codex_url(base_url):
|
||||
# 与运行时路径保持一致:使用 Codex envelope 的 best-effort headers。
|
||||
from src.services.codex.envelope import codex_oauth_envelope
|
||||
from src.services.provider.adapters.codex.envelope import codex_oauth_envelope
|
||||
|
||||
headers.update(codex_oauth_envelope.extra_headers() or {})
|
||||
|
||||
|
||||
Reference in New Issue
Block a user