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

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