mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-11 19:59:50 +08:00
优化 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:
@@ -204,6 +204,7 @@ logs/
|
||||
|
||||
# Git backup
|
||||
.git.backup/
|
||||
.worktrees/
|
||||
|
||||
# Database backups
|
||||
backups/
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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
@@ -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",
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)",
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user