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

- 为后台调度器注册失锁回调并只在 stop 成功后清空生命周期引用
- 停止调度器时移除定时 job,补充启动与任务协调器回归测试
- 降低多 worker 下重复调度风险,保持停机收口与统计聚合回归一致
This commit is contained in:
AAEE86
2026-03-19 22:25:06 +08:00
parent e4ebd5cca1
commit 1209c835c7
12 changed files with 567 additions and 25 deletions

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

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:

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()

View File

@@ -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:

View File

@@ -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:

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()
# ========== 实际任务实现 ==========

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:

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:

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

View File

@@ -1,12 +1,16 @@
from __future__ import annotations
import asyncio
import inspect
from collections.abc import Callable
from datetime import datetime, timedelta, timezone
from types import SimpleNamespace
from typing import Any, cast
from unittest.mock import MagicMock
import pytest
import src.main as main_module
import src.services.system.maintenance_scheduler as maintenance_scheduler_module
from src.config.settings import config
from src.services.system.maintenance_scheduler import MaintenanceScheduler
@@ -45,6 +49,85 @@ async def test_maintenance_scheduler_start_skips_startup_task_when_disabled(
assert created is False
@pytest.mark.asyncio
async def test_maintenance_scheduler_stop_cancels_startup_task_and_removes_jobs(
monkeypatch: pytest.MonkeyPatch,
) -> None:
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) == {
"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",
scheduler.CHECKIN_JOB_ID,
}
@pytest.mark.asyncio
async def test_stop_service_on_lock_lost_keeps_state_when_stop_fails() -> None:
state = main_module.LifecycleState()
service = SimpleNamespace()
state.quota_scheduler = cast(Any, service)
async def fail_stop() -> None:
raise RuntimeError("boom")
await main_module._stop_service_on_lock_lost(
state,
lock_name="quota_scheduler",
service_name="月卡额度重置调度器",
state_attr="quota_scheduler",
stop=fail_stop,
)
assert state.quota_scheduler is service
@pytest.mark.asyncio
async def test_stop_service_on_lock_lost_clears_state_after_success() -> None:
state = main_module.LifecycleState()
service = SimpleNamespace()
state.quota_scheduler = cast(Any, service)
stopped = False
async def stop() -> None:
nonlocal stopped
stopped = True
await main_module._stop_service_on_lock_lost(
state,
lock_name="quota_scheduler",
service_name="月卡额度重置调度器",
state_attr="quota_scheduler",
stop=stop,
)
assert stopped is True
assert state.quota_scheduler is None
def test_http_client_idle_cleanup_interval_env_invalid(
monkeypatch: pytest.MonkeyPatch,
) -> None:
@@ -276,7 +359,7 @@ async def test_perform_cleanup_deletes_first_and_uses_non_overlapping_windows(
class _FakeDateTime(datetime):
@classmethod
def now(cls, tz=None): # type: ignore[override]
def now(cls, tz: timezone | None = None) -> datetime: # type: ignore[override]
if tz is None:
return fixed_now.replace(tzinfo=None)
return fixed_now.astimezone(tz)
@@ -291,7 +374,7 @@ async def test_perform_cleanup_deletes_first_and_uses_non_overlapping_windows(
calls: list[tuple[str, datetime, int, datetime | None]] = []
def _record(name: str, count: int):
def _record(name: str, count: int) -> Callable[..., int]:
def _inner(
cutoff_time: datetime,
batch_size: int,
@@ -313,6 +396,10 @@ async def test_perform_cleanup_deletes_first_and_uses_non_overlapping_windows(
"auto_delete_expired_keys": False,
}
def _delete_old_records(cutoff_time: datetime, batch_size: int) -> int:
calls.append(("delete", cutoff_time, batch_size, None))
return 5
monkeypatch.setattr(maintenance_scheduler_module, "datetime", _FakeDateTime)
monkeypatch.setattr(
maintenance_scheduler_module.asyncio, "get_running_loop", lambda: _FakeLoop()
@@ -330,8 +417,7 @@ async def test_perform_cleanup_deletes_first_and_uses_non_overlapping_windows(
monkeypatch.setattr(
scheduler,
"_delete_old_records",
lambda cutoff_time, batch_size: calls.append(("delete", cutoff_time, batch_size, None))
or 5,
_delete_old_records,
)
monkeypatch.setattr(
scheduler,

View File

@@ -5,6 +5,7 @@ from types import SimpleNamespace
from typing import Any, cast
import pytest
from sqlalchemy.exc import IntegrityError
from src.models.database import StatsDaily, StatsUserDaily
from src.services.system.stats_aggregator import (
@@ -61,6 +62,20 @@ class _BatchUserStatsSession:
self.commit_count += 1
class _RetryCommitSession:
def __init__(self) -> None:
self.commit_count = 0
self.rollback_count = 0
def commit(self) -> None:
self.commit_count += 1
if self.commit_count == 1:
raise IntegrityError("insert", {}, Exception("duplicate key"))
def rollback(self) -> None:
self.rollback_count += 1
def test_query_stats_hybrid_batches_statsdaily_lookup_and_merges_realtime_ranges(
monkeypatch: pytest.MonkeyPatch,
) -> None:
@@ -167,6 +182,79 @@ def test_aggregate_user_daily_stats_batch_updates_all_users_in_two_queries() ->
assert user_two.total_cost == 0.0
def test_aggregate_daily_stats_bundle_retries_after_integrity_error(
monkeypatch: pytest.MonkeyPatch,
) -> None:
target_day = datetime(2026, 3, 1, tzinfo=timezone.utc)
db = _RetryCommitSession()
stats_calls: list[SimpleNamespace] = []
stage_calls: list[str] = []
def fake_aggregate_daily_stats(*_args: object, **_kwargs: object) -> SimpleNamespace:
stats = SimpleNamespace(is_complete=False, aggregated_at=None)
stats_calls.append(stats)
stage_calls.append("daily")
return stats
def fake_model_stats(*_args: object, **_kwargs: object) -> list[object]:
stage_calls.append("model")
return []
def fake_provider_stats(*_args: object, **_kwargs: object) -> list[object]:
stage_calls.append("provider")
return []
def fake_api_key_stats(*_args: object, **_kwargs: object) -> list[object]:
stage_calls.append("api_key")
return []
def fake_error_stats(*_args: object, **_kwargs: object) -> list[object]:
stage_calls.append("error")
return []
def fake_user_daily_stats(*_args: object, **_kwargs: object) -> list[object]:
stage_calls.append("user")
return []
monkeypatch.setattr(StatsAggregatorService, "aggregate_daily_stats", fake_aggregate_daily_stats)
monkeypatch.setattr(StatsAggregatorService, "aggregate_daily_model_stats", fake_model_stats)
monkeypatch.setattr(
StatsAggregatorService, "aggregate_daily_provider_stats", fake_provider_stats
)
monkeypatch.setattr(StatsAggregatorService, "aggregate_daily_api_key_stats", fake_api_key_stats)
monkeypatch.setattr(StatsAggregatorService, "aggregate_daily_error_stats", fake_error_stats)
monkeypatch.setattr(
StatsAggregatorService, "aggregate_user_daily_stats_batch", fake_user_daily_stats
)
result = StatsAggregatorService.aggregate_daily_stats_bundle(
cast(Any, db),
target_day,
user_ids=["user-1"],
)
assert db.commit_count == 2
assert db.rollback_count == 1
assert len(stats_calls) == 2
assert stage_calls == [
"daily",
"model",
"provider",
"api_key",
"error",
"user",
"daily",
"model",
"provider",
"api_key",
"error",
"user",
]
assert result is stats_calls[-1]
assert result.is_complete is True
assert result.aggregated_at is not None
def test_compute_percentiles_by_local_day_returns_sqlite_fallback_without_queries() -> None:
db = SimpleNamespace(bind=SimpleNamespace(dialect=SimpleNamespace(name="sqlite")))
time_range = TimeRangeParams(

View File

@@ -0,0 +1,106 @@
from __future__ import annotations
from typing import Any
import pytest
from src.utils.task_coordinator import StartupTaskCoordinator
class _FakeRedis:
def __init__(self, eval_results: list[int] | None = None, set_result: bool = True) -> None:
self.eval_results = list(eval_results or [])
self.eval_calls: list[tuple[Any, ...]] = []
self.set_result = set_result
self.set_calls: list[tuple[Any, ...]] = []
async def eval(self, script: str, numkeys: int, *args: Any) -> int:
self.eval_calls.append((script, numkeys, *args))
if self.eval_results:
return self.eval_results.pop(0)
return 1
async def set(self, key: str, value: str, *, nx: bool, ex: int) -> bool:
self.set_calls.append((key, value, nx, ex))
return self.set_result
class _FakeTask:
def __init__(self) -> None:
self.cancelled = False
def cancel(self) -> None:
self.cancelled = True
@pytest.mark.asyncio
async def test_startup_task_coordinator_acquire_starts_refresh_for_redis_lock(
monkeypatch: pytest.MonkeyPatch,
) -> None:
redis = _FakeRedis(eval_results=[1])
coordinator = StartupTaskCoordinator(redis)
started: list[tuple[str, int]] = []
monkeypatch.setattr(
coordinator,
"_start_refresh_task",
lambda name, ttl: started.append((name, ttl)),
)
acquired = await coordinator.acquire("maintenance_scheduler", ttl=120)
assert acquired is True
assert "maintenance_scheduler" in coordinator._tokens
assert started == [("maintenance_scheduler", 120)]
@pytest.mark.asyncio
async def test_startup_task_coordinator_refresh_lock_extends_matching_token() -> None:
redis = _FakeRedis(eval_results=[1])
coordinator = StartupTaskCoordinator(redis)
coordinator._tokens["maintenance_scheduler"] = "token-1"
refreshed = await coordinator._refresh_lock("maintenance_scheduler", ttl=180)
assert refreshed is True
assert len(redis.eval_calls) == 1
_script, numkeys, key, token, ttl = redis.eval_calls[0]
assert numkeys == 1
assert key == "task_lock:maintenance_scheduler"
assert token == "token-1"
assert ttl == 180
@pytest.mark.asyncio
async def test_startup_task_coordinator_release_cancels_refresh_task() -> None:
redis = _FakeRedis(eval_results=[1])
coordinator = StartupTaskCoordinator(redis)
coordinator._tokens["maintenance_scheduler"] = "token-1"
refresh_task = _FakeTask()
coordinator._refresh_tasks["maintenance_scheduler"] = refresh_task # type: ignore[assignment]
await coordinator.release("maintenance_scheduler")
assert refresh_task.cancelled is True
assert "maintenance_scheduler" not in coordinator._refresh_tasks
assert "maintenance_scheduler" not in coordinator._tokens
assert len(redis.eval_calls) == 1
_script, numkeys, key, token = redis.eval_calls[0]
assert numkeys == 1
assert key == "task_lock:maintenance_scheduler"
assert token == "token-1"
@pytest.mark.asyncio
async def test_startup_task_coordinator_notify_lock_lost_runs_registered_callback() -> None:
coordinator = StartupTaskCoordinator()
called: list[str] = []
async def on_lock_lost(name: str) -> None:
called.append(name)
coordinator.register_lock_lost_callback("maintenance_scheduler", on_lock_lost)
await coordinator._notify_lock_lost("maintenance_scheduler")
assert called == ["maintenance_scheduler"]