feat: OAuth 导入完成后自动触发配额刷新,统一配额常量定义

- 单个导入、批量导入、Kiro device flow 完成后后台异步刷新配额
- 将 CODEX_WHAM_USAGE_URL 和 QUOTA_REFRESH_PROVIDER_TYPES 统一到 key_quota_service.py,消除 3 处重复硬编码
This commit is contained in:
fawney19
2026-03-10 09:36:47 +08:00
parent 3b0dbadb1e
commit afd0dcf2ff
4 changed files with 105 additions and 33 deletions

View File

@@ -35,6 +35,9 @@ from src.services.provider_keys import (
reveal_endpoint_key_payload,
update_endpoint_key_response,
)
from src.services.provider_keys.key_quota_service import (
CODEX_WHAM_USAGE_URL as _CODEX_WHAM_USAGE_URL,
)
from src.utils.auth_utils import require_admin
router = APIRouter(tags=["Provider Keys"])
@@ -376,13 +379,7 @@ class AdminCreateProviderKeyAdapter(AdminApiAdapter):
)
# ========== Codex Quota Refresh API ==========
# Codex wham/usage API 地址(用于查询限额信息)
CODEX_WHAM_USAGE_URL = "https://chatgpt.com/backend-api/wham/usage"
# ========== Kiro Quota Refresh API ==========
# ========== Quota Refresh API ==========
class RefreshProviderQuotaRequest(BaseModel):
@@ -432,6 +429,6 @@ class AdminRefreshProviderQuotaAdapter(AdminApiAdapter):
return await refresh_provider_quota_for_provider(
db=context.db,
provider_id=self.provider_id,
codex_wham_usage_url=CODEX_WHAM_USAGE_URL,
codex_wham_usage_url=_CODEX_WHAM_USAGE_URL,
key_ids=self.key_ids,
)

View File

@@ -1290,7 +1290,7 @@ async def complete_provider_oauth(
proxy=key_proxy,
)
return ProviderCompleteOAuthResponse(
response = ProviderCompleteOAuthResponse(
key_id=str(new_key.id),
provider_type=provider_type,
expires_at=expires_at,
@@ -1299,6 +1299,11 @@ async def complete_provider_oauth(
replaced=replaced,
)
# 单个导入完成后,后台触发一次配额刷新
asyncio.create_task(_refresh_quota_after_import(provider_id, provider_type, [str(new_key.id)]))
return response
# ==============================================================================
# Import Refresh Token (从导出文件导入)
@@ -1755,7 +1760,7 @@ async def import_refresh_token(
proxy=key_proxy,
)
return ProviderCompleteOAuthResponse(
response = ProviderCompleteOAuthResponse(
key_id=str(new_key.id),
provider_type=provider_type,
expires_at=new_cfg.expires_at or None,
@@ -1764,6 +1769,13 @@ async def import_refresh_token(
replaced=replaced,
)
# 单个导入完成后,后台触发一次配额刷新
asyncio.create_task(
_refresh_quota_after_import(provider_id, provider_type, [str(new_key.id)])
)
return response
template = _require_oauth_template(provider_type)
# 用 refresh_token 换取 access_token
@@ -1879,7 +1891,7 @@ async def import_refresh_token(
proxy=key_proxy,
)
return ProviderCompleteOAuthResponse(
response = ProviderCompleteOAuthResponse(
key_id=str(new_key.id),
provider_type=provider_type,
expires_at=expires_at,
@@ -1888,6 +1900,50 @@ async def import_refresh_token(
replaced=replaced,
)
# 单个导入完成后,后台触发一次配额刷新
asyncio.create_task(_refresh_quota_after_import(provider_id, provider_type, [str(new_key.id)]))
return response
# ==============================================================================
# 导入后配额刷新
# ==============================================================================
def _extract_success_key_ids(result: BatchImportResponse) -> list[str]:
"""从批量导入结果中提取成功导入的 key_id 列表。"""
return [r.key_id for r in result.results if r.status == "success" and r.key_id]
async def _refresh_quota_after_import(
provider_id: str,
provider_type: str,
key_ids: list[str],
) -> None:
"""导入完成后触发一次配额刷新(使用独立 db session"""
from src.services.provider_keys.key_quota_service import (
CODEX_WHAM_USAGE_URL,
QUOTA_REFRESH_PROVIDER_TYPES,
refresh_provider_quota_for_provider,
)
if not key_ids or provider_type not in QUOTA_REFRESH_PROVIDER_TYPES:
return
try:
db = create_session()
try:
await refresh_provider_quota_for_provider(
db=db,
provider_id=provider_id,
codex_wham_usage_url=CODEX_WHAM_USAGE_URL,
key_ids=key_ids,
)
finally:
db.close()
except Exception as exc:
logger.warning("[BATCH_IMPORT] 导入后配额刷新失败 (provider={}): {}", provider_id, exc)
# ==============================================================================
# 通用批量导入(支持所有 OAuth Provider
@@ -2235,7 +2291,7 @@ async def batch_import_oauth(
batch_concurrency = (_pool_cfg.batch_concurrency or 8) if _pool_cfg else 8
if provider_type == ProviderType.KIRO.value:
return await _batch_import_kiro_internal(
result = await _batch_import_kiro_internal(
provider_id=provider_id,
provider=provider,
raw_credentials=payload.credentials,
@@ -2244,17 +2300,26 @@ async def batch_import_oauth(
key_proxy=key_proxy,
concurrency=batch_concurrency,
)
else:
result = await _batch_import_standard_oauth_internal(
provider_id=provider_id,
provider_type=provider_type,
provider=provider,
raw_credentials=payload.credentials,
db=db,
proxy_config=proxy_config,
key_proxy=key_proxy,
concurrency=batch_concurrency,
)
return await _batch_import_standard_oauth_internal(
provider_id=provider_id,
provider_type=provider_type,
provider=provider,
raw_credentials=payload.credentials,
db=db,
proxy_config=proxy_config,
key_proxy=key_proxy,
concurrency=batch_concurrency,
)
# 导入完成后,后台触发一次配额刷新
success_key_ids = _extract_success_key_ids(result)
if success_key_ids:
asyncio.create_task(
_refresh_quota_after_import(provider_id, provider_type, success_key_ids)
)
return result
async def _run_batch_import_task(
@@ -2351,6 +2416,11 @@ async def _run_batch_import_task(
state["message"] = f"导入完成:成功 {result.success},失败 {result.failed}"
await _save_batch_task_state(task_id, state, redis=redis)
# 导入完成后触发一次配额刷新
success_key_ids = _extract_success_key_ids(result)
if success_key_ids:
await _refresh_quota_after_import(provider_id, provider_type, success_key_ids)
except Exception as exc:
try:
db.rollback()
@@ -3056,6 +3126,10 @@ async def device_poll(
session["replaced"] = replaced
await redis.setex(redis_key, 60, json.dumps(session))
# 单个导入完成后,后台触发一次配额刷新
provider_type = ProviderType.KIRO.value
asyncio.create_task(_refresh_quota_after_import(provider_id, provider_type, [str(new_key.id)]))
return DevicePollResponse(
status="authorized",
key_id=str(new_key.id),

View File

@@ -26,12 +26,16 @@ from src.services.provider_keys.quota_refresh import (
QuotaRefreshHandler = Callable[..., Awaitable[dict]]
CODEX_WHAM_USAGE_URL = "https://chatgpt.com/backend-api/wham/usage"
_QUOTA_REFRESH_HANDLERS: dict[str, QuotaRefreshHandler] = {
ProviderType.CODEX: refresh_codex_key_quota,
ProviderType.ANTIGRAVITY: refresh_antigravity_key_quota,
ProviderType.KIRO: refresh_kiro_key_quota,
}
QUOTA_REFRESH_PROVIDER_TYPES: frozenset[str] = frozenset(_QUOTA_REFRESH_HANDLERS.keys())
def _normalize_api_format(api_format: Any) -> str:
"""规范化 api_format兼容大小写和首尾空白。"""
@@ -80,7 +84,7 @@ async def refresh_provider_quota_for_provider(
raise NotFoundException(f"Provider {provider_id} 不存在")
provider_type = normalize_provider_type(getattr(provider, "provider_type", ""))
if provider_type not in {ProviderType.CODEX, ProviderType.ANTIGRAVITY, ProviderType.KIRO}:
if provider_type not in QUOTA_REFRESH_PROVIDER_TYPES:
raise InvalidRequestException("仅支持 Codex / Antigravity / Kiro 类型的 Provider 刷新限额")
pool_cfg = parse_pool_config(getattr(provider, "config", None))
auto_remove_banned_keys = bool(pool_cfg and pool_cfg.auto_remove_banned_keys)

View File

@@ -21,21 +21,18 @@ from src.core.provider_types import ProviderType, normalize_provider_type
from src.database import create_session
from src.models.database import Provider, ProviderAPIKey
from src.services.provider.pool.config import parse_pool_config
from src.services.provider_keys.key_quota_service import refresh_provider_quota_for_provider
from src.services.provider_keys.key_quota_service import (
CODEX_WHAM_USAGE_URL,
QUOTA_REFRESH_PROVIDER_TYPES,
refresh_provider_quota_for_provider,
)
from src.services.system.scheduler import get_scheduler
# 与 admin 刷新额度 API 保持一致
_CODEX_WHAM_USAGE_URL = "https://chatgpt.com/backend-api/wham/usage"
_REDIS_PREFIX = "ap:quota_probe:last"
_DEFAULT_INTERVAL_MINUTES = 10
_DEFAULT_SCAN_INTERVAL_SECONDS = 60
_DEFAULT_MAX_KEYS_PER_PROVIDER = 50
_MAX_INTERVAL_MINUTES = 1440
_SUPPORTED_PROVIDER_TYPES = {
ProviderType.CODEX.value,
ProviderType.KIRO.value,
ProviderType.ANTIGRAVITY.value,
}
def _probe_stamp_key(provider_id: str, key_id: str) -> str:
@@ -260,7 +257,7 @@ class PoolQuotaProbeScheduler:
for provider in providers:
provider_id = str(getattr(provider, "id", "") or "")
provider_type = normalize_provider_type(getattr(provider, "provider_type", ""))
if not provider_id or provider_type not in _SUPPORTED_PROVIDER_TYPES:
if not provider_id or provider_type not in QUOTA_REFRESH_PROVIDER_TYPES:
continue
pool_cfg = parse_pool_config(getattr(provider, "config", None))
@@ -340,7 +337,7 @@ class PoolQuotaProbeScheduler:
result = await refresh_provider_quota_for_provider(
db=probe_db,
provider_id=task.provider_id,
codex_wham_usage_url=_CODEX_WHAM_USAGE_URL,
codex_wham_usage_url=CODEX_WHAM_USAGE_URL,
key_ids=task.probe_key_ids,
)
logger.info(