fix: 加固续租失败处理、verify_auth 异常捕获及调度器注册追踪

- task_coordinator: 续租连续失败 5 次后主动触发 lock_lost 回调,失败间加指数退避
- proxy_nodes: lock_lost 回调由 lambda 改为具名 async 函数,确保异步停止逻辑正确执行
- provider_ops: 将 prepare_verify_config 纳入外层 try,捕获 ValueError 并返回失败响应
- maintenance_scheduler: 用 _registered_job_ids 动态追踪已注册任务,stop 时按列表清理
- stats_aggregator: 内联 _do_aggregate 为 for/range(2) 循环,消除内嵌函数
This commit is contained in:
fawney19
2026-03-20 01:22:07 +08:00
parent 913ce2dbcb
commit aa83b4a7a7
7 changed files with 223 additions and 127 deletions

View File

@@ -109,10 +109,11 @@ async def _on_startup() -> None:
if active:
logger.info("启动 ProxyNode 心跳检测调度器...")
await proxy_node_health_scheduler.start()
task_coordinator.register_lock_lost_callback(
"proxy_node_health",
lambda _name: _stop_proxy_node_scheduler_on_lock_lost(logger),
)
async def _on_lock_lost(_name: str) -> None:
await _stop_proxy_node_scheduler_on_lock_lost(logger)
task_coordinator.register_lock_lost_callback("proxy_node_health", _on_lock_lost)
async def _on_shutdown() -> None:

View File

@@ -1011,41 +1011,41 @@ class ProviderOpsService:
# 使用架构的方法构建请求
verify_endpoint = f"{base_url}{architecture.get_verify_endpoint()}"
# 执行异步预处理(如获取动态 Cookie、登录获取 Token
# 返回值可以是 dict仅额外配置或 tuple[dict, dict](额外配置 + 凭据更新)
prepare_result = await architecture.prepare_verify_config(base_url, config, credentials)
if isinstance(prepare_result, tuple):
extra_config, updated_creds = prepare_result
else:
extra_config = prepare_result
updated_creds = {}
merged_config = {**config, **extra_config}
# Token Rotation: prepare_verify_config 可能已消耗旧 refresh_token 并获取新值,
# 无论后续验证是否成功都需要立即持久化,否则旧 token 已失效但数据库未更新。
if updated_creds and provider_id:
logger.info(
"验证过程检测到凭据变更: provider_id={}, updated_keys={}",
provider_id,
list(updated_creds.keys()),
)
self._persist_updated_credentials(provider_id, updated_creds)
headers = architecture.build_verify_headers(merged_config, credentials)
logger.debug(
"验证认证: architecture={}, endpoint={}, headers={}",
architecture_id,
verify_endpoint,
list(headers.keys()),
)
# 获取代理配置(支持 proxy_node_id、tunnel 模式和旧的 proxy URL
from src.services.proxy_node.resolver import resolve_ops_proxy_config_async
proxy, tunnel_node_id = await resolve_ops_proxy_config_async(config)
try:
# 执行异步预处理(如获取动态 Cookie、登录获取 Token
# 返回值可以是 dict仅额外配置或 tuple[dict, dict](额外配置 + 凭据更新)
prepare_result = await architecture.prepare_verify_config(base_url, config, credentials)
if isinstance(prepare_result, tuple):
extra_config, updated_creds = prepare_result
else:
extra_config = prepare_result
updated_creds = {}
merged_config = {**config, **extra_config}
# Token Rotation: prepare_verify_config 可能已消耗旧 refresh_token 并获取新值,
# 无论后续验证是否成功都需要立即持久化,否则旧 token 已失效但数据库未更新。
if updated_creds and provider_id:
logger.info(
"验证过程检测到凭据变更: provider_id={}, updated_keys={}",
provider_id,
list(updated_creds.keys()),
)
self._persist_updated_credentials(provider_id, updated_creds)
headers = architecture.build_verify_headers(merged_config, credentials)
logger.debug(
"验证认证: architecture={}, endpoint={}, headers={}",
architecture_id,
verify_endpoint,
list(headers.keys()),
)
# 获取代理配置(支持 proxy_node_id、tunnel 模式和旧的 proxy URL
from src.services.proxy_node.resolver import resolve_ops_proxy_config_async
proxy, tunnel_node_id = await resolve_ops_proxy_config_async(config)
# 构建 httpx client 参数
client_kwargs: dict[str, Any] = {
"timeout": 30.0,
@@ -1105,6 +1105,8 @@ class ProviderOpsService:
return result_dict
except ValueError as e:
return {"success": False, "message": str(e)}
except httpx.TimeoutException:
return {"success": False, "message": "连接超时"}
except httpx.ConnectError as e:

View File

@@ -50,6 +50,7 @@ class MaintenanceScheduler:
self._stats_aggregation_lock = asyncio.Lock()
self._wallet_daily_usage_lock = asyncio.Lock()
self._startup_task: asyncio.Task[Any] | None = None
self._registered_job_ids: list[str] = []
@staticmethod
def _get_http_client_idle_cleanup_interval_minutes() -> int:
@@ -140,120 +141,129 @@ class MaintenanceScheduler:
return
self.running = True
self._registered_job_ids.clear()
logger.info("系统维护调度器已启动")
scheduler = get_scheduler()
def _add_cron(job_id: str, **kwargs: Any) -> None:
scheduler.add_cron_job(job_id=job_id, **kwargs)
self._registered_job_ids.append(job_id)
def _add_interval(job_id: str, **kwargs: Any) -> None:
scheduler.add_interval_job(job_id=job_id, **kwargs)
self._registered_job_ids.append(job_id)
# 注册定时任务
# 统计聚合任务 - UTC 00:05 执行
scheduler.add_cron_job(
self._scheduled_stats_aggregation,
_add_cron(
"stats_aggregation",
func=self._scheduled_stats_aggregation,
hour=0,
minute=5,
job_id="stats_aggregation",
name="统计数据聚合",
timezone="UTC",
)
# 小时统计聚合任务 - 每小时 05 分执行UTC
scheduler.add_cron_job(
self._scheduled_hourly_stats_aggregation,
_add_cron(
"stats_hourly_aggregation",
func=self._scheduled_hourly_stats_aggregation,
hour="*",
minute=5,
job_id="stats_hourly_aggregation",
name="统计小时数据聚合",
timezone="UTC",
)
scheduler.add_cron_job(
self._scheduled_wallet_daily_usage_aggregation,
_add_cron(
"wallet_daily_usage_aggregation",
func=self._scheduled_wallet_daily_usage_aggregation,
hour=0,
minute=10,
job_id="wallet_daily_usage_aggregation",
name="钱包每日消费汇总",
)
# 清理任务 - 凌晨 3 点执行
scheduler.add_cron_job(
self._scheduled_cleanup,
_add_cron(
"usage_cleanup",
func=self._scheduled_cleanup,
hour=3,
minute=0,
job_id="usage_cleanup",
name="使用记录清理",
)
# 连接池监控 - 每 5 分钟
scheduler.add_interval_job(
self._scheduled_monitor,
_add_interval(
"pool_monitor",
func=self._scheduled_monitor,
minutes=5,
job_id="pool_monitor",
name="连接池监控",
)
# HTTP 代理/Tunnel 客户端空闲清理 - 默认每 5 分钟
scheduler.add_interval_job(
self._scheduled_http_client_idle_cleanup,
_add_interval(
"http_client_idle_cleanup",
func=self._scheduled_http_client_idle_cleanup,
minutes=self._get_http_client_idle_cleanup_interval_minutes(),
job_id="http_client_idle_cleanup",
name="HTTP客户端空闲清理",
)
# Pending 状态清理 - 每 5 分钟
scheduler.add_interval_job(
self._scheduled_pending_cleanup,
_add_interval(
"pending_cleanup",
func=self._scheduled_pending_cleanup,
minutes=5,
job_id="pending_cleanup",
name="Pending状态清理",
)
# 审计日志清理 - 凌晨 4 点执行
scheduler.add_cron_job(
self._scheduled_audit_cleanup,
_add_cron(
"audit_cleanup",
func=self._scheduled_audit_cleanup,
hour=4,
minute=0,
job_id="audit_cleanup",
name="审计日志清理",
)
# Gemini 文件映射清理 - 每小时执行
scheduler.add_interval_job(
self._scheduled_gemini_file_mapping_cleanup,
_add_interval(
"gemini_file_mapping_cleanup",
func=self._scheduled_gemini_file_mapping_cleanup,
hours=1,
job_id="gemini_file_mapping_cleanup",
name="Gemini文件映射清理",
)
# 请求候选记录清理 - 凌晨 3:30 执行
scheduler.add_cron_job(
self._scheduled_candidate_cleanup,
_add_cron(
"candidate_cleanup",
func=self._scheduled_candidate_cleanup,
hour=3,
minute=30,
job_id="candidate_cleanup",
name="请求候选记录清理",
)
# 数据库表维护 - 每周日凌晨 5 点执行 VACUUM ANALYZE
scheduler.add_cron_job(
self._scheduled_db_maintenance,
_add_cron(
"db_maintenance",
func=self._scheduled_db_maintenance,
day_of_week="sun",
hour=5,
minute=0,
job_id="db_maintenance",
name="数据库表维护",
)
# Antigravity User-Agent 版本刷新 - 每 6 小时
scheduler.add_interval_job(
self._scheduled_antigravity_ua_refresh,
_add_interval(
"antigravity_ua_refresh",
func=self._scheduled_antigravity_ua_refresh,
hours=6,
job_id="antigravity_ua_refresh",
name="Antigravity UA版本刷新",
)
# Provider 签到任务 - 根据配置时间执行
checkin_hour, checkin_minute = self._get_checkin_time()
scheduler.add_cron_job(
self._scheduled_provider_checkin,
_add_cron(
self.CHECKIN_JOB_ID,
func=self._scheduled_provider_checkin,
hour=checkin_hour,
minute=checkin_minute,
job_id=self.CHECKIN_JOB_ID,
name="Provider签到",
)
@@ -302,22 +312,9 @@ class MaintenanceScheduler:
self._startup_task = None
scheduler = get_scheduler()
for job_id in (
"stats_aggregation",
"stats_hourly_aggregation",
"wallet_daily_usage_aggregation",
"usage_cleanup",
"pool_monitor",
"http_client_idle_cleanup",
"pending_cleanup",
"audit_cleanup",
"gemini_file_mapping_cleanup",
"candidate_cleanup",
"db_maintenance",
"antigravity_ua_refresh",
self.CHECKIN_JOB_ID,
):
for job_id in self._registered_job_ids:
scheduler.remove_job(job_id)
self._registered_job_ids.clear()
logger.info("系统维护调度器已停止")

View File

@@ -873,30 +873,32 @@ class StatsAggregatorService:
db: Session, date: datetime, user_ids: list[str] | None = None
) -> StatsDaily:
"""聚合单日所有统计(原子提交)"""
for attempt in range(2):
if attempt > 0:
db.rollback()
logger.warning("每日统计聚合冲突,重试更新: {}", date)
def _do_aggregate() -> StatsDaily:
stats = StatsAggregatorService.aggregate_daily_stats(db, date, commit=False)
StatsAggregatorService.aggregate_daily_model_stats(db, date, commit=False)
StatsAggregatorService.aggregate_daily_provider_stats(db, date, commit=False)
StatsAggregatorService.aggregate_daily_api_key_stats(db, date, commit=False)
StatsAggregatorService.aggregate_daily_error_stats(db, date, commit=False)
try:
stats = StatsAggregatorService.aggregate_daily_stats(db, date, commit=False)
StatsAggregatorService.aggregate_daily_model_stats(db, date, commit=False)
StatsAggregatorService.aggregate_daily_provider_stats(db, date, commit=False)
StatsAggregatorService.aggregate_daily_api_key_stats(db, date, commit=False)
StatsAggregatorService.aggregate_daily_error_stats(db, date, commit=False)
if user_ids:
StatsAggregatorService.aggregate_user_daily_stats_batch(
db, date, user_ids, commit=False
)
if user_ids:
StatsAggregatorService.aggregate_user_daily_stats_batch(
db, date, user_ids, commit=False
)
stats.is_complete = True
stats.aggregated_at = datetime.now(timezone.utc)
db.commit()
return stats
stats.is_complete = True
stats.aggregated_at = datetime.now(timezone.utc)
db.commit()
return stats
except IntegrityError:
if attempt == 1:
raise
try:
return _do_aggregate()
except IntegrityError:
db.rollback()
logger.warning("每日统计聚合冲突,重试更新: {}", date)
return _do_aggregate()
raise RuntimeError("unreachable") # type hint helper
@staticmethod
def aggregate_hourly_stats(db: Session, hour_utc: datetime, commit: bool = True) -> StatsHourly:

View File

@@ -181,6 +181,8 @@ class StartupTaskCoordinator:
async def _refresh_lock_loop(self, name: str, ttl: int) -> None:
interval = max(1, ttl // 3)
max_consecutive_failures = 5
consecutive_failures = 0
try:
while name in self._tokens:
await asyncio.sleep(interval)
@@ -190,10 +192,24 @@ class StartupTaskCoordinator:
refreshed = await self._refresh_lock(name, ttl)
except asyncio.CancelledError:
raise
except Exception as exc: # pragma: no cover - 续租失败后下轮重试
logger.warning(f"续租任务锁 {name} 失败: {exc}")
except Exception as exc:
consecutive_failures += 1
if consecutive_failures >= max_consecutive_failures:
logger.error(
f"续租任务锁 {name} 连续失败 {consecutive_failures} 次,视为锁丢失: {exc}"
)
self._tokens.pop(name, None)
await self._notify_lock_lost(name)
break
backoff = min(interval, 2**consecutive_failures)
logger.warning(
f"续租任务锁 {name} 失败 ({consecutive_failures}/{max_consecutive_failures})"
f"{backoff}s 后重试: {exc}"
)
await asyncio.sleep(backoff)
continue
consecutive_failures = 0
if not refreshed:
logger.warning(f"任务 {name} 的 Redis 锁已失效,停止续租")
self._tokens.pop(name, None)

View File

@@ -0,0 +1,74 @@
from __future__ import annotations
from typing import Any
import pytest
from src.services.provider_ops.service import ProviderOpsService
from src.services.provider_ops.types import ConnectorAuthType
class _FakeDB:
new: tuple[Any, ...] = ()
dirty: tuple[Any, ...] = ()
deleted: tuple[Any, ...] = ()
def in_transaction(self) -> bool:
return False
def commit(self) -> None:
pass
def rollback(self) -> None:
pass
class _FailingArchitecture:
def get_verify_endpoint(self) -> str:
return "/verify"
async def prepare_verify_config(
self,
_base_url: str,
_config: dict[str, Any],
_credentials: dict[str, Any],
) -> dict[str, Any]:
raise ValueError("invalid refresh token")
def build_verify_headers(
self,
_config: dict[str, Any],
_credentials: dict[str, Any],
) -> dict[str, str]:
raise AssertionError("build_verify_headers should not be reached")
class _FakeRegistry:
def __init__(self, architecture: Any) -> None:
self._architecture = architecture
def get_or_default(self, _architecture_id: str) -> Any:
return self._architecture
@pytest.mark.asyncio
async def test_verify_auth_returns_failure_when_prepare_verify_config_raises_value_error(
monkeypatch: pytest.MonkeyPatch,
) -> None:
service = ProviderOpsService(_FakeDB())
architecture = _FailingArchitecture()
monkeypatch.setattr(
"src.services.provider_ops.service.get_registry",
lambda: _FakeRegistry(architecture),
)
result = await service.verify_auth(
base_url="https://example.com",
architecture_id="sub2api",
auth_type=ConnectorAuthType.SESSION_LOGIN,
config={},
credentials={"refresh_token": "stale-token"},
)
assert result == {"success": False, "message": "invalid refresh token"}

View File

@@ -56,19 +56,8 @@ async def test_maintenance_scheduler_stop_cancels_startup_task_and_removes_jobs(
scheduler = MaintenanceScheduler()
scheduler.running = True
scheduler._startup_task = asyncio.create_task(asyncio.sleep(3600))
removed_jobs: list[str] = []
monkeypatch.setattr(
maintenance_scheduler_module,
"get_scheduler",
lambda: SimpleNamespace(remove_job=lambda job_id: removed_jobs.append(job_id)),
)
await scheduler.stop()
assert scheduler.running is False
assert scheduler._startup_task is None
assert set(removed_jobs) == {
expected_job_ids = [
"stats_aggregation",
"stats_hourly_aggregation",
"wallet_daily_usage_aggregation",
@@ -82,7 +71,22 @@ async def test_maintenance_scheduler_stop_cancels_startup_task_and_removes_jobs(
"db_maintenance",
"antigravity_ua_refresh",
scheduler.CHECKIN_JOB_ID,
}
]
scheduler._registered_job_ids = list(expected_job_ids)
removed_jobs: list[str] = []
monkeypatch.setattr(
maintenance_scheduler_module,
"get_scheduler",
lambda: SimpleNamespace(remove_job=lambda job_id: removed_jobs.append(job_id)),
)
await scheduler.stop()
assert scheduler.running is False
assert scheduler._startup_task is None
assert set(removed_jobs) == set(expected_job_ids)
assert scheduler._registered_job_ids == []
@pytest.mark.asyncio