feat(oauth): 账号封禁前置 OAuth 验证、抽取 provider_context、完善账号状态分类

- 新增 verify_oauth_before_account_block:在标记账号封禁前先尝试刷新 token,
  区分 OAuth 过期与真正的账号级封禁,避免误标
- 抽取 provider_context.py 统一解析 provider_type,解决 ORM detached 访问问题
- account_state 新增 workspace_deactivated 分类和 auto-removable 状态集合,
  补充中文验证关键词匹配
- OAuth refresh 成功后仅清除可恢复的 token 错误,不再自动清除账号级 block
- deploy.sh 依赖指纹改用纯 shell 实现,移除对 Python tomllib 的依赖
- 前端 Pool 管理页面新增筛选和批量操作优化
- 补充对应测试用例
This commit is contained in:
fawney19
2026-03-20 16:50:59 +08:00
parent aa83b4a7a7
commit 25d38ae632
44 changed files with 1527 additions and 370 deletions
+19 -20
View File
@@ -108,27 +108,26 @@ if [ -n "$HUB_TAG" ]; then
esac esac
fi fi
# 提取 pyproject.toml 中"会影响运行时依赖安装"的最小指纹(与 CI 保持一致): # 提取 pyproject.toml 中会影响运行时依赖安装的字段指纹(纯 shell,无需 Python)
# - [build-system] requires / build-backend # 用 sed 提取 dependencies / requires 数组块和单值字段,排序后输出稳定文本
# - [project] requires-python / dependencies
# 使用 Python tomllib 解析,不受 TOML 格式变化影响。
pyproject_deps_fingerprint() { pyproject_deps_fingerprint() {
python3 - <<'PY' local file="pyproject.toml"
import json, pathlib, tomllib # 提取 "key = [..." 多行数组块(从 key 行到 ] 行)
extract_array() {
data = tomllib.loads(pathlib.Path("pyproject.toml").read_text("utf-8")) sed -n "/^$1[[:space:]]*=[[:space:]]*\[/,/\]/p" "$file" | grep '"' | sed 's/.*"\(.*\)".*/\1/' | sort
project = data.get("project") or {} }
build = data.get("build-system") or {} # 提取 "key = "value"" 单行值
extract_value() {
fingerprint = { grep -m1 "^$1[[:space:]]*=" "$file" 2>/dev/null | sed 's/.*"\(.*\)".*/\1/'
"requires-python": project.get("requires-python"), }
"dependencies": sorted(project.get("dependencies") or []), {
"build-backend": build.get("build-backend"), echo "requires-python=$(extract_value requires-python)"
"build-requires": sorted(build.get("requires") or []), echo "build-backend=$(extract_value build-backend)"
} echo "dependencies:"
extract_array dependencies
print(json.dumps(fingerprint, sort_keys=True, separators=(",", ":"))) echo "build-requires:"
PY extract_array requires
}
} }
# 计算依赖文件的哈希值(包含 Dockerfile.base.local) # 计算依赖文件的哈希值(包含 Dockerfile.base.local)
+6
View File
@@ -106,6 +106,12 @@ export interface PoolKeyDetail {
oauth_account_user_id?: string | null oauth_account_user_id?: string | null
oauth_account_name?: string | null oauth_account_name?: string | null
oauth_organizations?: OAuthOrganizationInfo[] | null oauth_organizations?: OAuthOrganizationInfo[] | null
account_status_code?: string | null
account_status_label?: string | null
account_status_reason?: string | null
account_status_blocked?: boolean
account_status_recoverable?: boolean
account_status_source?: string | null
quota_updated_at?: number | null quota_updated_at?: number | null
health_score?: number health_score?: number
circuit_breaker_open?: boolean circuit_breaker_open?: boolean
@@ -18,6 +18,8 @@ export interface ProviderOAuthCompleteResponse {
expires_at?: number | null expires_at?: number | null
has_refresh_token: boolean has_refresh_token: boolean
email?: string | null email?: string | null
account_state_recheck_attempted?: boolean
account_state_recheck_error?: string | null
} }
export interface ProviderOAuthCompleteResponseWithKey { export interface ProviderOAuthCompleteResponseWithKey {
@@ -107,15 +107,15 @@
<div class="flex items-center gap-1.5"> <div class="flex items-center gap-1.5">
<span class="text-xs font-medium truncate">{{ key.key_name || '未命名' }}</span> <span class="text-xs font-medium truncate">{{ key.key_name || '未命名' }}</span>
<Badge <Badge
v-if="isOAuthInvalid(key)"
variant="destructive"
class="text-[10px] px-1 py-0 h-4 shrink-0"
>OAuth失效</Badge>
<Badge
v-else
variant="outline" variant="outline"
class="text-[10px] px-1 py-0 h-4 shrink-0" class="text-[10px] px-1 py-0 h-4 shrink-0"
>{{ normalizeAuthTypeLabel(key.auth_type) }}</Badge> >{{ normalizeAuthTypeLabel(key.auth_type) }}</Badge>
<Badge
v-if="getStatusBadgeLabel(key)"
variant="destructive"
class="text-[10px] px-1 py-0 h-4 shrink-0"
:title="getStatusBadgeTitle(key)"
>{{ getStatusBadgeLabel(key) }}</Badge>
<Badge <Badge
v-if="key.oauth_plan_type" v-if="key.oauth_plan_type"
variant="outline" variant="outline"
@@ -127,11 +127,6 @@
class="text-[10px] px-1 py-0 h-4 shrink-0" class="text-[10px] px-1 py-0 h-4 shrink-0"
:title="getOAuthOrgBadge(key)?.title" :title="getOAuthOrgBadge(key)?.title"
>{{ getOAuthOrgBadge(key)?.label }}</Badge> >{{ getOAuthOrgBadge(key)?.label }}</Badge>
<Badge
v-if="isBannedKey(key)"
variant="destructive"
class="text-[10px] px-1 py-0 h-4 shrink-0"
>封号</Badge>
</div> </div>
<div class="flex items-center gap-1.5 mt-0.5 text-[11px] text-muted-foreground flex-wrap"> <div class="flex items-center gap-1.5 mt-0.5 text-[11px] text-muted-foreground flex-wrap">
<span :class="key.is_active ? '' : 'text-destructive'">{{ key.is_active ? '启用' : '禁用' }}</span> <span :class="key.is_active ? '' : 'text-destructive'">{{ key.is_active ? '启用' : '禁用' }}</span>
@@ -283,6 +278,7 @@ import {
import { exportKey, refreshProviderQuota } from '@/api/endpoints/keys' import { exportKey, refreshProviderQuota } from '@/api/endpoints/keys'
import { refreshProviderOAuth } from '@/api/endpoints/provider_oauth' import { refreshProviderOAuth } from '@/api/endpoints/provider_oauth'
import { useProxyNodesStore } from '@/stores/proxy-nodes' import { useProxyNodesStore } from '@/stores/proxy-nodes'
import { classifyAccountBlockLabel, cleanAccountBlockReason, isAccountLevelBlockReason, isRefreshFailedReason } from '@/utils/accountBlock'
import { getOAuthOrgBadge } from '@/utils/oauthIdentity' import { getOAuthOrgBadge } from '@/utils/oauthIdentity'
type QuickSelectorValue = type QuickSelectorValue =
@@ -321,12 +317,12 @@ const emit = defineEmits<{
}>() }>()
const QUICK_SELECT_OPTIONS: Array<{ value: QuickSelectorValue; label: string }> = [ const QUICK_SELECT_OPTIONS: Array<{ value: QuickSelectorValue; label: string }> = [
{ value: 'banned', label: '已封号' }, { value: 'banned', label: '账号异常' },
{ value: 'no_5h_limit', label: '无5H限额' }, { value: 'no_5h_limit', label: '无5H限额' },
{ value: 'no_weekly_limit', label: '无周限额' }, { value: 'no_weekly_limit', label: '无周限额' },
{ value: 'plan_free', label: '全部 Free' }, { value: 'plan_free', label: '全部 Free' },
{ value: 'plan_team', label: '全部 Team' }, { value: 'plan_team', label: '全部 Team' },
{ value: 'oauth_invalid', label: 'OAuth 失效' }, { value: 'oauth_invalid', label: 'Token 异常' },
{ value: 'proxy_unset', label: '未配置代理' }, { value: 'proxy_unset', label: '未配置代理' },
{ value: 'proxy_set', label: '已配置独立代理' }, { value: 'proxy_set', label: '已配置独立代理' },
{ value: 'disabled', label: '已禁用' }, { value: 'disabled', label: '已禁用' },
@@ -426,25 +422,43 @@ function normalizeAuthTypeLabel(authType: string): string {
return 'API Key' return 'API Key'
} }
function isBannedKey(key: PoolKeyDetail): boolean { function getStatusBadgeLabel(key: PoolKeyDetail): string | null {
const reason = normalizeText(key.oauth_invalid_reason) const explicitLabel = String(key.account_status_label || '').trim()
if (reason && /(banned|forbidden|blocked|suspend|封|禁|受限)/.test(reason)) return true if (explicitLabel) return explicitLabel
if (Array.isArray(key.scheduling_reasons)) {
return key.scheduling_reasons.some((item) => { const reason = String(key.oauth_invalid_reason || '').trim()
const code = normalizeText(item.code) if (isAccountLevelBlockReason(reason)) {
return code === 'account_banned' || code === 'account_forbidden' || code === 'account_blocked' const cleaned = cleanAccountBlockReason(reason)
}) return classifyAccountBlockLabel(cleaned || reason)
} }
return false
if (normalizeText(key.auth_type) !== 'oauth') return null
if (isRefreshFailedReason(reason)) return '续期失败'
if (key.oauth_invalid_at != null || normalizeText(reason)) return 'Token 失效'
if (typeof key.oauth_expires_at === 'number' && key.oauth_expires_at > 0) {
return key.oauth_expires_at * 1000 <= Date.now() ? 'Token 过期' : null
}
return null
} }
function isOAuthInvalid(key: PoolKeyDetail): boolean { function getStatusBadgeTitle(key: PoolKeyDetail): string {
if (normalizeText(key.auth_type) !== 'oauth') return false const label = getStatusBadgeLabel(key)
if (key.oauth_invalid_at != null || normalizeText(key.oauth_invalid_reason)) return true if (!label) return ''
if (typeof key.oauth_expires_at === 'number' && key.oauth_expires_at > 0) {
return key.oauth_expires_at * 1000 <= Date.now() const explicitReason = String(key.account_status_reason || '').trim()
if (explicitReason) return `${label}: ${explicitReason}`
const reason = String(key.oauth_invalid_reason || '').trim()
if (!reason) return label
if (isAccountLevelBlockReason(reason)) {
const cleaned = cleanAccountBlockReason(reason)
return cleaned ? `${label}: ${cleaned}` : label
} }
return false if (isRefreshFailedReason(reason)) {
const cleaned = reason.replace(/^\[REFRESH_FAILED\]\s*/i, '').trim()
return cleaned ? `${label}: ${cleaned}` : label
}
return `${label}: ${reason}`
} }
function formatRelativeTime(value: string): string { function formatRelativeTime(value: string): string {
@@ -28,7 +28,7 @@
<div class="space-y-0.5"> <div class="space-y-0.5">
<span class="text-sm font-medium">主动探测</span> <span class="text-sm font-medium">主动探测</span>
<p class="text-xs text-muted-foreground"> <p class="text-xs text-muted-foreground">
定期检查 Key 可用性,提前发现异常 按固定间隔主动刷新 Key 的账号状态与额度
</p> </p>
</div> </div>
<Switch <Switch
@@ -57,9 +57,9 @@
</div> </div>
<div class="flex items-center justify-between p-3 border rounded-lg bg-muted/50"> <div class="flex items-center justify-between p-3 border rounded-lg bg-muted/50">
<div class="space-y-0.5"> <div class="space-y-0.5">
<span class="text-sm font-medium">封号自动清除</span> <span class="text-sm font-medium">异常自动清除</span>
<p class="text-xs text-muted-foreground"> <p class="text-xs text-muted-foreground">
检测到账号被封禁时自动从号池中移除 仅在检测到不可恢复的账号异常时自动从号池中移除,不处理纯 Token 失效
</p> </p>
</div> </div>
<Switch <Switch
@@ -1158,7 +1158,7 @@ const emit = defineEmits<{
(e: 'refresh'): void (e: 'refresh'): void
}>() }>()
const { error: showError, success: showSuccess } = useToast() const { error: showError, success: showSuccess, warning: showWarning } = useToast()
const { confirm } = useConfirm() const { confirm } = useConfirm()
const { copyToClipboard } = useClipboard() const { copyToClipboard } = useClipboard()
const { tick: countdownTick, start: startCountdownTimer, stop: stopCountdownTimer } = useCountdownTimer() const { tick: countdownTick, start: startCountdownTimer, stop: stopCountdownTimer } = useCountdownTimer()
@@ -1643,7 +1643,15 @@ async function handleRefreshOAuth(key: EndpointAPIKey) {
refreshingOAuthKeyId.value = key.id refreshingOAuthKeyId.value = key.id
try { try {
const result = await refreshProviderOAuth(key.id) const result = await refreshProviderOAuth(key.id)
showSuccess('Token 刷新成功') if (result.account_state_recheck_attempted) {
if (result.account_state_recheck_error) {
showWarning('Token 刷新成功,但账号状态复检失败')
} else {
showSuccess('Token 刷新成功,已复检账号状态')
}
} else {
showSuccess('Token 刷新成功')
}
// 更新本地数据 // 更新本地数据
const keyInList = providerKeys.value.find(k => k.id === key.id) const keyInList = providerKeys.value.find(k => k.id === key.id)
if (keyInList) { if (keyInList) {
@@ -1678,7 +1686,7 @@ async function handleClearOAuthInvalid(key: EndpointAPIKey) {
const confirmed = await confirm({ const confirmed = await confirm({
title: '清除账号异常标记', title: '清除账号异常标记',
message: `确认账号 "${key.name || key.id.slice(0, 8)}" 已手动完成验证?清除后该 Key 将恢复正常调度。`, message: `确认账号 "${key.name || key.id.slice(0, 8)}" 已手动完成验证?清除后系统会按当前手动开关和调度状态重新评估该 Key。`,
confirmText: '确认清除', confirmText: '确认清除',
variant: 'default', variant: 'default',
}) })
@@ -1687,13 +1695,12 @@ async function handleClearOAuthInvalid(key: EndpointAPIKey) {
clearingOAuthInvalidKeyId.value = key.id clearingOAuthInvalidKeyId.value = key.id
try { try {
await clearOAuthInvalid(key.id) await clearOAuthInvalid(key.id)
showSuccess('已清除 OAuth 异常标记,Key 已自动启用') showSuccess('已清除 OAuth 异常标记')
// 更新本地数据 // 更新本地数据
const keyInList = providerKeys.value.find(k => k.id === key.id) const keyInList = providerKeys.value.find(k => k.id === key.id)
if (keyInList) { if (keyInList) {
keyInList.oauth_invalid_at = null keyInList.oauth_invalid_at = null
keyInList.oauth_invalid_reason = null keyInList.oauth_invalid_reason = null
keyInList.is_active = true
} }
await loadEndpoints() await loadEndpoints()
} catch (err: unknown) { } catch (err: unknown) {
+3
View File
@@ -37,6 +37,9 @@ const KEYWORDS_TOKEN_INVALID = [
const KEYWORDS_VERIFICATION = [ const KEYWORDS_VERIFICATION = [
'validation_required', 'validation_required',
'verify your account', 'verify your account',
'需要验证',
'验证账号',
'验证身份',
] ]
// 合并的完整列表 // 合并的完整列表
+26 -8
View File
@@ -127,7 +127,7 @@
全部 全部
</SelectItem> </SelectItem>
<SelectItem value="active"> <SelectItem value="active">
活跃 可调度
</SelectItem> </SelectItem>
<SelectItem value="cooldown"> <SelectItem value="cooldown">
冷却中 冷却中
@@ -213,7 +213,7 @@
全部状态 全部状态
</SelectItem> </SelectItem>
<SelectItem value="active"> <SelectItem value="active">
活跃 可调度
</SelectItem> </SelectItem>
<SelectItem value="cooldown"> <SelectItem value="cooldown">
冷却中 冷却中
@@ -2006,10 +2006,16 @@ async function handleRefreshOAuth(key: PoolKeyDetail) {
const target = keyPage.value.keys.find(k => k.key_id === key.key_id) const target = keyPage.value.keys.find(k => k.key_id === key.key_id)
if (target) { if (target) {
target.oauth_expires_at = result.expires_at ?? null target.oauth_expires_at = result.expires_at ?? null
target.oauth_invalid_at = null
target.oauth_invalid_reason = null
} }
success('Token 刷新成功') if (result.account_state_recheck_attempted) {
if (result.account_state_recheck_error) {
showWarning('Token 刷新成功,但账号状态复检失败')
} else {
success('Token 刷新成功,已复检账号状态')
}
} else {
success('Token 刷新成功')
}
await loadKeys() await loadKeys()
} catch (err) { } catch (err) {
showError(parseApiError(err, 'Token 刷新失败')) showError(parseApiError(err, 'Token 刷新失败'))
@@ -2359,6 +2365,11 @@ function getOAuthStatusTitle(key: PoolKeyDetail): string {
const status = getKeyOAuthExpires(key) const status = getKeyOAuthExpires(key)
if (!status) return '' if (!status) return ''
if (status.isInvalid) { if (status.isInvalid) {
const accountLabel = String(key.account_status_label || '').trim()
const accountReason = String(key.account_status_reason || '').trim()
if (accountLabel) {
return accountReason ? `${accountLabel}: ${accountReason}` : accountLabel
}
const cleaned = status.invalidReason && isAccountLevelBlockReason(status.invalidReason) const cleaned = status.invalidReason && isAccountLevelBlockReason(status.invalidReason)
? cleanAccountBlockReason(status.invalidReason) ? cleanAccountBlockReason(status.invalidReason)
: status.invalidReason : status.invalidReason
@@ -2377,11 +2388,15 @@ function getAccountAlertLabel(key: PoolKeyDetail): string | null {
if (cached !== undefined) return cached if (cached !== undefined) return cached
let result: string | null = null let result: string | null = null
const explicitLabel = String(key.account_status_label || '').trim()
if (key.account_status_blocked && explicitLabel) {
result = explicitLabel
}
const quotaText = String(key.account_quota || '').trim() const quotaText = String(key.account_quota || '').trim()
// 后端 _build_account_quota 返回的确切文本: "账号已封禁" / "访问受限" // 后端 _build_account_quota 返回的确切文本: "账号已封禁" / "访问受限"
if (quotaText === '账号已封禁' || quotaText === '封禁') result = '账号封禁' if (!result && (quotaText === '账号已封禁' || quotaText === '封禁')) result = '账号封禁'
else if (quotaText === '访问受限') result = '访问受限' else if (!result && quotaText === '访问受限') result = '访问受限'
else if (isAccountLevelBlockReason(key.oauth_invalid_reason)) { else if (!result && isAccountLevelBlockReason(key.oauth_invalid_reason)) {
const reason = String(key.oauth_invalid_reason || '').trim() const reason = String(key.oauth_invalid_reason || '').trim()
const cleaned = cleanAccountBlockReason(reason) const cleaned = cleanAccountBlockReason(reason)
result = classifyAccountBlockLabel(cleaned || reason) result = classifyAccountBlockLabel(cleaned || reason)
@@ -2395,6 +2410,9 @@ function getAccountAlertTitle(key: PoolKeyDetail): string {
const label = getAccountAlertLabel(key) const label = getAccountAlertLabel(key)
if (!label) return '' if (!label) return ''
const explicitReason = String(key.account_status_reason || '').trim()
if (explicitReason) return `${label}: ${explicitReason}`
const reason = String(key.oauth_invalid_reason || '').trim() const reason = String(key.oauth_invalid_reason || '').trim()
if (reason) { if (reason) {
if (isAccountLevelBlockReason(reason)) { if (isAccountLevelBlockReason(reason)) {
+110 -125
View File
@@ -82,9 +82,7 @@ async def pool_overview(
) -> PoolOverviewResponse: ) -> PoolOverviewResponse:
"""Return all pool-enabled providers with summary stats.""" """Return all pool-enabled providers with summary stats."""
adapter = AdminPoolOverviewAdapter() adapter = AdminPoolOverviewAdapter()
return await pipeline.run( return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
adapter=adapter, http_request=request, db=db, mode=adapter.mode
)
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -109,9 +107,7 @@ async def list_scheduling_presets(
"""Return scheduling preset definitions for frontend rendering.""" """Return scheduling preset definitions for frontend rendering."""
adapter = AdminListSchedulingPresetsAdapter() adapter = AdminListSchedulingPresetsAdapter()
return await pipeline.run( return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
adapter=adapter, http_request=request, db=db, mode=adapter.mode
)
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -143,9 +139,7 @@ async def list_pool_keys(
quick_selectors=quick_selectors.split(",") if quick_selectors else [], quick_selectors=quick_selectors.split(",") if quick_selectors else [],
search_scope=search_scope, search_scope=search_scope,
) )
return await pipeline.run( return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
adapter=adapter, http_request=request, db=db, mode=adapter.mode
)
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -162,9 +156,7 @@ async def batch_import_keys(
) -> BatchImportResponse: ) -> BatchImportResponse:
"""Batch import keys into a provider's pool.""" """Batch import keys into a provider's pool."""
adapter = AdminBatchImportKeysAdapter(provider_id=provider_id, body=body) adapter = AdminBatchImportKeysAdapter(provider_id=provider_id, body=body)
return await pipeline.run( return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
adapter=adapter, http_request=request, db=db, mode=adapter.mode
)
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -195,9 +187,7 @@ def _iter_batches(items: list[str], batch_size: int) -> list[list[str]]:
def _resolve_delete_batch_size(db: Session) -> int: def _resolve_delete_batch_size(db: Session) -> int:
try: try:
bind = db.get_bind() bind = db.get_bind()
dialect_name = str( dialect_name = str(getattr(getattr(bind, "dialect", None), "name", "") or "").lower()
getattr(getattr(bind, "dialect", None), "name", "") or ""
).lower()
except Exception: except Exception:
dialect_name = "" dialect_name = ""
@@ -211,6 +201,7 @@ _COOLDOWN_REASON_LABELS: dict[str, str] = {
"forbidden_403": "403 禁止", "forbidden_403": "403 禁止",
"overloaded_529": "529 过载", "overloaded_529": "529 过载",
"auth_failed_401": "401 认证失败", "auth_failed_401": "401 认证失败",
"account_deactivated_401": "401 账号停用",
"payment_required_402": "402 欠费", "payment_required_402": "402 欠费",
"server_error_500": "500 错误", "server_error_500": "500 错误",
"request_timeout_408": "408 超时", "request_timeout_408": "408 超时",
@@ -244,14 +235,17 @@ def _serialize_money(value: Any) -> str:
def _is_known_banned_key(key: ProviderAPIKey, provider_type: str) -> bool: def _is_known_banned_key(key: ProviderAPIKey, provider_type: str) -> bool:
from src.services.provider.pool.account_state import resolve_pool_account_state from src.services.provider.pool.account_state import (
resolve_pool_account_state,
should_auto_remove_account_state,
)
state = resolve_pool_account_state( state = resolve_pool_account_state(
provider_type=provider_type, provider_type=provider_type,
upstream_metadata=getattr(key, "upstream_metadata", None), upstream_metadata=getattr(key, "upstream_metadata", None),
oauth_invalid_reason=getattr(key, "oauth_invalid_reason", None), oauth_invalid_reason=getattr(key, "oauth_invalid_reason", None),
) )
return state.blocked return should_auto_remove_account_state(state)
def _build_account_quota(provider_type: str, upstream_metadata: Any) -> str | None: def _build_account_quota(provider_type: str, upstream_metadata: Any) -> str | None:
@@ -313,11 +307,7 @@ def _derive_oauth_expires_at(
if str(getattr(key, "auth_type", "") or "").strip().lower() != "oauth": if str(getattr(key, "auth_type", "") or "").strip().lower() != "oauth":
return None return None
cfg = ( cfg = auth_config if isinstance(auth_config, dict) else _extract_oauth_auth_config(key)
auth_config
if isinstance(auth_config, dict)
else _extract_oauth_auth_config(key)
)
if cfg: if cfg:
for field in ("expires_at", "expiresAt", "expiry", "exp"): for field in ("expires_at", "expiresAt", "expiry", "exp"):
expires_at = _normalize_oauth_expires_at(cfg.get(field)) expires_at = _normalize_oauth_expires_at(cfg.get(field))
@@ -337,9 +327,7 @@ def _derive_oauth_plan_type(
auth_config: dict[str, Any] | None = None, auth_config: dict[str, Any] | None = None,
) -> str | None: ) -> str | None:
# Prefer persisted normalized field # Prefer persisted normalized field
persisted = _normalize_oauth_plan_type( persisted = _normalize_oauth_plan_type(getattr(key, "oauth_plan_type", None), provider_type)
getattr(key, "oauth_plan_type", None), provider_type
)
if persisted: if persisted:
return persisted return persisted
@@ -347,11 +335,7 @@ def _derive_oauth_plan_type(
return None return None
# Fallback 1: encrypted auth_config (common for Codex/Antigravity) # Fallback 1: encrypted auth_config (common for Codex/Antigravity)
cfg = ( cfg = auth_config if isinstance(auth_config, dict) else _extract_oauth_auth_config(key)
auth_config
if isinstance(auth_config, dict)
else _extract_oauth_auth_config(key)
)
if cfg: if cfg:
for plan_key in ("plan_type", "tier", "plan", "subscription_plan"): for plan_key in ("plan_type", "tier", "plan", "subscription_plan"):
normalized = _normalize_oauth_plan_type(cfg.get(plan_key), provider_type) normalized = _normalize_oauth_plan_type(cfg.get(plan_key), provider_type)
@@ -430,9 +414,7 @@ def _compute_health_aggregate(
) -> tuple[float, bool]: ) -> tuple[float, bool]:
"""从按格式健康数据聚合出列表展示字段。""" """从按格式健康数据聚合出列表展示字段。"""
health_map = health_by_format if isinstance(health_by_format, dict) else {} health_map = health_by_format if isinstance(health_by_format, dict) else {}
circuit_map = ( circuit_map = circuit_breaker_by_format if isinstance(circuit_breaker_by_format, dict) else {}
circuit_breaker_by_format if isinstance(circuit_breaker_by_format, dict) else {}
)
if health_map: if health_map:
scores = [ scores = [
@@ -445,9 +427,7 @@ def _compute_health_aggregate(
health_score = 1.0 health_score = 1.0
any_circuit_open = any( any_circuit_open = any(
bool(item.get("open", False)) bool(item.get("open", False)) for item in circuit_map.values() if isinstance(item, dict)
for item in circuit_map.values()
if isinstance(item, dict)
) )
return health_score, any_circuit_open return health_score, any_circuit_open
@@ -536,14 +516,10 @@ async def batch_action_keys(
) -> BatchActionResponse: ) -> BatchActionResponse:
"""Batch enable/disable/delete/clear_cooldown/reset_cost/regenerate_fingerprint on pool keys.""" """Batch enable/disable/delete/clear_cooldown/reset_cost/regenerate_fingerprint on pool keys."""
adapter = AdminBatchActionKeysAdapter(provider_id=provider_id, body=body) adapter = AdminBatchActionKeysAdapter(provider_id=provider_id, body=body)
return await pipeline.run( return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
adapter=adapter, http_request=request, db=db, mode=adapter.mode
)
@router.post( @router.post("/{provider_id}/keys/resolve-selection", response_model=PoolKeySelectionResponse)
"/{provider_id}/keys/resolve-selection", response_model=PoolKeySelectionResponse
)
async def resolve_pool_key_selection( async def resolve_pool_key_selection(
provider_id: str, provider_id: str,
body: PoolKeySelectionRequest, body: PoolKeySelectionRequest,
@@ -552,9 +528,7 @@ async def resolve_pool_key_selection(
) -> PoolKeySelectionResponse: ) -> PoolKeySelectionResponse:
"""Resolve all key ids matching the current batch dialog filters.""" """Resolve all key ids matching the current batch dialog filters."""
adapter = AdminResolvePoolKeySelectionAdapter(provider_id=provider_id, body=body) adapter = AdminResolvePoolKeySelectionAdapter(provider_id=provider_id, body=body)
return await pipeline.run( return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
adapter=adapter, http_request=request, db=db, mode=adapter.mode
)
@router.get( @router.get(
@@ -568,12 +542,8 @@ async def get_batch_delete_task_status(
db: Session = Depends(get_db), db: Session = Depends(get_db),
) -> BatchDeleteTaskResponse: ) -> BatchDeleteTaskResponse:
"""Query the progress of an async batch-delete task.""" """Query the progress of an async batch-delete task."""
adapter = AdminBatchDeleteTaskStatusAdapter( adapter = AdminBatchDeleteTaskStatusAdapter(provider_id=provider_id, task_id=task_id)
provider_id=provider_id, task_id=task_id return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
)
return await pipeline.run(
adapter=adapter, http_request=request, db=db, mode=adapter.mode
)
@router.post("/{provider_id}/keys/cleanup-banned", response_model=BatchActionResponse) @router.post("/{provider_id}/keys/cleanup-banned", response_model=BatchActionResponse)
@@ -582,11 +552,9 @@ async def cleanup_banned_keys(
request: Request, request: Request,
db: Session = Depends(get_db), db: Session = Depends(get_db),
) -> BatchActionResponse: ) -> BatchActionResponse:
"""Delete known banned/suspended accounts for the provider.""" """Delete known hard-blocked abnormal accounts for the provider."""
adapter = AdminCleanupBannedKeysAdapter(provider_id=provider_id) adapter = AdminCleanupBannedKeysAdapter(provider_id=provider_id)
return await pipeline.run( return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
adapter=adapter, http_request=request, db=db, mode=adapter.mode
)
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -658,9 +626,7 @@ class AdminPoolOverviewAdapter(AdminApiAdapter):
db = context.db db = context.db
providers = ( providers = (
db.query(Provider) db.query(Provider)
.options( .options(load_only(*cast(tuple[Any, ...], _PROVIDER_OVERVIEW_LOAD_ONLY_ATTRS)))
load_only(*cast(tuple[Any, ...], _PROVIDER_OVERVIEW_LOAD_ONLY_ATTRS))
)
.order_by(Provider.provider_priority.asc()) .order_by(Provider.provider_priority.asc())
.all() .all()
) )
@@ -681,9 +647,7 @@ class AdminPoolOverviewAdapter(AdminApiAdapter):
ProviderAPIKey.provider_id, ProviderAPIKey.provider_id,
func.count(ProviderAPIKey.id).label("total"), func.count(ProviderAPIKey.id).label("total"),
func.coalesce( func.coalesce(
func.sum( func.sum(case((ProviderAPIKey.is_active.is_(True), 1), else_=0)),
case((ProviderAPIKey.is_active.is_(True), 1), else_=0)
),
0, 0,
).label("active"), ).label("active"),
) )
@@ -706,8 +670,8 @@ class AdminPoolOverviewAdapter(AdminApiAdapter):
if key_stats_by_provider.get(pid, {}).get("total", 0) > 0 if key_stats_by_provider.get(pid, {}).get("total", 0) > 0
] ]
if cooldown_targets: if cooldown_targets:
cooldown_count_by_provider = ( cooldown_count_by_provider = await pool_redis.batch_count_provider_cooldowns(
await pool_redis.batch_count_provider_cooldowns(cooldown_targets) cooldown_targets
) )
items: list[PoolOverviewItem] = [] items: list[PoolOverviewItem] = []
@@ -719,9 +683,7 @@ class AdminPoolOverviewAdapter(AdminApiAdapter):
PoolOverviewItem( PoolOverviewItem(
provider_id=pid, provider_id=pid,
provider_name=str(getattr(p, "name", "") or ""), provider_name=str(getattr(p, "name", "") or ""),
provider_type=str( provider_type=str(getattr(p, "provider_type", "custom") or "custom"),
getattr(p, "provider_type", "custom") or "custom"
),
total_keys=key_stats["total"], total_keys=key_stats["total"],
active_keys=key_stats["active"], active_keys=key_stats["active"],
cooldown_count=cooldown_count_by_provider.get(pid, 0), cooldown_count=cooldown_count_by_provider.get(pid, 0),
@@ -748,8 +710,17 @@ _ALLOWED_POOL_KEY_QUICK_SELECTORS = frozenset(
} }
) )
_ACCOUNT_BANNED_CODES = frozenset( _ACCOUNT_BANNED_CODES = frozenset(
{"account_banned", "account_forbidden", "account_blocked"} {
"account_banned",
"account_forbidden",
"account_blocked",
"account_suspended",
"account_disabled",
"workspace_deactivated",
"account_verification",
}
) )
_TOKEN_ISSUE_CODES = frozenset({"oauth_expired", "oauth_refresh_failed"})
_BANNED_REASON_PATTERN = re.compile(r"(banned|forbidden|blocked|suspend|封|禁|受限)") _BANNED_REASON_PATTERN = re.compile(r"(banned|forbidden|blocked|suspend|封|禁|受限)")
_PROVIDER_OVERVIEW_LOAD_ONLY_ATTRS: tuple[Any, ...] = ( _PROVIDER_OVERVIEW_LOAD_ONLY_ATTRS: tuple[Any, ...] = (
@@ -801,11 +772,7 @@ def _normalize_batch_text(value: Any) -> str:
def _normalize_pool_search_scope(value: Any) -> str: def _normalize_pool_search_scope(value: Any) -> str:
return ( return _FULL_SEARCH_SCOPE if _normalize_batch_text(value) == _FULL_SEARCH_SCOPE else "name"
_FULL_SEARCH_SCOPE
if _normalize_batch_text(value) == _FULL_SEARCH_SCOPE
else "name"
)
def _normalize_pool_quick_selectors(values: Any) -> list[str]: def _normalize_pool_quick_selectors(values: Any) -> list[str]:
@@ -837,8 +804,7 @@ def _get_quota_segments(account_quota: Any) -> list[str]:
return [ return [
segment segment
for segment in ( for segment in (
_normalize_quota_segment(part) _normalize_quota_segment(part) for part in str(account_quota or "").split("|")
for part in str(account_quota or "").split("|")
) )
if segment if segment
] ]
@@ -846,9 +812,7 @@ def _get_quota_segments(account_quota: Any) -> list[str]:
def _quota_segment_has_depleted_keyword(segment: str) -> bool: def _quota_segment_has_depleted_keyword(segment: str) -> bool:
return bool( return bool(
re.search( re.search(r"(无额度|额度不足|已耗尽|耗尽|depleted|exhausted|insufficient)", segment)
r"(无额度|额度不足|已耗尽|耗尽|depleted|exhausted|insufficient)", segment
)
) )
@@ -900,19 +864,27 @@ def _has_no_weekly_limit(account_quota: Any) -> bool:
def _detail_is_oauth_invalid(detail: PoolKeyDetail) -> bool: def _detail_is_oauth_invalid(detail: PoolKeyDetail) -> bool:
if _normalize_batch_text(detail.auth_type) != "oauth": if _normalize_batch_text(detail.auth_type) != "oauth":
return False return False
if detail.oauth_invalid_at is not None or _normalize_batch_text( status_code = _normalize_batch_text(detail.account_status_code)
detail.oauth_invalid_reason if status_code in _TOKEN_ISSUE_CODES:
): return True
if status_code in _ACCOUNT_BANNED_CODES or status_code == "oauth_request_failed":
return False
reason = _normalize_batch_text(detail.oauth_invalid_reason)
if reason.startswith("[oauth_expired]") or reason.startswith("[refresh_failed]"):
return True
if reason.startswith("[account_block]") or reason.startswith("[request_failed]"):
return False
if detail.oauth_invalid_at is not None or reason:
return True return True
expires_at = detail.oauth_expires_at expires_at = detail.oauth_expires_at
return ( return isinstance(expires_at, int) and expires_at > 0 and expires_at <= int(time.time())
isinstance(expires_at, int)
and expires_at > 0
and expires_at <= int(time.time())
)
def _detail_is_banned(detail: PoolKeyDetail) -> bool: def _detail_is_banned(detail: PoolKeyDetail) -> bool:
if _normalize_batch_text(detail.account_status_code) in _ACCOUNT_BANNED_CODES:
return True
reason = _normalize_batch_text(detail.oauth_invalid_reason) reason = _normalize_batch_text(detail.oauth_invalid_reason)
if reason and _BANNED_REASON_PATTERN.search(reason): if reason and _BANNED_REASON_PATTERN.search(reason):
return True return True
@@ -944,6 +916,8 @@ def _matches_pool_key_search(
detail.key_name, detail.key_name,
detail.auth_type, detail.auth_type,
detail.oauth_plan_type, detail.oauth_plan_type,
detail.account_status_label,
detail.account_status_reason,
detail.account_quota, detail.account_quota,
"独立代理" if _detail_has_proxy(detail) else "未配置代理", "独立代理" if _detail_has_proxy(detail) else "未配置代理",
"已启用" if detail.is_active else "已禁用", "已启用" if detail.is_active else "已禁用",
@@ -983,6 +957,7 @@ def _filter_pool_key_details(
quick_selectors: list[str] | None = None, quick_selectors: list[str] | None = None,
search_scope: str = _FULL_SEARCH_SCOPE, search_scope: str = _FULL_SEARCH_SCOPE,
require_cooldown: bool = False, require_cooldown: bool = False,
require_schedulable: bool = False,
) -> list[PoolKeyDetail]: ) -> list[PoolKeyDetail]:
normalized_selectors = _normalize_pool_quick_selectors(quick_selectors) normalized_selectors = _normalize_pool_quick_selectors(quick_selectors)
normalized_search_scope = _normalize_pool_search_scope(search_scope) normalized_search_scope = _normalize_pool_search_scope(search_scope)
@@ -990,19 +965,40 @@ def _filter_pool_key_details(
for detail in details: for detail in details:
if require_cooldown and not detail.cooldown_reason: if require_cooldown and not detail.cooldown_reason:
continue continue
if not _matches_pool_key_search( if require_schedulable and not _detail_is_schedulable(detail):
detail, search, search_scope=normalized_search_scope continue
): if not _matches_pool_key_search(detail, search, search_scope=normalized_search_scope):
continue continue
if normalized_selectors and not any( if normalized_selectors and not any(
_matches_pool_key_quick_selector(detail, selector) _matches_pool_key_quick_selector(detail, selector) for selector in normalized_selectors
for selector in normalized_selectors
): ):
continue continue
filtered.append(detail) filtered.append(detail)
return filtered return filtered
def _detail_is_schedulable(detail: PoolKeyDetail) -> bool:
status = str(getattr(detail, "scheduling_status", "") or "").strip().lower()
if status:
return status in {"available", "degraded"}
if not detail.is_active:
return False
if detail.account_status_blocked:
return False
if detail.cooldown_reason:
return False
if detail.circuit_breaker_open:
return False
if (
detail.cost_limit is not None
and detail.cost_limit > 0
and detail.cost_window_usage >= detail.cost_limit
):
return False
return True
def _build_pool_keys_base_query(db: Session, provider_id: str) -> Any: def _build_pool_keys_base_query(db: Session, provider_id: str) -> Any:
return ( return (
db.query(ProviderAPIKey) db.query(ProviderAPIKey)
@@ -1106,9 +1102,7 @@ async def _serialize_pool_key_details(
circuit_breaker_open=any_circuit_open, circuit_breaker_open=any_circuit_open,
cost_window_usage=cost_usage, cost_window_usage=cost_usage,
cost_limit=cost_limit, cost_limit=cost_limit,
cost_soft_threshold_percent=( cost_soft_threshold_percent=(pcfg.cost_soft_threshold_percent if pcfg else 80),
pcfg.cost_soft_threshold_percent if pcfg else 80
),
health_score=health_score, health_score=health_score,
) )
@@ -1167,9 +1161,7 @@ async def _serialize_pool_key_details(
key_name=str(getattr(k, "name", "") or ""), key_name=str(getattr(k, "name", "") or ""),
is_active=bool(k.is_active), is_active=bool(k.is_active),
auth_type=str(getattr(k, "auth_type", "api_key") or "api_key"), auth_type=str(getattr(k, "auth_type", "api_key") or "api_key"),
oauth_expires_at=_derive_oauth_expires_at( oauth_expires_at=_derive_oauth_expires_at(k, auth_config=oauth_auth_config),
k, auth_config=oauth_auth_config
),
oauth_invalid_at=( oauth_invalid_at=(
int(k.oauth_invalid_at.timestamp()) int(k.oauth_invalid_at.timestamp())
if getattr(k, "oauth_invalid_at", None) if getattr(k, "oauth_invalid_at", None)
@@ -1183,6 +1175,12 @@ async def _serialize_pool_key_details(
oauth_account_name=_derive_oauth_account_name(oauth_auth_config), oauth_account_name=_derive_oauth_account_name(oauth_auth_config),
oauth_account_user_id=_derive_oauth_account_user_id(oauth_auth_config), oauth_account_user_id=_derive_oauth_account_user_id(oauth_auth_config),
oauth_organizations=_derive_oauth_organizations(oauth_auth_config), oauth_organizations=_derive_oauth_organizations(oauth_auth_config),
account_status_code=account_state.code,
account_status_label=account_state.label,
account_status_reason=account_state.reason,
account_status_blocked=account_state.blocked,
account_status_recoverable=bool(getattr(account_state, "recoverable", False)),
account_status_source=getattr(account_state, "source", None),
quota_updated_at=_extract_quota_updated_at( quota_updated_at=_extract_quota_updated_at(
provider_type, provider_type,
getattr(k, "upstream_metadata", None), getattr(k, "upstream_metadata", None),
@@ -1197,9 +1195,7 @@ async def _serialize_pool_key_details(
v if (v := getattr(k, "cache_ttl_minutes", None)) is not None else 5 v if (v := getattr(k, "cache_ttl_minutes", None)) is not None else 5
), ),
max_probe_interval_minutes=( max_probe_interval_minutes=(
v v if (v := getattr(k, "max_probe_interval_minutes", None)) is not None else 32
if (v := getattr(k, "max_probe_interval_minutes", None)) is not None
else 32
), ),
note=getattr(k, "note", None), note=getattr(k, "note", None),
allowed_models=allowed_models, allowed_models=allowed_models,
@@ -1227,12 +1223,8 @@ async def _serialize_pool_key_details(
total_cost_usd=key_total_cost_usd, total_cost_usd=key_total_cost_usd,
sticky_sessions=sticky_counts.get(kid, 0), sticky_sessions=sticky_counts.get(kid, 0),
lru_score=lru_scores.get(kid), lru_score=lru_scores.get(kid),
created_at=( created_at=(k.created_at.isoformat() if getattr(k, "created_at", None) else None),
k.created_at.isoformat() if getattr(k, "created_at", None) else None last_used_at=(key_last_used_at.isoformat() if key_last_used_at else None),
),
last_used_at=(
key_last_used_at.isoformat() if key_last_used_at else None
),
scheduling_status=scheduling_status, scheduling_status=scheduling_status,
scheduling_reason=scheduling_reason, scheduling_reason=scheduling_reason,
scheduling_label=scheduling_label, scheduling_label=scheduling_label,
@@ -1257,6 +1249,7 @@ async def _resolve_filtered_pool_key_details(
quick_selectors: list[str], quick_selectors: list[str],
search_scope: str, search_scope: str,
require_cooldown: bool, require_cooldown: bool,
require_schedulable: bool,
max_scan: int = _DEFAULT_POOL_KEY_SCAN_LIMIT, max_scan: int = _DEFAULT_POOL_KEY_SCAN_LIMIT,
) -> tuple[list[PoolKeyDetail], float, float, float]: ) -> tuple[list[PoolKeyDetail], float, float, float]:
keys_query_started_at = time.perf_counter() keys_query_started_at = time.perf_counter()
@@ -1275,6 +1268,7 @@ async def _resolve_filtered_pool_key_details(
quick_selectors=quick_selectors, quick_selectors=quick_selectors,
search_scope=search_scope, search_scope=search_scope,
require_cooldown=require_cooldown, require_cooldown=require_cooldown,
require_schedulable=require_schedulable,
) )
return filtered_details, keys_query_ms, redis_state_ms, serialize_ms return filtered_details, keys_query_ms, redis_state_ms, serialize_ms
@@ -1304,18 +1298,12 @@ class AdminListPoolKeysAdapter(AdminApiAdapter):
pcfg = parse_pool_config(getattr(provider, "config", None)) pcfg = parse_pool_config(getattr(provider, "config", None))
pid = str(provider.id) pid = str(provider.id)
provider_type = str(getattr(provider, "provider_type", "custom") or "custom") provider_type = str(getattr(provider, "provider_type", "custom") or "custom")
normalized_quick_selectors = _normalize_pool_quick_selectors( normalized_quick_selectors = _normalize_pool_quick_selectors(self.quick_selectors)
self.quick_selectors
)
normalized_search_scope = _normalize_pool_search_scope(self.search_scope) normalized_search_scope = _normalize_pool_search_scope(self.search_scope)
q = _build_pool_keys_base_query(db, pid) q = _build_pool_keys_base_query(db, pid)
if self.search and normalized_search_scope != _FULL_SEARCH_SCOPE: if self.search and normalized_search_scope != _FULL_SEARCH_SCOPE:
escaped = ( escaped = self.search.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_")
self.search.replace("\\", "\\\\")
.replace("%", "\\%")
.replace("_", "\\_")
)
q = q.filter(ProviderAPIKey.name.ilike(f"%{escaped}%")) q = q.filter(ProviderAPIKey.name.ilike(f"%{escaped}%"))
if self.status == "active": if self.status == "active":
@@ -1326,6 +1314,7 @@ class AdminListPoolKeysAdapter(AdminApiAdapter):
total = 0 total = 0
if ( if (
normalized_quick_selectors normalized_quick_selectors
or self.status == "active"
or self.status == "cooldown" or self.status == "cooldown"
or (bool(self.search) and normalized_search_scope == _FULL_SEARCH_SCOPE) or (bool(self.search) and normalized_search_scope == _FULL_SEARCH_SCOPE)
): ):
@@ -1343,6 +1332,7 @@ class AdminListPoolKeysAdapter(AdminApiAdapter):
quick_selectors=normalized_quick_selectors, quick_selectors=normalized_quick_selectors,
search_scope=normalized_search_scope, search_scope=normalized_search_scope,
require_cooldown=self.status == "cooldown", require_cooldown=self.status == "cooldown",
require_schedulable=self.status == "active",
) )
total = len(filtered_details) total = len(filtered_details)
offset = (self.page - 1) * self.page_size offset = (self.page - 1) * self.page_size
@@ -1416,6 +1406,7 @@ class AdminResolvePoolKeySelectionAdapter(AdminApiAdapter):
quick_selectors=_normalize_pool_quick_selectors(self.body.quick_selectors), quick_selectors=_normalize_pool_quick_selectors(self.body.quick_selectors),
search_scope=_FULL_SEARCH_SCOPE, search_scope=_FULL_SEARCH_SCOPE,
require_cooldown=False, require_cooldown=False,
require_schedulable=False,
max_scan=_RESOLVE_SELECTION_SCAN_LIMIT, max_scan=_RESOLVE_SELECTION_SCAN_LIMIT,
) )
@@ -1435,9 +1426,7 @@ class AdminResolvePoolKeySelectionAdapter(AdminApiAdapter):
@dataclass @dataclass
class AdminBatchImportKeysAdapter(AdminApiAdapter): class AdminBatchImportKeysAdapter(AdminApiAdapter):
provider_id: str = "" provider_id: str = ""
body: BatchImportRequest = field( body: BatchImportRequest = field(default_factory=lambda: BatchImportRequest(keys=[]))
default_factory=lambda: BatchImportRequest(keys=[])
)
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override] async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db db = context.db
@@ -1611,9 +1600,7 @@ class AdminBatchActionKeysAdapter(AdminApiAdapter):
except Exception as exc: except Exception as exc:
db.rollback() db.rollback()
logger.error("batch action commit failed: {}", exc) logger.error("batch action commit failed: {}", exc)
return BatchActionResponse( return BatchActionResponse(affected=0, message=f"commit failed: {exc}")
affected=0, message=f"commit failed: {exc}"
)
admin_name = context.user.username if context.user else "admin" admin_name = context.user.username if context.user else "admin"
affected_ids = [str(k.id)[:8] for k in keys] affected_ids = [str(k.id)[:8] for k in keys]
@@ -1654,15 +1641,13 @@ class AdminCleanupBannedKeysAdapter(AdminApiAdapter):
raise NotFoundException("Provider not found", "provider") raise NotFoundException("Provider not found", "provider")
pid = str(provider.id) pid = str(provider.id)
provider_type = ( provider_type = str(getattr(provider, "provider_type", "") or "").strip().lower()
str(getattr(provider, "provider_type", "") or "").strip().lower()
)
keys = db.query(ProviderAPIKey).filter(ProviderAPIKey.provider_id == pid).all() keys = db.query(ProviderAPIKey).filter(ProviderAPIKey.provider_id == pid).all()
banned_keys = [key for key in keys if _is_known_banned_key(key, provider_type)] banned_keys = [key for key in keys if _is_known_banned_key(key, provider_type)]
if not banned_keys: if not banned_keys:
return BatchActionResponse(affected=0, message="未发现已知封号账号") return BatchActionResponse(affected=0, message="未发现可清理的异常账号")
banned_key_ids = [str(key.id) for key in banned_keys] banned_key_ids = [str(key.id) for key in banned_keys]
try: try:
@@ -1689,7 +1674,7 @@ class AdminCleanupBannedKeysAdapter(AdminApiAdapter):
admin_name = context.user.username if context.user else "admin" admin_name = context.user.username if context.user else "admin"
logger.warning( logger.warning(
"Pool cleanup banned by {}: provider={}, affected={}, key_ids={}", "Pool cleanup abnormal keys by {}: provider={}, affected={}, key_ids={}",
admin_name, admin_name,
self.provider_id[:8], self.provider_id[:8],
len(banned_key_ids), len(banned_key_ids),
@@ -1698,5 +1683,5 @@ class AdminCleanupBannedKeysAdapter(AdminApiAdapter):
return BatchActionResponse( return BatchActionResponse(
affected=len(banned_key_ids), affected=len(banned_key_ids),
message=f"已清理 {len(banned_key_ids)} 个已知封号账号", message=f"已清理 {len(banned_key_ids)} 个异常账号",
) )
+6
View File
@@ -88,6 +88,12 @@ class PoolKeyDetail(BaseModel):
oauth_account_name: str | None = None oauth_account_name: str | None = None
oauth_account_user_id: str | None = None oauth_account_user_id: str | None = None
oauth_organizations: list[OAuthOrganizationSummary] = Field(default_factory=list) oauth_organizations: list[OAuthOrganizationSummary] = Field(default_factory=list)
account_status_code: str | None = None
account_status_label: str | None = None
account_status_reason: str | None = None
account_status_blocked: bool = False
account_status_recoverable: bool = False
account_status_source: str | None = None
quota_updated_at: int | None = None quota_updated_at: int | None = None
# 健康度聚合字段(与 Provider Key 列表口径一致) # 健康度聚合字段(与 Provider Key 列表口径一致)
health_score: float = 1.0 health_score: float = 1.0
+107 -7
View File
@@ -70,10 +70,25 @@ def _mark_refresh_failed_sync(key_id: str, reason: str) -> None:
key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == key_id).first() key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == key_id).first()
if not key: if not key:
raise NotFoundException("Key 不存在", "key") raise NotFoundException("Key 不存在", "key")
current_reason = str(getattr(key, "oauth_invalid_reason", None) or "").strip()
if _should_preserve_refresh_failure_reason(current_reason):
return
key.oauth_invalid_at = datetime.now(timezone.utc) key.oauth_invalid_at = datetime.now(timezone.utc)
key.oauth_invalid_reason = reason key.oauth_invalid_reason = reason
def _should_preserve_refresh_failure_reason(reason: str | None) -> bool:
from src.services.provider.oauth_token import is_account_level_block
from src.services.provider.pool.account_state import OAUTH_EXPIRED_PREFIX
text = str(reason or "").strip()
if not text:
return False
if is_account_level_block(text):
return True
return text.startswith(OAUTH_EXPIRED_PREFIX)
def _store_refreshed_oauth_sync( def _store_refreshed_oauth_sync(
key_id: str, key_id: str,
access_token: str, access_token: str,
@@ -86,12 +101,15 @@ def _store_refreshed_oauth_sync(
key.api_key = crypto_service.encrypt(access_token) key.api_key = crypto_service.encrypt(access_token)
key.auth_config = crypto_service.encrypt(json.dumps(parsed_auth_config)) key.auth_config = crypto_service.encrypt(json.dumps(parsed_auth_config))
# 刷新成功 => 清除所有 oauth_invalid 标记(包括 [ACCOUNT_BLOCK])。 from src.services.provider.oauth_token import is_account_level_block
# Token 能成功刷新说明账号可用,之前的 block 标记应视为过时。
if getattr(key, "oauth_invalid_at", None) is not None: current_reason = str(getattr(key, "oauth_invalid_reason", None) or "").strip()
# 手动 refresh 只清除可恢复的 token 类异常,不自动清账号级 block。
if getattr(key, "oauth_invalid_at", None) is not None and not is_account_level_block(
current_reason
):
key.oauth_invalid_at = None key.oauth_invalid_at = None
key.oauth_invalid_reason = None key.oauth_invalid_reason = None
key.is_active = True
# ============================================================================== # ==============================================================================
@@ -278,6 +296,8 @@ class CompleteOAuthResponse(BaseModel):
expires_at: int | None = None expires_at: int | None = None
has_refresh_token: bool = False has_refresh_token: bool = False
email: str | None = None email: str | None = None
account_state_recheck_attempted: bool = False
account_state_recheck_error: str | None = None
class ProviderCompleteOAuthRequest(BaseModel): class ProviderCompleteOAuthRequest(BaseModel):
@@ -985,12 +1005,19 @@ async def complete_oauth(
access_token, access_token,
auth_config, auth_config,
) )
recheck_attempted, recheck_error = await _refresh_account_state_after_oauth_update(
provider_id=str(provider.id),
provider_type=provider_type,
key_ids=[key_id],
)
return CompleteOAuthResponse( return CompleteOAuthResponse(
provider_type=provider_type, provider_type=provider_type,
expires_at=expires_at, expires_at=expires_at,
has_refresh_token=bool(refresh_token), has_refresh_token=bool(refresh_token),
email=auth_config.get("email"), email=auth_config.get("email"),
account_state_recheck_attempted=recheck_attempted,
account_state_recheck_error=recheck_error,
) )
@@ -1044,10 +1071,17 @@ async def refresh_oauth(
try: try:
access_token, new_cfg = await refresh_access_token(cfg, proxy_config=proxy_config) access_token, new_cfg = await refresh_access_token(cfg, proxy_config=proxy_config)
except Exception as e: except Exception as e:
refresh_error = str(e) or type(e).__name__
await run_in_threadpool( await run_in_threadpool(
_mark_refresh_failed_sync, _mark_refresh_failed_sync,
key_id, key_id,
f"[REFRESH_FAILED] Token 续期失败: {e}", f"[REFRESH_FAILED] Token 续期失败: {refresh_error}",
)
await _recheck_account_state_after_failed_refresh(
provider_id=str(provider.id),
provider_type=provider_type,
key_id=key_id,
refresh_error=refresh_error,
) )
logger.warning("Kiro Key {} token 刷新失败,已标记为刷新失效: {}", key_id, e) logger.warning("Kiro Key {} token 刷新失败,已标记为刷新失效: {}", key_id, e)
raise InvalidRequestException("Kiro token refresh 失败,请检查凭据是否有效") raise InvalidRequestException("Kiro token refresh 失败,请检查凭据是否有效")
@@ -1058,12 +1092,19 @@ async def refresh_oauth(
access_token, access_token,
new_cfg.to_dict(), new_cfg.to_dict(),
) )
recheck_attempted, recheck_error = await _refresh_account_state_after_oauth_update(
provider_id=str(provider.id),
provider_type=provider_type,
key_ids=[key_id],
)
return CompleteOAuthResponse( return CompleteOAuthResponse(
provider_type=provider_type, provider_type=provider_type,
expires_at=new_cfg.expires_at or None, expires_at=new_cfg.expires_at or None,
has_refresh_token=bool(new_cfg.refresh_token), has_refresh_token=bool(new_cfg.refresh_token),
email=None, email=None,
account_state_recheck_attempted=recheck_attempted,
account_state_recheck_error=recheck_error,
) )
template = _require_oauth_template(provider_type) template = _require_oauth_template(provider_type)
@@ -1146,6 +1187,12 @@ async def refresh_oauth(
logger.warning( logger.warning(
"Key {} OAuth token 刷新失败,已标记为刷新失效: {}", key_id, error_reason "Key {} OAuth token 刷新失败,已标记为刷新失效: {}", key_id, error_reason
) )
await _recheck_account_state_after_failed_refresh(
provider_id=str(provider.id),
provider_type=provider_type,
key_id=key_id,
refresh_error=error_reason,
)
raise InvalidRequestException(f"token refresh 失败: {error_reason}") raise InvalidRequestException(f"token refresh 失败: {error_reason}")
@@ -1186,12 +1233,19 @@ async def refresh_oauth(
) )
await run_in_threadpool(_store_refreshed_oauth_sync, key_id, access_token, parsed) await run_in_threadpool(_store_refreshed_oauth_sync, key_id, access_token, parsed)
recheck_attempted, recheck_error = await _refresh_account_state_after_oauth_update(
provider_id=str(provider.id),
provider_type=provider_type,
key_ids=[key_id],
)
return CompleteOAuthResponse( return CompleteOAuthResponse(
provider_type=provider_type, provider_type=provider_type,
expires_at=expires_at, expires_at=expires_at,
has_refresh_token=bool(parsed.get("refresh_token")), has_refresh_token=bool(parsed.get("refresh_token")),
email=parsed.get("email"), email=parsed.get("email"),
account_state_recheck_attempted=recheck_attempted,
account_state_recheck_error=recheck_error,
) )
finally: finally:
if got_lock: if got_lock:
@@ -2055,6 +2109,22 @@ async def _refresh_quota_after_import(
key_ids: list[str], key_ids: list[str],
) -> None: ) -> None:
"""导入完成后触发一次配额刷新(使用独立 db session)。""" """导入完成后触发一次配额刷新(使用独立 db session)。"""
attempted, error = await _refresh_account_state_after_oauth_update(
provider_id=provider_id,
provider_type=provider_type,
key_ids=key_ids,
)
if attempted and error:
logger.warning("[BATCH_IMPORT] 导入后配额刷新失败 (provider={}): {}", provider_id, error)
async def _refresh_account_state_after_oauth_update(
*,
provider_id: str,
provider_type: str,
key_ids: list[str],
) -> tuple[bool, str | None]:
"""OAuth 更新成功后,立即复检账号额度/状态。"""
from src.services.provider_keys.key_quota_service import ( from src.services.provider_keys.key_quota_service import (
CODEX_WHAM_USAGE_URL, CODEX_WHAM_USAGE_URL,
QUOTA_REFRESH_PROVIDER_TYPES, QUOTA_REFRESH_PROVIDER_TYPES,
@@ -2062,7 +2132,7 @@ async def _refresh_quota_after_import(
) )
if not key_ids or provider_type not in QUOTA_REFRESH_PROVIDER_TYPES: if not key_ids or provider_type not in QUOTA_REFRESH_PROVIDER_TYPES:
return return False, None
try: try:
db = create_session() db = create_session()
try: try:
@@ -2074,8 +2144,38 @@ async def _refresh_quota_after_import(
) )
finally: finally:
db.close() db.close()
return True, None
except Exception as exc: except Exception as exc:
logger.warning("[BATCH_IMPORT] 导入后配额刷新失败 (provider={}): {}", provider_id, exc) brief = str(exc)[:120] if str(exc) else type(exc).__name__
return True, f"{type(exc).__name__}: {brief}"
async def _recheck_account_state_after_failed_refresh(
*,
provider_id: str,
provider_type: str,
key_id: str,
refresh_error: str,
) -> None:
attempted, error = await _refresh_account_state_after_oauth_update(
provider_id=provider_id,
provider_type=provider_type,
key_ids=[key_id],
)
if not attempted:
return
if error:
logger.warning(
"[OAUTH_REFRESH] Key {} 刷新失败后复检账号状态失败: {} (refresh_error={})",
key_id,
error,
refresh_error,
)
return
logger.info(
"[OAUTH_REFRESH] Key {} 刷新失败后已使用现有 access token 复检账号状态",
key_id,
)
# ============================================================================== # ==============================================================================
+58 -38
View File
@@ -44,7 +44,10 @@ from src.services.model.upstream_fetcher import (
build_format_to_config, build_format_to_config,
fetch_models_for_key, fetch_models_for_key,
) )
from src.services.provider.oauth_token import resolve_oauth_access_token from src.services.provider.oauth_token import (
resolve_oauth_access_token,
verify_oauth_before_account_block,
)
from src.services.proxy_node.resolver import resolve_effective_proxy from src.services.proxy_node.resolver import resolve_effective_proxy
from src.services.request.candidate import RequestCandidateService from src.services.request.candidate import RequestCandidateService
from src.services.request.model_test_debug import ( from src.services.request.model_test_debug import (
@@ -1224,37 +1227,45 @@ async def test_model(
or "permission" in str(error_obj.get("status", "")).lower() or "permission" in str(error_obj.get("status", "")).lower()
) )
): ):
from datetime import datetime, timezone should_mark = await verify_oauth_before_account_block(
endpoint=endpoint,
from src.services.provider.oauth_token import ( key=api_key,
OAUTH_ACCOUNT_BLOCK_PREFIX, candidate_reason="Google 要求验证账号",
request_id="test-model",
key_display=f"test-model:{api_key.id}",
) )
if should_mark:
api_key.oauth_invalid_at = datetime.now(timezone.utc) from src.services.provider.oauth_token import (
api_key.oauth_invalid_reason = ( OAUTH_ACCOUNT_BLOCK_PREFIX,
f"{OAUTH_ACCOUNT_BLOCK_PREFIX}Google 要求验证账号"
)
api_key.is_active = False
db.commit()
oauth_email = None
if getattr(api_key, "auth_config", None):
try:
decrypted = crypto_service.decrypt(api_key.auth_config)
parsed = json.loads(decrypted)
if isinstance(parsed, dict):
email_val = parsed.get("email")
if isinstance(email_val, str) and email_val.strip():
oauth_email = email_val.strip()
except Exception:
oauth_email = None
if oauth_email:
logger.warning(
"[test-model] Key {} (email={}) 因 403 verify 已标记为异常",
api_key.id,
oauth_email,
) )
else:
logger.warning("[test-model] Key {} 因 403 verify 已标记为异常", api_key.id) api_key.oauth_invalid_at = datetime.now(timezone.utc)
api_key.oauth_invalid_reason = (
f"{OAUTH_ACCOUNT_BLOCK_PREFIX}Google 要求验证账号"
)
db.commit()
oauth_email = None
if getattr(api_key, "auth_config", None):
try:
decrypted = crypto_service.decrypt(api_key.auth_config)
parsed = json.loads(decrypted)
if isinstance(parsed, dict):
email_val = parsed.get("email")
if isinstance(email_val, str) and email_val.strip():
oauth_email = email_val.strip()
except Exception:
oauth_email = None
if oauth_email:
logger.warning(
"[test-model] Key {} (email={}) 因 403 verify 已标记为异常",
api_key.id,
oauth_email,
)
else:
logger.warning(
"[test-model] Key {} 因 403 verify 已标记为异常",
api_key.id,
)
upstream_status = int( upstream_status = int(
response.get("status_code", 0) or error_obj.get("code", 0) or 500 response.get("status_code", 0) or error_obj.get("code", 0) or 500
@@ -1918,7 +1929,7 @@ async def _run_concurrent_test(
.filter(ProviderAPIKey.id == str(getattr(local_key, "id", "") or "")) .filter(ProviderAPIKey.id == str(getattr(local_key, "id", "") or ""))
.first() .first()
) )
parsed = _extract_test_response_or_raise( parsed = await _extract_test_response_or_raise(
response=response, response=response,
endpoint=local_endpoint, endpoint=local_endpoint,
provider_name=str(local_provider.name), provider_name=str(local_provider.name),
@@ -2078,14 +2089,15 @@ async def _run_concurrent_test(
} }
def _maybe_mark_test_oauth_key_invalid( async def _maybe_mark_test_oauth_key_invalid(
*, *,
db: Session, db: Session,
endpoint: Any,
key: Any, key: Any,
auth_type: str, auth_type: str,
error_payload: Any, error_payload: Any,
) -> None: ) -> None:
if auth_type != "oauth" or not isinstance(error_payload, dict): if auth_type != "oauth" or key is None or not isinstance(error_payload, dict):
return return
error_obj = error_payload.get("error") error_obj = error_payload.get("error")
@@ -2101,17 +2113,24 @@ def _maybe_mark_test_oauth_key_invalid(
): ):
return return
from datetime import datetime, timezone should_mark = await verify_oauth_before_account_block(
endpoint=endpoint,
key=key,
candidate_reason="Google 要求验证账号",
request_id="provider-query-test",
key_display=f"provider-query-test:{getattr(key, 'id', '?')}",
)
if not should_mark:
return
from src.services.provider.oauth_token import OAUTH_ACCOUNT_BLOCK_PREFIX from src.services.provider.oauth_token import OAUTH_ACCOUNT_BLOCK_PREFIX
key.oauth_invalid_at = datetime.now(timezone.utc) key.oauth_invalid_at = datetime.now(timezone.utc)
key.oauth_invalid_reason = f"{OAUTH_ACCOUNT_BLOCK_PREFIX}Google 要求验证账号" key.oauth_invalid_reason = f"{OAUTH_ACCOUNT_BLOCK_PREFIX}Google 要求验证账号"
key.is_active = False
db.commit() db.commit()
def _extract_test_response_or_raise( async def _extract_test_response_or_raise(
*, *,
response: dict[str, Any], response: dict[str, Any],
endpoint: Any, endpoint: Any,
@@ -2127,8 +2146,9 @@ def _extract_test_response_or_raise(
parsed_payload = _parse_jsonish(parsed_payload.get("response_body")) parsed_payload = _parse_jsonish(parsed_payload.get("response_body"))
if isinstance(parsed_payload, dict) and parsed_payload.get("error"): if isinstance(parsed_payload, dict) and parsed_payload.get("error"):
_maybe_mark_test_oauth_key_invalid( await _maybe_mark_test_oauth_key_invalid(
db=db, db=db,
endpoint=endpoint,
key=api_key, key=api_key,
auth_type=auth_type, auth_type=auth_type,
error_payload=parsed_payload, error_payload=parsed_payload,
@@ -2411,7 +2431,7 @@ async def test_model_failover(
db=db, db=db,
) )
set_candidate_model_test_debug(candidate, _extract_test_debug_payload(response)) set_candidate_model_test_debug(candidate, _extract_test_debug_payload(response))
return _extract_test_response_or_raise( return await _extract_test_response_or_raise(
response=response, response=response,
endpoint=endpoint, endpoint=endpoint,
provider_name=str(provider_obj.name), provider_name=str(provider_obj.name),
+2 -2
View File
@@ -45,6 +45,7 @@ from src.core.crypto import crypto_service
from src.core.logger import logger from src.core.logger import logger
from src.models.database import ApiKey, ProviderAPIKey, ProviderEndpoint, User, VideoTask from src.models.database import ApiKey, ProviderAPIKey, ProviderEndpoint, User, VideoTask
from src.services.billing.rule_service import BillingRuleLookupResult, BillingRuleService from src.services.billing.rule_service import BillingRuleLookupResult, BillingRuleService
from src.services.provider.provider_context import resolve_provider_proxy
from src.services.scheduling.aware_scheduler import ProviderCandidate from src.services.scheduling.aware_scheduler import ProviderCandidate
from src.services.usage.service import UsageService from src.services.usage.service import UsageService
@@ -478,9 +479,8 @@ class GeminiVeoHandler(VideoHandlerBase):
resolve_effective_proxy, resolve_effective_proxy,
) )
provider = getattr(endpoint, "provider", None) if endpoint else None
eff_proxy = resolve_effective_proxy( eff_proxy = resolve_effective_proxy(
getattr(provider, "proxy", None) if provider else None, resolve_provider_proxy(endpoint=endpoint, key=key),
getattr(key, "proxy", None), getattr(key, "proxy", None),
) )
+7 -2
View File
@@ -290,14 +290,19 @@ class PoolAdvancedConfig(BaseModel):
le=32, le=32,
description="批量操作并发数(前端批量刷新 OAuth/额度等)。默认 8", description="批量操作并发数(前端批量刷新 OAuth/额度等)。默认 8",
) )
probing_enabled: bool = Field(False, description="启用主动探测(定期检查 Key 可用性)") probing_enabled: bool = Field(
False, description="启用主动探测(定期刷新 Key 的账号状态与额度)"
)
probing_interval_minutes: int | None = Field( probing_interval_minutes: int | None = Field(
None, None,
ge=1, ge=1,
le=1440, le=1440,
description="主动探测间隔(分钟)。默认 10", description="主动探测间隔(分钟)。默认 10",
) )
auto_remove_banned_keys: bool = Field(False, description="检测到封号时自动清除账号") auto_remove_banned_keys: bool = Field(
False,
description="检测到不可恢复账号异常时自动清除账号(不处理纯 Token 失效)",
)
class ClaudeCodeAdvancedConfig(BaseModel): class ClaudeCodeAdvancedConfig(BaseModel):
+61 -13
View File
@@ -26,6 +26,7 @@ from src.core.logger import logger
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
from src.services.health.monitor import get_health_monitor from src.services.health.monitor import get_health_monitor
from src.services.provider.format import normalize_endpoint_signature from src.services.provider.format import normalize_endpoint_signature
from src.services.provider.oauth_token import verify_oauth_before_account_block
from src.services.provider.pool.config import parse_pool_config from src.services.provider.pool.config import parse_pool_config
from src.services.rate_limit.adaptive_rpm import get_adaptive_rpm_manager from src.services.rate_limit.adaptive_rpm import get_adaptive_rpm_manager
from src.services.rate_limit.detector import RateLimitType, detect_rate_limit_type from src.services.rate_limit.detector import RateLimitType, detect_rate_limit_type
@@ -171,7 +172,14 @@ class ErrorHandlerService:
and str(getattr(key, "auth_type", "") or "").lower() == "oauth" and str(getattr(key, "auth_type", "") or "").lower() == "oauth"
and self._is_account_validation_required(error_response_text) and self._is_account_validation_required(error_response_text)
): ):
self._mark_oauth_key_blocked(key, request_id, provider=provider) should_mark = await self._verify_oauth_before_account_block(
endpoint=endpoint,
key=key,
request_id=request_id,
candidate_reason="Google 要求验证账号",
)
if should_mark:
self._mark_oauth_key_blocked(key, request_id, provider=provider)
# 403 suspended -> 标记 OAuth key 为账号被暂停 # 403 suspended -> 标记 OAuth key 为账号被暂停
elif ( elif (
status_code == 403 status_code == 403
@@ -179,12 +187,19 @@ class ErrorHandlerService:
and str(getattr(key, "auth_type", "") or "").lower() == "oauth" and str(getattr(key, "auth_type", "") or "").lower() == "oauth"
and self._is_account_suspended(error_response_text) and self._is_account_suspended(error_response_text)
): ):
self._mark_oauth_key_blocked( should_mark = await self._verify_oauth_before_account_block(
key, endpoint=endpoint,
request_id, key=key,
reason="AWS 账号被暂停", request_id=request_id,
provider=provider, candidate_reason="AWS 账号被暂停",
) )
if should_mark:
self._mark_oauth_key_blocked(
key,
request_id,
reason="AWS 账号被暂停",
provider=provider,
)
# 401 account_deactivated -> 标记 OAuth key 为账号被永久停用 # 401 account_deactivated -> 标记 OAuth key 为账号被永久停用
elif ( elif (
status_code == 401 status_code == 401
@@ -192,12 +207,19 @@ class ErrorHandlerService:
and str(getattr(key, "auth_type", "") or "").lower() == "oauth" and str(getattr(key, "auth_type", "") or "").lower() == "oauth"
and self._is_account_deactivated(error_response_text) and self._is_account_deactivated(error_response_text)
): ):
self._mark_oauth_key_blocked( should_mark = await self._verify_oauth_before_account_block(
key, endpoint=endpoint,
request_id, key=key,
reason="账号已被停用 (account_deactivated)", request_id=request_id,
provider=provider, candidate_reason="账号已被停用 (account_deactivated)",
) )
if should_mark:
self._mark_oauth_key_blocked(
key,
request_id,
reason="账号已被停用 (account_deactivated)",
provider=provider,
)
return return
# 限流错误 # 限流错误
@@ -404,6 +426,10 @@ class ErrorHandlerService:
from datetime import datetime, timezone from datetime import datetime, timezone
from src.services.provider.oauth_token import OAUTH_ACCOUNT_BLOCK_PREFIX from src.services.provider.oauth_token import OAUTH_ACCOUNT_BLOCK_PREFIX
from src.services.provider.pool.account_state import (
resolve_pool_account_state,
should_auto_remove_account_state,
)
key.oauth_invalid_at = datetime.now(timezone.utc) key.oauth_invalid_at = datetime.now(timezone.utc)
key.oauth_invalid_reason = f"{OAUTH_ACCOUNT_BLOCK_PREFIX}{reason}" key.oauth_invalid_reason = f"{OAUTH_ACCOUNT_BLOCK_PREFIX}{reason}"
@@ -412,8 +438,13 @@ class ErrorHandlerService:
pool_cfg = parse_pool_config(getattr(provider, "config", None)) pool_cfg = parse_pool_config(getattr(provider, "config", None))
auto_remove_enabled = bool(pool_cfg and pool_cfg.auto_remove_banned_keys) auto_remove_enabled = bool(pool_cfg and pool_cfg.auto_remove_banned_keys)
account_state = resolve_pool_account_state(
provider_type=str(getattr(provider, "provider_type", "") or ""),
upstream_metadata=getattr(key, "upstream_metadata", None),
oauth_invalid_reason=getattr(key, "oauth_invalid_reason", None),
)
if auto_remove_enabled: if auto_remove_enabled and should_auto_remove_account_state(account_state):
key_id = str(getattr(key, "id", "") or "") key_id = str(getattr(key, "id", "") or "")
provider_id = str(getattr(key, "provider_id", "") or "") provider_id = str(getattr(key, "provider_id", "") or "")
display = self._format_key_display(key) display = self._format_key_display(key)
@@ -431,7 +462,7 @@ class ErrorHandlerService:
self.db.commit() self.db.commit()
logger.warning( logger.warning(
" [{}] {} 因 {} 已标记为账号异常并自动停用", " [{}] {} 因 {} 已标记为账号异常并阻止调度",
request_id, request_id,
self._format_key_display(key), self._format_key_display(key),
reason, reason,
@@ -439,6 +470,23 @@ class ErrorHandlerService:
except Exception as mark_exc: except Exception as mark_exc:
logger.debug(" [{}] 标记 oauth_invalid 失败: {}", request_id, mark_exc) logger.debug(" [{}] 标记 oauth_invalid 失败: {}", request_id, mark_exc)
async def _verify_oauth_before_account_block(
self,
*,
endpoint: ProviderEndpoint,
key: ProviderAPIKey,
request_id: str | None,
candidate_reason: str,
) -> bool:
"""Before applying an account-level block, distinguish it from OAuth expiry."""
return await verify_oauth_before_account_block(
endpoint=endpoint,
key=key,
candidate_reason=candidate_reason,
request_id=request_id,
key_display=self._format_key_display(key),
)
@staticmethod @staticmethod
def _schedule_auto_cleanup_after_delete(*, provider_id: str, key_id: str) -> None: def _schedule_auto_cleanup_after_delete(*, provider_id: str, key_id: str) -> None:
if not provider_id or not key_id: if not provider_id or not key_id:
+10 -7
View File
@@ -18,6 +18,7 @@ from src.core.crypto import crypto_service
from src.core.logger import logger from src.core.logger import logger
from src.core.provider_auth_types import ProviderAuthInfo from src.core.provider_auth_types import ProviderAuthInfo
from src.core.provider_oauth_utils import enrich_auth_config, post_oauth_token from src.core.provider_oauth_utils import enrich_auth_config, post_oauth_token
from src.services.provider.provider_context import resolve_provider_proxy
if TYPE_CHECKING: if TYPE_CHECKING:
from src.models.database import ProviderAPIKey, ProviderEndpoint from src.models.database import ProviderAPIKey, ProviderEndpoint
@@ -101,9 +102,14 @@ def _persist_refreshed_token(
key.api_key = crypto_service.encrypt(access_token) key.api_key = crypto_service.encrypt(access_token)
key.auth_config = crypto_service.encrypt(json.dumps(token_meta)) key.auth_config = crypto_service.encrypt(json.dumps(token_meta))
# 刷新成功 => 清除所有 oauth_invalid 标记(包括 [ACCOUNT_BLOCK])。 # 刷新成功只清除可恢复的 token 类异常。
# Token 能成功刷新说明账号可用,之前的 block 标记应视为过时。 # 账号级 block(如验证要求/工作区停用)不能靠 token refresh 自动恢复。
if getattr(key, "oauth_invalid_at", None) is not None: from src.services.provider.oauth_token import is_account_level_block
current_reason = str(getattr(key, "oauth_invalid_reason", None) or "").strip()
if getattr(key, "oauth_invalid_at", None) is not None and not is_account_level_block(
current_reason
):
key.oauth_invalid_at = None key.oauth_invalid_at = None
key.oauth_invalid_reason = None key.oauth_invalid_reason = None
@@ -213,10 +219,7 @@ def _get_proxy_config(key: Any, endpoint: Any = None) -> Any:
try: try:
from src.services.proxy_node.resolver import resolve_effective_proxy from src.services.proxy_node.resolver import resolve_effective_proxy
provider = getattr(key, "provider", None) or ( provider_proxy = resolve_provider_proxy(endpoint=endpoint, key=key)
getattr(endpoint, "provider", None) if endpoint else None
)
provider_proxy = getattr(provider, "proxy", None)
key_proxy = getattr(key, "proxy", None) key_proxy = getattr(key, "proxy", None)
return resolve_effective_proxy(provider_proxy, key_proxy) return resolve_effective_proxy(provider_proxy, key_proxy)
except Exception: except Exception:
+56 -5
View File
@@ -22,6 +22,10 @@ from typing import Any
from src.core.logger import logger from src.core.logger import logger
from src.database import create_session from src.database import create_session
from src.models.database import ProviderAPIKey from src.models.database import ProviderAPIKey
from src.services.provider.pool.account_state import (
OAUTH_EXPIRED_PREFIX,
OAUTH_REFRESH_FAILED_PREFIX,
)
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Account-level block 结构化标记 # Account-level block 结构化标记
@@ -76,6 +80,51 @@ def is_account_level_block(reason: str | None) -> bool:
) and not _is_refresh_recoverable_account_block(text) ) and not _is_refresh_recoverable_account_block(text)
async def verify_oauth_before_account_block(
*,
endpoint: Any,
key: Any,
candidate_reason: str,
request_id: str | None = None,
key_display: str | None = None,
) -> bool:
"""Before applying an account-level block, distinguish it from OAuth expiry."""
display = key_display or str(getattr(key, "id", "?") or "?")
try:
from src.services.provider.auth import get_provider_auth
await get_provider_auth(endpoint, key, force_refresh=True, refresh_skew=0)
except Exception as exc:
logger.debug(
"[OAUTH_VERIFY] [{}] {} account-block precheck failed: {}",
request_id,
display,
exc,
)
latest_reason = str(getattr(key, "oauth_invalid_reason", None) or "").strip()
if latest_reason.startswith(OAUTH_EXPIRED_PREFIX) or latest_reason.startswith(
OAUTH_REFRESH_FAILED_PREFIX
):
logger.info(
"[OAUTH_VERIFY] [{}] {} candidate account block ({}) skipped due to {}",
request_id,
display,
candidate_reason,
latest_reason[:120],
)
return False
logger.debug(
"[OAUTH_VERIFY] [{}] {} proceeding with account block ({}), post-refresh reason: {}",
request_id,
display,
candidate_reason,
latest_reason[:120] if latest_reason else "<none>",
)
return True
@dataclass(frozen=True, slots=True) @dataclass(frozen=True, slots=True)
class OAuthAccessTokenResult: class OAuthAccessTokenResult:
access_token: str access_token: str
@@ -135,13 +184,14 @@ async def resolve_oauth_access_token(
if row is not None: if row is not None:
row.api_key = key_obj.api_key row.api_key = key_obj.api_key
row.auth_config = key_obj.auth_config row.auth_config = key_obj.auth_config
# Refresh succeeded => clear all invalid markers (including # Refresh succeeded => only clear recoverable token errors.
# account-level blocks). A successful token refresh proves the # True account-level blocks must be cleared explicitly.
# account is usable; stale block marks should not persist. current_reason = str(getattr(row, "oauth_invalid_reason", None) or "")
if row.oauth_invalid_at is not None: if row.oauth_invalid_at is not None and not is_account_level_block(
current_reason
):
row.oauth_invalid_at = None row.oauth_invalid_at = None
row.oauth_invalid_reason = None row.oauth_invalid_reason = None
row.is_active = True
db.commit() db.commit()
except Exception as e: except Exception as e:
# Don't fail caller path; token is still usable for this request. # Don't fail caller path; token is still usable for this request.
@@ -157,6 +207,7 @@ async def resolve_oauth_access_token(
__all__ = [ __all__ = [
"OAuthAccessTokenResult", "OAuthAccessTokenResult",
"TOKEN_INVALIDATED_KEYWORDS", "TOKEN_INVALIDATED_KEYWORDS",
"verify_oauth_before_account_block",
"looks_like_token_invalidated", "looks_like_token_invalidated",
"resolve_oauth_access_token", "resolve_oauth_access_token",
] ]
+73 -2
View File
@@ -54,6 +54,9 @@ _TOKEN_INVALID_KEYWORDS: tuple[str, ...] = (
_KEYWORDS_VERIFICATION: tuple[str, ...] = ( _KEYWORDS_VERIFICATION: tuple[str, ...] = (
"validation_required", "validation_required",
"verify your account", "verify your account",
"需要验证",
"验证账号",
"验证身份",
) )
# 合并的完整列表(用于 is_account_level_block_reason 快速判断) # 合并的完整列表(用于 is_account_level_block_reason 快速判断)
@@ -64,6 +67,16 @@ ACCOUNT_BLOCK_REASON_KEYWORDS: tuple[str, ...] = (
*_KEYWORDS_VERIFICATION, *_KEYWORDS_VERIFICATION,
) )
AUTO_REMOVABLE_ACCOUNT_STATE_CODES: frozenset[str] = frozenset(
{
"account_banned",
"account_suspended",
"account_disabled",
"workspace_deactivated",
"account_forbidden",
}
)
def _classify_block_reason(text: str) -> tuple[str, str]: def _classify_block_reason(text: str) -> tuple[str, str]:
"""Return (code, label) based on the oauth_invalid_reason text.""" """Return (code, label) based on the oauth_invalid_reason text."""
@@ -89,6 +102,8 @@ class PoolAccountState:
code: str | None = None # account_banned / account_forbidden / account_blocked code: str | None = None # account_banned / account_forbidden / account_blocked
label: str | None = None label: str | None = None
reason: str | None = None reason: str | None = None
source: str | None = None # metadata / oauth_invalid / oauth_refresh / oauth_request
recoverable: bool = False
def _is_truthy_flag(value: Any) -> bool: def _is_truthy_flag(value: Any) -> bool:
@@ -119,6 +134,11 @@ def _extract_reason(source: dict[str, Any] | None, *fields: str) -> str | None:
return None return None
def _is_workspace_deactivated_reason(reason: str | None) -> bool:
text = _clean_text(reason)
return bool(text and "deactivated_workspace" in text.lower())
def _resolve_from_metadata( def _resolve_from_metadata(
provider_type: str | None, provider_type: str | None,
upstream_metadata: Any, upstream_metadata: Any,
@@ -140,6 +160,7 @@ def _resolve_from_metadata(
code=quota_block.code, code=quota_block.code,
label=quota_block.label, label=quota_block.label,
reason=quota_block.reason, reason=quota_block.reason,
source="metadata",
) )
for source in (provider_bucket, upstream_metadata): for source in (provider_bucket, upstream_metadata):
@@ -152,16 +173,26 @@ def _resolve_from_metadata(
code="account_banned", code="account_banned",
label="账号封禁", label="账号封禁",
reason=reason or "账号已封禁", reason=reason or "账号已封禁",
source="metadata",
) )
if _is_truthy_flag(source.get("is_forbidden")) or _is_truthy_flag( if _is_truthy_flag(source.get("is_forbidden")) or _is_truthy_flag(
source.get("account_disabled") source.get("account_disabled")
): ):
reason = _extract_reason(source, "forbidden_reason", "ban_reason", "reason", "message") reason = _extract_reason(source, "forbidden_reason", "ban_reason", "reason", "message")
if _is_workspace_deactivated_reason(reason):
return PoolAccountState(
blocked=True,
code="workspace_deactivated",
label="工作区停用",
reason=reason or "工作区已停用",
source="metadata",
)
return PoolAccountState( return PoolAccountState(
blocked=True, blocked=True,
code="account_forbidden", code="account_forbidden",
label="访问受限", label="访问受限",
reason=reason or "账号访问受限", reason=reason or "账号访问受限",
source="metadata",
) )
return None return None
@@ -182,6 +213,7 @@ def _resolve_from_oauth_invalid_reason(reason: str | None) -> PoolAccountState |
code=code, code=code,
label=label, label=label,
reason=cleaned or "账号异常", reason=cleaned or "账号异常",
source="oauth_invalid",
) )
if text.startswith(OAUTH_EXPIRED_PREFIX): if text.startswith(OAUTH_EXPIRED_PREFIX):
@@ -191,10 +223,31 @@ def _resolve_from_oauth_invalid_reason(reason: str | None) -> PoolAccountState |
code="oauth_expired", code="oauth_expired",
label="Token 失效", label="Token 失效",
reason=cleaned or "OAuth Token 已过期且无法续期", reason=cleaned or "OAuth Token 已过期且无法续期",
source="oauth_invalid",
recoverable=True,
) )
if text.startswith(OAUTH_REFRESH_FAILED_PREFIX) or text.startswith(OAUTH_REQUEST_FAILED_PREFIX): if text.startswith(OAUTH_REFRESH_FAILED_PREFIX):
return None cleaned = text[len(OAUTH_REFRESH_FAILED_PREFIX) :].strip()
return PoolAccountState(
blocked=False,
code="oauth_refresh_failed",
label="续期失败",
reason=cleaned or "OAuth Token 续期失败",
source="oauth_refresh",
recoverable=True,
)
if text.startswith(OAUTH_REQUEST_FAILED_PREFIX):
cleaned = text[len(OAUTH_REQUEST_FAILED_PREFIX) :].strip()
return PoolAccountState(
blocked=False,
code="oauth_request_failed",
label="请求失败",
reason=cleaned or "账号状态检查失败",
source="oauth_request",
recoverable=True,
)
if text.startswith("["): if text.startswith("["):
return None return None
@@ -207,6 +260,7 @@ def _resolve_from_oauth_invalid_reason(reason: str | None) -> PoolAccountState |
code=code, code=code,
label=label, label=label,
reason=text, reason=text,
source="oauth_invalid",
) )
return None return None
@@ -231,12 +285,29 @@ def resolve_pool_account_state(
return PoolAccountState(blocked=False) return PoolAccountState(blocked=False)
def should_auto_remove_account_state(state: PoolAccountState) -> bool:
"""Whether a resolved account state is safe to auto-remove.
Auto-removal is limited to hard, non-recoverable account abnormalities.
Pure token failures (`oauth_expired`, `oauth_refresh_failed`) and
softer/manual-recoverable states like `account_verification` are excluded.
"""
return bool(
state.blocked
and not state.recoverable
and str(state.code or "").strip().lower() in AUTO_REMOVABLE_ACCOUNT_STATE_CODES
)
__all__ = [ __all__ = [
"ACCOUNT_BLOCK_REASON_KEYWORDS", "ACCOUNT_BLOCK_REASON_KEYWORDS",
"AUTO_REMOVABLE_ACCOUNT_STATE_CODES",
"OAUTH_ACCOUNT_BLOCK_PREFIX", "OAUTH_ACCOUNT_BLOCK_PREFIX",
"OAUTH_EXPIRED_PREFIX", "OAUTH_EXPIRED_PREFIX",
"OAUTH_REFRESH_FAILED_PREFIX", "OAUTH_REFRESH_FAILED_PREFIX",
"OAUTH_REQUEST_FAILED_PREFIX", "OAUTH_REQUEST_FAILED_PREFIX",
"PoolAccountState", "PoolAccountState",
"resolve_pool_account_state", "resolve_pool_account_state",
"should_auto_remove_account_state",
] ]
@@ -91,6 +91,12 @@ class _AccountStateDimension:
blocked_label = snapshot.account_block_label or "账号异常" blocked_label = snapshot.account_block_label or "账号异常"
if blocked_label == "账号封禁": if blocked_label == "账号封禁":
blocked_code = "account_banned" blocked_code = "account_banned"
elif blocked_label == "工作区停用":
blocked_code = "workspace_deactivated"
elif blocked_label == "账号停用":
blocked_code = "account_disabled"
elif blocked_label == "需要验证":
blocked_code = "account_verification"
elif blocked_label == "访问受限": elif blocked_label == "访问受限":
blocked_code = "account_forbidden" blocked_code = "account_forbidden"
else: else:
+116
View File
@@ -0,0 +1,116 @@
"""Helpers for resolving provider metadata without touching detached ORM relations."""
from __future__ import annotations
import time
from typing import Any
def _safe_getattr(obj: Any, attr: str) -> Any:
if obj is None:
return None
try:
return getattr(obj, attr)
except Exception:
# Intentionally broad: ORM objects may raise DetachedInstanceError,
# MissingGreenlet, or other SQLAlchemy errors when accessing
# lazy-loaded attributes on expired/detached objects.
return None
def _normalize_text(value: Any) -> str:
return str(value or "").strip()
def _extract_provider_id(*, endpoint: Any | None = None, key: Any | None = None) -> str | None:
for source in (key, endpoint):
provider_id = _normalize_text(_safe_getattr(source, "provider_id"))
if provider_id:
return provider_id
for provider_obj in (_safe_getattr(endpoint, "provider"), _safe_getattr(key, "provider")):
provider_id = _normalize_text(_safe_getattr(provider_obj, "id"))
if provider_id:
return provider_id
return None
_SNAPSHOT_CACHE: dict[str, tuple[float, dict[str, Any]]] = {}
_SNAPSHOT_TTL: float = 30.0 # seconds
def _load_provider_snapshot(provider_id: str | None) -> dict[str, Any] | None:
normalized_id = _normalize_text(provider_id)
if not normalized_id:
return None
now = time.monotonic()
cached = _SNAPSHOT_CACHE.get(normalized_id)
if cached is not None and (now - cached[0]) < _SNAPSHOT_TTL:
return cached[1]
from src.database import create_session
from src.models.database import Provider
with create_session() as db:
row = db.query(Provider).filter(Provider.id == normalized_id).first()
if row is None:
return None
snapshot = {
"provider_type": _safe_getattr(row, "provider_type"),
"proxy": _safe_getattr(row, "proxy"),
}
_SNAPSHOT_CACHE[normalized_id] = (now, snapshot)
return snapshot
def resolve_provider_type(
*,
endpoint: Any | None = None,
key: Any | None = None,
explicit_provider_type: str | None = None,
decrypted_auth_config: dict[str, Any] | None = None,
) -> str | None:
provider_type = _normalize_text(explicit_provider_type).lower()
if provider_type:
return provider_type
for source in (endpoint, key):
provider_type = _normalize_text(_safe_getattr(source, "provider_type")).lower()
if provider_type:
return provider_type
for provider_obj in (_safe_getattr(endpoint, "provider"), _safe_getattr(key, "provider")):
provider_type = _normalize_text(_safe_getattr(provider_obj, "provider_type")).lower()
if provider_type:
return provider_type
if isinstance(decrypted_auth_config, dict):
provider_type = _normalize_text(decrypted_auth_config.get("provider_type")).lower()
if provider_type:
return provider_type
snapshot = _load_provider_snapshot(_extract_provider_id(endpoint=endpoint, key=key))
provider_type = _normalize_text((snapshot or {}).get("provider_type")).lower()
return provider_type or None
def resolve_provider_proxy(
*,
endpoint: Any | None = None,
key: Any | None = None,
) -> dict[str, Any] | None:
for source in (endpoint, key):
provider_proxy = _safe_getattr(source, "provider_proxy")
if isinstance(provider_proxy, dict):
return provider_proxy
for provider_obj in (_safe_getattr(endpoint, "provider"), _safe_getattr(key, "provider")):
provider_proxy = _safe_getattr(provider_obj, "proxy")
if isinstance(provider_proxy, dict):
return provider_proxy
snapshot = _load_provider_snapshot(_extract_provider_id(endpoint=endpoint, key=key))
provider_proxy = (snapshot or {}).get("proxy")
return provider_proxy if isinstance(provider_proxy, dict) else None
+9 -2
View File
@@ -18,6 +18,7 @@ from typing import Any
from src.core.api_format.metadata import resolve_endpoint_definition from src.core.api_format.metadata import resolve_endpoint_definition
from src.core.provider_types import ProviderType from src.core.provider_types import ProviderType
from src.services.provider.adapters.codex.context import is_codex_compact_request from src.services.provider.adapters.codex.context import is_codex_compact_request
from src.services.provider.provider_context import resolve_provider_type
class UpstreamStreamPolicy(str, Enum): class UpstreamStreamPolicy(str, Enum):
@@ -59,8 +60,14 @@ def get_upstream_stream_policy(
- Codex + openai:compact: follow endpoint/client policy (no hard force). - Codex + openai:compact: follow endpoint/client policy (no hard force).
""" """
provider_obj = getattr(endpoint, "provider", None) # Prefer the caller-supplied provider_type; fall back to detached-safe provider lookup.
pt = str(provider_type or getattr(provider_obj, "provider_type", "") or "").strip().lower() pt = (
resolve_provider_type(
endpoint=endpoint,
explicit_provider_type=provider_type,
)
or ""
)
sig = str(endpoint_sig or getattr(endpoint, "api_format", "") or "").strip().lower() sig = str(endpoint_sig or getattr(endpoint, "api_format", "") or "").strip().lower()
is_codex_cli = pt == ProviderType.CODEX and sig == "openai:cli" is_codex_cli = pt == ProviderType.CODEX and sig == "openai:cli"
is_codex_compact = pt == ProviderType.CODEX and sig == "openai:compact" is_codex_compact = pt == ProviderType.CODEX and sig == "openai:compact"
+10 -25
View File
@@ -21,6 +21,7 @@ from src.core.api_format import (
from src.core.logger import logger from src.core.logger import logger
from src.core.provider_types import ProviderType, normalize_provider_type from src.core.provider_types import ProviderType, normalize_provider_type
from src.services.provider.format import normalize_endpoint_signature from src.services.provider.format import normalize_endpoint_signature
from src.services.provider.provider_context import resolve_provider_type
from src.services.provider.request_context import ( from src.services.provider.request_context import (
get_selected_base_url, get_selected_base_url,
set_selected_base_url, set_selected_base_url,
@@ -137,34 +138,18 @@ def _get_provider_type(
"""尽力获取 Provider.provider_type(用于 Antigravity 等 Provider 特判)。 """尽力获取 Provider.provider_type(用于 Antigravity 等 Provider 特判)。
优先级: 优先级:
1. endpoint.provider.provider_type 1. endpoint.provider_type(如果调用方已做扁平化注入)
2. key.provider.provider_type 2. key.provider.provider_type
3. decrypted_auth_config["provider_type"](OAuth 导入的凭证) 3. decrypted_auth_config["provider_type"](OAuth 导入的凭证)
4. provider_id 对应的 Provider 记录
""" """
try: resolved = resolve_provider_type(
provider = getattr(endpoint, "provider", None) endpoint=endpoint,
if provider is not None: key=key,
pt = getattr(provider, "provider_type", None) decrypted_auth_config=decrypted_auth_config,
if pt: )
return str(pt).lower() if resolved:
except Exception: return resolved
pass
try:
if key is not None:
provider = getattr(key, "provider", None)
if provider is not None:
pt = getattr(provider, "provider_type", None)
if pt:
return str(pt).lower()
except Exception:
pass
# Fallback: OAuth 导入的凭证可能包含 provider_type(如 Kiro)
if decrypted_auth_config:
pt = decrypted_auth_config.get("provider_type")
if isinstance(pt, str) and pt.strip():
return pt.strip().lower()
# Fallback: 历史 Vertex 数据可能缺少 provider_type,但 base_url 已固定到 aiplatform。 # Fallback: 历史 Vertex 数据可能缺少 provider_type,但 base_url 已固定到 aiplatform。
try: try:
@@ -216,14 +216,11 @@ def _clear_oauth_invalid_marker(db: Session, key_id: str) -> dict[str, str]:
old_reason = key.oauth_invalid_reason old_reason = key.oauth_invalid_reason
key.oauth_invalid_at = None key.oauth_invalid_at = None
key.oauth_invalid_reason = None key.oauth_invalid_reason = None
key.is_active = True
db.commit() db.commit()
_run_async_with_fallback(_invalidate_cache_after_clear_oauth_invalid(key_id)) _run_async_with_fallback(_invalidate_cache_after_clear_oauth_invalid(key_id))
logger.info( logger.info("[OK] 手动清除 Key {}... 的 OAuth 失效标记 (原因: {})", key_id[:8], old_reason)
"[OK] 手动清除 Key {}... 的 OAuth 失效标记并自动启用 (原因: {})", key_id[:8], old_reason return {"message": "已清除 OAuth 失效标记"}
)
return {"message": "已清除 OAuth 失效标记,Key 已自动启用"}
def clear_oauth_invalid_response(db: Session, key_id: str) -> dict[str, str]: def clear_oauth_invalid_response(db: Session, key_id: str) -> dict[str, str]:
@@ -15,7 +15,10 @@ from src.core.provider_types import ProviderType, normalize_provider_type
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
from src.services.model.upstream_fetcher import merge_upstream_metadata from src.services.model.upstream_fetcher import merge_upstream_metadata
from src.services.provider.pool import redis_ops as pool_redis from src.services.provider.pool import redis_ops as pool_redis
from src.services.provider.pool.account_state import resolve_pool_account_state from src.services.provider.pool.account_state import (
resolve_pool_account_state,
should_auto_remove_account_state,
)
from src.services.provider.pool.config import parse_pool_config from src.services.provider.pool.config import parse_pool_config
from src.services.provider_keys.key_side_effects import run_delete_key_side_effects from src.services.provider_keys.key_side_effects import run_delete_key_side_effects
from src.services.provider_keys.quota_refresh import ( from src.services.provider_keys.quota_refresh import (
@@ -88,7 +91,7 @@ async def refresh_provider_quota_for_provider(
if provider_type not in QUOTA_REFRESH_PROVIDER_TYPES: if provider_type not in QUOTA_REFRESH_PROVIDER_TYPES:
raise InvalidRequestException("仅支持 Codex / Antigravity / Kiro 类型的 Provider 刷新限额") raise InvalidRequestException("仅支持 Codex / Antigravity / Kiro 类型的 Provider 刷新限额")
pool_cfg = parse_pool_config(getattr(provider, "config", None)) pool_cfg = parse_pool_config(getattr(provider, "config", None))
auto_remove_banned_keys = bool(pool_cfg and pool_cfg.auto_remove_banned_keys) auto_remove_abnormal_keys = bool(pool_cfg and pool_cfg.auto_remove_banned_keys)
selected_key_ids: list[str] | None = None selected_key_ids: list[str] | None = None
if key_ids is not None: if key_ids is not None:
@@ -201,7 +204,7 @@ async def refresh_provider_quota_for_provider(
if rid: if rid:
result_index_by_key_id[rid] = result result_index_by_key_id[rid] = result
if metadata_updates or state_updates or auto_remove_banned_keys: if metadata_updates or state_updates or auto_remove_abnormal_keys:
for key in keys: for key in keys:
key_dirty = False key_dirty = False
if key.id in metadata_updates: if key.id in metadata_updates:
@@ -216,13 +219,13 @@ async def refresh_provider_quota_for_provider(
setattr(key, field_name, field_value) setattr(key, field_name, field_value)
key_dirty = True key_dirty = True
if auto_remove_banned_keys: if auto_remove_abnormal_keys:
account_state = resolve_pool_account_state( account_state = resolve_pool_account_state(
provider_type=provider_type, provider_type=provider_type,
upstream_metadata=getattr(key, "upstream_metadata", None), upstream_metadata=getattr(key, "upstream_metadata", None),
oauth_invalid_reason=getattr(key, "oauth_invalid_reason", None), oauth_invalid_reason=getattr(key, "oauth_invalid_reason", None),
) )
if account_state.blocked: if should_auto_remove_account_state(account_state):
key_id = str(getattr(key, "id", "") or "") key_id = str(getattr(key, "id", "") or "")
auto_removed_contexts.append( auto_removed_contexts.append(
( (
@@ -261,7 +264,7 @@ async def refresh_provider_quota_for_provider(
deleted_key_allowed_models=allowed_models, deleted_key_allowed_models=allowed_models,
) )
logger.warning( logger.warning(
"[QUOTA_REFRESH] Provider {}: auto removed {} banned key(s): {}", "[QUOTA_REFRESH] Provider {}: auto removed {} abnormal key(s): {}",
provider_id, provider_id,
len(auto_removed_contexts), len(auto_removed_contexts),
[ctx[0][:8] for ctx in auto_removed_contexts if ctx[0]], [ctx[0][:8] for ctx in auto_removed_contexts if ctx[0]],
@@ -3,15 +3,14 @@
行为: 行为:
- 当 provider.pool_advanced.probing_enabled=true 时启用 - 当 provider.pool_advanced.probing_enabled=true 时启用
- Key 在静默超过 probing_interval_minutes 后,主动触发额度刷新 - Key 以固定间隔主动触发额度刷新,用于检查 OAuth / 额度状态
- Key 一旦被实际请求使用(last_used_at 变新),探测冷却自动重置 - 实际请求使用不会跳过定期探测;探测节流仅由刷新时间与主动探测时间控制
""" """
from __future__ import annotations from __future__ import annotations
import os import os
import time import time
from datetime import datetime, timezone
from typing import Any from typing import Any
from sqlalchemy.orm import load_only from sqlalchemy.orm import load_only
@@ -40,16 +39,6 @@ def _probe_stamp_key(provider_id: str, key_id: str) -> str:
return f"{_REDIS_PREFIX}:{provider_id}:{key_id}" return f"{_REDIS_PREFIX}:{provider_id}:{key_id}"
def _to_unix_seconds(value: datetime | None) -> int | None:
if not isinstance(value, datetime):
return None
dt = value if value.tzinfo else value.replace(tzinfo=timezone.utc)
try:
return int(dt.timestamp())
except Exception:
return None
def _to_float(value: Any) -> float | None: def _to_float(value: Any) -> float | None:
if isinstance(value, bool): if isinstance(value, bool):
return None return None
@@ -121,13 +110,12 @@ def _select_probe_key_ids(
key_id = str(getattr(key, "id", "") or "") key_id = str(getattr(key, "id", "") or "")
if not key_id: if not key_id:
continue continue
last_used_ts = _to_unix_seconds(getattr(key, "last_used_at", None))
quota_updated_ts = _extract_quota_updated_at( quota_updated_ts = _extract_quota_updated_at(
provider_type, provider_type,
getattr(key, "upstream_metadata", None), getattr(key, "upstream_metadata", None),
) )
last_probe_ts = last_probe_timestamps.get(key_id) last_probe_ts = last_probe_timestamps.get(key_id)
anchor_ts = max(last_used_ts or 0, quota_updated_ts or 0, last_probe_ts or 0) anchor_ts = max(quota_updated_ts or 0, last_probe_ts or 0)
if anchor_ts <= 0 or (now_ts - anchor_ts) >= interval_seconds: if anchor_ts <= 0 or (now_ts - anchor_ts) >= interval_seconds:
stale.append((anchor_ts, key_id)) stale.append((anchor_ts, key_id))
@@ -338,7 +326,6 @@ class PoolQuotaProbeScheduler:
load_only( load_only(
ProviderAPIKey.id, ProviderAPIKey.id,
ProviderAPIKey.provider_id, ProviderAPIKey.provider_id,
ProviderAPIKey.last_used_at,
ProviderAPIKey.upstream_metadata, ProviderAPIKey.upstream_metadata,
) )
) )
@@ -50,6 +50,15 @@ def _extract_reason(source: dict[str, Any], *fields: str) -> str | None:
return None return None
def _is_workspace_deactivated_reason(reason: str | None) -> bool:
if not reason:
return False
lowered = reason.strip().lower()
if not lowered:
return False
return "deactivated_workspace" in lowered
def _pct_is_exhausted(value: Any) -> bool: def _pct_is_exhausted(value: Any) -> bool:
pct = _to_float(value) pct = _to_float(value)
if pct is None: if pct is None:
@@ -218,6 +227,13 @@ class CodexQuotaReader(PoolQuotaReader):
if not _is_truthy_flag(self._data.get("account_disabled")): if not _is_truthy_flag(self._data.get("account_disabled")):
return AccountBlockResult(blocked=False) return AccountBlockResult(blocked=False)
reason = _extract_reason(self._data, "forbidden_reason", "ban_reason", "reason", "message") reason = _extract_reason(self._data, "forbidden_reason", "ban_reason", "reason", "message")
if _is_workspace_deactivated_reason(reason):
return AccountBlockResult(
blocked=True,
code="workspace_deactivated",
label="工作区停用",
reason=reason or "工作区已停用",
)
return AccountBlockResult( return AccountBlockResult(
blocked=True, blocked=True,
code="account_forbidden", code="account_forbidden",
@@ -75,9 +75,8 @@ async def refresh_antigravity_key_quota(
fetch_ctx, timeout_seconds=10.0 fetch_ctx, timeout_seconds=10.0
) )
except AntigravityAccountForbiddenException as e: except AntigravityAccountForbiddenException as e:
# 对齐 AM:所有 403 一律标记 is_forbidden 并停用 # 对齐 AM:所有 403 一律标记 is_forbidden;手动启用状态保持不变。
state_updates[key.id] = { state_updates[key.id] = {
"is_active": False,
"oauth_invalid_at": datetime.now(timezone.utc), "oauth_invalid_at": datetime.now(timezone.utc),
"oauth_invalid_reason": f"账户访问被禁止: {e.reason or e.message}", "oauth_invalid_reason": f"账户访问被禁止: {e.reason or e.message}",
} }
@@ -91,7 +90,7 @@ async def refresh_antigravity_key_quota(
} }
} }
logger.warning( logger.warning(
"[ANTIGRAVITY_QUOTA] Key {} 账户访问被禁止,已自动停用: {}", "[ANTIGRAVITY_QUOTA] Key {} 账户访问被禁止,已更新账号状态: {}",
key.id, key.id,
e.reason or e.message, e.reason or e.message,
) )
@@ -101,7 +100,7 @@ async def refresh_antigravity_key_quota(
"status": "forbidden", "status": "forbidden",
"message": f"账户访问被禁止: {e.reason or e.message}", "message": f"账户访问被禁止: {e.reason or e.message}",
"is_forbidden": True, "is_forbidden": True,
"auto_disabled": True, "auto_disabled": False,
} }
if ok and upstream_meta: if ok and upstream_meta:
@@ -114,7 +113,6 @@ async def refresh_antigravity_key_quota(
state_updates[key.id] = { state_updates[key.id] = {
"oauth_invalid_at": None, "oauth_invalid_at": None,
"oauth_invalid_reason": None, "oauth_invalid_reason": None,
"is_active": True,
} }
return { return {
"key_id": key.id, "key_id": key.id,
@@ -280,7 +280,6 @@ async def refresh_codex_key_quota(
state_updates[key.id] = { state_updates[key.id] = {
"oauth_invalid_at": None, "oauth_invalid_at": None,
"oauth_invalid_reason": None, "oauth_invalid_reason": None,
"is_active": True,
} }
return { return {
"key_id": key.id, "key_id": key.id,
@@ -352,7 +351,6 @@ async def refresh_codex_key_quota(
state_updates[key.id] = { state_updates[key.id] = {
"oauth_invalid_at": None, "oauth_invalid_at": None,
"oauth_invalid_reason": None, "oauth_invalid_reason": None,
"is_active": True,
} }
return { return {
"key_id": key.id, "key_id": key.id,
@@ -76,9 +76,8 @@ async def refresh_kiro_key_quota(
proxy_config=proxy_config, proxy_config=proxy_config,
) )
except KiroAccountBannedException as e: except KiroAccountBannedException as e:
# 账户被封禁,自动停用并标记 # 账户被封禁,记录账号状态;手动启用状态保持不变。
state_updates[key.id] = { state_updates[key.id] = {
"is_active": False,
"oauth_invalid_at": datetime.now(timezone.utc), "oauth_invalid_at": datetime.now(timezone.utc),
"oauth_invalid_reason": f"账户已封禁: {e.reason or e.message}", "oauth_invalid_reason": f"账户已封禁: {e.reason or e.message}",
} }
@@ -92,7 +91,7 @@ async def refresh_kiro_key_quota(
} }
} }
logger.warning( logger.warning(
"[KIRO_QUOTA] Key {} 账户已封禁,已自动停用: {}", "[KIRO_QUOTA] Key {} 账户已封禁,已更新账号状态: {}",
key.id, key.id,
e.reason or e.message, e.reason or e.message,
) )
@@ -102,18 +101,17 @@ async def refresh_kiro_key_quota(
"status": "banned", "status": "banned",
"message": f"账户已封禁: {e.reason or e.message}", "message": f"账户已封禁: {e.reason or e.message}",
"is_banned": True, "is_banned": True,
"auto_disabled": True, "auto_disabled": False,
} }
except RuntimeError as e: except RuntimeError as e:
error_msg = str(e) error_msg = str(e)
# 检查是否需要标记账号异常 # 检查是否需要标记账号异常
if "401" in error_msg or "认证失败" in error_msg: if "401" in error_msg or "认证失败" in error_msg:
state_updates[key.id] = { state_updates[key.id] = {
"is_active": False,
"oauth_invalid_at": datetime.now(timezone.utc), "oauth_invalid_at": datetime.now(timezone.utc),
"oauth_invalid_reason": "Kiro Token 无效或已过期", "oauth_invalid_reason": "Kiro Token 无效或已过期",
} }
logger.warning("[KIRO_QUOTA] Key {} Token 无效,已标记为异常并自动停用", key.id) logger.warning("[KIRO_QUOTA] Key {} Token 无效,已标记为异常", key.id)
return { return {
"key_id": key.id, "key_id": key.id,
"key_name": key.name, "key_name": key.name,
@@ -137,7 +135,6 @@ async def refresh_kiro_key_quota(
state_updates[key.id] = { state_updates[key.id] = {
"oauth_invalid_at": None, "oauth_invalid_at": None,
"oauth_invalid_reason": None, "oauth_invalid_reason": None,
"is_active": True,
} }
# 如果 auth_config 有更新(例如 token 刷新),也需要更新 # 如果 auth_config 有更新(例如 token 刷新),也需要更新
@@ -1,8 +1,12 @@
from __future__ import annotations from __future__ import annotations
import sys
import types
from types import SimpleNamespace from types import SimpleNamespace
from typing import Any, cast from typing import Any, cast
import pytest
from src.services.orchestration.error_handler import ErrorHandlerService from src.services.orchestration.error_handler import ErrorHandlerService
@@ -31,11 +35,16 @@ def _build_key() -> SimpleNamespace:
) )
def test_mark_oauth_key_blocked_auto_remove_enabled(monkeypatch: Any) -> None: def test_mark_oauth_key_blocked_auto_remove_enabled_skips_verification_state(
monkeypatch: Any,
) -> None:
db = _FakeDB() db = _FakeDB()
service = ErrorHandlerService(db=cast(Any, db)) service = ErrorHandlerService(db=cast(Any, db))
key = _build_key() key = _build_key()
provider = SimpleNamespace(config={"pool_advanced": {"auto_remove_banned_keys": True}}) provider = SimpleNamespace(
provider_type="codex",
config={"pool_advanced": {"auto_remove_banned_keys": True}},
)
monkeypatch.setattr( monkeypatch.setattr(
ErrorHandlerService, ErrorHandlerService,
@@ -46,8 +55,8 @@ def test_mark_oauth_key_blocked_auto_remove_enabled(monkeypatch: Any) -> None:
service._mark_oauth_key_blocked(cast(Any, key), "req-1", provider=cast(Any, provider)) service._mark_oauth_key_blocked(cast(Any, key), "req-1", provider=cast(Any, provider))
assert db.commit_count == 1 assert db.commit_count == 1
assert db.deleted == [key] assert db.deleted == []
assert key.is_active is False assert key.is_active is True
assert str(key.oauth_invalid_reason).startswith("[ACCOUNT_BLOCK] ") assert str(key.oauth_invalid_reason).startswith("[ACCOUNT_BLOCK] ")
@@ -61,5 +70,62 @@ def test_mark_oauth_key_blocked_auto_remove_disabled() -> None:
assert db.commit_count == 1 assert db.commit_count == 1
assert db.deleted == [] assert db.deleted == []
assert key.is_active is False assert key.is_active is True
assert str(key.oauth_invalid_reason).startswith("[ACCOUNT_BLOCK] ") assert str(key.oauth_invalid_reason).startswith("[ACCOUNT_BLOCK] ")
def test_mark_oauth_key_blocked_auto_remove_enabled_for_deactivated_account(
monkeypatch: Any,
) -> None:
db = _FakeDB()
service = ErrorHandlerService(db=cast(Any, db))
key = _build_key()
provider = SimpleNamespace(
provider_type="codex",
config={"pool_advanced": {"auto_remove_banned_keys": True}},
)
monkeypatch.setattr(
ErrorHandlerService,
"_schedule_auto_cleanup_after_delete",
staticmethod(lambda **kwargs: None),
)
service._mark_oauth_key_blocked(
cast(Any, key),
"req-1",
reason="account has been deactivated",
provider=cast(Any, provider),
)
assert db.commit_count == 1
assert db.deleted == [key]
assert key.is_active is True
assert key.oauth_invalid_reason == "[ACCOUNT_BLOCK] account has been deactivated"
@pytest.mark.asyncio
async def test_verify_oauth_before_account_block_skips_when_refresh_marks_token_expired(
monkeypatch: Any,
) -> None:
db = _FakeDB()
service = ErrorHandlerService(db=cast(Any, db))
key = _build_key()
endpoint = SimpleNamespace()
fake_module = types.ModuleType("src.services.provider.auth")
async def _fake_get_provider_auth(*_args: Any, **_kwargs: Any) -> None:
key.oauth_invalid_reason = "[OAUTH_EXPIRED] token expired"
fake_module.get_provider_auth = _fake_get_provider_auth
monkeypatch.setitem(sys.modules, "src.services.provider.auth", fake_module)
should_mark = await service._verify_oauth_before_account_block(
endpoint=cast(Any, endpoint),
key=cast(Any, key),
request_id="req-1",
candidate_reason="Google 要求验证账号",
)
assert should_mark is False
+38 -1
View File
@@ -2,7 +2,10 @@
from __future__ import annotations from __future__ import annotations
from src.services.provider.pool.account_state import resolve_pool_account_state from src.services.provider.pool.account_state import (
resolve_pool_account_state,
should_auto_remove_account_state,
)
def test_resolve_from_kiro_banned_metadata() -> None: def test_resolve_from_kiro_banned_metadata() -> None:
@@ -41,6 +44,18 @@ def test_resolve_from_structured_oauth_reason_verification() -> None:
assert state.reason == "Google requires verification" assert state.reason == "Google requires verification"
def test_resolve_from_structured_oauth_reason_verification_chinese() -> None:
state = resolve_pool_account_state(
provider_type="codex",
upstream_metadata=None,
oauth_invalid_reason="[ACCOUNT_BLOCK] Google 要求验证账号",
)
assert state.blocked is True
assert state.code == "account_verification"
assert state.label == "需要验证"
assert state.reason == "Google 要求验证账号"
def test_resolve_from_structured_oauth_reason_suspended() -> None: def test_resolve_from_structured_oauth_reason_suspended() -> None:
state = resolve_pool_account_state( state = resolve_pool_account_state(
provider_type="codex", provider_type="codex",
@@ -155,3 +170,25 @@ def test_request_failed_prefix_does_not_block() -> None:
oauth_invalid_reason="[REQUEST_FAILED] Codex 账户访问受限 (403)", oauth_invalid_reason="[REQUEST_FAILED] Codex 账户访问受限 (403)",
) )
assert state.blocked is False assert state.blocked is False
def test_auto_remove_state_excludes_token_expired_and_verification() -> None:
expired = resolve_pool_account_state(
provider_type="codex",
upstream_metadata=None,
oauth_invalid_reason="[OAUTH_EXPIRED] token invalidated",
)
verification = resolve_pool_account_state(
provider_type="codex",
upstream_metadata=None,
oauth_invalid_reason="[ACCOUNT_BLOCK] Google 要求验证账号",
)
disabled = resolve_pool_account_state(
provider_type="codex",
upstream_metadata=None,
oauth_invalid_reason="[ACCOUNT_BLOCK] account has been deactivated",
)
assert should_auto_remove_account_state(expired) is False
assert should_auto_remove_account_state(verification) is False
assert should_auto_remove_account_state(disabled) is True
@@ -25,7 +25,7 @@ def test_select_probe_key_ids_selects_silent_keys_only() -> None:
keys = [ keys = [
_key("k1"), # never used, should be probed _key("k1"), # never used, should be probed
_key("k2", last_used_at=now - timedelta(minutes=2)), # recently used, skip _key("k2", last_used_at=now - timedelta(minutes=2)), # recently used,仍可定期探测
_key( _key(
"k3", "k3",
upstream_metadata={"codex": {"updated_at": now_ts - (20 * 60)}}, upstream_metadata={"codex": {"updated_at": now_ts - (20 * 60)}},
@@ -40,10 +40,10 @@ def test_select_probe_key_ids_selects_silent_keys_only() -> None:
last_probe_timestamps={}, last_probe_timestamps={},
limit=0, limit=0,
) )
assert selected == ["k1", "k3"] assert selected == ["k1", "k2", "k3"]
def test_select_probe_key_ids_resets_probe_window_after_key_usage() -> None: def test_select_probe_key_ids_keeps_periodic_probe_even_after_recent_usage() -> None:
now = datetime(2026, 3, 5, 12, 0, 0, tzinfo=timezone.utc) now = datetime(2026, 3, 5, 12, 0, 0, tzinfo=timezone.utc)
now_ts = int(now.timestamp()) now_ts = int(now.timestamp())
@@ -55,7 +55,7 @@ def test_select_probe_key_ids_resets_probe_window_after_key_usage() -> None:
) )
] ]
# 上一次主动探测非常早,但 key 刚刚被真实流量使用,应跳过本次探测 # 即使 key 刚刚被真实流量使用,只要上次额度刷新/主动探测已过窗口,仍应继续定期探测
selected = _select_probe_key_ids( selected = _select_probe_key_ids(
keys=keys, # type: ignore[arg-type] keys=keys, # type: ignore[arg-type]
provider_type="codex", provider_type="codex",
@@ -64,7 +64,7 @@ def test_select_probe_key_ids_resets_probe_window_after_key_usage() -> None:
last_probe_timestamps={"k1": now_ts - (25 * 60)}, last_probe_timestamps={"k1": now_ts - (25 * 60)},
limit=0, limit=0,
) )
assert selected == [] assert selected == ["k1"]
def test_select_probe_key_ids_applies_limit_by_oldest_anchor_first() -> None: def test_select_probe_key_ids_applies_limit_by_oldest_anchor_first() -> None:
@@ -72,9 +72,9 @@ def test_select_probe_key_ids_applies_limit_by_oldest_anchor_first() -> None:
now_ts = int(now.timestamp()) now_ts = int(now.timestamp())
keys = [ keys = [
_key("k1", last_used_at=now - timedelta(minutes=60)), _key("k1", upstream_metadata={"codex": {"updated_at": now_ts - (60 * 60)}}),
_key("k2", last_used_at=now - timedelta(minutes=50)), _key("k2", upstream_metadata={"codex": {"updated_at": now_ts - (50 * 60)}}),
_key("k3", last_used_at=now - timedelta(minutes=40)), _key("k3", upstream_metadata={"codex": {"updated_at": now_ts - (40 * 60)}}),
] ]
selected = _select_probe_key_ids( selected = _select_probe_key_ids(
@@ -83,6 +83,20 @@ def test_account_state_takes_priority_over_manual_disabled() -> None:
assert summary.reason == "account_forbidden" assert summary.reason == "account_forbidden"
def test_workspace_deactivated_uses_specific_reason_code() -> None:
dimensions = evaluate_pool_scheduling_dimensions(
_snapshot(
account_blocked=True,
account_block_label="工作区停用",
account_block_reason="deactivated_workspace",
)
)
summary = summarize_pool_scheduling_dimensions(dimensions)
assert summary.status == "blocked"
assert summary.reason == "workspace_deactivated"
def test_summary_degraded_when_cost_reaches_soft_threshold() -> None: def test_summary_degraded_when_cost_reaches_soft_threshold() -> None:
dimensions = evaluate_pool_scheduling_dimensions( dimensions = evaluate_pool_scheduling_dimensions(
_snapshot(cost_window_usage=8200, cost_limit=10000, cost_soft_threshold_percent=80) _snapshot(cost_window_usage=8200, cost_limit=10000, cost_soft_threshold_percent=80)
+26 -6
View File
@@ -47,9 +47,7 @@ class _FakeSessionCtx:
return False return False
def _install_module( def _install_module(monkeypatch: pytest.MonkeyPatch, name: str, attrs: dict[str, Any]) -> None:
monkeypatch: pytest.MonkeyPatch, name: str, attrs: dict[str, Any]
) -> None:
fake_module = types.ModuleType(name) fake_module = types.ModuleType(name)
for key, value in attrs.items(): for key, value in attrs.items():
setattr(fake_module, key, value) setattr(fake_module, key, value)
@@ -110,9 +108,7 @@ def test_mark_refresh_token_invalid_persists_detached_key(
assert fake_db.committed is True assert fake_db.committed is True
assert key.oauth_invalid_at is not None assert key.oauth_invalid_at is not None
assert row.oauth_invalid_at is not None assert row.oauth_invalid_at is not None
assert str(key.oauth_invalid_reason).startswith( assert str(key.oauth_invalid_reason).startswith("[REFRESH_FAILED] Token 续期失败 (401)")
"[REFRESH_FAILED] Token 续期失败 (401)"
)
assert "refresh_token_reused" in str(row.oauth_invalid_reason) assert "refresh_token_reused" in str(row.oauth_invalid_reason)
@@ -153,6 +149,30 @@ def test_persist_refreshed_token_clears_legacy_token_invalidated_account_block(
assert key.oauth_invalid_reason is None assert key.oauth_invalid_reason is None
def test_persist_refreshed_token_preserves_true_account_block(
monkeypatch: pytest.MonkeyPatch,
) -> None:
key = SimpleNamespace(
id="key-1",
api_key="old-api",
auth_config="old-config",
oauth_invalid_at=datetime.now(timezone.utc),
oauth_invalid_reason="[ACCOUNT_BLOCK] Google requires verification",
)
monkeypatch.setattr(
module, "object_session", lambda _key: (_ for _ in ()).throw(RuntimeError())
)
monkeypatch.setattr(module.crypto_service, "encrypt", lambda value: f"enc:{value}")
module._persist_refreshed_token(key, "new-token", {"refresh_token": "rt-2"})
assert key.api_key == "enc:new-token"
assert key.auth_config == 'enc:{"refresh_token": "rt-2"}'
assert key.oauth_invalid_at is not None
assert key.oauth_invalid_reason == "[ACCOUNT_BLOCK] Google requires verification"
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_refresh_generic_oauth_token_persists_enriched_account_name( async def test_refresh_generic_oauth_token_persists_enriched_account_name(
monkeypatch: pytest.MonkeyPatch, monkeypatch: pytest.MonkeyPatch,
@@ -1,6 +1,8 @@
from __future__ import annotations from __future__ import annotations
import json import json
import sys
import types
from types import SimpleNamespace from types import SimpleNamespace
import pytest import pytest
@@ -9,6 +11,32 @@ from src.core.vertex_auth import VertexAuthService
from src.services.provider.auth import get_provider_auth from src.services.provider.auth import get_provider_auth
class _FakeQuery:
def __init__(self, row: object | None) -> None:
self._row = row
def filter(self, *_args: object, **_kwargs: object) -> "_FakeQuery":
return self
def first(self) -> object | None:
return self._row
class _FakeSessionCtx:
def __init__(self, row: object | None) -> None:
self._row = row
def __enter__(self) -> "_FakeSessionCtx":
return self
def __exit__(self, exc_type: object, exc: object, tb: object) -> bool:
_ = exc_type, exc, tb
return False
def query(self, _model: object) -> _FakeQuery:
return _FakeQuery(self._row)
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_get_provider_auth_vertex_service_account_uses_provider_proxy( async def test_get_provider_auth_vertex_service_account_uses_provider_proxy(
monkeypatch: pytest.MonkeyPatch, monkeypatch: pytest.MonkeyPatch,
@@ -124,3 +152,76 @@ async def test_get_provider_auth_vertex_service_account_prefers_key_proxy(
assert auth is not None assert auth is not None
assert captured["proxy_config"] == key_proxy assert captured["proxy_config"] == key_proxy
@pytest.mark.asyncio
async def test_get_provider_auth_vertex_service_account_uses_provider_id_lookup_without_touching_endpoint_provider(
monkeypatch: pytest.MonkeyPatch,
) -> None:
sa_json = {
"client_email": "[email protected]",
"private_key": "-----BEGIN PRIVATE KEY-----\nTEST\n-----END PRIVATE KEY-----\n",
"project_id": "demo-project",
}
provider_proxy = {"node_id": "provider-node", "enabled": True}
class _DetachedEndpoint:
provider_id = "provider-1"
@property
def provider(self) -> object:
raise RuntimeError("detached endpoint provider should not be lazy-loaded")
fake_provider = SimpleNamespace(
id="provider-1", provider_type="vertex_ai", proxy=provider_proxy
)
fake_database = types.ModuleType("src.database")
fake_database.create_session = lambda: _FakeSessionCtx(fake_provider)
fake_models = types.ModuleType("src.models.database")
fake_models.Provider = type("Provider", (), {"id": "id"})
monkeypatch.setitem(sys.modules, "src.database", fake_database)
monkeypatch.setitem(sys.modules, "src.models.database", fake_models)
monkeypatch.setattr(
"src.core.crypto.crypto_service.decrypt",
lambda value: json.dumps(sa_json) if value == "enc_cfg" else "",
)
captured: dict[str, object] = {}
def _fake_build_proxy_client_kwargs(
proxy_config: dict[str, object] | None = None,
*,
timeout: float = 30.0,
**_: object,
) -> dict[str, object]:
captured["proxy_config"] = proxy_config
return {"timeout": timeout}
async def _fake_get_access_token(
self: VertexAuthService,
*,
httpx_client_kwargs: dict[str, object] | None = None,
) -> str:
captured["httpx_client_kwargs"] = httpx_client_kwargs
return "ya29.test-token"
monkeypatch.setattr(
"src.services.proxy_node.resolver.build_proxy_client_kwargs",
_fake_build_proxy_client_kwargs,
)
monkeypatch.setattr(VertexAuthService, "get_access_token", _fake_get_access_token)
key = SimpleNamespace(
auth_type="service_account",
auth_config="enc_cfg",
api_key="enc_key",
provider_id="provider-1",
proxy=None,
)
auth = await get_provider_auth(_DetachedEndpoint(), key) # type: ignore[arg-type]
assert auth is not None
assert captured["proxy_config"] == provider_proxy
assert captured["httpx_client_kwargs"] == {"timeout": 30}
@@ -244,10 +244,10 @@ def test_clear_oauth_invalid_response_invalidates_caches(
result = command_module.clear_oauth_invalid_response(cast(Any, db), key_id="key-1") result = command_module.clear_oauth_invalid_response(cast(Any, db), key_id="key-1")
assert result["message"] == "已清除 OAuth 失效标记,Key 已自动启用" assert result["message"] == "已清除 OAuth 失效标记"
assert key.oauth_invalid_at is None assert key.oauth_invalid_at is None
assert key.oauth_invalid_reason is None assert key.oauth_invalid_reason is None
assert key.is_active is True assert key.is_active is False
assert db.commit_count == 1 assert db.commit_count == 1
assert cache_calls == [("key", "key-1"), ("models", None)] assert cache_calls == [("key", "key-1"), ("models", None)]
@@ -771,10 +771,10 @@ async def test_antigravity_refresher_forbidden_collects_updates_without_commit(
) )
assert result["status"] == "forbidden" assert result["status"] == "forbidden"
assert result["auto_disabled"] is True assert result["auto_disabled"] is False
assert key.is_active is True assert key.is_active is True
assert key.oauth_invalid_reason is None assert key.oauth_invalid_reason is None
assert state_updates["k1"]["is_active"] is False assert "is_active" not in state_updates["k1"]
assert state_updates["k1"]["oauth_invalid_reason"].startswith("账户访问被禁止") assert state_updates["k1"]["oauth_invalid_reason"].startswith("账户访问被禁止")
assert metadata_updates["k1"]["antigravity"]["is_forbidden"] is True assert metadata_updates["k1"]["antigravity"]["is_forbidden"] is True
assert db.commit_count == 0 assert db.commit_count == 0
@@ -911,7 +911,7 @@ async def test_kiro_refresher_runtime_401_marks_key_invalid(
assert "401" in result["message"] assert "401" in result["message"]
assert key.is_active is True assert key.is_active is True
assert key.oauth_invalid_reason is None assert key.oauth_invalid_reason is None
assert state_updates["k1"]["is_active"] is False assert "is_active" not in state_updates["k1"]
assert state_updates["k1"]["oauth_invalid_reason"] == "Kiro Token 无效或已过期" assert state_updates["k1"]["oauth_invalid_reason"] == "Kiro Token 无效或已过期"
assert db.commit_count == 0 assert db.commit_count == 0
@@ -414,3 +414,53 @@ async def test_refresh_provider_quota_auto_removes_banned_keys_when_enabled(
assert result["results"][0]["auto_removed"] is True assert result["results"][0]["auto_removed"] is True
assert deleted_side_effect_calls == [("p1", ["gpt-4o"])] assert deleted_side_effect_calls == [("p1", ["gpt-4o"])]
assert redis_cleared == [("p1", "k1"), ("p1", "k1")] assert redis_cleared == [("p1", "k1"), ("p1", "k1")]
@pytest.mark.asyncio
async def test_refresh_provider_quota_does_not_auto_remove_oauth_expired_keys(
monkeypatch: pytest.MonkeyPatch,
) -> None:
provider = SimpleNamespace(
id="p1",
provider_type=ProviderType.CODEX,
endpoints=[SimpleNamespace(api_format="openai:cli", is_active=True)],
config={"pool_advanced": {"auto_remove_banned_keys": True}},
)
key = SimpleNamespace(
id="k1",
name="K1",
provider_id="p1",
allowed_models=None,
upstream_metadata={},
is_active=True,
oauth_invalid_at=None,
oauth_invalid_reason=None,
)
db = _FakeDB(provider=provider, keys=[key])
async def _fake_handler(**kwargs: Any) -> dict[str, Any]:
state_updates = kwargs["state_updates"]
state_updates["k1"] = {
"oauth_invalid_at": "expired-at",
"oauth_invalid_reason": "[OAUTH_EXPIRED] token invalidated",
}
return {"key_id": "k1", "key_name": "K1", "status": "error", "message": "expired"}
monkeypatch.setattr(
quota_service_module,
"_select_refresh_endpoint",
lambda provider, provider_type: provider.endpoints[0],
)
monkeypatch.setattr(
quota_service_module, "_resolve_quota_refresh_handler", lambda _: _fake_handler
)
result = await refresh_provider_quota_for_provider(
db=cast(Any, db),
provider_id="p1",
codex_wham_usage_url="https://example.test/wham/usage",
)
assert result["auto_removed"] == 0
assert db.deleted == []
assert key.oauth_invalid_reason == "[OAUTH_EXPIRED] token invalidated"
@@ -1,5 +1,7 @@
from __future__ import annotations from __future__ import annotations
import sys
import types
from dataclasses import dataclass from dataclasses import dataclass
from types import SimpleNamespace from types import SimpleNamespace
@@ -16,6 +18,33 @@ class _DummyEndpoint:
api_format: str api_format: str
custom_path: str | None = None custom_path: str | None = None
provider: object | None = None provider: object | None = None
provider_id: str | None = None
class _FakeQuery:
def __init__(self, row: object | None) -> None:
self._row = row
def filter(self, *_args: object, **_kwargs: object) -> "_FakeQuery":
return self
def first(self) -> object | None:
return self._row
class _FakeSessionCtx:
def __init__(self, row: object | None) -> None:
self._row = row
def __enter__(self) -> "_FakeSessionCtx":
return self
def __exit__(self, exc_type: object, exc: object, tb: object) -> bool:
_ = exc_type, exc, tb
return False
def query(self, _model: object) -> _FakeQuery:
return _FakeQuery(self._row)
def test_codex_openai_cli_uses_responses_path_without_v1_prefix() -> None: def test_codex_openai_cli_uses_responses_path_without_v1_prefix() -> None:
@@ -56,14 +85,16 @@ def test_codex_openai_cli_uses_compact_suffix_when_context_marked_compact() -> N
api_format="openai:cli", api_format="openai:cli",
provider=SimpleNamespace(provider_type="codex"), provider=SimpleNamespace(provider_type="codex"),
) )
set_codex_request_context(CodexRequestContext(is_compact=True)) try:
url = build_provider_url( set_codex_request_context(CodexRequestContext(is_compact=True))
endpoint, # type: ignore[arg-type] url = build_provider_url(
path_params={"model": "ignored"}, endpoint, # type: ignore[arg-type]
is_stream=False, path_params={"model": "ignored"},
) is_stream=False,
assert url == "https://chatgpt.com/backend-api/codex/responses/compact" )
set_codex_request_context(None) assert url == "https://chatgpt.com/backend-api/codex/responses/compact"
finally:
set_codex_request_context(None)
def test_codex_openai_compact_uses_compact_path_without_v1_prefix() -> None: def test_codex_openai_compact_uses_compact_path_without_v1_prefix() -> None:
@@ -78,3 +109,34 @@ def test_codex_openai_compact_uses_compact_path_without_v1_prefix() -> None:
is_stream=False, is_stream=False,
) )
assert url == "https://chatgpt.com/backend-api/codex/responses/compact" assert url == "https://chatgpt.com/backend-api/codex/responses/compact"
def test_codex_openai_cli_uses_provider_id_lookup_without_touching_endpoint_provider(
monkeypatch,
) -> None:
class _DetachedEndpoint:
base_url = "https://chatgpt.com/backend-api/codex"
api_format = "openai:cli"
custom_path = None
provider_id = "provider-1"
@property
def provider(self) -> object:
raise RuntimeError("detached endpoint provider should not be lazy-loaded")
fake_provider = SimpleNamespace(id="provider-1", provider_type="codex", proxy=None)
fake_database = types.ModuleType("src.database")
fake_database.create_session = lambda: _FakeSessionCtx(fake_provider)
fake_models = types.ModuleType("src.models.database")
fake_models.Provider = type("Provider", (), {"id": "id"})
monkeypatch.setitem(sys.modules, "src.database", fake_database)
monkeypatch.setitem(sys.modules, "src.models.database", fake_models)
url = build_provider_url(
_DetachedEndpoint(), # type: ignore[arg-type]
path_params={"model": "ignored"},
is_stream=True,
)
assert url == "https://chatgpt.com/backend-api/codex/responses"
+2 -2
View File
@@ -192,8 +192,8 @@ def test_resolve_pool_account_state_keeps_codex_metadata_block() -> None:
) )
assert state.blocked is True assert state.blocked is True
assert state.code == "account_forbidden" assert state.code == "workspace_deactivated"
assert state.label == "访问受限" assert state.label == "工作区停用"
assert state.reason == "deactivated_workspace" assert state.reason == "deactivated_workspace"
@@ -1,5 +1,7 @@
from __future__ import annotations from __future__ import annotations
import sys
import types
from dataclasses import dataclass from dataclasses import dataclass
from types import SimpleNamespace from types import SimpleNamespace
@@ -19,6 +21,33 @@ class _DummyEndpoint:
api_format: str api_format: str
config: dict | None = None config: dict | None = None
provider: object | None = None provider: object | None = None
provider_id: str | None = None
class _FakeQuery:
def __init__(self, row: object | None) -> None:
self._row = row
def filter(self, *_args: object, **_kwargs: object) -> "_FakeQuery":
return self
def first(self) -> object | None:
return self._row
class _FakeSessionCtx:
def __init__(self, row: object | None) -> None:
self._row = row
def __enter__(self) -> "_FakeSessionCtx":
return self
def __exit__(self, exc_type: object, exc: object, tb: object) -> bool:
_ = exc_type, exc, tb
return False
def query(self, _model: object) -> _FakeQuery:
return _FakeQuery(self._row)
def test_get_upstream_stream_policy_defaults_to_auto() -> None: def test_get_upstream_stream_policy_defaults_to_auto() -> None:
@@ -59,6 +88,24 @@ def test_get_upstream_stream_policy_codex_compact_forces_non_stream() -> None:
set_codex_request_context(None) set_codex_request_context(None)
def test_get_upstream_stream_policy_uses_explicit_provider_type_without_touching_endpoint_provider() -> (
None
):
class _DetachedEndpoint:
api_format = "openai:cli"
config = None
@property
def provider(self) -> object:
raise RuntimeError("detached endpoint provider should not be lazy-loaded")
ep = _DetachedEndpoint()
assert (
get_upstream_stream_policy(ep, provider_type="codex") == UpstreamStreamPolicy.FORCE_STREAM
)
def test_get_upstream_stream_policy_codex_openai_compact_defaults_to_auto() -> None: def test_get_upstream_stream_policy_codex_openai_compact_defaults_to_auto() -> None:
ep = _DummyEndpoint( ep = _DummyEndpoint(
api_format="openai:compact", api_format="openai:compact",
@@ -68,6 +115,30 @@ def test_get_upstream_stream_policy_codex_openai_compact_defaults_to_auto() -> N
assert get_upstream_stream_policy(ep) == UpstreamStreamPolicy.AUTO assert get_upstream_stream_policy(ep) == UpstreamStreamPolicy.AUTO
def test_get_upstream_stream_policy_uses_provider_id_lookup_without_touching_endpoint_provider(
monkeypatch,
) -> None:
class _DetachedEndpoint:
api_format = "openai:cli"
config = None
provider_id = "provider-1"
@property
def provider(self) -> object:
raise RuntimeError("detached endpoint provider should not be lazy-loaded")
fake_provider = SimpleNamespace(id="provider-1", provider_type="codex", proxy=None)
fake_database = types.ModuleType("src.database")
fake_database.create_session = lambda: _FakeSessionCtx(fake_provider)
fake_models = types.ModuleType("src.models.database")
fake_models.Provider = type("Provider", (), {"id": "id"})
monkeypatch.setitem(sys.modules, "src.database", fake_database)
monkeypatch.setitem(sys.modules, "src.models.database", fake_models)
assert get_upstream_stream_policy(_DetachedEndpoint()) == UpstreamStreamPolicy.FORCE_STREAM
def test_enforce_stream_mode_for_upstream_openai_chat_sets_stream_options_usage() -> None: def test_enforce_stream_mode_for_upstream_openai_chat_sets_stream_options_usage() -> None:
body = {"stream": False} body = {"stream": False}
out = enforce_stream_mode_for_upstream( out = enforce_stream_mode_for_upstream(
+108
View File
@@ -0,0 +1,108 @@
from src.api.admin.pool.routes import _detail_is_oauth_invalid, _filter_pool_key_details
from src.api.admin.pool.schemas import PoolKeyDetail
def _detail(
key_id: str,
*,
is_active: bool = True,
scheduling_status: str = "available",
account_status_blocked: bool = False,
account_status_code: str | None = None,
account_status_label: str | None = None,
account_status_reason: str | None = None,
auth_type: str = "api_key",
oauth_invalid_at: int | None = None,
oauth_invalid_reason: str | None = None,
oauth_expires_at: int | None = None,
cooldown_reason: str | None = None,
circuit_breaker_open: bool = False,
cost_limit: int | None = None,
cost_window_usage: int = 0,
) -> PoolKeyDetail:
return PoolKeyDetail(
key_id=key_id,
key_name=key_id,
is_active=is_active,
auth_type=auth_type,
oauth_invalid_at=oauth_invalid_at,
oauth_invalid_reason=oauth_invalid_reason,
oauth_expires_at=oauth_expires_at,
scheduling_status=scheduling_status,
scheduling_reason=scheduling_status or "available",
scheduling_label=scheduling_status or "available",
account_status_blocked=account_status_blocked,
account_status_code=account_status_code,
account_status_label=account_status_label,
account_status_reason=account_status_reason,
cooldown_reason=cooldown_reason,
circuit_breaker_open=circuit_breaker_open,
cost_limit=cost_limit,
cost_window_usage=cost_window_usage,
)
def test_filter_pool_key_details_require_schedulable_keeps_available_and_degraded() -> None:
details = [
_detail("available", scheduling_status="available"),
_detail("degraded", scheduling_status="degraded"),
_detail("blocked", scheduling_status="blocked", account_status_blocked=True),
]
filtered = _filter_pool_key_details(details, require_schedulable=True)
assert [item.key_id for item in filtered] == ["available", "degraded"]
def test_filter_pool_key_details_require_schedulable_uses_fallback_when_status_missing() -> None:
details = [
_detail("manual-disabled", scheduling_status="", is_active=False),
_detail("cooldown", scheduling_status="", cooldown_reason="rate_limited_429"),
_detail("usable", scheduling_status="", is_active=True),
]
filtered = _filter_pool_key_details(details, require_schedulable=True)
assert [item.key_id for item in filtered] == ["usable"]
def test_detail_is_oauth_invalid_excludes_account_disabled_state() -> None:
detail = _detail(
"disabled-account",
auth_type="oauth",
oauth_invalid_at=1,
oauth_invalid_reason="[ACCOUNT_BLOCK] account has been deactivated",
account_status_blocked=True,
account_status_code="account_disabled",
account_status_label="账号停用",
)
assert _detail_is_oauth_invalid(detail) is False
def test_detail_is_oauth_invalid_accepts_token_expired_state() -> None:
detail = _detail(
"expired-token",
auth_type="oauth",
oauth_invalid_at=1,
oauth_invalid_reason="[OAUTH_EXPIRED] token invalidated",
account_status_blocked=True,
account_status_code="oauth_expired",
account_status_label="Token 失效",
)
assert _detail_is_oauth_invalid(detail) is True
def test_detail_is_oauth_invalid_accepts_refresh_failed_state() -> None:
detail = _detail(
"refresh-failed",
auth_type="oauth",
oauth_invalid_at=1,
oauth_invalid_reason="[REFRESH_FAILED] refresh_token_reused",
account_status_blocked=False,
account_status_code="oauth_refresh_failed",
account_status_label="续期失败",
)
assert _detail_is_oauth_invalid(detail) is True
@@ -1,7 +1,9 @@
from __future__ import annotations from __future__ import annotations
import json import json
from contextlib import contextmanager
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest import pytest
@@ -84,6 +86,30 @@ def _make_oauth_key(*, key_id: str, name: str, auth_config: dict[str, object]) -
) )
class _SingleKeyQuery:
def __init__(self, key: SimpleNamespace | None) -> None:
self._key = key
def filter(self, *_args: object, **_kwargs: object) -> "_SingleKeyQuery":
return self
def first(self) -> SimpleNamespace | None:
return self._key
class _SingleKeyDB:
def __init__(self, key: SimpleNamespace | None) -> None:
self._key = key
def query(self, _model: object) -> _SingleKeyQuery:
return _SingleKeyQuery(self._key)
@contextmanager
def _fake_db_context(db: _SingleKeyDB):
yield db
def test_check_duplicate_oauth_account_codex_allows_same_user_different_account_id( def test_check_duplicate_oauth_account_codex_allows_same_user_different_account_id(
monkeypatch: pytest.MonkeyPatch, monkeypatch: pytest.MonkeyPatch,
) -> None: ) -> None:
@@ -151,3 +177,87 @@ def test_check_duplicate_oauth_account_codex_rejects_same_account_user_identity(
"plan_type": "team", "plan_type": "team",
}, },
) )
def test_mark_refresh_failed_sync_preserves_existing_account_block(
monkeypatch: pytest.MonkeyPatch,
) -> None:
key = SimpleNamespace(
id="key-1",
oauth_invalid_at="old-invalid-at",
oauth_invalid_reason="[ACCOUNT_BLOCK] account has been deactivated",
)
db = _SingleKeyDB(key)
monkeypatch.setattr(module, "get_db_context", lambda: _fake_db_context(db))
module._mark_refresh_failed_sync(
"key-1",
"[REFRESH_FAILED] Token 续期失败 (400): refresh_token_reused",
)
assert key.oauth_invalid_at == "old-invalid-at"
assert key.oauth_invalid_reason == "[ACCOUNT_BLOCK] account has been deactivated"
@pytest.mark.asyncio
async def test_refresh_account_state_after_oauth_update_refreshes_supported_provider(
monkeypatch: pytest.MonkeyPatch,
) -> None:
fake_db = SimpleNamespace(close=MagicMock())
captured: dict[str, object] = {}
async def _fake_refresh_provider_quota_for_provider(**kwargs: object) -> dict[str, object]:
captured.update(kwargs)
return {"success": 1}
monkeypatch.setattr(module, "create_session", lambda: fake_db)
from src.services.provider_keys import key_quota_service as quota_module
monkeypatch.setattr(
quota_module,
"refresh_provider_quota_for_provider",
_fake_refresh_provider_quota_for_provider,
)
attempted, error = await module._refresh_account_state_after_oauth_update(
provider_id="provider-1",
provider_type="codex",
key_ids=["key-1"],
)
assert attempted is True
assert error is None
assert captured["provider_id"] == "provider-1"
assert captured["key_ids"] == ["key-1"]
fake_db.close.assert_called_once()
@pytest.mark.asyncio
async def test_refresh_account_state_after_oauth_update_returns_error_when_refresh_fails(
monkeypatch: pytest.MonkeyPatch,
) -> None:
fake_db = SimpleNamespace(close=MagicMock())
async def _fake_refresh_provider_quota_for_provider(**_kwargs: object) -> dict[str, object]:
raise RuntimeError("quota refresh failed")
monkeypatch.setattr(module, "create_session", lambda: fake_db)
from src.services.provider_keys import key_quota_service as quota_module
monkeypatch.setattr(
quota_module,
"refresh_provider_quota_for_provider",
_fake_refresh_provider_quota_for_provider,
)
attempted, error = await module._refresh_account_state_after_oauth_update(
provider_id="provider-1",
provider_type="codex",
key_ids=["key-1"],
)
assert attempted is True
assert "quota refresh failed" in error
fake_db.close.assert_called_once()
@@ -1,18 +1,23 @@
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest import pytest
from pydantic import ValidationError from pydantic import ValidationError
from src.api.admin.provider_query import TestModelFailoverRequest as FailoverRequestModel from src.api.admin import provider_query as provider_query_module
from src.api.admin.provider_query import ( from src.api.admin.provider_query import (
DEFAULT_MODEL_TEST_MESSAGE, DEFAULT_MODEL_TEST_MESSAGE,
)
from src.api.admin.provider_query import TestModelFailoverRequest as FailoverRequestModel
from src.api.admin.provider_query import (
_build_direct_test_candidates, _build_direct_test_candidates,
_build_test_attempts_from_candidate_keys, _build_test_attempts_from_candidate_keys,
_filter_test_candidates_by_endpoint, _filter_test_candidates_by_endpoint,
_flatten_test_candidates_for_concurrency, _flatten_test_candidates_for_concurrency,
_maybe_mark_test_oauth_key_invalid,
_require_test_endpoint_base_url, _require_test_endpoint_base_url,
_resolve_test_message,
_resolve_test_effective_model, _resolve_test_effective_model,
_resolve_test_message,
) )
from src.services.scheduling.schemas import PoolCandidate from src.services.scheduling.schemas import PoolCandidate
@@ -176,3 +181,68 @@ def test_require_test_endpoint_base_url_trims_whitespace() -> None:
) )
assert _require_test_endpoint_base_url(endpoint) == "https://api.anthropic.com" assert _require_test_endpoint_base_url(endpoint) == "https://api.anthropic.com"
@pytest.mark.asyncio
async def test_maybe_mark_test_oauth_key_invalid_skips_account_block_when_oauth_check_fails(
monkeypatch: pytest.MonkeyPatch,
) -> None:
key = SimpleNamespace(id="key-1", oauth_invalid_at=None, oauth_invalid_reason=None)
endpoint = SimpleNamespace(api_format="openai:chat")
db = MagicMock()
async def _fake_verify(**_: object) -> bool:
key.oauth_invalid_reason = "[OAUTH_EXPIRED] refresh token expired"
return False
monkeypatch.setattr(provider_query_module, "verify_oauth_before_account_block", _fake_verify)
await _maybe_mark_test_oauth_key_invalid(
db=db,
endpoint=endpoint,
key=key,
auth_type="oauth",
error_payload={
"error": {
"code": 403,
"message": "Please verify your account",
"status": "PERMISSION_DENIED",
}
},
)
assert key.oauth_invalid_at is None
assert key.oauth_invalid_reason == "[OAUTH_EXPIRED] refresh token expired"
db.commit.assert_not_called()
@pytest.mark.asyncio
async def test_maybe_mark_test_oauth_key_invalid_marks_account_block_after_oauth_check(
monkeypatch: pytest.MonkeyPatch,
) -> None:
key = SimpleNamespace(id="key-2", oauth_invalid_at=None, oauth_invalid_reason=None)
endpoint = SimpleNamespace(api_format="openai:chat")
db = MagicMock()
async def _fake_verify(**_: object) -> bool:
return True
monkeypatch.setattr(provider_query_module, "verify_oauth_before_account_block", _fake_verify)
await _maybe_mark_test_oauth_key_invalid(
db=db,
endpoint=endpoint,
key=key,
auth_type="oauth",
error_payload={
"error": {
"code": 403,
"message": "verify your account",
"status": "PERMISSION_DENIED",
}
},
)
assert key.oauth_invalid_at is not None
assert str(key.oauth_invalid_reason).startswith("[ACCOUNT_BLOCK] ")
db.commit.assert_called_once()