mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
feat: OAuth 账户管理、维护调度、端点健康检查增强及前端优化
- 新增 OAuth 账户管理对话框和提供商详情抽屉中的 OAuth 信息展示 - 新增维护调度器(maintenance_scheduler)支持定时清理和健康检查 - 增强端点健康检查器,支持更多检测策略 - 重构 codex 服务为 metadata_collectors 模块 - 优化 OpenAI CLI normalizer 代码结构 - 前端: 改进使用量表格、统计图表、指南页面和异步任务管理 - 扩展多个数据库字符串列为 TEXT 类型 - 新增倒计时 composable 和 provider OAuth API 端点
This commit is contained in:
@@ -242,9 +242,7 @@ class AdminUpdateEndpointKeyAdapter(AdminApiAdapter):
|
||||
if "auth_type" in update_data:
|
||||
if target_auth_type == "api_key":
|
||||
if current_auth_type in {"vertex_ai", "oauth"} and not update_data.get("api_key"):
|
||||
raise InvalidRequestException(
|
||||
"切换到 API Key 认证模式时,必须提供新的 API Key"
|
||||
)
|
||||
raise InvalidRequestException("切换到 API Key 认证模式时,必须提供新的 API Key")
|
||||
# 切换回 API Key:清理非本模式配置
|
||||
update_data["auth_config"] = None
|
||||
elif target_auth_type == "vertex_ai":
|
||||
@@ -629,16 +627,29 @@ def _build_key_response(
|
||||
key_dict.pop("_sa_instance_state", None)
|
||||
key_dict.pop("api_key", None) # 移除敏感字段,避免泄露
|
||||
|
||||
# 提取 OAuth expires_at(如果是 OAuth 类型)
|
||||
# 提取 OAuth 元数据(如果是 OAuth 类型)
|
||||
oauth_expires_at = None
|
||||
oauth_email = None
|
||||
oauth_plan_type = None
|
||||
oauth_account_id = None
|
||||
encrypted_auth_config = key_dict.pop("auth_config", None) # 移除敏感字段,避免泄露
|
||||
if auth_type == "oauth" and encrypted_auth_config:
|
||||
try:
|
||||
decrypted_config = crypto_service.decrypt(encrypted_auth_config)
|
||||
auth_config = json.loads(decrypted_config)
|
||||
oauth_expires_at = auth_config.get("expires_at")
|
||||
except Exception:
|
||||
pass
|
||||
oauth_email = auth_config.get("email")
|
||||
oauth_plan_type = auth_config.get("plan_type") # Codex: plus/free/team/enterprise
|
||||
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)
|
||||
|
||||
# 从 health_by_format 计算汇总字段(便于列表展示)
|
||||
health_by_format = key.health_by_format or {}
|
||||
@@ -685,6 +696,13 @@ def _build_key_response(
|
||||
"circuit_breaker_open": any_circuit_open,
|
||||
# OAuth 相关
|
||||
"oauth_expires_at": oauth_expires_at,
|
||||
"oauth_email": oauth_email,
|
||||
"oauth_plan_type": oauth_plan_type,
|
||||
"oauth_account_id": oauth_account_id,
|
||||
"oauth_invalid_at": (
|
||||
int(key.oauth_invalid_at.timestamp()) if key.oauth_invalid_at else None
|
||||
),
|
||||
"oauth_invalid_reason": key.oauth_invalid_reason,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@@ -12,14 +12,13 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import hashlib
|
||||
import json
|
||||
import secrets
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
import base64
|
||||
import hashlib
|
||||
|
||||
from urllib.parse import parse_qsl, urlencode, urlparse
|
||||
|
||||
import httpx
|
||||
@@ -32,12 +31,11 @@ from src.clients.redis_client import get_redis_client
|
||||
from src.core.crypto import crypto_service
|
||||
from src.core.exceptions import InvalidRequestException, NotFoundException
|
||||
from src.core.logger import logger
|
||||
from src.database.database import get_db
|
||||
from src.models.database import Provider, ProviderAPIKey
|
||||
|
||||
from src.core.provider_oauth_utils import enrich_auth_config, post_oauth_token
|
||||
from src.core.provider_templates.fixed_providers import FIXED_PROVIDERS
|
||||
from src.core.provider_templates.types import ProviderType
|
||||
from src.database.database import get_db
|
||||
from src.models.database import Provider, ProviderAPIKey
|
||||
|
||||
router = APIRouter(prefix="/api/admin/provider-oauth", tags=["Provider OAuth"])
|
||||
|
||||
@@ -66,7 +64,8 @@ def _state_key(nonce: str) -> str:
|
||||
@dataclass(frozen=True)
|
||||
class ProviderOAuthStateData:
|
||||
nonce: str
|
||||
key_id: str
|
||||
key_id: str # 可能为空(新流程)
|
||||
provider_id: str # 新增
|
||||
provider_type: str
|
||||
pkce_verifier: str | None
|
||||
created_at: int
|
||||
@@ -76,6 +75,7 @@ async def _create_state(
|
||||
redis: Redis,
|
||||
*,
|
||||
key_id: str,
|
||||
provider_id: str,
|
||||
provider_type: str,
|
||||
pkce_verifier: str | None,
|
||||
) -> str:
|
||||
@@ -83,6 +83,7 @@ async def _create_state(
|
||||
data = {
|
||||
"nonce": nonce,
|
||||
"key_id": key_id,
|
||||
"provider_id": provider_id,
|
||||
"provider_type": provider_type,
|
||||
"pkce_verifier": pkce_verifier,
|
||||
"created_at": int(time.time()),
|
||||
@@ -105,6 +106,7 @@ async def _consume_state(redis: Redis, nonce: str) -> ProviderOAuthStateData | N
|
||||
return ProviderOAuthStateData(
|
||||
nonce=str(parsed.get("nonce") or ""),
|
||||
key_id=str(parsed.get("key_id") or ""),
|
||||
provider_id=str(parsed.get("provider_id") or ""),
|
||||
provider_type=str(parsed.get("provider_type") or ""),
|
||||
pkce_verifier=parsed.get("pkce_verifier"),
|
||||
created_at=int(parsed.get("created_at") or 0),
|
||||
@@ -131,6 +133,20 @@ class CompleteOAuthResponse(BaseModel):
|
||||
provider_type: str
|
||||
expires_at: int | None = None
|
||||
has_refresh_token: bool = False
|
||||
email: str | None = None
|
||||
|
||||
|
||||
class ProviderCompleteOAuthRequest(BaseModel):
|
||||
callback_url: str = Field(..., min_length=5, description="浏览器地址栏中的完整回调 URL")
|
||||
name: str | None = Field(None, max_length=100, description="账号名称(可选)")
|
||||
|
||||
|
||||
class ProviderCompleteOAuthResponse(BaseModel):
|
||||
key_id: str
|
||||
provider_type: str
|
||||
expires_at: int | None = None
|
||||
has_refresh_token: bool = False
|
||||
email: str | None = None
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
@@ -179,7 +195,11 @@ async def supported_types() -> list[dict[str, Any]]:
|
||||
for provider_type, template in FIXED_PROVIDERS.items():
|
||||
result.append(
|
||||
{
|
||||
"provider_type": str(provider_type.value) if hasattr(provider_type, "value") else str(provider_type),
|
||||
"provider_type": (
|
||||
str(provider_type.value)
|
||||
if hasattr(provider_type, "value")
|
||||
else str(provider_type)
|
||||
),
|
||||
"display_name": template.display_name,
|
||||
"scopes": list(template.oauth.scopes),
|
||||
"redirect_uri": template.oauth.redirect_uri,
|
||||
@@ -228,6 +248,7 @@ async def start_oauth(
|
||||
state = await _create_state(
|
||||
redis,
|
||||
key_id=key_id,
|
||||
provider_id=str(provider.id),
|
||||
provider_type=provider_type,
|
||||
pkce_verifier=pkce_verifier,
|
||||
)
|
||||
@@ -336,7 +357,10 @@ async def complete_oauth(
|
||||
form["client_secret"] = template.oauth.client_secret
|
||||
if state_data.pkce_verifier:
|
||||
form["code_verifier"] = state_data.pkce_verifier
|
||||
headers = {"Content-Type": "application/x-www-form-urlencoded", "Accept": "application/json"}
|
||||
headers = {
|
||||
"Content-Type": "application/x-www-form-urlencoded",
|
||||
"Accept": "application/json",
|
||||
}
|
||||
data = form
|
||||
json_body = None
|
||||
|
||||
@@ -391,10 +415,19 @@ 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,
|
||||
has_refresh_token=bool(refresh_token),
|
||||
email=auth_config.get("email"),
|
||||
)
|
||||
|
||||
|
||||
@@ -452,7 +485,10 @@ async def refresh_oauth(
|
||||
}
|
||||
if template.oauth.client_secret:
|
||||
form["client_secret"] = template.oauth.client_secret
|
||||
headers = {"Content-Type": "application/x-www-form-urlencoded", "Accept": "application/json"}
|
||||
headers = {
|
||||
"Content-Type": "application/x-www-form-urlencoded",
|
||||
"Accept": "application/json",
|
||||
}
|
||||
data = form
|
||||
json_body = None
|
||||
|
||||
@@ -469,7 +505,25 @@ async def refresh_oauth(
|
||||
)
|
||||
|
||||
if resp.status_code < 200 or resp.status_code >= 300:
|
||||
raise InvalidRequestException("token refresh 失败")
|
||||
# 解析错误原因
|
||||
error_reason = f"HTTP {resp.status_code}"
|
||||
try:
|
||||
error_body = resp.json()
|
||||
if "error" in error_body:
|
||||
error_reason = str(error_body.get("error_description") or error_body.get("error"))
|
||||
except Exception:
|
||||
error_reason = resp.text[:100] if resp.text else f"HTTP {resp.status_code}"
|
||||
|
||||
# 标记为失效(400/401/403 通常表示永久性错误)
|
||||
if resp.status_code in (400, 401, 403):
|
||||
from datetime import datetime, timezone
|
||||
|
||||
key.oauth_invalid_at = datetime.now(timezone.utc)
|
||||
key.oauth_invalid_reason = error_reason
|
||||
db.commit()
|
||||
logger.warning("Key {} OAuth token 刷新失败,已标记为失效: {}", key_id, error_reason)
|
||||
|
||||
raise InvalidRequestException(f"token refresh 失败: {error_reason}")
|
||||
|
||||
token = resp.json()
|
||||
access_token = str(token.get("access_token") or "")
|
||||
@@ -503,10 +557,247 @@ async def refresh_oauth(
|
||||
)
|
||||
|
||||
key.auth_config = crypto_service.encrypt(json.dumps(parsed))
|
||||
# 刷新成功,清除失效标记
|
||||
key.oauth_invalid_at = None
|
||||
key.oauth_invalid_reason = None
|
||||
db.commit()
|
||||
|
||||
return CompleteOAuthResponse(
|
||||
provider_type=provider_type,
|
||||
expires_at=expires_at,
|
||||
has_refresh_token=bool(parsed.get("refresh_token")),
|
||||
email=parsed.get("email"),
|
||||
)
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
# Provider-level OAuth (不需要预先创建 key)
|
||||
# ==============================================================================
|
||||
|
||||
|
||||
@router.post("/providers/{provider_id}/start", response_model=StartOAuthResponse)
|
||||
async def start_provider_oauth(
|
||||
provider_id: str,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
) -> StartOAuthResponse:
|
||||
"""基于 Provider 启动 OAuth(不需要预先创建 key)。"""
|
||||
provider = db.query(Provider).filter(Provider.id == provider_id).first()
|
||||
if not provider:
|
||||
raise NotFoundException("Provider 不存在", "provider")
|
||||
provider_type = _require_fixed_provider(provider)
|
||||
|
||||
try:
|
||||
template = FIXED_PROVIDERS.get(ProviderType(provider_type))
|
||||
except Exception:
|
||||
template = None
|
||||
if not template:
|
||||
raise InvalidRequestException("不支持的 provider_type")
|
||||
|
||||
redis = await get_redis_client(require_redis=True)
|
||||
assert redis is not None
|
||||
|
||||
pkce_verifier: str | None = None
|
||||
code_challenge: str | None = None
|
||||
|
||||
if template.oauth.use_pkce:
|
||||
pkce_verifier = secrets.token_urlsafe(32)
|
||||
code_challenge = _pkce_s256(pkce_verifier)
|
||||
|
||||
state = await _create_state(
|
||||
redis,
|
||||
key_id="", # 空,complete 时创建
|
||||
provider_id=provider_id,
|
||||
provider_type=provider_type,
|
||||
pkce_verifier=pkce_verifier,
|
||||
)
|
||||
|
||||
params: dict[str, Any] = {
|
||||
"client_id": template.oauth.client_id,
|
||||
"response_type": "code",
|
||||
"redirect_uri": template.oauth.redirect_uri,
|
||||
"scope": " ".join(template.oauth.scopes),
|
||||
"state": state,
|
||||
}
|
||||
|
||||
if provider_type == ProviderType.CODEX.value:
|
||||
params.update(
|
||||
{
|
||||
"prompt": "login",
|
||||
"id_token_add_organizations": "true",
|
||||
"codex_cli_simplified_flow": "true",
|
||||
}
|
||||
)
|
||||
|
||||
if template.oauth.use_pkce and code_challenge:
|
||||
params["code_challenge"] = code_challenge
|
||||
params["code_challenge_method"] = "S256"
|
||||
|
||||
authorization_url = f"{template.oauth.authorize_url}?{urlencode(params)}"
|
||||
|
||||
return StartOAuthResponse(
|
||||
authorization_url=authorization_url,
|
||||
redirect_uri=template.oauth.redirect_uri,
|
||||
provider_type=provider_type,
|
||||
instructions=(
|
||||
"1) 打开 authorization_url 完成授权\n"
|
||||
"2) 授权后会跳转到 redirect_uri(localhost)\n"
|
||||
"3) 复制浏览器地址栏完整 URL,调用 complete 接口粘贴 callback_url"
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@router.post("/providers/{provider_id}/complete", response_model=ProviderCompleteOAuthResponse)
|
||||
async def complete_provider_oauth(
|
||||
provider_id: str,
|
||||
payload: ProviderCompleteOAuthRequest,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
) -> ProviderCompleteOAuthResponse:
|
||||
"""完成 Provider OAuth 并创建 key。"""
|
||||
redis = await get_redis_client(require_redis=True)
|
||||
assert redis is not None
|
||||
|
||||
params = _parse_callback_params(payload.callback_url)
|
||||
code = params.get("code")
|
||||
state = params.get("state")
|
||||
if not code or not state:
|
||||
raise InvalidRequestException("callback_url 缺少 code/state")
|
||||
|
||||
state_data = await _consume_state(redis, state)
|
||||
if not state_data or state_data.provider_id != provider_id:
|
||||
raise InvalidRequestException("state 无效或已过期")
|
||||
|
||||
provider = db.query(Provider).filter(Provider.id == provider_id).first()
|
||||
if not provider:
|
||||
raise NotFoundException("Provider 不存在", "provider")
|
||||
provider_type = _require_fixed_provider(provider)
|
||||
|
||||
try:
|
||||
template = FIXED_PROVIDERS.get(ProviderType(provider_type))
|
||||
except Exception:
|
||||
template = None
|
||||
if not template:
|
||||
raise InvalidRequestException("不支持的 provider_type")
|
||||
|
||||
# exchange token
|
||||
token_url = template.oauth.token_url
|
||||
is_json = "anthropic.com" in token_url
|
||||
|
||||
if is_json:
|
||||
body: dict[str, Any] = {
|
||||
"grant_type": "authorization_code",
|
||||
"client_id": template.oauth.client_id,
|
||||
"redirect_uri": template.oauth.redirect_uri,
|
||||
"code": code,
|
||||
"state": state,
|
||||
}
|
||||
if state_data.pkce_verifier:
|
||||
body["code_verifier"] = state_data.pkce_verifier
|
||||
headers = {"Content-Type": "application/json", "Accept": "application/json"}
|
||||
data = None
|
||||
json_body = body
|
||||
else:
|
||||
form: dict[str, str] = {
|
||||
"grant_type": "authorization_code",
|
||||
"client_id": template.oauth.client_id,
|
||||
"redirect_uri": template.oauth.redirect_uri,
|
||||
"code": code,
|
||||
}
|
||||
if template.oauth.client_secret:
|
||||
form["client_secret"] = template.oauth.client_secret
|
||||
if state_data.pkce_verifier:
|
||||
form["code_verifier"] = state_data.pkce_verifier
|
||||
headers = {
|
||||
"Content-Type": "application/x-www-form-urlencoded",
|
||||
"Accept": "application/json",
|
||||
}
|
||||
data = form
|
||||
json_body = None
|
||||
|
||||
proxy_config = getattr(provider, "proxy", None)
|
||||
|
||||
resp = await post_oauth_token(
|
||||
provider_type=provider_type,
|
||||
token_url=token_url,
|
||||
headers=headers,
|
||||
data=data,
|
||||
json_body=json_body,
|
||||
proxy_config=proxy_config,
|
||||
timeout_seconds=30.0,
|
||||
)
|
||||
|
||||
if resp.status_code < 200 or resp.status_code >= 300:
|
||||
raise InvalidRequestException("token exchange 失败")
|
||||
|
||||
token = resp.json()
|
||||
access_token = str(token.get("access_token") or "")
|
||||
refresh_token = str(token.get("refresh_token") or "")
|
||||
expires_in = token.get("expires_in")
|
||||
expires_at: int | None = None
|
||||
try:
|
||||
if expires_in is not None:
|
||||
expires_at = int(time.time()) + int(expires_in)
|
||||
except Exception:
|
||||
expires_at = None
|
||||
|
||||
if not access_token:
|
||||
raise InvalidRequestException("token exchange 返回缺少 access_token")
|
||||
|
||||
# 构建 auth_config
|
||||
auth_config: dict[str, Any] = {
|
||||
"provider_type": provider_type,
|
||||
"token_type": token.get("token_type"),
|
||||
"refresh_token": refresh_token or None,
|
||||
"expires_at": expires_at,
|
||||
"scope": token.get("scope"),
|
||||
"updated_at": int(time.time()),
|
||||
}
|
||||
|
||||
auth_config = await enrich_auth_config(
|
||||
provider_type=provider_type,
|
||||
auth_config=auth_config,
|
||||
token_response=token,
|
||||
access_token=access_token,
|
||||
proxy_config=proxy_config,
|
||||
)
|
||||
|
||||
# 确定账号名称
|
||||
name = (payload.name or "").strip()
|
||||
if not name:
|
||||
name = auth_config.get("email") or f"账号_{int(time.time())}"
|
||||
|
||||
# 从 Provider 的 endpoints 中提取所有 api_format 作为 Key 的支持格式
|
||||
api_formats = [ep.api_format for ep in provider.endpoints if ep.api_format and ep.is_active]
|
||||
|
||||
# 创建 key
|
||||
from src.models.database import ProviderAPIKey as ProviderAPIKeyModel
|
||||
|
||||
new_key = ProviderAPIKeyModel(
|
||||
provider_id=provider_id,
|
||||
name=name,
|
||||
api_key=crypto_service.encrypt(access_token),
|
||||
auth_type="oauth",
|
||||
auth_config=crypto_service.encrypt(json.dumps(auth_config)),
|
||||
api_formats=api_formats,
|
||||
is_active=True,
|
||||
)
|
||||
db.add(new_key)
|
||||
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,
|
||||
expires_at=expires_at,
|
||||
has_refresh_token=bool(refresh_token),
|
||||
email=auth_config.get("email"),
|
||||
)
|
||||
|
||||
@@ -438,12 +438,28 @@ async def test_model(
|
||||
raise HTTPException(status_code=500, detail="Failed to decrypt 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:
|
||||
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")
|
||||
if account_id:
|
||||
extra_headers["chatgpt-account-id"] = 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)
|
||||
|
||||
endpoint_config = {
|
||||
"api_key": api_key_value,
|
||||
"api_key_id": api_key.id, # 添加API Key ID用于用量记录
|
||||
"base_url": endpoint.base_url,
|
||||
"api_format": endpoint.api_format,
|
||||
"extra_headers": get_extra_headers_from_endpoint(endpoint),
|
||||
"extra_headers": extra_headers if extra_headers else None,
|
||||
"timeout": TimeoutDefaults.HTTP_REQUEST,
|
||||
}
|
||||
|
||||
@@ -478,8 +494,7 @@ async def test_model(
|
||||
async with httpx.AsyncClient(
|
||||
timeout=endpoint_config["timeout"], verify=get_ssl_context()
|
||||
) as client:
|
||||
# 非流式测试
|
||||
logger.debug(f"[test-model] 开始非流式测试...")
|
||||
logger.debug("[test-model] 开始端点测试...")
|
||||
|
||||
response = await adapter_class.check_endpoint(
|
||||
client,
|
||||
@@ -497,7 +512,7 @@ async def test_model(
|
||||
)
|
||||
|
||||
# 记录提供商返回信息
|
||||
logger.debug(f"[test-model] 非流式测试结果:")
|
||||
logger.debug("[test-model] 端点测试结果:")
|
||||
logger.debug(f"[test-model] Status Code: {response.get('status_code')}")
|
||||
logger.debug(f"[test-model] Response Headers: {response.get('headers', {})}")
|
||||
response_data = response.get("response", {})
|
||||
|
||||
@@ -68,7 +68,6 @@ from src.models.database import (
|
||||
User,
|
||||
)
|
||||
from src.services.cache.aware_scheduler import ProviderCandidate
|
||||
from src.services.provider.codex import maybe_patch_request_for_codex
|
||||
from src.services.provider.transport import (
|
||||
build_provider_url,
|
||||
get_vertex_ai_effective_format,
|
||||
@@ -720,6 +719,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
key_id=str(key.id),
|
||||
provider_api_format=str(endpoint.api_format) if endpoint.api_format else None,
|
||||
)
|
||||
ctx.provider_type = str(getattr(provider, "provider_type", "") or "")
|
||||
|
||||
# ctx.api_format 是枚举,需要取 value 作为字符串
|
||||
_api_format_str = (
|
||||
@@ -754,13 +754,18 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
else:
|
||||
request_body = dict(original_request_body)
|
||||
|
||||
# 确定目标变体(用于 Codex 等需要特殊处理的上游)
|
||||
provider_type = str(getattr(provider, "provider_type", "") or "").lower()
|
||||
target_variant = provider_type if provider_type == "codex" else None
|
||||
|
||||
# 跨格式:先做请求体转换(失败触发 failover)
|
||||
registry = get_format_converter_registry()
|
||||
if needs_conversion:
|
||||
registry = get_format_converter_registry()
|
||||
request_body = registry.convert_request(
|
||||
request_body,
|
||||
str(client_api_format),
|
||||
str(provider_api_format),
|
||||
target_variant=target_variant,
|
||||
)
|
||||
# 格式转换后,为需要 model 字段的格式设置模型名
|
||||
self._set_model_after_conversion(
|
||||
@@ -779,13 +784,14 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
else:
|
||||
# 同格式:按原逻辑做轻量清理(子类可覆盖以移除不需要的字段)
|
||||
request_body = self.prepare_provider_request_body(request_body)
|
||||
|
||||
# Provider-specific compatibility patches (e.g. Codex requires store=false and instructions).
|
||||
request_body = maybe_patch_request_for_codex(
|
||||
provider_type=str(getattr(provider, "provider_type", "") or ""),
|
||||
provider_api_format=str(provider_api_format),
|
||||
request_body=request_body,
|
||||
)
|
||||
# 同格式时也需要应用 target_variant 转换(如 Codex)
|
||||
if target_variant:
|
||||
request_body = registry.convert_request(
|
||||
request_body,
|
||||
str(provider_api_format),
|
||||
str(provider_api_format),
|
||||
target_variant=target_variant,
|
||||
)
|
||||
|
||||
# 构建请求(上游始终使用 header 认证,不跟随客户端的 query 方式)
|
||||
provider_payload, provider_headers = self._request_builder.build(
|
||||
@@ -1087,13 +1093,18 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
else:
|
||||
request_body = dict(request_body_ref["body"])
|
||||
|
||||
# 确定目标变体(用于 Codex 等需要特殊处理的上游)
|
||||
provider_type = str(getattr(provider, "provider_type", "") or "").lower()
|
||||
target_variant = provider_type if provider_type == "codex" else None
|
||||
|
||||
# 跨格式:先做请求体转换(失败触发 failover)
|
||||
registry = get_format_converter_registry()
|
||||
if needs_conversion:
|
||||
registry = get_format_converter_registry()
|
||||
request_body = registry.convert_request(
|
||||
request_body,
|
||||
client_api_format,
|
||||
provider_api_format,
|
||||
target_variant=target_variant,
|
||||
)
|
||||
# 格式转换后,为需要 model 字段的格式设置模型名
|
||||
self._set_model_after_conversion(
|
||||
@@ -1112,13 +1123,14 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
else:
|
||||
# 同格式:按原逻辑做轻量清理(子类可覆盖以移除不需要的字段)
|
||||
request_body = self.prepare_provider_request_body(request_body)
|
||||
|
||||
# Provider-specific compatibility patches (e.g. Codex requires store=false and instructions).
|
||||
request_body = maybe_patch_request_for_codex(
|
||||
provider_type=str(getattr(provider, "provider_type", "") or ""),
|
||||
provider_api_format=str(provider_api_format),
|
||||
request_body=request_body,
|
||||
)
|
||||
# 同格式时也需要应用 target_variant 转换(如 Codex)
|
||||
if target_variant:
|
||||
request_body = registry.convert_request(
|
||||
request_body,
|
||||
provider_api_format,
|
||||
provider_api_format,
|
||||
target_variant=target_variant,
|
||||
)
|
||||
|
||||
# 构建请求(上游始终使用 header 认证,不跟随客户端的 query 方式)
|
||||
provider_payload, provider_hdrs = self._request_builder.build(
|
||||
|
||||
@@ -638,13 +638,13 @@ class CliAdapterBase(ApiAdapter):
|
||||
url = cls.build_endpoint_url(base_url, request_data, model_name)
|
||||
|
||||
# 合并 CLI 额外头部到 extra_headers
|
||||
cli_extra = cls.get_cli_extra_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)
|
||||
|
||||
# 使用统一的头部构建函数
|
||||
headers = cls.build_headers_with_extra(api_key, merged_extra if merged_extra else None)
|
||||
body = cls.build_request_body(request_data)
|
||||
body = cls.build_request_body(request_data, base_url=base_url)
|
||||
|
||||
# 获取有效的模型名称
|
||||
effective_model_name = model_name or request_data.get("model")
|
||||
@@ -686,17 +686,25 @@ class CliAdapterBase(ApiAdapter):
|
||||
raise NotImplementedError(f"{cls.FORMAT_ID} adapter must implement build_endpoint_url")
|
||||
|
||||
@classmethod
|
||||
def build_request_body(cls, request_data: dict[str, Any] | None = None) -> dict[str, Any]:
|
||||
def build_request_body(
|
||||
cls,
|
||||
request_data: dict[str, Any] | None = None,
|
||||
*,
|
||||
base_url: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""构建测试请求体,使用转换器注册表自动处理格式转换
|
||||
|
||||
Args:
|
||||
request_data: 可选的请求数据,会与默认测试请求合并
|
||||
base_url: API 基础 URL,用于判断特殊端点(如 Codex)
|
||||
|
||||
Returns:
|
||||
转换为目标 API 格式的请求体
|
||||
"""
|
||||
from src.api.handlers.base.request_builder import build_test_request_body
|
||||
|
||||
# 基类不使用 base_url,子类可覆盖以支持特殊端点
|
||||
_ = base_url
|
||||
return build_test_request_body(cls.FORMAT_ID, request_data)
|
||||
|
||||
@classmethod
|
||||
@@ -710,13 +718,16 @@ class CliAdapterBase(ApiAdapter):
|
||||
return None
|
||||
|
||||
@classmethod
|
||||
def get_cli_extra_headers(cls) -> dict[str, str]:
|
||||
def get_cli_extra_headers(cls, *, base_url: str | None = None) -> dict[str, str]:
|
||||
"""
|
||||
获取CLI额外请求头 - 子类可覆盖
|
||||
|
||||
用于 check_endpoint 测试请求时添加额外的头部。
|
||||
默认实现只添加 User-Agent(如果有)。
|
||||
|
||||
Args:
|
||||
base_url: API 基础 URL,子类可据此判断特殊端点(如 Codex)
|
||||
|
||||
Returns:
|
||||
额外请求头字典
|
||||
"""
|
||||
|
||||
@@ -72,7 +72,6 @@ from src.models.database import (
|
||||
User,
|
||||
)
|
||||
from src.services.cache.aware_scheduler import ProviderCandidate
|
||||
from src.services.provider.codex import maybe_patch_request_for_codex
|
||||
from src.services.provider.transport import build_provider_url
|
||||
from src.services.system.config import SystemConfigService
|
||||
from src.utils.sse_parser import SSEEventParser
|
||||
@@ -425,6 +424,8 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
mapped_model: str | None,
|
||||
fallback_model: str,
|
||||
is_stream: bool,
|
||||
*,
|
||||
target_variant: str | None = None,
|
||||
) -> tuple[dict[str, Any], str]:
|
||||
"""
|
||||
跨格式请求转换的公共逻辑
|
||||
@@ -438,6 +439,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
mapped_model: 映射后的模型名
|
||||
fallback_model: 备用模型名(通常是原始请求的 model)
|
||||
is_stream: 是否流式请求
|
||||
target_variant: 目标变体(如 "codex"),用于同格式但有细微差异的上游
|
||||
|
||||
Returns:
|
||||
(转换后的请求体, 用于 URL 的模型名)
|
||||
@@ -447,6 +449,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
request_body,
|
||||
str(client_api_format),
|
||||
str(provider_api_format),
|
||||
target_variant=target_variant,
|
||||
)
|
||||
|
||||
# 先计算 URL 模型(在清理 body 中的 model 字段之前)
|
||||
@@ -697,6 +700,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
# 记录 Provider 信息
|
||||
ctx.provider_name = str(provider.name)
|
||||
ctx.provider_id = str(provider.id)
|
||||
ctx.provider_type = str(getattr(provider, "provider_type", "") or "")
|
||||
ctx.endpoint_id = str(endpoint.id)
|
||||
ctx.key_id = str(key.id)
|
||||
|
||||
@@ -730,6 +734,10 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
)
|
||||
ctx.needs_conversion = needs_conversion
|
||||
|
||||
# 确定目标变体(用于 Codex 等需要特殊处理的上游)
|
||||
provider_type = str(getattr(provider, "provider_type", "") or "").lower()
|
||||
target_variant = provider_type if provider_type == "codex" else None
|
||||
|
||||
# 跨格式:先做请求体转换(失败触发 failover)
|
||||
if needs_conversion and provider_api_format:
|
||||
request_body, url_model = self._convert_request_for_cross_format(
|
||||
@@ -739,6 +747,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
mapped_model,
|
||||
ctx.model,
|
||||
is_stream=True,
|
||||
target_variant=target_variant,
|
||||
)
|
||||
else:
|
||||
# 同格式:按原逻辑做轻量清理(子类可覆盖)
|
||||
@@ -746,13 +755,15 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
url_model = (
|
||||
self.get_model_for_url(request_body, mapped_model) or mapped_model or ctx.model
|
||||
)
|
||||
|
||||
# Provider-specific compatibility patches (e.g. Codex requires store=false and instructions).
|
||||
request_body = maybe_patch_request_for_codex(
|
||||
provider_type=str(getattr(provider, "provider_type", "") or ""),
|
||||
provider_api_format=str(provider_api_format),
|
||||
request_body=request_body,
|
||||
)
|
||||
# 同格式时也需要应用 target_variant 转换(如 Codex)
|
||||
if target_variant:
|
||||
registry = get_format_converter_registry()
|
||||
request_body = registry.convert_request(
|
||||
request_body,
|
||||
provider_api_format,
|
||||
provider_api_format,
|
||||
target_variant=target_variant,
|
||||
)
|
||||
|
||||
# 获取认证信息(处理 Service Account 等异步认证场景)
|
||||
auth_info = await get_provider_auth(endpoint, key)
|
||||
@@ -1922,12 +1933,19 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
try:
|
||||
from src.models.database import ApiKey as ApiKeyModel
|
||||
|
||||
# 采集上游元数据(仅成功请求)
|
||||
if ctx.is_success():
|
||||
self._collect_upstream_metadata(bg_db, ctx)
|
||||
|
||||
user = bg_db.query(User).filter(User.id == ctx.user_id).first()
|
||||
api_key = bg_db.query(ApiKeyModel).filter(ApiKeyModel.id == ctx.api_key_id).first()
|
||||
|
||||
if not user or not api_key:
|
||||
logger.warning(
|
||||
f"[{ctx.request_id}] 无法记录统计: user={user is not None}, api_key={api_key is not None}"
|
||||
"[{}] 无法记录统计: user={} api_key={}",
|
||||
ctx.request_id,
|
||||
user is not None,
|
||||
api_key is not None,
|
||||
)
|
||||
return
|
||||
|
||||
@@ -2153,6 +2171,19 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
except Exception as e:
|
||||
logger.exception("记录流式统计信息时出错")
|
||||
|
||||
@staticmethod
|
||||
def _collect_upstream_metadata(db: Session, ctx: StreamContext) -> None:
|
||||
"""采集上游元数据并更新 ProviderAPIKey.upstream_metadata(带节流)"""
|
||||
from src.services.provider.metadata_collectors import collect_and_save_upstream_metadata
|
||||
|
||||
collect_and_save_upstream_metadata(
|
||||
db,
|
||||
provider_type=ctx.provider_type or "",
|
||||
key_id=ctx.key_id or "",
|
||||
response_headers=ctx.response_headers or {},
|
||||
request_id=ctx.request_id or "",
|
||||
)
|
||||
|
||||
async def _record_stream_failure(
|
||||
self,
|
||||
ctx: StreamContext,
|
||||
@@ -2285,6 +2316,10 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
)
|
||||
needs_conversion = bool(getattr(candidate, "needs_conversion", False))
|
||||
|
||||
# 确定目标变体(用于 Codex 等需要特殊处理的上游)
|
||||
provider_type = str(getattr(provider, "provider_type", "") or "").lower()
|
||||
target_variant = provider_type if provider_type == "codex" else None
|
||||
|
||||
# 跨格式:先做请求体转换(失败触发 failover)
|
||||
if needs_conversion and provider_api_format:
|
||||
request_body, url_model = self._convert_request_for_cross_format(
|
||||
@@ -2294,6 +2329,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
mapped_model,
|
||||
model,
|
||||
is_stream=False,
|
||||
target_variant=target_variant,
|
||||
)
|
||||
else:
|
||||
# 同格式:按原逻辑做轻量清理(子类可覆盖)
|
||||
@@ -2301,13 +2337,15 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
url_model = (
|
||||
self.get_model_for_url(request_body, mapped_model) or mapped_model or model
|
||||
)
|
||||
|
||||
# Provider-specific compatibility patches (e.g. Codex requires store=false and instructions).
|
||||
request_body = maybe_patch_request_for_codex(
|
||||
provider_type=str(getattr(provider, "provider_type", "") or ""),
|
||||
provider_api_format=str(provider_api_format),
|
||||
request_body=request_body,
|
||||
)
|
||||
# 同格式时也需要应用 target_variant 转换(如 Codex)
|
||||
if target_variant:
|
||||
registry = get_format_converter_registry()
|
||||
request_body = registry.convert_request(
|
||||
request_body,
|
||||
provider_api_format,
|
||||
provider_api_format,
|
||||
target_variant=target_variant,
|
||||
)
|
||||
|
||||
# 获取认证信息(处理 Service Account 等异步认证场景)
|
||||
auth_info = await get_provider_auth(endpoint, key)
|
||||
|
||||
@@ -271,7 +271,7 @@ async def _calculate_and_record_usage(
|
||||
cache_read_input_tokens=cache_read_input_tokens,
|
||||
request_type="endpoint_test", # 使用特殊的请求类型标识测试
|
||||
api_format=api_format,
|
||||
is_stream=False,
|
||||
is_stream=request_data.get("stream", False) if request_data else False,
|
||||
response_time_ms=response_time_ms,
|
||||
first_byte_time_ms=response_time_ms,
|
||||
status_code=status_code,
|
||||
@@ -587,58 +587,191 @@ class HttpRequestExecutor:
|
||||
self.timeout = timeout
|
||||
|
||||
async def execute(self, request: EndpointCheckRequest) -> EndpointCheckResult:
|
||||
"""执行HTTP请求"""
|
||||
"""执行HTTP请求(支持流式和非流式响应)"""
|
||||
start_time = time.time()
|
||||
request_id = request.request_id or str(uuid.uuid4())[:8]
|
||||
|
||||
# 检查是否是流式请求
|
||||
is_stream = request.json_body.get("stream", False) if request.json_body else False
|
||||
|
||||
try:
|
||||
# 使用httpx进行异步请求
|
||||
async with httpx.AsyncClient(timeout=self.timeout, verify=get_ssl_context()) as client:
|
||||
response = await client.post(
|
||||
url=request.url, json=request.json_body, headers=request.headers
|
||||
)
|
||||
if is_stream:
|
||||
# 流式请求:读取 SSE 事件直到完成
|
||||
response_data = await self._execute_stream_request(client, request)
|
||||
end_time = time.time()
|
||||
response_time_ms = int((end_time - start_time) * 1000)
|
||||
|
||||
end_time = time.time()
|
||||
response_time_ms = int((end_time - start_time) * 1000)
|
||||
if response_data.get("error"):
|
||||
# 流式请求返回错误
|
||||
return EndpointCheckResult(
|
||||
status_code=response_data.get("status_code", 500),
|
||||
headers=response_data.get("headers", {}),
|
||||
response_time_ms=response_time_ms,
|
||||
request_id=request_id,
|
||||
response_data=None,
|
||||
error_message=response_data.get("error"),
|
||||
)
|
||||
|
||||
# 处理响应
|
||||
if response.status_code == 200:
|
||||
try:
|
||||
response_data = response.json()
|
||||
logger.debug(
|
||||
f"[{request.api_format}] check_endpoint | response | json={_truncate_repr(response_data)}"
|
||||
return EndpointCheckResult(
|
||||
status_code=200,
|
||||
headers=response_data.get("headers", {}),
|
||||
response_time_ms=response_time_ms,
|
||||
request_id=request_id,
|
||||
response_data=response_data.get("final_response"),
|
||||
)
|
||||
else:
|
||||
# 非流式请求:直接读取响应
|
||||
response = await client.post(
|
||||
url=request.url, json=request.json_body, headers=request.headers
|
||||
)
|
||||
except Exception:
|
||||
response_data = None
|
||||
logger.debug(f"[{request.api_format}] check_endpoint | response | invalid json")
|
||||
|
||||
return EndpointCheckResult(
|
||||
status_code=response.status_code,
|
||||
headers=dict(response.headers),
|
||||
response_time_ms=response_time_ms,
|
||||
request_id=request_id,
|
||||
response_data=response_data,
|
||||
)
|
||||
else:
|
||||
# 对于非200状态码,使用错误处理器
|
||||
error_body = response.text[:500] if response.text else "(empty)"
|
||||
logger.debug(
|
||||
f"[{request.api_format}] check_endpoint | response | error={error_body}"
|
||||
)
|
||||
end_time = time.time()
|
||||
response_time_ms = int((end_time - start_time) * 1000)
|
||||
|
||||
# 创建HTTPStatusError让错误处理器处理
|
||||
http_error = httpx.HTTPStatusError(
|
||||
message=f"HTTP {response.status_code}: {error_body}",
|
||||
request=None, # 我们不需要完整的request对象
|
||||
response=response,
|
||||
)
|
||||
if response.status_code == 200:
|
||||
try:
|
||||
response_data = response.json()
|
||||
logger.debug(
|
||||
f"[{request.api_format}] check_endpoint | response | json={_truncate_repr(response_data)}"
|
||||
)
|
||||
except Exception:
|
||||
response_data = None
|
||||
logger.debug(
|
||||
f"[{request.api_format}] check_endpoint | response | invalid json"
|
||||
)
|
||||
|
||||
return await ErrorHandler.handle_error(http_error, request)
|
||||
return EndpointCheckResult(
|
||||
status_code=response.status_code,
|
||||
headers=dict(response.headers),
|
||||
response_time_ms=response_time_ms,
|
||||
request_id=request_id,
|
||||
response_data=response_data,
|
||||
)
|
||||
else:
|
||||
error_body = response.text[:500] if response.text else "(empty)"
|
||||
logger.debug(
|
||||
f"[{request.api_format}] check_endpoint | response | error={error_body}"
|
||||
)
|
||||
http_error = httpx.HTTPStatusError(
|
||||
message=f"HTTP {response.status_code}: {error_body}",
|
||||
request=None,
|
||||
response=response,
|
||||
)
|
||||
return await ErrorHandler.handle_error(http_error, request)
|
||||
|
||||
except Exception as e:
|
||||
# 使用统一错误处理器处理异常
|
||||
return await ErrorHandler.handle_error(e, request)
|
||||
|
||||
async def _execute_stream_request(
|
||||
self, client: httpx.AsyncClient, request: EndpointCheckRequest
|
||||
) -> dict[str, Any]:
|
||||
"""执行流式请求并收集响应"""
|
||||
try:
|
||||
async with client.stream(
|
||||
"POST", request.url, json=request.json_body, headers=request.headers
|
||||
) as response:
|
||||
headers = dict(response.headers)
|
||||
|
||||
if response.status_code != 200:
|
||||
error_body = ""
|
||||
async for chunk in response.aiter_text():
|
||||
error_body += chunk
|
||||
if len(error_body) > 500:
|
||||
break
|
||||
logger.debug(
|
||||
"[{}] check_endpoint | stream error | {}",
|
||||
request.api_format,
|
||||
error_body[:500],
|
||||
)
|
||||
return {
|
||||
"error": f"HTTP {response.status_code}: {error_body[:500]}",
|
||||
"status_code": response.status_code,
|
||||
"headers": headers,
|
||||
}
|
||||
|
||||
# 收集 SSE 事件(兼容多种 API 格式)
|
||||
final_response: dict[str, Any] = {}
|
||||
collected_text = ""
|
||||
|
||||
async for line in response.aiter_lines():
|
||||
if not line or not line.startswith("data:"):
|
||||
continue
|
||||
|
||||
data_str = line[5:].strip()
|
||||
if data_str == "[DONE]":
|
||||
break
|
||||
|
||||
try:
|
||||
event = json.loads(data_str)
|
||||
event_type = event.get("type", "")
|
||||
|
||||
# OpenAI Responses API 事件
|
||||
if event_type == "response.output_text.delta":
|
||||
delta = event.get("delta", "")
|
||||
if isinstance(delta, str):
|
||||
collected_text += delta
|
||||
elif event_type == "response.completed":
|
||||
final_response = event.get("response", {})
|
||||
break
|
||||
|
||||
# OpenAI Chat Completions 格式
|
||||
elif "choices" in event:
|
||||
for choice in event.get("choices", []):
|
||||
delta = choice.get("delta", {})
|
||||
content = delta.get("content")
|
||||
if content:
|
||||
collected_text += content
|
||||
if choice.get("finish_reason"):
|
||||
final_response = event
|
||||
break
|
||||
|
||||
# Claude Messages API 格式
|
||||
elif event_type == "content_block_delta":
|
||||
delta = event.get("delta", {})
|
||||
text = delta.get("text", "")
|
||||
if text:
|
||||
collected_text += text
|
||||
elif event_type == "message_stop":
|
||||
break
|
||||
|
||||
# Gemini SSE 格式
|
||||
elif "candidates" in event:
|
||||
for candidate in event.get("candidates", []):
|
||||
content = candidate.get("content", {})
|
||||
for part in content.get("parts", []):
|
||||
text = part.get("text", "")
|
||||
if text:
|
||||
collected_text += text
|
||||
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
|
||||
# 如果没有收到最终响应事件,构建一个基本响应
|
||||
if not final_response:
|
||||
final_response = {
|
||||
"status": "completed",
|
||||
"output": [
|
||||
{
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": collected_text}],
|
||||
}
|
||||
],
|
||||
}
|
||||
|
||||
logger.debug(
|
||||
"[{}] check_endpoint | stream completed | text_length={}",
|
||||
request.api_format,
|
||||
len(collected_text),
|
||||
)
|
||||
|
||||
return {"final_response": final_response, "headers": headers}
|
||||
|
||||
except Exception as e:
|
||||
logger.warning("[{}] check_endpoint | stream error | {}", request.api_format, e)
|
||||
return {"error": str(e), "status_code": 500, "headers": {}}
|
||||
|
||||
|
||||
class UsageCalculator:
|
||||
"""用量计算器 - 专门负责Token计数和费用计算"""
|
||||
|
||||
@@ -15,12 +15,14 @@ from __future__ import annotations
|
||||
|
||||
import json
|
||||
import time
|
||||
|
||||
import httpx
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import httpx
|
||||
from sqlalchemy.orm import object_session
|
||||
|
||||
from src.clients.redis_client import get_redis_client
|
||||
from src.core.api_format import (
|
||||
UPSTREAM_DROP_HEADERS,
|
||||
HeaderBuilder,
|
||||
@@ -30,9 +32,6 @@ from src.core.api_format import (
|
||||
from src.core.crypto import crypto_service
|
||||
from src.core.logger import logger
|
||||
from src.core.provider_oauth_utils import enrich_auth_config, post_oauth_token
|
||||
from sqlalchemy.orm import object_session
|
||||
|
||||
from src.clients.redis_client import get_redis_client
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.models.database import ProviderAPIKey, ProviderEndpoint
|
||||
@@ -108,6 +107,8 @@ def get_test_request_data(request_data: dict[str, Any] | None = None) -> dict[st
|
||||
def build_test_request_body(
|
||||
format_id: str,
|
||||
request_data: dict[str, Any] | None = None,
|
||||
*,
|
||||
target_variant: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""构建测试请求体,自动处理格式转换
|
||||
|
||||
@@ -116,6 +117,7 @@ def build_test_request_body(
|
||||
Args:
|
||||
format_id: 目标 endpoint signature(如 "claude:chat", "gemini:chat", "openai:cli")
|
||||
request_data: 可选的请求数据,会与默认测试请求合并
|
||||
target_variant: 目标变体(如 "codex"),用于同格式但有细微差异的上游
|
||||
|
||||
Returns:
|
||||
转换为目标 API 格式的请求体
|
||||
@@ -124,21 +126,19 @@ def build_test_request_body(
|
||||
format_conversion_registry,
|
||||
register_default_normalizers,
|
||||
)
|
||||
from src.core.api_format.utils import get_base_format
|
||||
|
||||
register_default_normalizers()
|
||||
|
||||
# 获取测试请求数据(OpenAI 格式)
|
||||
source_data = get_test_request_data(request_data)
|
||||
|
||||
# CLI 格式使用基础格式进行转换(claude:cli -> claude:chat)
|
||||
target_format = get_base_format(format_id) or format_id
|
||||
|
||||
# 使用注册表进行格式转换 (openai:chat -> 目标基础格式)
|
||||
# 直接使用目标格式进行转换,不再转换为基础格式
|
||||
# 这样 openai:cli 会正确转换为 Responses API 格式
|
||||
return format_conversion_registry.convert_request(
|
||||
source_data,
|
||||
make_signature_key("openai", "chat"),
|
||||
target_format,
|
||||
format_id,
|
||||
target_variant=target_variant,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -39,6 +39,7 @@ class StreamContext:
|
||||
# Provider 信息(在请求执行时填充)
|
||||
provider_name: str | None = None
|
||||
provider_id: str | None = None
|
||||
provider_type: str | None = None # Provider 类型(如 codex),用于元数据采集
|
||||
endpoint_id: str | None = None
|
||||
key_id: str | None = None
|
||||
attempt_id: str | None = None
|
||||
|
||||
@@ -96,6 +96,10 @@ class StreamTelemetryRecorder:
|
||||
bg_db = next(db_gen)
|
||||
|
||||
try:
|
||||
# 采集上游元数据(仅成功请求,放在 writer 获取之前以确保执行)
|
||||
if ctx.is_success():
|
||||
self._collect_upstream_metadata(bg_db, ctx)
|
||||
|
||||
writer = await self._get_telemetry_writer(bg_db, ctx, response_time_ms)
|
||||
if writer is None:
|
||||
return
|
||||
@@ -502,6 +506,19 @@ class StreamTelemetryRecorder:
|
||||
error_message=ctx.error_message or f"HTTP {ctx.status_code}",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _collect_upstream_metadata(db: Session, ctx: StreamContext) -> None:
|
||||
"""采集上游元数据并更新 ProviderAPIKey.upstream_metadata(带节流)"""
|
||||
from src.services.provider.metadata_collectors import collect_and_save_upstream_metadata
|
||||
|
||||
collect_and_save_upstream_metadata(
|
||||
db,
|
||||
provider_type=ctx.provider_type or "",
|
||||
key_id=ctx.key_id or "",
|
||||
response_headers=ctx.response_headers or {},
|
||||
request_id=ctx.request_id or "",
|
||||
)
|
||||
|
||||
def _get_status_from_ctx(self, ctx: StreamContext) -> str:
|
||||
"""根据上下文获取状态字符串"""
|
||||
if ctx.is_success():
|
||||
|
||||
@@ -140,9 +140,9 @@ class ClaudeCliAdapter(CliAdapterBase):
|
||||
return config.internal_user_agent_claude_cli
|
||||
|
||||
@classmethod
|
||||
def get_cli_extra_headers(cls) -> dict[str, str]:
|
||||
def get_cli_extra_headers(cls, *, base_url: str | None = None) -> dict[str, str]:
|
||||
"""获取Claude CLI额外请求头,包含 x-app: cli 标识"""
|
||||
headers = super().get_cli_extra_headers()
|
||||
headers = super().get_cli_extra_headers(base_url=base_url)
|
||||
headers["x-app"] = "cli" # 标识 CLI 模式,让上游使用正确的认证方式
|
||||
return headers
|
||||
|
||||
|
||||
@@ -158,9 +158,9 @@ class GeminiCliAdapter(CliAdapterBase):
|
||||
return config.internal_user_agent_gemini_cli
|
||||
|
||||
@classmethod
|
||||
def get_cli_extra_headers(cls) -> dict[str, str]:
|
||||
def get_cli_extra_headers(cls, *, base_url: str | None = None) -> dict[str, str]:
|
||||
"""获取Gemini CLI额外请求头,包含 x-app: cli 标识"""
|
||||
headers = super().get_cli_extra_headers()
|
||||
headers = super().get_cli_extra_headers(base_url=base_url)
|
||||
headers["x-app"] = "cli" # 标识 CLI 模式,让上游使用正确的 adapter
|
||||
return headers
|
||||
|
||||
|
||||
@@ -6,6 +6,7 @@ OpenAI CLI Adapter - 基于通用 CLI Adapter 基类的简化实现
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
@@ -67,20 +68,74 @@ class OpenAICliAdapter(CliAdapterBase):
|
||||
def build_endpoint_url(
|
||||
cls, base_url: str, request_data: dict[str, Any], model_name: str | None = None
|
||||
) -> str:
|
||||
"""构建OpenAI CLI API端点URL"""
|
||||
"""构建OpenAI CLI API端点URL(使用 Responses API)
|
||||
|
||||
对于 Codex OAuth 端点(如 chatgpt.com/backend-api/codex),直接追加 /responses;
|
||||
对于标准 OpenAI API,使用 /v1/responses。
|
||||
"""
|
||||
base_url = base_url.rstrip("/")
|
||||
# Codex OAuth 端点:chatgpt.com/backend-api/codex -> /responses
|
||||
if cls._is_codex_url(base_url):
|
||||
return f"{base_url}/responses"
|
||||
# 标准 OpenAI API
|
||||
if base_url.endswith("/v1"):
|
||||
return f"{base_url}/chat/completions"
|
||||
return f"{base_url}/responses"
|
||||
else:
|
||||
return f"{base_url}/v1/chat/completions"
|
||||
return f"{base_url}/v1/responses"
|
||||
|
||||
@classmethod
|
||||
def _is_codex_url(cls, base_url: str) -> bool:
|
||||
"""判断是否是 Codex OAuth 端点"""
|
||||
return "/backend-api/codex" in base_url or base_url.endswith("/codex")
|
||||
|
||||
# build_request_body 使用基类实现
|
||||
# OPENAI -> OPENAI_CLI 无转换器,会直接透传原始请求
|
||||
# OpenAI CLI normalizer 会自动添加 instructions 字段
|
||||
|
||||
@classmethod
|
||||
def build_request_body(
|
||||
cls,
|
||||
request_data: dict[str, Any] | None = None,
|
||||
*,
|
||||
base_url: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""构建测试请求体(Codex 端点需要强制 stream=true 等特性)"""
|
||||
from src.api.handlers.base.request_builder import build_test_request_body
|
||||
|
||||
target_variant = "codex" if base_url and cls._is_codex_url(base_url) else None
|
||||
return build_test_request_body(
|
||||
cls.FORMAT_ID,
|
||||
request_data,
|
||||
target_variant=target_variant,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def get_cli_user_agent(cls) -> str | None:
|
||||
"""获取OpenAI CLI User-Agent"""
|
||||
return config.internal_user_agent_openai_cli
|
||||
|
||||
@classmethod
|
||||
def get_cli_extra_headers(cls, *, base_url: str | None = None) -> dict[str, str]:
|
||||
"""
|
||||
获取额外请求头
|
||||
|
||||
对于 Codex OAuth 端点,添加特定头部(缺少可能导致 Cloudflare 拦截)。
|
||||
对于标准 OpenAI API 端点,仅添加 User-Agent。
|
||||
"""
|
||||
headers: dict[str, str] = {}
|
||||
|
||||
# User-Agent
|
||||
cli_user_agent = cls.get_cli_user_agent()
|
||||
if cli_user_agent:
|
||||
headers["User-Agent"] = cli_user_agent
|
||||
|
||||
# 仅 Codex 端点添加特定头部
|
||||
if base_url and cls._is_codex_url(base_url):
|
||||
headers["x-oai-web-search-eligible"] = "true"
|
||||
headers["session_id"] = str(uuid.uuid4())
|
||||
headers["accept"] = "text/event-stream"
|
||||
headers["originator"] = "codex_cli_rs"
|
||||
|
||||
return headers
|
||||
|
||||
|
||||
__all__ = ["OpenAICliAdapter"]
|
||||
|
||||
@@ -28,8 +28,18 @@ class FormatNormalizer(ABC):
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def request_from_internal(self, internal: InternalRequest) -> dict[str, Any]:
|
||||
"""将内部表示转换为格式特定请求"""
|
||||
def request_from_internal(
|
||||
self,
|
||||
internal: InternalRequest,
|
||||
*,
|
||||
target_variant: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""将内部表示转换为格式特定请求
|
||||
|
||||
Args:
|
||||
internal: 内部请求表示
|
||||
target_variant: 目标变体(如 "codex"),用于同格式但有细微差异的上游
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
# ============ 响应转换 ============
|
||||
|
||||
@@ -154,7 +154,12 @@ class ClaudeNormalizer(FormatNormalizer):
|
||||
|
||||
return internal
|
||||
|
||||
def request_from_internal(self, internal: InternalRequest) -> dict[str, Any]:
|
||||
def request_from_internal(
|
||||
self,
|
||||
internal: InternalRequest,
|
||||
*,
|
||||
target_variant: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
system_text = internal.system or self._join_instructions(internal.instructions)
|
||||
|
||||
# Claude Messages API: messages[] 仅允许 user/assistant,且需要交替;这里做最小修复
|
||||
|
||||
@@ -189,7 +189,12 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
|
||||
return internal
|
||||
|
||||
def request_from_internal(self, internal: InternalRequest) -> dict[str, Any]:
|
||||
def request_from_internal(
|
||||
self,
|
||||
internal: InternalRequest,
|
||||
*,
|
||||
target_variant: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
system_text = internal.system or self._join_instructions(internal.instructions)
|
||||
|
||||
# tools/tool_choice
|
||||
|
||||
@@ -200,7 +200,12 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
|
||||
return internal
|
||||
|
||||
def request_from_internal(self, internal: InternalRequest) -> dict[str, Any]:
|
||||
def request_from_internal(
|
||||
self,
|
||||
internal: InternalRequest,
|
||||
*,
|
||||
target_variant: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
out_messages: list[dict[str, Any]] = []
|
||||
|
||||
if internal.instructions:
|
||||
|
||||
@@ -114,38 +114,58 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
|
||||
return internal
|
||||
|
||||
def request_from_internal(self, internal: InternalRequest) -> dict[str, Any]:
|
||||
# Codex 需要的 include 项
|
||||
_CODEX_REQUIRED_INCLUDE = "reasoning.encrypted_content"
|
||||
|
||||
def request_from_internal(
|
||||
self,
|
||||
internal: InternalRequest,
|
||||
*,
|
||||
target_variant: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
is_codex = str(target_variant or "").lower() == "codex"
|
||||
|
||||
result: dict[str, Any] = {
|
||||
"model": internal.model,
|
||||
"input": self._internal_messages_to_input(internal.messages),
|
||||
"input": self._internal_messages_to_input(
|
||||
internal.messages, system_to_developer=is_codex
|
||||
),
|
||||
}
|
||||
|
||||
instructions_text = self._join_instructions(internal)
|
||||
if instructions_text:
|
||||
result["instructions"] = instructions_text
|
||||
# 合并 instructions,如果没有则使用 system
|
||||
instructions_text = (
|
||||
self._join_instructions(internal.instructions)
|
||||
if internal.instructions
|
||||
else internal.system
|
||||
)
|
||||
# Responses API 兼容 instructions 字段,Codex 强制要求
|
||||
# 统一添加该字段以确保兼容性
|
||||
result["instructions"] = instructions_text or ""
|
||||
|
||||
# max_output_tokens/temperature/top_p: Codex 不支持,标准 API 可选
|
||||
if not is_codex:
|
||||
if internal.max_tokens is not None:
|
||||
# Responses API 使用 max_output_tokens
|
||||
result["max_output_tokens"] = internal.max_tokens
|
||||
if internal.temperature is not None:
|
||||
result["temperature"] = internal.temperature
|
||||
if internal.top_p is not None:
|
||||
result["top_p"] = internal.top_p
|
||||
|
||||
if internal.max_tokens is not None:
|
||||
# Responses API 使用 max_output_tokens;兼容层仍可能接受 max_tokens
|
||||
result["max_output_tokens"] = internal.max_tokens
|
||||
if internal.temperature is not None:
|
||||
result["temperature"] = internal.temperature
|
||||
if internal.top_p is not None:
|
||||
result["top_p"] = internal.top_p
|
||||
if internal.stop_sequences:
|
||||
result["stop"] = list(internal.stop_sequences)
|
||||
if internal.stream:
|
||||
result["stream"] = True
|
||||
# Codex 强制要求 stream=true;其他情况尊重客户端请求
|
||||
result["stream"] = True if is_codex else bool(internal.stream)
|
||||
|
||||
if internal.tools:
|
||||
# Responses API 使用扁平结构: {type, name, description, parameters}
|
||||
# 而非 Chat Completions 的嵌套结构: {type, function: {name, ...}}
|
||||
result["tools"] = [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": t.name,
|
||||
"description": t.description,
|
||||
"parameters": t.parameters or {},
|
||||
**(t.extra.get("openai_function") or {}),
|
||||
},
|
||||
"name": t.name,
|
||||
"description": t.description or "",
|
||||
"parameters": t.parameters or {},
|
||||
**(t.extra.get("openai_tool") or {}),
|
||||
}
|
||||
for t in internal.tools
|
||||
@@ -154,6 +174,48 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
if internal.tool_choice:
|
||||
result["tool_choice"] = self._tool_choice_to_openai(internal.tool_choice)
|
||||
|
||||
# 还原 OpenAI Responses API 的其他字段(黑名单:已单独处理的字段不还原)
|
||||
openai_cli_extra = internal.extra.get("openai_cli", {})
|
||||
handled_keys = {
|
||||
"model",
|
||||
"input",
|
||||
"instructions",
|
||||
"max_output_tokens",
|
||||
"max_tokens",
|
||||
"temperature",
|
||||
"top_p",
|
||||
"stop",
|
||||
"stream",
|
||||
"tools",
|
||||
"tool_choice",
|
||||
}
|
||||
for key, value in openai_cli_extra.items():
|
||||
if key not in handled_keys and key not in result:
|
||||
result[key] = value
|
||||
|
||||
# 统一设置 store=false(Codex 强制要求,标准 API 兼容)
|
||||
if "store" not in result:
|
||||
result["store"] = False
|
||||
|
||||
# Codex 特定设置(覆盖/删除不支持的字段)
|
||||
if is_codex:
|
||||
result["parallel_tool_calls"] = True
|
||||
# 添加 reasoning.encrypted_content 到 include
|
||||
include = result.get("include", [])
|
||||
if not isinstance(include, list):
|
||||
include = []
|
||||
if self._CODEX_REQUIRED_INCLUDE not in include:
|
||||
include.append(self._CODEX_REQUIRED_INCLUDE)
|
||||
result["include"] = include
|
||||
# 删除 Codex 不支持的字段
|
||||
for key in (
|
||||
"previous_response_id",
|
||||
"prompt_cache_key",
|
||||
"service_tier",
|
||||
"max_completion_tokens",
|
||||
):
|
||||
result.pop(key, None)
|
||||
|
||||
return result
|
||||
|
||||
# =========================
|
||||
@@ -922,7 +984,12 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
blocks.append(UnknownBlock(raw_type=ptype or "unknown", payload=part))
|
||||
return blocks
|
||||
|
||||
def _internal_messages_to_input(self, messages: list[InternalMessage]) -> list[dict[str, Any]]:
|
||||
def _internal_messages_to_input(
|
||||
self,
|
||||
messages: list[InternalMessage],
|
||||
*,
|
||||
system_to_developer: bool = False,
|
||||
) -> list[dict[str, Any]]:
|
||||
out: list[dict[str, Any]] = []
|
||||
for msg in messages:
|
||||
# ToolUseBlock -> function_call
|
||||
@@ -979,6 +1046,9 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
|
||||
# 普通 message(TextBlock)
|
||||
role = self._role_to_openai(msg.role)
|
||||
# Codex 不接受 system 角色,需要转换为 developer
|
||||
if system_to_developer and role == "system":
|
||||
role = "developer"
|
||||
content_items: list[dict[str, Any]] = []
|
||||
has_text = False
|
||||
|
||||
@@ -990,7 +1060,9 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
if isinstance(block, UnknownBlock):
|
||||
continue # 跳过其他未知块
|
||||
if isinstance(block, TextBlock) and block.text:
|
||||
content_items.append({"type": "input_text", "text": block.text})
|
||||
# assistant 角色使用 output_text,其他角色使用 input_text
|
||||
text_type = "output_text" if role == "assistant" else "input_text"
|
||||
content_items.append({"type": text_type, "text": block.text})
|
||||
has_text = True
|
||||
|
||||
if has_text:
|
||||
@@ -1083,7 +1155,8 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
if tool_choice.type == ToolChoiceType.REQUIRED:
|
||||
return "required"
|
||||
if tool_choice.type == ToolChoiceType.TOOL:
|
||||
return {"type": "function", "function": {"name": tool_choice.tool_name or ""}}
|
||||
# Responses API 使用扁平结构: {type, name}
|
||||
return {"type": "function", "name": tool_choice.tool_name or ""}
|
||||
return "auto"
|
||||
|
||||
def _role_from_value(self, role: Any) -> Role:
|
||||
@@ -1148,14 +1221,11 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
return {}
|
||||
return {k: v for k, v in payload.items() if k not in keep_keys}
|
||||
|
||||
def _join_instructions(self, internal: InternalRequest) -> str:
|
||||
if internal.instructions:
|
||||
parts: list[str] = []
|
||||
for seg in internal.instructions:
|
||||
if seg.text:
|
||||
parts.append(seg.text)
|
||||
return "\n\n".join(parts)
|
||||
return internal.system or ""
|
||||
def _join_instructions(self, instructions: list[InstructionSegment]) -> str | None:
|
||||
"""合并 instructions 为单一字符串,与其他 normalizer 保持一致"""
|
||||
parts = [seg.text for seg in instructions if seg.text]
|
||||
joined = "\n\n".join(parts)
|
||||
return joined or None
|
||||
|
||||
def _error_type_from_value(self, value: str) -> ErrorType:
|
||||
for t in ErrorType:
|
||||
|
||||
@@ -67,8 +67,10 @@ class FormatConversionRegistry:
|
||||
request: dict[str, Any],
|
||||
source_format: str,
|
||||
target_format: str,
|
||||
*,
|
||||
target_variant: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
if str(source_format).upper() == str(target_format).upper():
|
||||
if str(source_format).upper() == str(target_format).upper() and not target_variant:
|
||||
return request
|
||||
|
||||
src = self._require_normalizer(source_format)
|
||||
@@ -79,7 +81,7 @@ class FormatConversionRegistry:
|
||||
):
|
||||
try:
|
||||
internal = src.request_to_internal(request)
|
||||
return tgt.request_from_internal(internal)
|
||||
return tgt.request_from_internal(internal, target_variant=target_variant)
|
||||
except Exception as e:
|
||||
raise FormatConversionError(source_format, target_format, str(e)) from e
|
||||
|
||||
|
||||
@@ -10,7 +10,6 @@ import jwt
|
||||
from src.clients.http_client import HTTPClientPool, build_proxy_url
|
||||
from src.core.logger import logger
|
||||
|
||||
|
||||
_ANTHROPIC_TOKEN_URL = "https://console.anthropic.com/v1/oauth/token"
|
||||
_GOOGLE_USERINFO_URL = "https://www.googleapis.com/oauth2/v1/userinfo?alt=json"
|
||||
|
||||
@@ -206,18 +205,20 @@ async def post_oauth_token(
|
||||
)
|
||||
|
||||
|
||||
def parse_codex_id_token(id_token: str | None) -> tuple[str | None, str | None]:
|
||||
def parse_codex_id_token(id_token: str | None) -> dict[str, Any]:
|
||||
"""Parse Codex id_token WITHOUT signature verification.
|
||||
|
||||
Extract:
|
||||
Extract from claim `https://api.openai.com/auth`:
|
||||
- email: claim `email`
|
||||
- account_id: claim `https://api.openai.com/auth`.`chatgpt_account_id`
|
||||
- account_id: `chatgpt_account_id`
|
||||
- plan_type: `chatgpt_plan_type` (e.g. "plus", "free", "team", "enterprise")
|
||||
- user_id: `chatgpt_user_id`
|
||||
|
||||
Return (email, account_id). On any failure returns (None, None).
|
||||
Return dict with extracted fields. On any failure returns empty dict.
|
||||
"""
|
||||
|
||||
if not id_token:
|
||||
return (None, None)
|
||||
return {}
|
||||
try:
|
||||
claims = jwt.decode(
|
||||
id_token,
|
||||
@@ -226,17 +227,29 @@ def parse_codex_id_token(id_token: str | None) -> tuple[str | None, str | None]:
|
||||
"verify_aud": False,
|
||||
},
|
||||
)
|
||||
result: dict[str, Any] = {}
|
||||
|
||||
email = claims.get("email")
|
||||
if isinstance(email, str) and email:
|
||||
result["email"] = email
|
||||
|
||||
auth_info = claims.get("https://api.openai.com/auth") or {}
|
||||
account_id = None
|
||||
if isinstance(auth_info, dict):
|
||||
account_id = auth_info.get("chatgpt_account_id")
|
||||
return (
|
||||
str(email) if isinstance(email, str) and email else None,
|
||||
str(account_id) if isinstance(account_id, str) and account_id else None,
|
||||
)
|
||||
if isinstance(account_id, str) and account_id:
|
||||
result["account_id"] = account_id
|
||||
|
||||
plan_type = auth_info.get("chatgpt_plan_type")
|
||||
if isinstance(plan_type, str) and plan_type:
|
||||
result["plan_type"] = plan_type
|
||||
|
||||
user_id = auth_info.get("chatgpt_user_id")
|
||||
if isinstance(user_id, str) and user_id:
|
||||
result["user_id"] = user_id
|
||||
|
||||
return result
|
||||
except Exception:
|
||||
return (None, None)
|
||||
return {}
|
||||
|
||||
|
||||
async def fetch_google_email(
|
||||
@@ -306,11 +319,22 @@ async def enrich_auth_config(
|
||||
# Codex
|
||||
if provider_type == "codex":
|
||||
id_token = token_response.get("id_token")
|
||||
email, account_id = parse_codex_id_token(str(id_token) if id_token else None)
|
||||
if email:
|
||||
auth_config["email"] = email
|
||||
if account_id:
|
||||
auth_config["account_id"] = account_id
|
||||
logger.debug(
|
||||
"Codex enrich_auth_config: id_token_present={} token_keys={}",
|
||||
bool(id_token),
|
||||
list(token_response.keys()),
|
||||
)
|
||||
codex_info = parse_codex_id_token(str(id_token) if id_token else None)
|
||||
if codex_info:
|
||||
logger.debug("Codex parsed id_token fields: {}", list(codex_info.keys()))
|
||||
if codex_info.get("email"):
|
||||
auth_config["email"] = codex_info["email"]
|
||||
if codex_info.get("account_id"):
|
||||
auth_config["account_id"] = codex_info["account_id"]
|
||||
if codex_info.get("plan_type"):
|
||||
auth_config["plan_type"] = codex_info["plan_type"]
|
||||
if codex_info.get("user_id"):
|
||||
auth_config["user_id"] = codex_info["user_id"]
|
||||
return auth_config
|
||||
|
||||
# Gemini family (gemini_cli / antigravity)
|
||||
|
||||
@@ -1391,6 +1391,13 @@ class ProviderAPIKey(Base):
|
||||
model_include_patterns = Column(JSON, nullable=True) # 包含规则列表,空表示不过滤(包含所有)
|
||||
model_exclude_patterns = Column(JSON, nullable=True) # 排除规则列表,空表示不排除
|
||||
|
||||
# 上游元数据(由响应头解析器采集,如 Codex 额度信息)
|
||||
upstream_metadata = Column(JSON, nullable=True, default=dict)
|
||||
|
||||
# OAuth 失效状态(账号被封、授权撤销、刷新失败等)
|
||||
oauth_invalid_at = Column(DateTime(timezone=True), nullable=True) # 失效时间
|
||||
oauth_invalid_reason = Column(String(255), nullable=True) # 失效原因
|
||||
|
||||
# 时间戳
|
||||
created_at = Column(
|
||||
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
|
||||
|
||||
@@ -154,9 +154,7 @@ class ProviderEndpointResponse(BaseModel):
|
||||
header_rules: list[HeaderRule] | None = Field(default=None, description="请求头规则列表")
|
||||
|
||||
# 请求体配置
|
||||
body_rules: list[BodyRule] | None = Field(
|
||||
default=None, description="请求体规则列表"
|
||||
)
|
||||
body_rules: list[BodyRule] | None = Field(default=None, description="请求体规则列表")
|
||||
|
||||
max_retries: int
|
||||
|
||||
@@ -514,7 +512,18 @@ class EndpointAPIKeyResponse(BaseModel):
|
||||
capabilities: dict[str, bool] | None = Field(default=None, description="Key 能力标签")
|
||||
|
||||
# OAuth 相关
|
||||
oauth_expires_at: int | None = Field(default=None, description="OAuth Token 过期时间(Unix 时间戳)")
|
||||
oauth_expires_at: int | None = Field(
|
||||
default=None, description="OAuth Token 过期时间(Unix 时间戳)"
|
||||
)
|
||||
oauth_email: str | None = Field(default=None, description="OAuth 账号邮箱")
|
||||
oauth_plan_type: str | None = Field(
|
||||
default=None, description="OAuth 账号套餐类型(如 free/plus/team/enterprise)"
|
||||
)
|
||||
oauth_account_id: str | None = Field(default=None, description="OAuth 账号 ID")
|
||||
oauth_invalid_at: int | None = Field(
|
||||
default=None, description="OAuth Token 失效时间(Unix 时间戳),如账号被封、授权撤销等"
|
||||
)
|
||||
oauth_invalid_reason: str | None = Field(default=None, description="OAuth Token 失效原因")
|
||||
|
||||
# 缓存与熔断配置
|
||||
cache_ttl_minutes: int = Field(default=5, description="缓存 TTL(分钟),0=禁用")
|
||||
@@ -576,6 +585,11 @@ class EndpointAPIKeyResponse(BaseModel):
|
||||
model_include_patterns: list[str] | None = Field(None, description="模型包含规则")
|
||||
model_exclude_patterns: list[str] | None = Field(None, description="模型排除规则")
|
||||
|
||||
# 上游元数据(由响应头采集,如 Codex 额度信息)
|
||||
upstream_metadata: dict[str, Any] | None = Field(
|
||||
None, description="上游元数据(如 Codex 额度信息)"
|
||||
)
|
||||
|
||||
# 时间戳
|
||||
last_used_at: datetime | None = None
|
||||
created_at: datetime
|
||||
@@ -701,7 +715,9 @@ class ProviderWithEndpointsSummary(BaseModel):
|
||||
# Provider 基本信息
|
||||
id: str
|
||||
name: str
|
||||
provider_type: str | None = Field(default=None, description="Provider 类型(custom/claude_code/codex/gemini_cli/antigravity)")
|
||||
provider_type: str | None = Field(
|
||||
default=None, description="Provider 类型(custom/claude_code/codex/gemini_cli/antigravity)"
|
||||
)
|
||||
description: str | None = None
|
||||
website: str | None = None
|
||||
provider_priority: int = Field(default=100, description="提供商优先级(数字越小越优先)")
|
||||
|
||||
@@ -1,106 +0,0 @@
|
||||
"""
|
||||
Codex upstream request compatibility helpers.
|
||||
|
||||
The Codex upstream (https://chatgpt.com/backend-api/codex) is largely compatible with the
|
||||
OpenAI Responses (/responses, aka "openai:cli") schema, but enforces some extra constraints.
|
||||
|
||||
CLIProxyAPI's reference implementation applies a small set of mutations before forwarding.
|
||||
We replicate the same mutations here to keep Aether's routing compatible when
|
||||
Provider.provider_type == "codex".
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
_CODEX_REQUIRED_INCLUDE_ITEM = "reasoning.encrypted_content"
|
||||
|
||||
|
||||
def patch_openai_cli_request_for_codex(request_body: dict[str, Any]) -> dict[str, Any]:
|
||||
"""
|
||||
Mutate an OpenAI Responses (openai:cli) request into a Codex-compatible payload.
|
||||
|
||||
Notes (based on CLIProxyAPI translators):
|
||||
- `store` must be explicitly set to false.
|
||||
- `instructions` field must exist (Codex rejects missing instructions).
|
||||
- Codex rejects several generation params, so strip them.
|
||||
- Codex does not accept `system` role inside the `input` array.
|
||||
- Enable `parallel_tool_calls` and request encrypted reasoning content.
|
||||
"""
|
||||
if not isinstance(request_body, dict):
|
||||
return request_body
|
||||
|
||||
result: dict[str, Any] = dict(request_body)
|
||||
|
||||
# Required by Codex: explicitly disable storing.
|
||||
result["store"] = False
|
||||
|
||||
# Required by Codex: ensure instructions exists (can be empty).
|
||||
instructions = result.get("instructions")
|
||||
if instructions is None:
|
||||
result["instructions"] = ""
|
||||
elif not isinstance(instructions, str):
|
||||
result["instructions"] = str(instructions)
|
||||
|
||||
# Codex defaults/tooling expectations
|
||||
result["parallel_tool_calls"] = True
|
||||
|
||||
include_value = result.get("include")
|
||||
include: list[str] = []
|
||||
if isinstance(include_value, list):
|
||||
include = [v for v in include_value if isinstance(v, str) and v]
|
||||
if _CODEX_REQUIRED_INCLUDE_ITEM not in include:
|
||||
include.append(_CODEX_REQUIRED_INCLUDE_ITEM)
|
||||
result["include"] = include
|
||||
|
||||
# Codex Responses rejects token limit fields and some sampling params.
|
||||
for key in (
|
||||
"max_output_tokens",
|
||||
"max_completion_tokens",
|
||||
"max_tokens",
|
||||
"temperature",
|
||||
"top_p",
|
||||
"service_tier",
|
||||
):
|
||||
result.pop(key, None)
|
||||
|
||||
# Convert role "system" to "developer" in input array to comply with Codex API requirements.
|
||||
input_value = result.get("input")
|
||||
if isinstance(input_value, list):
|
||||
patched_input: list[Any] = []
|
||||
for item in input_value:
|
||||
if (
|
||||
isinstance(item, dict)
|
||||
and item.get("type") == "message"
|
||||
and item.get("role") == "system"
|
||||
):
|
||||
item = dict(item)
|
||||
item["role"] = "developer"
|
||||
patched_input.append(item)
|
||||
result["input"] = patched_input
|
||||
|
||||
return result
|
||||
|
||||
|
||||
def maybe_patch_request_for_codex(
|
||||
*,
|
||||
provider_type: str | None,
|
||||
provider_api_format: str | None,
|
||||
request_body: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Apply Codex compatibility patches only when the selected upstream is Codex and the
|
||||
endpoint uses the OpenAI Responses schema ("openai:cli").
|
||||
"""
|
||||
if str(provider_type or "").strip().lower() != "codex":
|
||||
return request_body
|
||||
if str(provider_api_format or "").strip().lower() != "openai:cli":
|
||||
return request_body
|
||||
return patch_openai_cli_request_for_codex(request_body)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"maybe_patch_request_for_codex",
|
||||
"patch_openai_cli_request_for_codex",
|
||||
]
|
||||
|
||||
150
src/services/provider/metadata_collectors/__init__.py
Normal file
150
src/services/provider/metadata_collectors/__init__.py
Normal file
@@ -0,0 +1,150 @@
|
||||
"""
|
||||
上游元数据采集器(MetadataCollector)
|
||||
|
||||
可扩展注册表模式:
|
||||
- 每个 Provider 类型可注册一个 MetadataCollector
|
||||
- 从响应头解析有价值的元数据(额度、限流等)
|
||||
- 解析结果存入 ProviderAPIKey.upstream_metadata
|
||||
|
||||
扩展方式:
|
||||
1. 创建新文件实现 MetadataCollector
|
||||
2. 在本文件底部注册
|
||||
"""
|
||||
|
||||
import time
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any, ClassVar
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.core.logger import logger
|
||||
|
||||
# 节流:每个 key_id 至少间隔 _THROTTLE_SECONDS 秒才写入一次
|
||||
_THROTTLE_SECONDS = 30
|
||||
_last_write_ts: dict[str, float] = {}
|
||||
|
||||
|
||||
class MetadataCollector(ABC):
|
||||
"""元数据采集器基类"""
|
||||
|
||||
# 支持的 provider_type 列表(小写)
|
||||
PROVIDER_TYPES: ClassVar[list[str]] = []
|
||||
|
||||
@abstractmethod
|
||||
def parse_headers(self, headers: dict[str, str]) -> dict[str, Any] | None:
|
||||
"""解析响应头,返回结构化元数据。返回 None 表示无可用数据。"""
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class MetadataCollectorRegistry:
|
||||
"""元数据采集器注册表"""
|
||||
|
||||
_collectors: ClassVar[list[MetadataCollector]] = []
|
||||
_type_index: ClassVar[dict[str, MetadataCollector]] = {}
|
||||
|
||||
@classmethod
|
||||
def register(cls, collector: MetadataCollector) -> None:
|
||||
cls._collectors.append(collector)
|
||||
for pt in collector.PROVIDER_TYPES:
|
||||
cls._type_index[pt.lower()] = collector
|
||||
logger.info(
|
||||
"[MetadataCollectorRegistry] 注册: {} -> {}",
|
||||
collector.__class__.__name__,
|
||||
collector.PROVIDER_TYPES,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def collect(cls, provider_type: str, headers: dict[str, str]) -> dict[str, Any] | None:
|
||||
"""根据 provider_type 查找采集器并解析响应头"""
|
||||
collector = cls._type_index.get(provider_type.lower())
|
||||
if collector is None:
|
||||
return None
|
||||
try:
|
||||
return collector.parse_headers(headers)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"[MetadataCollectorRegistry] {} 解析失败", collector.__class__.__name__
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
_initialized = False
|
||||
|
||||
|
||||
def _ensure_collectors_registered() -> None:
|
||||
"""惰性注册所有采集器(首次调用时执行,避免循环导入)"""
|
||||
global _initialized
|
||||
if _initialized:
|
||||
return
|
||||
_initialized = True
|
||||
|
||||
# 延迟导入,避免模块加载时的循环依赖
|
||||
from src.services.provider.metadata_collectors.codex import CodexMetadataCollector
|
||||
|
||||
MetadataCollectorRegistry.register(CodexMetadataCollector())
|
||||
|
||||
|
||||
def collect_and_save_upstream_metadata(
|
||||
db: Session,
|
||||
*,
|
||||
provider_type: str,
|
||||
key_id: str,
|
||||
response_headers: dict[str, str],
|
||||
request_id: str,
|
||||
) -> None:
|
||||
"""采集上游元数据并更新 ProviderAPIKey.upstream_metadata(带节流)
|
||||
|
||||
每个 key_id 至少间隔 _THROTTLE_SECONDS 秒才执行一次数据库写入,
|
||||
避免高并发时频繁更新同一行。
|
||||
|
||||
Args:
|
||||
db: 数据库 Session
|
||||
provider_type: Provider 类型(如 "codex")
|
||||
key_id: ProviderAPIKey.id
|
||||
response_headers: 上游响应头
|
||||
request_id: 请求 ID(用于日志)
|
||||
"""
|
||||
if not provider_type or not key_id or not response_headers:
|
||||
return
|
||||
|
||||
# 确保采集器已注册
|
||||
_ensure_collectors_registered()
|
||||
|
||||
# 节流检查
|
||||
now = time.monotonic()
|
||||
last_ts = _last_write_ts.get(key_id, 0.0)
|
||||
if now - last_ts < _THROTTLE_SECONDS:
|
||||
return
|
||||
|
||||
try:
|
||||
metadata = MetadataCollectorRegistry.collect(provider_type, response_headers)
|
||||
if metadata is None:
|
||||
return
|
||||
|
||||
from src.models.database import ProviderAPIKey
|
||||
|
||||
key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == key_id).first()
|
||||
if key is None:
|
||||
return
|
||||
|
||||
key.upstream_metadata = metadata
|
||||
db.commit()
|
||||
_last_write_ts[key_id] = now
|
||||
logger.debug(
|
||||
"[{}] 已更新 ProviderAPIKey({}) upstream_metadata",
|
||||
request_id,
|
||||
key_id,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("[{}] 采集上游元数据失败", request_id)
|
||||
try:
|
||||
db.rollback()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
__all__ = [
|
||||
"MetadataCollector",
|
||||
"MetadataCollectorRegistry",
|
||||
"collect_and_save_upstream_metadata",
|
||||
]
|
||||
95
src/services/provider/metadata_collectors/codex.py
Normal file
95
src/services/provider/metadata_collectors/codex.py
Normal file
@@ -0,0 +1,95 @@
|
||||
"""
|
||||
Codex Provider 元数据采集器
|
||||
|
||||
从响应头解析 Codex 额度/限流信息:
|
||||
- x-codex-plan-type: 套餐类型
|
||||
- x-codex-primary-*: 主限额窗口(通常 7 天)
|
||||
- x-codex-secondary-*: 次级限额窗口(通常 5 小时)
|
||||
- x-codex-credits-*: 积分信息
|
||||
"""
|
||||
|
||||
from typing import Any, ClassVar
|
||||
|
||||
from src.services.provider.metadata_collectors import MetadataCollector
|
||||
|
||||
|
||||
def _safe_float(value: str | None) -> float | None:
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
return float(value)
|
||||
except (ValueError, TypeError):
|
||||
return None
|
||||
|
||||
|
||||
def _safe_int(value: str | None) -> int | None:
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
return int(float(value))
|
||||
except (ValueError, TypeError):
|
||||
return None
|
||||
|
||||
|
||||
def _safe_bool(value: str | None) -> bool | None:
|
||||
if value is None:
|
||||
return None
|
||||
return value.lower() in ("true", "1", "yes")
|
||||
|
||||
|
||||
class CodexMetadataCollector(MetadataCollector):
|
||||
"""Codex 额度/限流元数据采集器"""
|
||||
|
||||
# 支持 codex 类型,通过响应头判断是否是 Codex
|
||||
PROVIDER_TYPES: ClassVar[list[str]] = ["codex"]
|
||||
|
||||
def parse_headers(self, headers: dict[str, str]) -> dict[str, Any] | None:
|
||||
# 大小写不敏感查找
|
||||
lower_headers = {k.lower(): v for k, v in headers.items()}
|
||||
|
||||
plan_type = lower_headers.get("x-codex-plan-type")
|
||||
if plan_type is None:
|
||||
# 没有 Codex 特征头,跳过
|
||||
return None
|
||||
|
||||
result: dict[str, Any] = {"plan_type": plan_type}
|
||||
|
||||
# 主限额窗口(7 天)
|
||||
primary_used = _safe_float(lower_headers.get("x-codex-primary-used-percent"))
|
||||
if primary_used is not None:
|
||||
result["primary_used_percent"] = primary_used
|
||||
primary_reset_seconds = _safe_int(lower_headers.get("x-codex-primary-reset-after-seconds"))
|
||||
if primary_reset_seconds is not None:
|
||||
result["primary_reset_seconds"] = primary_reset_seconds
|
||||
primary_reset_at = _safe_int(lower_headers.get("x-codex-primary-reset-at"))
|
||||
if primary_reset_at is not None:
|
||||
result["primary_reset_at"] = primary_reset_at
|
||||
primary_window = _safe_int(lower_headers.get("x-codex-primary-window-minutes"))
|
||||
if primary_window is not None:
|
||||
result["primary_window_minutes"] = primary_window
|
||||
|
||||
# 次级限额窗口(5 小时)
|
||||
secondary_used = _safe_float(lower_headers.get("x-codex-secondary-used-percent"))
|
||||
if secondary_used is not None:
|
||||
result["secondary_used_percent"] = secondary_used
|
||||
secondary_reset_seconds = _safe_int(
|
||||
lower_headers.get("x-codex-secondary-reset-after-seconds")
|
||||
)
|
||||
if secondary_reset_seconds is not None:
|
||||
result["secondary_reset_seconds"] = secondary_reset_seconds
|
||||
secondary_reset_at = _safe_int(lower_headers.get("x-codex-secondary-reset-at"))
|
||||
if secondary_reset_at is not None:
|
||||
result["secondary_reset_at"] = secondary_reset_at
|
||||
secondary_window = _safe_int(lower_headers.get("x-codex-secondary-window-minutes"))
|
||||
if secondary_window is not None:
|
||||
result["secondary_window_minutes"] = secondary_window
|
||||
|
||||
# 积分信息
|
||||
has_credits = _safe_bool(lower_headers.get("x-codex-credits-has-credits"))
|
||||
if has_credits is not None:
|
||||
result["has_credits"] = has_credits
|
||||
credits_balance = _safe_float(lower_headers.get("x-codex-credits-balance"))
|
||||
if credits_balance is not None:
|
||||
result["credits_balance"] = credits_balance
|
||||
|
||||
return result
|
||||
@@ -189,6 +189,11 @@ class SystemConfigService:
|
||||
"value": "Aether",
|
||||
"description": "发件人名称",
|
||||
},
|
||||
# OAuth Token 刷新配置
|
||||
"enable_oauth_token_refresh": {
|
||||
"value": True,
|
||||
"description": "是否启用 OAuth Token 自动刷新任务,主动刷新即将过期的 OAuth token",
|
||||
},
|
||||
}
|
||||
|
||||
@classmethod
|
||||
|
||||
@@ -9,6 +9,7 @@
|
||||
- 连接池监控:定期检查数据库连接池状态
|
||||
- Pending 状态清理:清理异常的 Pending 状态记录
|
||||
- Gemini 文件映射清理:清理过期的 Gemini 文件→Key 映射
|
||||
- OAuth Token 刷新:主动刷新即将过期的 OAuth token
|
||||
|
||||
使用 APScheduler 进行任务调度,支持时区配置。
|
||||
"""
|
||||
@@ -38,12 +39,26 @@ class MaintenanceScheduler:
|
||||
|
||||
# 签到任务的 job_id
|
||||
CHECKIN_JOB_ID = "provider_checkin"
|
||||
# OAuth 刷新任务的 job_id
|
||||
OAUTH_REFRESH_JOB_ID = "oauth_token_refresh"
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.running = False
|
||||
self._interval_tasks = []
|
||||
self._stats_aggregation_lock = asyncio.Lock()
|
||||
|
||||
def trigger_oauth_refresh_check(self) -> None:
|
||||
"""
|
||||
触发 OAuth Token 刷新检查
|
||||
|
||||
当新增 OAuth Key 时调用此方法,重新调度刷新任务。
|
||||
会取消当前的调度,并立即重新计算下次执行时间。
|
||||
"""
|
||||
if not self.running:
|
||||
return
|
||||
|
||||
asyncio.create_task(self._schedule_next_oauth_refresh())
|
||||
|
||||
def _get_checkin_time(self) -> tuple[int, int]:
|
||||
"""获取签到任务的执行时间
|
||||
|
||||
@@ -203,6 +218,11 @@ class MaintenanceScheduler:
|
||||
name="Provider签到",
|
||||
)
|
||||
|
||||
# OAuth Token 刷新任务 - 动态调度
|
||||
# 根据最近即将过期的 token 时间来调度,避免固定间隔频繁查询
|
||||
# 启动时先执行一次,计算下次执行时间
|
||||
asyncio.create_task(self._schedule_next_oauth_refresh())
|
||||
|
||||
# 启动时执行一次初始化任务
|
||||
asyncio.create_task(self._run_startup_tasks())
|
||||
|
||||
@@ -274,6 +294,152 @@ class MaintenanceScheduler:
|
||||
"""Provider 签到任务(定时调用)"""
|
||||
await self._perform_provider_checkin()
|
||||
|
||||
async def _scheduled_oauth_token_refresh(self) -> None:
|
||||
"""OAuth Token 刷新任务(定时调用)"""
|
||||
await self._perform_oauth_token_refresh()
|
||||
# 执行完成后,调度下次执行
|
||||
await self._schedule_next_oauth_refresh()
|
||||
|
||||
async def _schedule_next_oauth_refresh(self) -> None:
|
||||
"""
|
||||
动态调度下次 OAuth Token 刷新任务
|
||||
|
||||
策略:
|
||||
- 查询所有 OAuth Key 的 expires_at
|
||||
- 找到最近即将过期的 token(在 refresh_threshold 内)
|
||||
- 设置下次执行时间为:最近过期时间 - 提前量(如提前 1 小时刷新)
|
||||
- 如果没有即将过期的 token,设置默认间隔(如 6 小时后)
|
||||
"""
|
||||
import json
|
||||
import time
|
||||
|
||||
from src.core.crypto import crypto_service
|
||||
from src.models.database import ProviderAPIKey
|
||||
|
||||
# 延迟启动,等待系统初始化
|
||||
await asyncio.sleep(5)
|
||||
|
||||
scheduler = get_scheduler()
|
||||
job_id = "oauth_token_refresh"
|
||||
|
||||
try:
|
||||
db = create_session()
|
||||
try:
|
||||
# 检查配置开关
|
||||
if not SystemConfigService.get_config(db, "enable_oauth_token_refresh", True):
|
||||
logger.info("OAuth Token 自动刷新已禁用,不调度任务")
|
||||
return
|
||||
# 查找所有活跃的 OAuth 类型 Key
|
||||
oauth_keys = (
|
||||
db.query(ProviderAPIKey)
|
||||
.filter(
|
||||
ProviderAPIKey.auth_type == "oauth",
|
||||
ProviderAPIKey.is_active == True, # noqa: E712
|
||||
)
|
||||
.all()
|
||||
)
|
||||
|
||||
if not oauth_keys:
|
||||
# 没有 OAuth Key,6 小时后再检查
|
||||
next_run = datetime.now(timezone.utc) + timedelta(hours=6)
|
||||
scheduler.add_date_job(
|
||||
self._scheduled_oauth_token_refresh,
|
||||
run_date=next_run,
|
||||
job_id=job_id,
|
||||
name="OAuth Token刷新",
|
||||
)
|
||||
logger.info("没有 OAuth Key,下次检查时间: {}", next_run.isoformat())
|
||||
return
|
||||
|
||||
now = int(time.time())
|
||||
# 24 小时内过期的都需要刷新(含提前量)
|
||||
refresh_window = 24 * 3600
|
||||
# 提前 1 小时执行刷新
|
||||
refresh_advance = 1 * 3600
|
||||
refresh_threshold_seconds = refresh_window + refresh_advance
|
||||
refresh_threshold = now + refresh_threshold_seconds
|
||||
|
||||
earliest_expires_at: int | None = None
|
||||
|
||||
for key in oauth_keys:
|
||||
if not key.auth_config:
|
||||
continue
|
||||
|
||||
try:
|
||||
decrypted_config = crypto_service.decrypt(key.auth_config)
|
||||
token_meta = json.loads(decrypted_config)
|
||||
expires_at = token_meta.get("expires_at")
|
||||
|
||||
if expires_at is None:
|
||||
continue
|
||||
|
||||
expires_at_int = int(expires_at)
|
||||
|
||||
# 已经过期或在阈值内,需要立即刷新
|
||||
if expires_at_int <= refresh_threshold:
|
||||
# 立即执行
|
||||
next_run = datetime.now(timezone.utc) + timedelta(seconds=10)
|
||||
scheduler.add_date_job(
|
||||
self._scheduled_oauth_token_refresh,
|
||||
run_date=next_run,
|
||||
job_id=job_id,
|
||||
name="OAuth Token刷新",
|
||||
)
|
||||
logger.info(
|
||||
"发现即将过期的 OAuth Token,立即执行刷新: {}",
|
||||
next_run.isoformat(),
|
||||
)
|
||||
return
|
||||
|
||||
# 记录最近的过期时间
|
||||
if earliest_expires_at is None or expires_at_int < earliest_expires_at:
|
||||
earliest_expires_at = expires_at_int
|
||||
|
||||
except Exception:
|
||||
continue
|
||||
|
||||
# 计算下次执行时间
|
||||
if earliest_expires_at is not None:
|
||||
# 在最近过期时间前 24 小时 + 提前量执行
|
||||
next_run_ts = earliest_expires_at - refresh_threshold_seconds
|
||||
# 确保不会是过去的时间
|
||||
if next_run_ts <= now:
|
||||
next_run_ts = now + 60 # 1 分钟后
|
||||
next_run = datetime.fromtimestamp(next_run_ts, tz=timezone.utc)
|
||||
else:
|
||||
# 没有有效的过期时间,6 小时后再检查
|
||||
next_run = datetime.now(timezone.utc) + timedelta(hours=6)
|
||||
|
||||
# 限制最大间隔为 24 小时
|
||||
max_next_run = datetime.now(timezone.utc) + timedelta(hours=24)
|
||||
if next_run > max_next_run:
|
||||
next_run = max_next_run
|
||||
|
||||
scheduler.add_date_job(
|
||||
self._scheduled_oauth_token_refresh,
|
||||
run_date=next_run,
|
||||
job_id=job_id,
|
||||
name="OAuth Token刷新",
|
||||
)
|
||||
logger.info("OAuth Token 刷新任务已调度,下次执行时间: {}", next_run.isoformat())
|
||||
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
except Exception as e:
|
||||
logger.exception("调度 OAuth Token 刷新任务失败: {}", e)
|
||||
# 出错时 1 小时后重试
|
||||
next_run = datetime.now(timezone.utc) + timedelta(hours=1)
|
||||
try:
|
||||
scheduler.add_date_job(
|
||||
self._scheduled_oauth_token_refresh,
|
||||
run_date=next_run,
|
||||
job_id=job_id,
|
||||
name="OAuth Token刷新",
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# ========== 实际任务实现 ==========
|
||||
|
||||
async def _perform_stats_aggregation(self, backfill: bool = False) -> None:
|
||||
@@ -1031,6 +1197,294 @@ class MaintenanceScheduler:
|
||||
|
||||
return total_deleted
|
||||
|
||||
async def _perform_oauth_token_refresh(self) -> None:
|
||||
"""
|
||||
主动刷新即将过期的 OAuth token
|
||||
|
||||
策略:
|
||||
- 查找所有 auth_type='oauth' 且 is_active=True 的 Key
|
||||
- 检查 auth_config 中的 expires_at,如果在 24 小时内过期则刷新
|
||||
- 使用 refresh_token 换取新的 access_token
|
||||
- 更新数据库中的 token 信息
|
||||
"""
|
||||
import json
|
||||
import time
|
||||
|
||||
from src.core.crypto import crypto_service
|
||||
from src.core.provider_oauth_utils import enrich_auth_config, post_oauth_token
|
||||
from src.core.provider_templates.fixed_providers import FIXED_PROVIDERS
|
||||
from src.core.provider_templates.types import ProviderType
|
||||
from src.models.database import ProviderAPIKey
|
||||
|
||||
# 检查配置开关
|
||||
check_db = create_session()
|
||||
try:
|
||||
if not SystemConfigService.get_config(check_db, "enable_oauth_token_refresh", True):
|
||||
logger.info("OAuth Token 自动刷新已禁用,跳过任务")
|
||||
return
|
||||
finally:
|
||||
check_db.close()
|
||||
|
||||
logger.info("开始执行 OAuth Token 刷新任务...")
|
||||
|
||||
db = create_session()
|
||||
refreshed_count = 0
|
||||
failed_count = 0
|
||||
skipped_count = 0
|
||||
|
||||
try:
|
||||
# 查找所有活跃的 OAuth 类型 Key
|
||||
oauth_keys = (
|
||||
db.query(ProviderAPIKey)
|
||||
.filter(
|
||||
ProviderAPIKey.auth_type == "oauth",
|
||||
ProviderAPIKey.is_active == True, # noqa: E712
|
||||
)
|
||||
.all()
|
||||
)
|
||||
|
||||
if not oauth_keys:
|
||||
logger.info("没有找到需要刷新的 OAuth Key")
|
||||
return
|
||||
|
||||
logger.info("找到 {} 个 OAuth Key,开始检查过期状态...", len(oauth_keys))
|
||||
|
||||
now = int(time.time())
|
||||
# 24 小时内过期的都刷新(含提前量)
|
||||
refresh_window = 24 * 3600
|
||||
# 提前 1 小时执行刷新
|
||||
refresh_advance = 1 * 3600
|
||||
refresh_threshold = now + refresh_window + refresh_advance
|
||||
|
||||
for key in oauth_keys:
|
||||
try:
|
||||
# 解密 auth_config
|
||||
if not key.auth_config:
|
||||
skipped_count += 1
|
||||
continue
|
||||
|
||||
try:
|
||||
decrypted_config = crypto_service.decrypt(key.auth_config)
|
||||
token_meta = json.loads(decrypted_config)
|
||||
except Exception:
|
||||
logger.warning("Key {} auth_config 解密失败,跳过", key.id)
|
||||
skipped_count += 1
|
||||
continue
|
||||
|
||||
expires_at = token_meta.get("expires_at")
|
||||
refresh_token = token_meta.get("refresh_token")
|
||||
provider_type = str(token_meta.get("provider_type") or "")
|
||||
|
||||
# 检查是否需要刷新
|
||||
if expires_at is None:
|
||||
skipped_count += 1
|
||||
continue
|
||||
|
||||
try:
|
||||
expires_at_int = int(expires_at)
|
||||
except (ValueError, TypeError):
|
||||
skipped_count += 1
|
||||
continue
|
||||
|
||||
if expires_at_int > refresh_threshold:
|
||||
# 还没到刷新时间
|
||||
skipped_count += 1
|
||||
continue
|
||||
|
||||
if not refresh_token or not provider_type:
|
||||
logger.warning(
|
||||
"Key {} 缺少 refresh_token 或 provider_type,无法刷新", key.id
|
||||
)
|
||||
skipped_count += 1
|
||||
continue
|
||||
|
||||
# 获取 provider 模板
|
||||
try:
|
||||
provider_type_enum = ProviderType(provider_type)
|
||||
except ValueError:
|
||||
logger.warning("Key {} 未知的 provider_type: {}", key.id, provider_type)
|
||||
skipped_count += 1
|
||||
continue
|
||||
|
||||
template = FIXED_PROVIDERS.get(provider_type_enum)
|
||||
if not template or not template.oauth:
|
||||
logger.warning("Key {} provider {} 不支持 OAuth", key.id, provider_type)
|
||||
skipped_count += 1
|
||||
continue
|
||||
|
||||
# 获取代理配置
|
||||
proxy_config = None
|
||||
if key.provider and key.provider.endpoints:
|
||||
for endpoint in key.provider.endpoints:
|
||||
if endpoint.proxy:
|
||||
proxy_config = endpoint.proxy
|
||||
break
|
||||
|
||||
# 执行刷新
|
||||
token_url = template.oauth.token_url
|
||||
is_json = "anthropic.com" in token_url
|
||||
|
||||
if is_json:
|
||||
body = {
|
||||
"grant_type": "refresh_token",
|
||||
"client_id": template.oauth.client_id,
|
||||
"refresh_token": str(refresh_token),
|
||||
}
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"Accept": "application/json",
|
||||
}
|
||||
data = None
|
||||
json_body = body
|
||||
else:
|
||||
form = {
|
||||
"grant_type": "refresh_token",
|
||||
"client_id": template.oauth.client_id,
|
||||
"refresh_token": str(refresh_token),
|
||||
}
|
||||
if template.oauth.client_secret:
|
||||
form["client_secret"] = template.oauth.client_secret
|
||||
headers = {
|
||||
"Content-Type": "application/x-www-form-urlencoded",
|
||||
"Accept": "application/json",
|
||||
}
|
||||
data = form
|
||||
json_body = None
|
||||
|
||||
logger.info(
|
||||
"刷新 Key {} ({}) 的 OAuth token,当前过期时间: {}",
|
||||
key.id,
|
||||
key.name,
|
||||
datetime.fromtimestamp(expires_at_int, tz=timezone.utc).isoformat(),
|
||||
)
|
||||
|
||||
resp = await post_oauth_token(
|
||||
provider_type=provider_type,
|
||||
token_url=token_url,
|
||||
headers=headers,
|
||||
data=data,
|
||||
json_body=json_body,
|
||||
proxy_config=proxy_config,
|
||||
timeout_seconds=30.0,
|
||||
)
|
||||
|
||||
if 200 <= resp.status_code < 300:
|
||||
token = resp.json()
|
||||
access_token = str(token.get("access_token") or "")
|
||||
new_refresh_token = str(token.get("refresh_token") or "")
|
||||
expires_in = token.get("expires_in")
|
||||
new_expires_at = None
|
||||
|
||||
try:
|
||||
if expires_in is not None:
|
||||
new_expires_at = int(time.time()) + int(expires_in)
|
||||
except Exception:
|
||||
new_expires_at = None
|
||||
|
||||
if access_token:
|
||||
# 更新 token_meta
|
||||
token_meta["token_type"] = token.get("token_type")
|
||||
if new_refresh_token:
|
||||
token_meta["refresh_token"] = new_refresh_token
|
||||
token_meta["expires_at"] = new_expires_at
|
||||
token_meta["scope"] = token.get("scope")
|
||||
token_meta["updated_at"] = int(time.time())
|
||||
|
||||
# 提取额外信息
|
||||
token_meta = await enrich_auth_config(
|
||||
provider_type=provider_type,
|
||||
auth_config=token_meta,
|
||||
token_response=token,
|
||||
access_token=access_token,
|
||||
proxy_config=proxy_config,
|
||||
)
|
||||
|
||||
# 更新数据库
|
||||
encrypted_token = crypto_service.encrypt(access_token)
|
||||
encrypted_config = crypto_service.encrypt(json.dumps(token_meta))
|
||||
|
||||
key.api_key = encrypted_token
|
||||
key.auth_config = encrypted_config
|
||||
# 刷新成功,清除失效标记
|
||||
key.oauth_invalid_at = None
|
||||
key.oauth_invalid_reason = None
|
||||
db.commit()
|
||||
|
||||
refreshed_count += 1
|
||||
new_expires_str = (
|
||||
datetime.fromtimestamp(new_expires_at, tz=timezone.utc).isoformat()
|
||||
if new_expires_at
|
||||
else "unknown"
|
||||
)
|
||||
logger.info(
|
||||
"Key {} ({}) OAuth token 刷新成功,新过期时间: {}",
|
||||
key.id,
|
||||
key.name,
|
||||
new_expires_str,
|
||||
)
|
||||
else:
|
||||
failed_count += 1
|
||||
logger.warning(
|
||||
"Key {} ({}) 刷新响应中没有 access_token", key.id, key.name
|
||||
)
|
||||
else:
|
||||
failed_count += 1
|
||||
# 解析错误原因
|
||||
error_reason = "HTTP {}".format(resp.status_code)
|
||||
try:
|
||||
error_body = resp.json()
|
||||
if "error" in error_body:
|
||||
error_reason = str(
|
||||
error_body.get("error_description") or error_body.get("error")
|
||||
)
|
||||
except Exception:
|
||||
error_reason = (
|
||||
resp.text[:100] if resp.text else "HTTP {}".format(resp.status_code)
|
||||
)
|
||||
|
||||
# 标记为失效(400/401/403 通常表示永久性错误)
|
||||
if resp.status_code in (400, 401, 403):
|
||||
key.oauth_invalid_at = datetime.now(timezone.utc)
|
||||
key.oauth_invalid_reason = error_reason
|
||||
db.commit()
|
||||
logger.warning(
|
||||
"Key {} ({}) OAuth token 刷新失败,已标记为失效: {}",
|
||||
key.id,
|
||||
key.name,
|
||||
error_reason,
|
||||
)
|
||||
else:
|
||||
logger.warning(
|
||||
"Key {} ({}) OAuth token 刷新失败,状态码: {},响应: {}",
|
||||
key.id,
|
||||
key.name,
|
||||
resp.status_code,
|
||||
resp.text[:200],
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
failed_count += 1
|
||||
logger.exception("Key {} OAuth token 刷新出错: {}", key.id, e)
|
||||
try:
|
||||
db.rollback()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# 避免请求过于频繁
|
||||
await asyncio.sleep(1)
|
||||
|
||||
except Exception as e:
|
||||
logger.exception("OAuth Token 刷新任务执行出错: {}", e)
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
logger.info(
|
||||
"OAuth Token 刷新任务完成: 刷新 {} 个,失败 {} 个,跳过 {} 个",
|
||||
refreshed_count,
|
||||
failed_count,
|
||||
skipped_count,
|
||||
)
|
||||
|
||||
|
||||
# 全局单例
|
||||
_maintenance_scheduler = None
|
||||
|
||||
@@ -14,6 +14,7 @@ from typing import Any, Callable
|
||||
|
||||
from apscheduler.schedulers.asyncio import AsyncIOScheduler
|
||||
from apscheduler.triggers.cron import CronTrigger
|
||||
from apscheduler.triggers.date import DateTrigger
|
||||
from apscheduler.triggers.interval import IntervalTrigger
|
||||
|
||||
from src.core.logger import logger
|
||||
@@ -136,6 +137,40 @@ class TaskScheduler:
|
||||
|
||||
logger.info(f"已注册间隔任务: {display_name}, 执行间隔: {interval_desc}")
|
||||
|
||||
def add_date_job(
|
||||
self,
|
||||
func: Callable[..., Any],
|
||||
run_date: datetime,
|
||||
job_id: str | None = None,
|
||||
name: str | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Any:
|
||||
"""
|
||||
添加一次性定时任务(在指定时间执行一次)
|
||||
|
||||
Args:
|
||||
func: 要执行的函数
|
||||
run_date: 执行时间(datetime 对象)
|
||||
job_id: 任务ID
|
||||
name: 任务名称(用于日志)
|
||||
**kwargs: 传递给任务函数的参数
|
||||
"""
|
||||
trigger = DateTrigger(run_date=run_date)
|
||||
|
||||
job_id = job_id or func.__name__
|
||||
display_name = name or job_id
|
||||
|
||||
self.scheduler.add_job(
|
||||
func,
|
||||
trigger,
|
||||
id=job_id,
|
||||
name=display_name,
|
||||
replace_existing=True,
|
||||
kwargs=kwargs,
|
||||
)
|
||||
|
||||
logger.info("已注册一次性任务: {}, 执行时间: {}", display_name, run_date.isoformat())
|
||||
|
||||
def start(self) -> Any:
|
||||
"""启动调度器"""
|
||||
if self._started:
|
||||
|
||||
Reference in New Issue
Block a user