优化 OAuth Token 导入解析与号池管理直刷空列表问题 (#330)

* docs: add aether-proxy C++ rewrite design

* chore: ignore local worktrees directory

* 修复 OAuth Token 导入识别与账号信息解析

* 修复号池管理直刷账号列表为空

---------

Co-authored-by: Your Name <[email protected]>
This commit is contained in:
Dao
2026-04-30 18:07:00 +08:00
committed by GitHub
co-authored by Your Name
parent c34565b02b
commit 392c557831
14 changed files with 750 additions and 308 deletions
+1
View File
@@ -204,6 +204,7 @@ logs/
# Git backup
.git.backup/
.worktrees/
# Database backups
backups/
+1 -1
View File
@@ -94,7 +94,7 @@ export async function completeProviderLevelOAuth(
export async function importProviderRefreshToken(
providerId: string,
data: { refresh_token: string; name?: string; proxy_node_id?: string }
data: { refresh_token?: string; access_token?: string; name?: string; proxy_node_id?: string }
): Promise<ProviderOAuthCompleteResponseWithKey> {
const resp = await client.post(`/api/admin/provider-oauth/providers/${providerId}/import-refresh-token`, data)
return resp.data
@@ -292,6 +292,7 @@ export interface EndpointAPIKey {
oauth_account_user_id?: string | null // Codex ChatGPT account-user 联合 ID
oauth_account_name?: string | null
oauth_organizations?: OAuthOrganizationInfo[] | null // OAuth 关联组织/工作区摘要
oauth_temporary?: boolean | null // 是否为仅 Access Token 导入的临时 OAuth 账号
oauth_invalid_at?: number | null // 兼容字段;优先使用 status_snapshot.oauth
oauth_invalid_reason?: string | null // 兼容字段;优先使用 status_snapshot.oauth
status_snapshot?: ProviderKeyStatusSnapshot | null
@@ -410,8 +410,8 @@
:reset-key="importInputResetKey"
drop-title="拖入授权文件或点击选择"
drop-hint="支持 .json / .txt,可多选"
manual-placeholder="粘贴 Refresh Token 或 JSON 内容"
paste-toggle-text="或手动粘贴 Refresh Token"
manual-placeholder="粘贴 Refresh Token / Access Token 或 JSON 内容"
paste-toggle-text="或手动粘贴 Token"
file-toggle-text="或选择 JSON 文件导入"
textarea-class="min-h-[200px] text-xs font-mono break-all !rounded-xl"
@error="handleImportInputError"
@@ -893,7 +893,7 @@ function isBatchImport(text: string): boolean {
return lines.length > 1
}
function parseImportText(text: string): { refresh_token: string; name?: string } | null {
function parseImportText(text: string): { refresh_token?: string; access_token?: string; name?: string } | null {
const trimmed = text.trim()
if (!trimmed) return null
@@ -907,9 +907,19 @@ function parseImportText(text: string): { refresh_token: string; name?: string }
if (typeof parsed === 'object' && parsed !== null) {
const obj = parsed as Record<string, unknown>
const refreshToken = obj.refresh_token
if (typeof refreshToken === 'string' && refreshToken.trim()) {
const refreshTokenCamel = obj.refreshToken
const accessToken = obj.access_token
const accessTokenCamel = obj.accessToken
const normalizedRefreshToken = typeof refreshToken === 'string' && refreshToken.trim()
? refreshToken.trim()
: (typeof refreshTokenCamel === 'string' && refreshTokenCamel.trim() ? refreshTokenCamel.trim() : undefined)
const normalizedAccessToken = typeof accessToken === 'string' && accessToken.trim()
? accessToken.trim()
: (typeof accessTokenCamel === 'string' && accessTokenCamel.trim() ? accessTokenCamel.trim() : undefined)
if (normalizedRefreshToken || normalizedAccessToken) {
return {
refresh_token: refreshToken.trim(),
refresh_token: normalizedRefreshToken,
access_token: normalizedAccessToken,
name: (typeof obj.name === 'string' ? obj.name : undefined) || (typeof obj.oauth_email === 'string' ? obj.oauth_email : undefined),
}
}
@@ -919,9 +929,34 @@ function parseImportText(text: string): { refresh_token: string; name?: string }
// Not JSON: treat as raw token.
}
if (isLikelyJwtToken(trimmed)) {
return { access_token: trimmed }
}
return { refresh_token: trimmed }
}
function isLikelyJwtToken(token: string): boolean {
const parts = token.trim().split('.')
if (parts.length !== 3 || parts.some(part => !part)) return false
try {
const header = JSON.parse(decodeBase64Url(parts[0])) as Record<string, unknown>
const payload = JSON.parse(decodeBase64Url(parts[1])) as Record<string, unknown>
const tokenType = typeof header.typ === 'string' ? header.typ.toLowerCase() : ''
if (tokenType && tokenType !== 'jwt' && tokenType !== 'at+jwt') return false
return ['exp', 'aud', 'iss', 'scope', 'scp'].some(key => key in payload)
} catch {
return false
}
}
function decodeBase64Url(value: string): string {
const normalized = value.replace(/-/g, '+').replace(/_/g, '/')
const padded = normalized.padEnd(normalized.length + ((4 - (normalized.length % 4)) % 4), '=')
return atob(padded)
}
function handleImportInputError(payload: { message: string; title?: string }) {
showError(payload.message, payload.title)
}
@@ -370,7 +370,16 @@
>
{{ getKeyOAuthExpires(key)?.text }}
</span>
<Badge
v-if="key.oauth_temporary"
variant="outline"
class="text-[10px] px-1.5 py-0 shrink-0"
title="仅通过 Access Token 导入,无法自动刷新,到期后需要重新导入"
>
临时
</Badge>
<Button
v-else
variant="ghost"
size="icon"
class="h-4 w-4 shrink-0"
+16 -6
View File
@@ -1013,7 +1013,7 @@
<!-- Empty keys -->
<div
v-if="keyPage.keys.length === 0 && !keysLoading"
v-if="keyPage.keys.length === 0 && !keysLoading && keysLoadedOnce"
class="flex flex-col items-center justify-center py-16 text-center"
>
<div class="mx-auto flex h-16 w-16 items-center justify-center rounded-full bg-muted">
@@ -1232,23 +1232,28 @@ async function loadOverview() {
const enabledProviders = allProviders.filter(item => item.pool_enabled)
poolProviders.value = enabledProviders
// Keep selected provider aligned with dropdown options.
const selectedId = selectedProviderId.value
const queryProviderId = getQueryValue('providerId')
const queryProviderExists = Boolean(
queryProviderId && enabledProviders.some(item => item.provider_id === queryProviderId),
)
const selectedId = selectedProviderId.value || (queryProviderExists ? queryProviderId : null)
const selectedStillExists = Boolean(
selectedId && enabledProviders.some(item => item.provider_id === selectedId),
)
if (!selectedStillExists) {
if (enabledProviders.length > 0) {
// Do not block overview loading on key list fetch; keys area has its own loader.
void selectProvider(enabledProviders[0].provider_id)
await selectProvider(enabledProviders[0].provider_id)
} else {
selectedProviderId.value = null
selectedProviderData.value = null
keysLoadedOnce.value = false
showAccountBatchDialog.value = false
closeProviderProxyPopovers()
resetKeyPage()
}
} else if (selectedId && selectedId !== selectedProviderId.value) {
await selectProvider(selectedId, { preserveSearch: true })
}
} catch (err) {
if (requestId !== overviewRequestId) return
@@ -1455,6 +1460,7 @@ async function selectProvider(id: string, options: { preserveSearch?: boolean }
clearTimeout(keysSearchDebounceTimer)
keysSearchDebounceTimer = null
}
keysLoadedOnce.value = false
resetKeyPage(1, pageSize.value)
const keysTask = loadKeys()
// Provider summary is non-blocking for key list rendering.
@@ -1486,6 +1492,7 @@ function createEmptyKeyPage(page = 1, pageSizeValue = 50): PoolKeysPageResponse
const keyPage = ref<PoolKeysPageResponse>(createEmptyKeyPage())
const keysLoading = ref(false)
const keysLoadedOnce = ref(false)
const refreshingCurrentPageQuota = ref(false)
const searchQuery = ref('')
const statusFilter = ref('all')
@@ -1524,10 +1531,11 @@ watch(searchQuery, (value) => {
watch(
() => getQueryValue('providerId'),
(value) => {
if (overviewLoading.value) return
if (!value || value === selectedProviderId.value) return
if (!poolProviders.value.some(item => item.provider_id === value)) return
void selectProvider(value, { preserveSearch: true })
},
{ immediate: true },
)
watch(selectedProviderId, (value) => {
@@ -1700,9 +1708,11 @@ async function loadKeys() {
})
if (requestId !== keysRequestId || selectedProviderId.value !== providerId) return
keyPage.value = nextPage
keysLoadedOnce.value = true
} catch (err) {
if (requestId !== keysRequestId || selectedProviderId.value !== providerId) return
resetKeyPage(page, pageSizeValue)
keysLoadedOnce.value = true
showError(parseApiError(err))
} finally {
if (requestId === keysRequestId) {
+425 -295
View File
@@ -36,7 +36,7 @@ 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.core.provider_oauth_utils import enrich_auth_config, post_oauth_token
from src.core.provider_oauth_utils import enrich_auth_config, parse_codex_id_token, post_oauth_token
from src.core.provider_templates.fixed_providers import FIXED_PROVIDERS
from src.core.provider_templates.types import ProviderType
from src.database import get_db_context
@@ -150,6 +150,10 @@ def _mark_refresh_failed_sync(key_id: str, reason: str) -> None:
if not key:
raise NotFoundException("Key 不存在", "key")
current_reason = str(getattr(key, "oauth_invalid_reason", None) or "").strip()
from src.services.provider.oauth_token import is_account_level_block
if is_account_level_block(current_reason):
return
merged_reason = _merge_refresh_failure_reason(current_reason, reason)
if merged_reason is None:
return
@@ -1577,11 +1581,62 @@ def _coerce_import_str(value: Any) -> str | None:
return normalized or None
def _decode_unverified_jwt_json_part(part: str) -> dict[str, Any] | None:
try:
padding = "=" * (-len(part) % 4)
decoded = base64.urlsafe_b64decode(f"{part}{padding}")
payload = json.loads(decoded.decode("utf-8"))
except Exception:
return None
return payload if isinstance(payload, dict) else None
def _looks_like_access_token(token: str) -> bool:
"""识别手工粘贴的 JWT Access Token,避免误当 Refresh Token 去刷新验证。"""
parts = token.strip().split(".")
if len(parts) != 3 or not all(parts):
return False
header = _decode_unverified_jwt_json_part(parts[0])
payload = _decode_unverified_jwt_json_part(parts[1])
if not header or not payload:
return False
token_type = str(header.get("typ") or "").lower()
if token_type and token_type not in {"jwt", "at+jwt"}:
return False
return any(key in payload for key in ("exp", "aud", "iss", "scope", "scp"))
def _build_standard_oauth_import_entry_from_token(token: str) -> dict[str, Any]:
if _looks_like_access_token(token):
return {"access_token": token}
return {"refresh_token": token}
def _normalize_single_import_tokens(
*,
refresh_token: str | None,
access_token: str | None,
) -> tuple[str, str]:
refresh_token_value = (refresh_token or "").strip()
access_token_value = (access_token or "").strip()
if (
refresh_token_value
and not access_token_value
and _looks_like_access_token(refresh_token_value)
):
access_token_value = refresh_token_value
refresh_token_value = ""
return refresh_token_value, access_token_value
def _extract_standard_oauth_import_entry(item: Any) -> dict[str, Any] | None:
if isinstance(item, str):
token = _coerce_import_str(item)
if token:
return {"refresh_token": token}
return _build_standard_oauth_import_entry_from_token(token)
return None
if not isinstance(item, dict):
return None
@@ -1589,10 +1644,17 @@ def _extract_standard_oauth_import_entry(item: Any) -> dict[str, Any] | None:
refresh_token = _coerce_import_str(item.get("refresh_token")) or _coerce_import_str(
item.get("refreshToken")
)
if not refresh_token:
access_token = _coerce_import_str(item.get("access_token")) or _coerce_import_str(
item.get("accessToken")
)
if not refresh_token and not access_token:
return None
entry: dict[str, Any] = {"refresh_token": refresh_token}
entry: dict[str, Any] = {}
if refresh_token:
entry["refresh_token"] = refresh_token
if access_token:
entry["access_token"] = access_token
account_id = (
_coerce_import_str(item.get("account_id"))
@@ -1678,14 +1740,18 @@ def _parse_standard_oauth_import_entries(raw_input: str) -> list[dict[str, Any]]
for line in raw.splitlines():
token = line.strip()
if token and not token.startswith("#"):
result.append({"refresh_token": token})
result.append(_build_standard_oauth_import_entry_from_token(token))
return result
def _parse_tokens_input(raw_input: str) -> list[str]:
"""兼容旧逻辑:仅返回 refresh_token 列表。"""
return [entry["refresh_token"] for entry in _parse_standard_oauth_import_entries(raw_input)]
return [
str(entry["refresh_token"])
for entry in _parse_standard_oauth_import_entries(raw_input)
if entry.get("refresh_token")
]
def _parse_kiro_import_input(raw_input: str) -> list[dict[str, Any]]:
@@ -1788,7 +1854,8 @@ def _parse_kiro_import_input(raw_input: str) -> list[dict[str, Any]]:
class ImportRefreshTokenRequest(BaseModel):
refresh_token: str = Field(..., min_length=1, description="Refresh Token")
refresh_token: str | None = Field(None, min_length=1, description="Refresh Token")
access_token: str | None = Field(None, min_length=1, description="Access Token(Codex 可选)")
name: str | None = Field(None, max_length=100, description="账号名称(可选)")
proxy_node_id: str | None = Field(
None,
@@ -1944,6 +2011,243 @@ def _apply_codex_import_hints(auth_config: dict[str, Any], import_entry: dict[st
auth_config[field] = value
def _decode_access_token_expires_at(access_token: str) -> int | None:
"""从 access token JWT payload 中解析 exp。"""
try:
import jwt
claims = jwt.decode(
access_token,
options={
"verify_signature": False,
"verify_aud": False,
},
)
except Exception:
return None
if not isinstance(claims, dict):
return None
try:
exp = int(claims.get("exp"))
except Exception:
return None
return exp if exp > 0 else None
def _build_access_token_import_payload(
*,
provider_type: str,
access_token: str,
refresh_token: str | None = None,
refresh_error: str | None = None,
) -> tuple[dict[str, Any], dict[str, Any]]:
expires_at = _decode_access_token_expires_at(access_token)
token_data: dict[str, Any] = {
"access_token": access_token,
"token_type": "Bearer",
}
if expires_at is not None:
token_data["expires_at"] = expires_at
auth_config: dict[str, Any] = {
"provider_type": provider_type,
"token_type": "Bearer",
"refresh_token": refresh_token or None,
"expires_at": expires_at,
"scope": None,
"updated_at": int(time.time()),
"imported_from_access_token": True,
"access_token_import_temporary": not bool(refresh_token),
}
if refresh_error:
auth_config["refresh_token_import_error"] = refresh_error[:300]
return token_data, auth_config
async def _build_codex_access_token_import_auth_config(
*,
provider_type: str,
access_token: str,
refresh_token: str | None,
proxy_config: dict[str, Any] | None,
import_entry: dict[str, Any] | None = None,
refresh_error: str | None = None,
) -> tuple[str, dict[str, Any], int | None]:
if provider_type != ProviderType.CODEX.value:
raise InvalidRequestException("仅 Codex 支持 Access Token 导入")
access_token = access_token.strip()
if len(access_token) < 10:
raise InvalidRequestException("Access Token 无效或过短")
token_data, auth_config = _build_access_token_import_payload(
provider_type=provider_type,
access_token=access_token,
refresh_token=refresh_token,
refresh_error=refresh_error,
)
try:
auth_config = await enrich_auth_config(
provider_type=provider_type,
auth_config=auth_config,
token_response=token_data,
access_token=access_token,
proxy_config=proxy_config,
)
except Exception as exc:
logger.warning("Codex Access Token 导入: enrich_auth_config 失败: {}", exc)
parsed = parse_codex_id_token(access_token)
for key, value in parsed.items():
if value and not auth_config.get(key):
auth_config[key] = value
if import_entry:
_apply_codex_import_hints(auth_config, import_entry)
return access_token, auth_config, auth_config.get("expires_at")
async def _exchange_refresh_token_for_import(
*,
provider_type: str,
template: Any,
refresh_token: str,
proxy_config: dict[str, Any] | None,
timeout_seconds: float,
) -> tuple[str, dict[str, Any], int | None]:
token_url = template.oauth.token_url
is_json = "anthropic.com" in token_url
scope_str = " ".join(template.oauth.scopes) if template.oauth.scopes else ""
if is_json:
body: dict[str, Any] = {
"grant_type": "refresh_token",
"client_id": template.oauth.client_id,
"refresh_token": refresh_token,
}
if scope_str:
body["scope"] = scope_str
headers = {"Content-Type": "application/json", "Accept": "application/json"}
data = None
json_body = body
else:
form: dict[str, str] = {
"grant_type": "refresh_token",
"client_id": template.oauth.client_id,
"refresh_token": refresh_token,
}
if scope_str:
form["scope"] = scope_str
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
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=timeout_seconds,
)
if resp.status_code < 200 or resp.status_code >= 300:
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}"
raise InvalidRequestException(f"Refresh Token 验证失败: {error_reason}")
token_data = resp.json()
access_token = str(token_data.get("access_token") or "")
new_refresh_token = str(token_data.get("refresh_token") or "") or refresh_token
if not access_token:
raise InvalidRequestException("token refresh 返回缺少 access_token")
expires_in = token_data.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
auth_config: dict[str, Any] = {
"provider_type": provider_type,
"token_type": token_data.get("token_type"),
"refresh_token": new_refresh_token or None,
"expires_at": expires_at,
"scope": token_data.get("scope"),
"updated_at": int(time.time()),
}
auth_config = await enrich_auth_config(
provider_type=provider_type,
auth_config=auth_config,
token_response=token_data,
access_token=access_token,
proxy_config=proxy_config,
)
return access_token, auth_config, expires_at
async def _build_standard_oauth_import_credentials(
*,
provider_type: str,
template: Any,
import_entry: dict[str, Any],
proxy_config: dict[str, Any] | None,
timeout_seconds: float,
) -> tuple[str, dict[str, Any], int | None]:
refresh_token = str(import_entry.get("refresh_token") or "").strip()
access_token_input = str(import_entry.get("access_token") or "").strip()
if refresh_token:
try:
access_token, auth_config, expires_at = await _exchange_refresh_token_for_import(
provider_type=provider_type,
template=template,
refresh_token=refresh_token,
proxy_config=proxy_config,
timeout_seconds=timeout_seconds,
)
if provider_type == ProviderType.CODEX.value:
_apply_codex_import_hints(auth_config, import_entry)
return access_token, auth_config, expires_at
except Exception as exc:
if access_token_input and provider_type == ProviderType.CODEX.value:
return await _build_codex_access_token_import_auth_config(
provider_type=provider_type,
access_token=access_token_input,
refresh_token=refresh_token,
proxy_config=proxy_config,
import_entry=import_entry,
refresh_error=str(exc),
)
raise
if access_token_input:
if provider_type != ProviderType.CODEX.value:
raise InvalidRequestException("仅 Codex 支持 Access Token 导入")
return await _build_codex_access_token_import_auth_config(
provider_type=provider_type,
access_token=access_token_input,
refresh_token=None,
proxy_config=proxy_config,
import_entry=import_entry,
)
raise InvalidRequestException("Token 无效或过短")
@router.post(
"/providers/{provider_id}/import-refresh-token",
response_model=ProviderCompleteOAuthResponse,
@@ -1969,9 +2273,9 @@ async def import_refresh_token(
)
if provider_type == ProviderType.KIRO.value:
raw_import = payload.refresh_token.strip()
raw_import = (payload.refresh_token or "").strip()
if not raw_import:
raise InvalidRequestException("Refresh Token 不能为空")
raise InvalidRequestException("Kiro 导入需要 Refresh Token 或完整凭据 JSON")
# 使用统一的解析函数
credentials = _parse_kiro_import_input(raw_import)
@@ -2045,93 +2349,43 @@ async def import_refresh_token(
template = _require_oauth_template(provider_type)
# 用 refresh_token 换取 access_token
refresh_token = payload.refresh_token.strip()
token_url = template.oauth.token_url
is_json = "anthropic.com" in token_url
scope_str = " ".join(template.oauth.scopes) if template.oauth.scopes else ""
if is_json:
body: dict[str, Any] = {
"grant_type": "refresh_token",
"client_id": template.oauth.client_id,
"refresh_token": refresh_token,
}
if scope_str:
body["scope"] = scope_str
headers = {"Content-Type": "application/json", "Accept": "application/json"}
data = None
json_body = body
else:
form: dict[str, str] = {
"grant_type": "refresh_token",
"client_id": template.oauth.client_id,
"refresh_token": refresh_token,
}
if scope_str:
form["scope"] = scope_str
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
# proxy_config 和 key_proxy 已在上方 Kiro 分支之前统一解析
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,
refresh_token, access_token_input = _normalize_single_import_tokens(
refresh_token=payload.refresh_token,
access_token=payload.access_token,
)
if not refresh_token and not access_token_input:
raise InvalidRequestException("缺少 refresh_token 或 access_token")
if access_token_input and provider_type != ProviderType.CODEX.value:
raise InvalidRequestException("仅 Codex 支持 Access Token 导入")
if resp.status_code < 200 or resp.status_code >= 300:
error_reason = f"HTTP {resp.status_code}"
refresh_error: str | None = None
if refresh_token:
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}"
raise InvalidRequestException(f"Refresh Token 验证失败: {error_reason}")
token = resp.json()
access_token = str(token.get("access_token") or "")
new_refresh_token = str(token.get("refresh_token") or "") or refresh_token
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 refresh 返回缺少 access_token")
# 构建 auth_config
auth_config: dict[str, Any] = {
"provider_type": provider_type,
"token_type": token.get("token_type"),
"refresh_token": new_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,
)
access_token, auth_config, expires_at = await _exchange_refresh_token_for_import(
provider_type=provider_type,
template=template,
refresh_token=refresh_token,
proxy_config=proxy_config,
timeout_seconds=30.0,
)
except Exception as exc:
if not access_token_input or provider_type != ProviderType.CODEX.value:
raise
refresh_error = str(exc)
access_token, auth_config, expires_at = await _build_codex_access_token_import_auth_config(
provider_type=provider_type,
access_token=access_token_input,
refresh_token=refresh_token,
proxy_config=proxy_config,
refresh_error=refresh_error,
)
else:
access_token, auth_config, expires_at = await _build_codex_access_token_import_auth_config(
provider_type=provider_type,
access_token=access_token_input,
refresh_token=None,
proxy_config=proxy_config,
)
# 检查是否存在重复的 OAuth 账号(失效账号允许覆盖)
existing_key = _check_duplicate_oauth_account(db, provider_id, auth_config)
@@ -2162,7 +2416,7 @@ async def import_refresh_token(
key_id=str(new_key.id),
provider_type=provider_type,
expires_at=expires_at,
has_refresh_token=bool(new_refresh_token),
has_refresh_token=bool(auth_config.get("refresh_token")),
email=auth_config.get("email"),
replaced=replaced,
)
@@ -2284,10 +2538,6 @@ async def _batch_import_standard_oauth_internal(
raise InvalidRequestException("未找到有效的 Token 数据")
api_formats = _get_provider_api_formats(provider)
token_url = template.oauth.token_url
is_json = "anthropic.com" in token_url
scope_str = " ".join(template.oauth.scopes) if template.oauth.scopes else ""
total = len(import_entries)
results: list[BatchImportResultItem] = [None] * total # type: ignore[list-item]
success_count = 0
@@ -2305,215 +2555,95 @@ async def _batch_import_standard_oauth_internal(
async with sem:
try:
refresh_token = import_entry.get("refresh_token", "")
if not refresh_token or len(refresh_token) < 10:
result_item = BatchImportResultItem(
index=idx,
status="error",
error="Token 无效或过短",
)
failed_count += 1
else:
if is_json:
body: dict[str, Any] = {
"grant_type": "refresh_token",
"client_id": template.oauth.client_id,
"refresh_token": refresh_token,
}
if scope_str:
body["scope"] = scope_str
headers = {
"Content-Type": "application/json",
"Accept": "application/json",
}
data = None
json_body = body
else:
form: dict[str, str] = {
"grant_type": "refresh_token",
"client_id": template.oauth.client_id,
"refresh_token": refresh_token,
}
if scope_str:
form["scope"] = scope_str
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
try:
_release_batch_import_db_connection_before_await(db)
resp = await post_oauth_token(
try:
_release_batch_import_db_connection_before_await(db)
access_token, auth_config, _expires_at = (
await _build_standard_oauth_import_credentials(
provider_type=provider_type,
token_url=token_url,
headers=headers,
data=data,
json_body=json_body,
template=template,
import_entry=import_entry,
proxy_config=proxy_config,
timeout_seconds=timeout_seconds,
)
except Exception as exc:
result_item = BatchImportResultItem(
index=idx,
status="error",
error=f"Token 刷新请求失败: {exc}",
)
failed_count += 1
processed_count += 1
results[idx] = result_item
if progress_hook is not None:
_release_batch_import_db_connection_before_await(db)
await progress_hook(
total, processed_count, success_count, failed_count, result_item
)
return
if resp.status_code < 200 or resp.status_code >= 300:
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}"
)
result_item = BatchImportResultItem(
index=idx,
status="error",
error=f"Token 验证失败: {error_reason}",
)
failed_count += 1
processed_count += 1
results[idx] = result_item
if progress_hook is not None:
_release_batch_import_db_connection_before_await(db)
await progress_hook(
total, processed_count, success_count, failed_count, result_item
)
return
token_data = resp.json()
access_token = str(token_data.get("access_token") or "")
new_refresh_token = str(token_data.get("refresh_token") or "") or refresh_token
if not access_token:
result_item = BatchImportResultItem(
index=idx,
status="error",
error="Token 刷新返回缺少 access_token",
)
failed_count += 1
processed_count += 1
results[idx] = result_item
if progress_hook is not None:
_release_batch_import_db_connection_before_await(db)
await progress_hook(
total, processed_count, success_count, failed_count, result_item
)
return
expires_in = token_data.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
auth_config: dict[str, Any] = {
"provider_type": provider_type,
"token_type": token_data.get("token_type"),
"refresh_token": new_refresh_token or None,
"expires_at": expires_at,
"scope": token_data.get("scope"),
"updated_at": int(time.time()),
}
try:
)
except Exception as exc:
result_item = BatchImportResultItem(
index=idx,
status="error",
error=f"Token 验证失败: {exc}",
)
failed_count += 1
processed_count += 1
results[idx] = result_item
if progress_hook is not None:
_release_batch_import_db_connection_before_await(db)
auth_config = await enrich_auth_config(
provider_type=provider_type,
auth_config=auth_config,
token_response=token_data,
access_token=access_token,
proxy_config=proxy_config,
await progress_hook(
total, processed_count, success_count, failed_count, result_item
)
except Exception as exc:
logger.warning("批量导入: enrich_auth_config 失败 (index={}): {}", idx, exc)
return
if provider_type == ProviderType.CODEX.value:
_apply_codex_import_hints(auth_config, import_entry)
async with db_lock:
try:
existing_key = _check_duplicate_oauth_account(
db, provider_id, auth_config
)
except InvalidRequestException as exc:
result_item = BatchImportResultItem(
index=idx,
status="error",
error=str(exc),
)
failed_count += 1
processed_count += 1
results[idx] = result_item
if progress_hook is not None:
_release_batch_import_db_connection_before_await(db)
await progress_hook(
total,
processed_count,
success_count,
failed_count,
result_item,
)
return
async with db_lock:
try:
existing_key = _check_duplicate_oauth_account(
db, provider_id, auth_config
)
except InvalidRequestException as exc:
result_item = BatchImportResultItem(
index=idx,
status="error",
error=str(exc),
)
failed_count += 1
processed_count += 1
results[idx] = result_item
if progress_hook is not None:
_release_batch_import_db_connection_before_await(db)
await progress_hook(
total,
processed_count,
success_count,
failed_count,
result_item,
)
return
replaced = False
if existing_key:
new_key = _update_existing_oauth_key(
db,
existing_key,
access_token,
auth_config,
flush_only=True,
proxy=key_proxy,
)
name = existing_key.name
replaced = True
replaced = False
if existing_key:
new_key = _update_existing_oauth_key(
db,
existing_key,
access_token,
auth_config,
flush_only=True,
proxy=key_proxy,
)
name = existing_key.name
replaced = True
else:
email = auth_config.get("email")
if email:
name = f"{provider_type}_{email}"
else:
email = auth_config.get("email")
if email:
name = f"{provider_type}_{email}"
else:
name = f"{provider_type}_{int(time.time())}_{idx}"
if len(name) > 100:
name = name[:100]
name = f"{provider_type}_{int(time.time())}_{idx}"
if len(name) > 100:
name = name[:100]
new_key = _create_oauth_key(
db,
provider_id=provider_id,
name=name,
access_token=access_token,
auth_config=auth_config,
api_formats=api_formats,
flush_only=True,
proxy=key_proxy,
)
pending_success_writes += 1
pending_success_writes = _commit_batch_import_writes_if_needed(
db, pending_success_writes
new_key = _create_oauth_key(
db,
provider_id=provider_id,
name=name,
access_token=access_token,
auth_config=auth_config,
api_formats=api_formats,
flush_only=True,
proxy=key_proxy,
)
pending_success_writes += 1
pending_success_writes = _commit_batch_import_writes_if_needed(
db, pending_success_writes
)
result_item = BatchImportResultItem(
index=idx,
status="success",
+3
View File
@@ -311,6 +311,8 @@ def _decode_unverified_jwt_payload(token: str) -> dict[str, Any] | None:
def _extract_codex_fields_from_claims(claims: dict[str, Any]) -> dict[str, Any]:
auth_info = claims.get("https://api.openai.com/auth")
auth = auth_info if isinstance(auth_info, dict) else {}
profile_info = claims.get("https://api.openai.com/profile")
profile = profile_info if isinstance(profile_info, dict) else {}
result: dict[str, Any] = {}
@@ -318,6 +320,7 @@ def _extract_codex_fields_from_claims(claims: dict[str, Any]) -> dict[str, Any]:
[
claims.get("email"),
auth.get("email"),
profile.get("email"),
]
)
if email:
+4
View File
@@ -863,6 +863,10 @@ class EndpointAPIKeyResponse(BaseModel):
default_factory=list,
description="OAuth 关联的组织/工作区摘要列表",
)
oauth_temporary: bool = Field(
default=False,
description="是否为仅 Access Token 导入、不可自动刷新的临时 OAuth 账号",
)
oauth_invalid_at: int | None = Field(
default=None,
description="OAuth Token 失效时间(Unix 时间戳,兼容字段;优先使用 status_snapshot.oauth)",
+7 -1
View File
@@ -516,7 +516,13 @@ async def get_provider_auth(
# Refresh 失败(非锁竞争)且 access token 已过期 → 升级标记为 [OAUTH_EXPIRED]
# 注意:未获取到锁说明其他实例正在刷新,不应在此标记为过期
if should_refresh and not _refreshed and not _lost_lock and expires_at is not None:
if (
should_refresh
and refresh_token
and not _refreshed
and not _lost_lock
and expires_at is not None
):
try:
token_truly_expired = int(time.time()) >= int(expires_at)
except Exception:
@@ -64,6 +64,7 @@ def build_key_response(
oauth_account_user_id = None
auth_config: dict[str, Any] | None = None
oauth_organizations: list[dict[str, object]] = []
oauth_temporary = False
encrypted_auth_config = key_dict.pop("auth_config", None) # 移除敏感字段,避免泄露
if auth_type == "oauth" and isinstance(encrypted_auth_config, str) and encrypted_auth_config:
try:
@@ -81,6 +82,7 @@ def build_key_response(
oauth_account_name = auth_config.get("account_name")
oauth_account_user_id = auth_config.get("account_user_id")
oauth_organizations = normalize_oauth_organizations(auth_config.get("organizations"))
oauth_temporary = bool(auth_config.get("access_token_import_temporary"))
except Exception as e:
logger.error("Failed to decrypt auth_config for key {}: {}", key.id, e)
@@ -154,6 +156,7 @@ def build_key_response(
"oauth_account_name": oauth_account_name,
"oauth_account_user_id": oauth_account_user_id,
"oauth_organizations": oauth_organizations,
"oauth_temporary": oauth_temporary,
"oauth_invalid_at": status_snapshot.oauth.invalid_at,
"oauth_invalid_reason": getattr(key, "oauth_invalid_reason", None),
"status_snapshot": asdict(status_snapshot),
@@ -1,5 +1,7 @@
from __future__ import annotations
import base64
import json
from itertools import count
from types import SimpleNamespace
from unittest.mock import MagicMock
@@ -10,6 +12,14 @@ import pytest
from src.api.admin import provider_oauth as oauthmod
def _unsigned_jwt(payload: dict[str, object]) -> str:
def _encode(data: dict[str, object]) -> str:
raw = json.dumps(data, separators=(",", ":")).encode()
return base64.urlsafe_b64encode(raw).decode().rstrip("=")
return f"{_encode({'alg': 'RS256', 'typ': 'JWT'})}.{_encode(payload)}.signature"
@pytest.mark.asyncio
async def test_standard_batch_import_releases_db_connection_before_network_await(
monkeypatch: pytest.MonkeyPatch,
@@ -228,3 +238,152 @@ async def test_kiro_batch_import_releases_db_connection_before_refresh(
assert result.failed == 1
assert release_calls
db.commit.assert_not_called()
@pytest.mark.asyncio
async def test_codex_batch_import_falls_back_to_access_token_when_refresh_fails(
monkeypatch: pytest.MonkeyPatch,
) -> None:
created_auth_configs: list[dict[str, object]] = []
monkeypatch.setattr(
oauthmod,
"_require_oauth_template",
lambda _provider_type: SimpleNamespace(
oauth=SimpleNamespace(
token_url="https://example.com/oauth/token",
client_id="client-id",
client_secret=None,
scopes=[],
)
),
)
monkeypatch.setattr(
oauthmod,
"_parse_standard_oauth_import_entries",
lambda _raw: [{"refresh_token": "r" * 120, "access_token": "a" * 120}],
)
monkeypatch.setattr(oauthmod, "_get_provider_api_formats", lambda _provider: ["openai:cli"])
monkeypatch.setattr(oauthmod, "_release_batch_import_db_connection_before_await", lambda _db: None)
async def _fake_exchange_refresh_token_for_import(**_kwargs: object) -> tuple[str, dict[str, object], int | None]:
raise RuntimeError("invalid refresh")
async def _fake_access_import(**kwargs: object) -> tuple[str, dict[str, object], int | None]:
auth_config = {
"provider_type": "codex",
"refresh_token": kwargs["refresh_token"],
"expires_at": 2_000_000_000,
"access_token_import_temporary": False,
"refresh_token_import_error": kwargs.get("refresh_error"),
"email": "[email protected]",
}
return str(kwargs["access_token"]), auth_config, 2_000_000_000
monkeypatch.setattr(oauthmod, "_exchange_refresh_token_for_import", _fake_exchange_refresh_token_for_import)
monkeypatch.setattr(oauthmod, "_build_codex_access_token_import_auth_config", _fake_access_import)
monkeypatch.setattr(oauthmod, "_check_duplicate_oauth_account", lambda *_args, **_kwargs: None)
def _fake_create_oauth_key(*_args: object, **kwargs: object) -> SimpleNamespace:
created_auth_configs.append(dict(kwargs["auth_config"]))
return SimpleNamespace(id="key-1")
monkeypatch.setattr(oauthmod, "_create_oauth_key", _fake_create_oauth_key)
db = MagicMock()
result = await oauthmod._batch_import_standard_oauth_internal(
provider_id="provider-1",
provider_type="codex",
provider=SimpleNamespace(endpoints=[]), # type: ignore[arg-type]
raw_credentials="ignored",
db=db,
concurrency=1,
)
assert result.success == 1
assert created_auth_configs[0]["refresh_token"] == "r" * 120
assert created_auth_configs[0]["access_token_import_temporary"] is False
assert "invalid refresh" in str(created_auth_configs[0]["refresh_token_import_error"])
def test_standard_oauth_import_parses_plain_jwt_as_access_token() -> None:
token = _unsigned_jwt(
{
"iss": "https://auth.openai.com",
"aud": ["https://api.openai.com/v1"],
"exp": 2_000_000_000,
"scp": ["openid", "offline_access"],
}
)
entries = oauthmod._parse_standard_oauth_import_entries(token)
assert entries == [{"access_token": token}]
@pytest.mark.asyncio
async def test_codex_plain_access_token_import_does_not_refresh(
monkeypatch: pytest.MonkeyPatch,
) -> None:
token = _unsigned_jwt(
{
"iss": "https://auth.openai.com",
"aud": ["https://api.openai.com/v1"],
"exp": 2_000_000_000,
"https://api.openai.com/profile": {"email": "[email protected]"},
}
)
created_auth_configs: list[dict[str, object]] = []
refresh_calls = 0
monkeypatch.setattr(
oauthmod,
"_require_oauth_template",
lambda _provider_type: SimpleNamespace(
oauth=SimpleNamespace(
token_url="https://example.com/oauth/token",
client_id="client-id",
client_secret=None,
scopes=[],
)
),
)
monkeypatch.setattr(oauthmod, "_get_provider_api_formats", lambda _provider: ["openai:cli"])
monkeypatch.setattr(oauthmod, "_release_batch_import_db_connection_before_await", lambda _db: None)
async def _fake_exchange_refresh_token_for_import(**_kwargs: object) -> tuple[str, dict[str, object], int | None]:
nonlocal refresh_calls
refresh_calls += 1
raise AssertionError("Access Token must not be refreshed as Refresh Token")
async def _fake_enrich_auth_config(**kwargs: object) -> dict[str, object]:
auth_config = dict(kwargs["auth_config"]) # type: ignore[call-overload]
auth_config["email"] = "[email protected]"
return auth_config
monkeypatch.setattr(oauthmod, "_exchange_refresh_token_for_import", _fake_exchange_refresh_token_for_import)
monkeypatch.setattr(oauthmod, "enrich_auth_config", _fake_enrich_auth_config)
monkeypatch.setattr(oauthmod, "_check_duplicate_oauth_account", lambda *_args, **_kwargs: None)
def _fake_create_oauth_key(*_args: object, **kwargs: object) -> SimpleNamespace:
created_auth_configs.append(dict(kwargs["auth_config"]))
return SimpleNamespace(id="key-1")
monkeypatch.setattr(oauthmod, "_create_oauth_key", _fake_create_oauth_key)
db = MagicMock()
result = await oauthmod._batch_import_standard_oauth_internal(
provider_id="provider-1",
provider_type="codex",
provider=SimpleNamespace(endpoints=[]), # type: ignore[arg-type]
raw_credentials=token,
db=db,
concurrency=1,
)
assert result.success == 1
assert refresh_calls == 0
assert created_auth_configs[0]["refresh_token"] is None
assert created_auth_configs[0]["access_token_import_temporary"] is True
@@ -45,6 +45,25 @@ def test_parse_codex_id_token_extracts_auth_claim_fields() -> None:
}
def test_parse_codex_id_token_extracts_profile_email() -> None:
token = _encode_unsigned_jwt(
{
"https://api.openai.com/profile": {
"email": "[email protected]",
"email_verified": True,
},
"https://api.openai.com/auth": {
"chatgpt_account_id": "acc-1",
},
}
)
parsed = parse_codex_id_token(token)
assert parsed["email"] == "[email protected]"
assert parsed["account_id"] == "acc-1"
def test_parse_codex_id_token_accepts_json_payload_string() -> None:
payload = {
"email": "[email protected]",
@@ -1,6 +1,8 @@
from __future__ import annotations
import json
import jwt
from contextlib import contextmanager
from types import SimpleNamespace
from unittest.mock import MagicMock
@@ -261,3 +263,63 @@ async def test_refresh_account_state_after_oauth_update_returns_error_when_refre
assert attempted is True
assert "quota refresh failed" in error
fake_db.close.assert_called_once()
def test_parse_standard_oauth_import_entries_accepts_codex_access_token_only() -> None:
entries = module._parse_standard_oauth_import_entries(
'[{"access_token":"at_1","accountId":"acc-1","email":"[email protected]"}]'
)
assert entries == [
{
"access_token": "at_1",
"account_id": "acc-1",
"email": "[email protected]",
}
]
def test_parse_standard_oauth_import_entries_keeps_refresh_and_access_token() -> None:
entries = module._parse_standard_oauth_import_entries(
'{"refreshToken":"rt_1","accessToken":"at_1","planType":"PLUS"}'
)
assert entries == [
{
"refresh_token": "rt_1",
"access_token": "at_1",
"plan_type": "plus",
}
]
def test_decode_access_token_expires_at_reads_jwt_exp() -> None:
token = jwt.encode({"exp": 2_000_000_000}, key="", algorithm="none")
if isinstance(token, bytes):
token = token.decode("utf-8")
assert module._decode_access_token_expires_at(token) == 2_000_000_000
def test_normalize_single_import_tokens_treats_plain_jwt_as_access_token() -> None:
token = jwt.encode(
{
"iss": "https://auth.openai.com",
"aud": ["https://api.openai.com/v1"],
"exp": 2_000_000_000,
"scp": ["openid", "offline_access"],
},
key="",
algorithm="none",
)
if isinstance(token, bytes):
token = token.decode("utf-8")
token = f"{token}signature"
refresh_token, access_token = module._normalize_single_import_tokens(
refresh_token=token,
access_token=None,
)
assert refresh_token == ""
assert access_token == token