mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
feat: OAuth 导入完成后自动触发配额刷新,统一配额常量定义
- 单个导入、批量导入、Kiro device flow 完成后后台异步刷新配额 - 将 CODEX_WHAM_USAGE_URL 和 QUOTA_REFRESH_PROVIDER_TYPES 统一到 key_quota_service.py,消除 3 处重复硬编码
This commit is contained in:
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user