fix(startup): 收口 leader 失锁后的后台任务

- 为后台调度器注册失锁回调并只在 stop 成功后清空生命周期引用
- 停止调度器时移除定时 job,补充启动与任务协调器回归测试
- 降低多 worker 下重复调度风险,保持停机收口与统计聚合回归一致
This commit is contained in:
AAEE86
2026-03-19 22:26:51 +08:00
parent e4ebd5cca1
commit 1209c835c7
12 changed files with 567 additions and 25 deletions
+89
View File
@@ -7,6 +7,7 @@ from __future__ import annotations
import asyncio
import time
from collections.abc import Awaitable, Callable
from contextlib import asynccontextmanager
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any
@@ -114,6 +115,44 @@ class LifecycleState:
warmup_task: asyncio.Task[None] | None = None
async def _stop_service_on_lock_lost(
state: LifecycleState,
*,
lock_name: str,
service_name: str,
state_attr: str,
stop: Callable[[], Awaitable[Any]],
) -> None:
logger.warning("检测到 {} 的 leader 锁已丢失,停止本实例的 {}", lock_name, service_name)
try:
await stop()
except Exception:
logger.exception("丢失 {} leader 锁后停止 {} 失败", lock_name, service_name)
return
setattr(state, state_attr, None)
def _make_lock_lost_callback(
state: LifecycleState,
*,
lock_name: str,
service_name: str,
state_attr: str,
stop: Callable[[], Awaitable[Any]],
) -> Callable[[str], Awaitable[None]]:
async def _callback(_name: str) -> None:
await _stop_service_on_lock_lost(
state,
lock_name=lock_name,
service_name=service_name,
state_attr=state_attr,
stop=stop,
)
return _callback
def _configure_uvicorn_access_log() -> None:
"""禁用 uvicorn access 日志(在子进程中执行)。"""
import logging
@@ -388,6 +427,16 @@ async def _start_background_services(state: LifecycleState) -> None:
if quota_scheduler_active:
state.quota_scheduler = get_quota_scheduler()
await state.quota_scheduler.start()
state.task_coordinator.register_lock_lost_callback(
"quota_scheduler",
_make_lock_lost_callback(
state,
lock_name="quota_scheduler",
service_name="月卡额度重置调度器",
state_attr="quota_scheduler",
stop=state.quota_scheduler.stop,
),
)
else:
logger.info("检测到其他 worker 已运行额度调度器,本实例跳过")
state.quota_scheduler = None
@@ -398,6 +447,16 @@ async def _start_background_services(state: LifecycleState) -> None:
state.maintenance_scheduler = get_maintenance_scheduler()
logger.info("启动系统维护调度器...")
await state.maintenance_scheduler.start()
state.task_coordinator.register_lock_lost_callback(
"maintenance_scheduler",
_make_lock_lost_callback(
state,
lock_name="maintenance_scheduler",
service_name="系统维护调度器",
state_attr="maintenance_scheduler",
stop=state.maintenance_scheduler.stop,
),
)
else:
logger.info("检测到其他 worker 已运行维护调度器,本实例跳过")
state.maintenance_scheduler = None
@@ -408,6 +467,16 @@ async def _start_background_services(state: LifecycleState) -> None:
state.model_fetch_scheduler = get_model_fetch_scheduler()
logger.info("启动模型自动获取调度器...")
await state.model_fetch_scheduler.start()
state.task_coordinator.register_lock_lost_callback(
"model_fetch_scheduler",
_make_lock_lost_callback(
state,
lock_name="model_fetch_scheduler",
service_name="模型自动获取调度器",
state_attr="model_fetch_scheduler",
stop=state.model_fetch_scheduler.stop,
),
)
else:
logger.info("检测到其他 worker 已运行模型获取调度器,本实例跳过")
state.model_fetch_scheduler = None
@@ -420,6 +489,16 @@ async def _start_background_services(state: LifecycleState) -> None:
state.pool_quota_probe_scheduler = get_pool_quota_probe_scheduler()
logger.info("启动号池额度主动探测调度器...")
await state.pool_quota_probe_scheduler.start()
state.task_coordinator.register_lock_lost_callback(
"pool_quota_probe_scheduler",
_make_lock_lost_callback(
state,
lock_name="pool_quota_probe_scheduler",
service_name="号池额度主动探测调度器",
state_attr="pool_quota_probe_scheduler",
stop=state.pool_quota_probe_scheduler.stop,
),
)
else:
logger.info("检测到其他 worker 已运行号池额度主动探测调度器,本实例跳过")
state.pool_quota_probe_scheduler = None
@@ -430,6 +509,16 @@ async def _start_background_services(state: LifecycleState) -> None:
state.task_poller = get_task_poller()
logger.info("启动 TaskPollervideo...")
await state.task_poller.start()
state.task_coordinator.register_lock_lost_callback(
"task_poller:video",
_make_lock_lost_callback(
state,
lock_name="task_poller:video",
service_name="TaskPollervideo",
state_attr="task_poller",
stop=state.task_poller.stop,
),
)
else:
logger.info("检测到其他 worker 已运行 TaskPollervideo),本实例跳过")
state.task_poller = None
+26 -3
View File
@@ -20,6 +20,9 @@ if TYPE_CHECKING:
from sqlalchemy.orm import Session
_proxy_node_task_coordinator: Any | None = None
def _reset_tunnel_connected_on_startup() -> None:
"""服务端启动时将所有 tunnel_connected=True 的节点重置为 False/OFFLINE。
@@ -89,8 +92,10 @@ async def _on_startup() -> None:
from src.clients import get_redis_client
global _proxy_node_task_coordinator
redis_client = await get_redis_client()
task_coordinator = StartupTaskCoordinator(redis_client)
_proxy_node_task_coordinator = task_coordinator
proxy_node_health_scheduler = get_proxy_node_health_scheduler()
active = await task_coordinator.acquire("proxy_node_health")
@@ -104,6 +109,10 @@ 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_shutdown() -> None:
@@ -117,14 +126,28 @@ async def _on_shutdown() -> None:
from src.clients import get_redis_client
redis_client = await get_redis_client()
task_coordinator = StartupTaskCoordinator(redis_client)
global _proxy_node_task_coordinator
if _proxy_node_task_coordinator is None:
redis_client = await get_redis_client()
_proxy_node_task_coordinator = StartupTaskCoordinator(redis_client)
scheduler = get_proxy_node_health_scheduler()
if scheduler.running:
logger.info("停止 ProxyNode 心跳检测调度器...")
await scheduler.stop()
await task_coordinator.release("proxy_node_health")
await _proxy_node_task_coordinator.release("proxy_node_health")
_proxy_node_task_coordinator = None
async def _stop_proxy_node_scheduler_on_lock_lost(logger: Any) -> None:
from src.services.proxy_node.health_scheduler import get_proxy_node_health_scheduler
scheduler = get_proxy_node_health_scheduler()
if not scheduler.running:
return
logger.warning("检测到 proxy_node_health 的 leader 锁已丢失,停止本实例的 ProxyNode 心跳检测")
await scheduler.stop()
async def _health_check() -> ModuleHealth:
+5 -3
View File
@@ -54,9 +54,7 @@ KEY_FETCH_TIMEOUT_SECONDS = 120
MODEL_FETCH_HTTP_TIMEOUT = 10.0
# 启动时首次自动获取开关与延迟
MODEL_FETCH_STARTUP_ENABLED = (
os.getenv("MODEL_FETCH_STARTUP_ENABLED", "true").lower() == "true"
)
MODEL_FETCH_STARTUP_ENABLED = os.getenv("MODEL_FETCH_STARTUP_ENABLED", "true").lower() == "true"
MODEL_FETCH_STARTUP_DELAY_SECONDS = max(
0,
int(os.getenv("MODEL_FETCH_STARTUP_DELAY_SECONDS", "10")),
@@ -315,6 +313,8 @@ class ModelFetchScheduler:
async def stop(self) -> None:
"""停止调度器"""
self._running = False
scheduler = get_scheduler()
scheduler.remove_job("model_auto_fetch")
# 取消并等待启动任务完成
if self._startup_task and not self._startup_task.done():
@@ -348,6 +348,8 @@ class ModelFetchScheduler:
async def _scheduled_fetch_models(self) -> None:
"""定时任务入口"""
if not self._running:
return
async with self._lock:
await self._perform_fetch_all_keys()
@@ -183,6 +183,8 @@ class PoolQuotaProbeScheduler:
if not self.running:
return
self.running = False
scheduler = get_scheduler()
scheduler.remove_job("pool_quota_probe_check")
logger.info("PoolQuotaProbeScheduler stopped")
async def _scheduled_probe_check(self) -> None:
@@ -74,9 +74,13 @@ class ProxyNodeHealthScheduler:
if not self.running:
return
self.running = False
scheduler = get_scheduler()
scheduler.remove_job("proxy_node_health_check")
logger.info("ProxyNodeHealthScheduler stopped")
async def _scheduled_check(self) -> None:
if not self.running:
return
await self._check_heartbeats()
self._check_count = (self._check_count + 1) % _EVENT_CLEANUP_INTERVAL
if self._check_count == 0:
+54 -2
View File
@@ -49,6 +49,7 @@ class MaintenanceScheduler:
self._interval_tasks = []
self._stats_aggregation_lock = asyncio.Lock()
self._wallet_daily_usage_lock = asyncio.Lock()
self._startup_task: asyncio.Task[Any] | None = None
@staticmethod
def _get_http_client_idle_cleanup_interval_minutes() -> int:
@@ -260,8 +261,9 @@ class MaintenanceScheduler:
if config.maintenance_startup_tasks_enabled:
from src.utils.async_utils import safe_create_task
safe_create_task(self._run_startup_tasks())
self._startup_task = safe_create_task(self._run_startup_tasks())
else:
self._startup_task = None
logger.info("维护调度器启动任务已禁用(MAINTENANCE_STARTUP_TASKS_ENABLED=false")
async def _run_startup_tasks(self) -> None:
@@ -290,8 +292,32 @@ class MaintenanceScheduler:
return
self.running = False
if self._startup_task and not self._startup_task.done():
self._startup_task.cancel()
try:
await self._startup_task
except asyncio.CancelledError:
pass
self._startup_task = None
scheduler = get_scheduler()
scheduler.stop()
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,
):
scheduler.remove_job(job_id)
logger.info("系统维护调度器已停止")
@@ -299,22 +325,32 @@ class MaintenanceScheduler:
async def _scheduled_stats_aggregation(self, backfill: bool = False) -> None:
"""统计聚合任务(定时调用)"""
if not self.running:
return
await self._perform_stats_aggregation(backfill=backfill)
async def _scheduled_wallet_daily_usage_aggregation(self) -> None:
"""钱包每日消费汇总任务(定时调用)"""
if not self.running:
return
await self._perform_wallet_daily_usage_aggregation()
async def _scheduled_hourly_stats_aggregation(self) -> None:
"""小时统计聚合任务(定时调用)"""
if not self.running:
return
await self._perform_hourly_stats_aggregation()
async def _scheduled_cleanup(self) -> None:
"""清理任务(定时调用)"""
if not self.running:
return
await self._perform_cleanup()
async def _scheduled_monitor(self) -> None:
"""监控任务(定时调用)"""
if not self.running:
return
try:
from src.database import log_pool_status
@@ -324,6 +360,8 @@ class MaintenanceScheduler:
async def _scheduled_http_client_idle_cleanup(self) -> None:
"""HTTP 客户端空闲清理任务(定时调用)。"""
if not self.running:
return
try:
stats = await HTTPClientPool.cleanup_idle_clients()
if stats.get("proxy_closed", 0) or stats.get("tunnel_closed", 0):
@@ -337,26 +375,38 @@ class MaintenanceScheduler:
async def _scheduled_pending_cleanup(self) -> None:
"""Pending 清理任务(定时调用)"""
if not self.running:
return
await self._perform_pending_cleanup()
async def _scheduled_audit_cleanup(self) -> None:
"""审计日志清理任务(定时调用)"""
if not self.running:
return
await self._perform_audit_cleanup()
async def _scheduled_candidate_cleanup(self) -> None:
"""请求候选记录清理任务(定时调用)"""
if not self.running:
return
await self._perform_candidate_cleanup()
async def _scheduled_db_maintenance(self) -> None:
"""数据库表维护任务(定时调用)"""
if not self.running:
return
await self._perform_db_maintenance()
async def _scheduled_gemini_file_mapping_cleanup(self) -> None:
"""Gemini 文件映射清理任务(定时调用)"""
if not self.running:
return
await self._perform_gemini_file_mapping_cleanup()
async def _scheduled_antigravity_ua_refresh(self) -> None:
"""Antigravity User-Agent 版本刷新(定时调用)"""
if not self.running:
return
try:
from src.services.provider.adapters.antigravity.client import refresh_user_agent
@@ -366,6 +416,8 @@ class MaintenanceScheduler:
async def _scheduled_provider_checkin(self) -> None:
"""Provider 签到任务(定时调用)"""
if not self.running:
return
await self._perform_provider_checkin()
# ========== 实际任务实现 ==========
+22 -13
View File
@@ -873,21 +873,30 @@ class StatsAggregatorService:
db: Session, date: datetime, user_ids: list[str] | None = None
) -> 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)
if user_ids:
StatsAggregatorService.aggregate_user_daily_stats_batch(
db, date, user_ids, commit=False
)
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)
stats.is_complete = True
stats.aggregated_at = datetime.now(timezone.utc)
db.commit()
return stats
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
try:
return _do_aggregate()
except IntegrityError:
db.rollback()
logger.warning("每日统计聚合冲突,重试更新: {}", date)
return _do_aggregate()
@staticmethod
def aggregate_hourly_stats(db: Session, hour_utc: datetime, commit: bool = True) -> StatsHourly:
+4
View File
@@ -54,10 +54,14 @@ class QuotaScheduler:
return
self.running = False
scheduler = get_scheduler()
scheduler.remove_job("quota_reset_check")
logger.info("Quota scheduler stopped")
async def _scheduled_quota_check(self) -> None:
"""额度检查任务(定时调用)"""
if not self.running:
return
await self._check_and_reset_quotas()
async def _check_and_reset_quotas(self) -> None:
+77
View File
@@ -11,9 +11,11 @@
from __future__ import annotations
import asyncio
import os
import pathlib
import uuid
from collections.abc import Awaitable, Callable
from typing import Any
from src.core.logger import logger
@@ -35,6 +37,8 @@ class StartupTaskCoordinator:
self.redis = redis_client
self._tokens: dict[str, str] = {}
self._file_handles: dict[str, object] = {}
self._refresh_tasks: dict[str, asyncio.Task[None]] = {}
self._lock_lost_callbacks: dict[str, Callable[[str], Awaitable[None] | None]] = {}
self._lock_dir = pathlib.Path(lock_dir or os.getenv("TASK_LOCK_DIR", "./.locks"))
if not self._lock_dir.exists():
self._lock_dir.mkdir(parents=True, exist_ok=True)
@@ -80,6 +84,7 @@ class StartupTaskCoordinator:
)
if result == 1:
self._tokens[name] = token
self._start_refresh_task(name, ttl)
logger.info(f"任务 {name} 通过 Redis 锁独占执行")
return True
return False
@@ -88,6 +93,7 @@ class StartupTaskCoordinator:
acquired = await self.redis.set(self._redis_key(name), token, nx=True, ex=ttl)
if acquired:
self._tokens[name] = token
self._start_refresh_task(name, ttl)
logger.info(f"任务 {name} 通过 Redis 锁独占执行")
return True
return False
@@ -97,6 +103,11 @@ class StartupTaskCoordinator:
return await self._acquire_file_lock(name)
async def release(self, name: str) -> Any:
refresh_task = self._refresh_tasks.pop(name, None)
if refresh_task is not None:
refresh_task.cancel()
self._lock_lost_callbacks.pop(name, None)
if self.redis and name in self._tokens:
token = self._tokens.pop(name)
script = """
@@ -137,6 +148,72 @@ class StartupTaskCoordinator:
handle.close()
return False
def register_lock_lost_callback(
self,
name: str,
callback: Callable[[str], Awaitable[None] | None],
) -> None:
self._lock_lost_callbacks[name] = callback
def _start_refresh_task(self, name: str, ttl: int) -> None:
if not self.redis:
return
existing_task = self._refresh_tasks.pop(name, None)
if existing_task is not None:
existing_task.cancel()
self._refresh_tasks[name] = asyncio.create_task(self._refresh_lock_loop(name, ttl))
async def _refresh_lock(self, name: str, ttl: int) -> bool:
token = self._tokens.get(name)
if not self.redis or token is None:
return False
script = """
if redis.call('GET', KEYS[1]) == ARGV[1] then
return redis.call('EXPIRE', KEYS[1], tonumber(ARGV[2]))
end
return 0
"""
result = await self.redis.eval(script, 1, self._redis_key(name), token, ttl)
return result == 1
async def _refresh_lock_loop(self, name: str, ttl: int) -> None:
interval = max(1, ttl // 3)
try:
while name in self._tokens:
await asyncio.sleep(interval)
if name not in self._tokens:
break
try:
refreshed = await self._refresh_lock(name, ttl)
except asyncio.CancelledError:
raise
except Exception as exc: # pragma: no cover - 续租失败后下轮重试
logger.warning(f"续租任务锁 {name} 失败: {exc}")
continue
if not refreshed:
logger.warning(f"任务 {name} 的 Redis 锁已失效,停止续租")
self._tokens.pop(name, None)
await self._notify_lock_lost(name)
break
except asyncio.CancelledError:
raise
async def _notify_lock_lost(self, name: str) -> None:
callback = self._lock_lost_callbacks.pop(name, None)
if callback is None:
return
try:
result = callback(name)
if result is not None:
await result
except Exception as exc: # pragma: no cover - 回调失败仅记录日志
logger.exception("任务 {} 的失锁回调执行失败: {}", name, exc)
async def ensure_singleton_task(
name: str, redis_client: Any | None = None, ttl: int | None = None