mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-12 04:09:48 +08:00
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:
@@ -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)
|
||||||
|
|||||||
@@ -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) {
|
||||||
|
|||||||
@@ -37,6 +37,9 @@ const KEYWORDS_TOKEN_INVALID = [
|
|||||||
const KEYWORDS_VERIFICATION = [
|
const KEYWORDS_VERIFICATION = [
|
||||||
'validation_required',
|
'validation_required',
|
||||||
'verify your account',
|
'verify your account',
|
||||||
|
'需要验证',
|
||||||
|
'验证账号',
|
||||||
|
'验证身份',
|
||||||
]
|
]
|
||||||
|
|
||||||
// 合并的完整列表
|
// 合并的完整列表
|
||||||
|
|||||||
@@ -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
@@ -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)} 个异常账号",
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
# ==============================================================================
|
# ==============================================================================
|
||||||
|
|||||||
@@ -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),
|
||||||
|
|||||||
@@ -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),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -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):
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -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"
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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"
|
||||||
|
|||||||
@@ -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(
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
Reference in New Issue
Block a user