feat: OAuth 账户管理、维护调度、端点健康检查增强及前端优化

- 新增 OAuth 账户管理对话框和提供商详情抽屉中的 OAuth 信息展示
- 新增维护调度器(maintenance_scheduler)支持定时清理和健康检查
- 增强端点健康检查器,支持更多检测策略
- 重构 codex 服务为 metadata_collectors 模块
- 优化 OpenAI CLI normalizer 代码结构
- 前端: 改进使用量表格、统计图表、指南页面和异步任务管理
- 扩展多个数据库字符串列为 TEXT 类型
- 新增倒计时 composable 和 provider OAuth API 端点
This commit is contained in:
fawney19
2026-02-04 23:59:45 +08:00
parent 24c9105628
commit 4d6e7c094f
64 changed files with 3885 additions and 930 deletions

View File

@@ -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,
}
)

View File

@@ -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_urilocalhost\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"),
)

View File

@@ -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", {})

View File

@@ -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(

View File

@@ -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:
额外请求头字典
"""

View File

@@ -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)

View File

@@ -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计数和费用计算"""

View File

@@ -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,
)

View File

@@ -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

View File

@@ -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():

View File

@@ -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

View File

@@ -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

View File

@@ -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"]

View File

@@ -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
# ============ 响应转换 ============

View File

@@ -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且需要交替这里做最小修复

View File

@@ -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

View File

@@ -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:

View File

@@ -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=falseCodex 强制要求,标准 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):
# 普通 messageTextBlock
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:

View File

@@ -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

View File

@@ -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)

View File

@@ -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

View File

@@ -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="提供商优先级(数字越小越优先)")

View File

@@ -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",
]

View 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",
]

View 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

View File

@@ -189,6 +189,11 @@ class SystemConfigService:
"value": "Aether",
"description": "发件人名称",
},
# OAuth Token 刷新配置
"enable_oauth_token_refresh": {
"value": True,
"description": "是否启用 OAuth Token 自动刷新任务,主动刷新即将过期的 OAuth token",
},
}
@classmethod

View File

@@ -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 Key6 小时后再检查
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

View File

@@ -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: