mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 09:20:22 +08:00
- 为后台调度器注册失锁回调并只在 stop 成功后清空生命周期引用 - 停止调度器时移除定时 job,补充启动与任务协调器回归测试 - 降低多 worker 下重复调度风险,保持停机收口与统计聚合回归一致
271 lines
8.8 KiB
Python
271 lines
8.8 KiB
Python
from __future__ import annotations
|
|
|
|
from datetime import date, datetime, time, timedelta, timezone
|
|
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 (
|
|
AggregatedStats,
|
|
StatsAggregatorService,
|
|
query_stats_hybrid,
|
|
)
|
|
from src.services.system.time_range import TimeRangeParams
|
|
|
|
|
|
class _FakeQuery:
|
|
def __init__(self, *, all_result: list[Any] | None = None) -> None:
|
|
self._all_result = all_result if all_result is not None else []
|
|
|
|
def filter(self, *_args: object, **_kwargs: object) -> _FakeQuery:
|
|
return self
|
|
|
|
def group_by(self, *_args: object, **_kwargs: object) -> _FakeQuery:
|
|
return self
|
|
|
|
def all(self) -> list[Any]:
|
|
return self._all_result
|
|
|
|
|
|
class _HybridQuerySession:
|
|
def __init__(self, stats_daily_rows: list[SimpleNamespace]) -> None:
|
|
self._stats_daily_rows = stats_daily_rows
|
|
self.stats_daily_query_count = 0
|
|
|
|
def query(self, entity: object) -> _FakeQuery:
|
|
if entity is StatsDaily:
|
|
self.stats_daily_query_count += 1
|
|
return _FakeQuery(all_result=self._stats_daily_rows)
|
|
raise AssertionError(f"Unexpected query entity: {entity}")
|
|
|
|
|
|
class _BatchUserStatsSession:
|
|
def __init__(
|
|
self, existing_rows: list[StatsUserDaily], aggregated_rows: list[SimpleNamespace]
|
|
) -> None:
|
|
self._responses: list[list[Any]] = [list(existing_rows), list(aggregated_rows)]
|
|
self.added: list[StatsUserDaily] = []
|
|
self.commit_count = 0
|
|
|
|
def query(self, *_entities: object) -> _FakeQuery:
|
|
if not self._responses:
|
|
raise AssertionError("Unexpected extra query")
|
|
return _FakeQuery(all_result=self._responses.pop(0))
|
|
|
|
def add(self, row: StatsUserDaily) -> None:
|
|
self.added.append(row)
|
|
|
|
def commit(self) -> None:
|
|
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:
|
|
today = datetime.now(timezone.utc).date()
|
|
historical_cached_day = today - timedelta(days=4)
|
|
historical_missing_day = today - timedelta(days=3)
|
|
realtime_day = today
|
|
|
|
cached_row = SimpleNamespace(
|
|
date=datetime.combine(historical_cached_day, time.min, tzinfo=timezone.utc),
|
|
total_requests=10,
|
|
success_requests=9,
|
|
error_requests=1,
|
|
input_tokens=100,
|
|
output_tokens=50,
|
|
cache_creation_tokens=5,
|
|
cache_read_tokens=3,
|
|
cache_creation_cost=1.2,
|
|
cache_read_cost=0.8,
|
|
total_cost=3.5,
|
|
actual_total_cost=3.0,
|
|
avg_response_time_ms=200.0,
|
|
)
|
|
db = _HybridQuerySession(stats_daily_rows=[cached_row])
|
|
|
|
calls: list[tuple[datetime, datetime]] = []
|
|
|
|
def _fake_aggregate_usage_range(
|
|
_db: object,
|
|
start_utc: datetime,
|
|
end_utc: datetime,
|
|
filters: object | None = None, # noqa: ARG001
|
|
) -> AggregatedStats:
|
|
calls.append((start_utc, end_utc))
|
|
return AggregatedStats(total_requests=1, success_requests=1)
|
|
|
|
class _FakeParams:
|
|
def get_complete_utc_dates(self) -> tuple[list[date], None, None]:
|
|
return [historical_cached_day, historical_missing_day, realtime_day], None, None
|
|
|
|
monkeypatch.setattr(
|
|
"src.services.system.stats_aggregator.aggregate_usage_range",
|
|
_fake_aggregate_usage_range,
|
|
)
|
|
|
|
result = query_stats_hybrid(cast(Any, db), cast(Any, _FakeParams()))
|
|
|
|
assert db.stats_daily_query_count == 1
|
|
assert calls == [
|
|
(
|
|
datetime.combine(historical_missing_day, time.min, tzinfo=timezone.utc),
|
|
datetime.combine(
|
|
historical_missing_day + timedelta(days=1), time.min, tzinfo=timezone.utc
|
|
),
|
|
),
|
|
(
|
|
datetime.combine(realtime_day, time.min, tzinfo=timezone.utc),
|
|
datetime.combine(realtime_day + timedelta(days=1), time.min, tzinfo=timezone.utc),
|
|
),
|
|
]
|
|
assert result.total_requests == 12
|
|
assert result.success_requests == 11
|
|
|
|
|
|
def test_aggregate_user_daily_stats_batch_updates_all_users_in_two_queries() -> None:
|
|
target_day = datetime(2026, 3, 1, tzinfo=timezone.utc)
|
|
aggregated_rows = [
|
|
SimpleNamespace(
|
|
user_id="user-1",
|
|
username="alice",
|
|
total_requests=4,
|
|
error_requests=1,
|
|
input_tokens=20,
|
|
output_tokens=8,
|
|
cache_creation_tokens=2,
|
|
cache_read_tokens=1,
|
|
total_cost=1.5,
|
|
)
|
|
]
|
|
db = _BatchUserStatsSession(existing_rows=[], aggregated_rows=aggregated_rows)
|
|
|
|
result = StatsAggregatorService.aggregate_user_daily_stats_batch(
|
|
cast(Any, db),
|
|
target_day,
|
|
["user-1", "user-2"],
|
|
commit=True,
|
|
)
|
|
|
|
assert len(result) == 2
|
|
assert db.commit_count == 1
|
|
assert len(db.added) == 2
|
|
|
|
user_one = next(row for row in result if row.user_id == "user-1")
|
|
user_two = next(row for row in result if row.user_id == "user-2")
|
|
|
|
assert user_one.username == "alice"
|
|
assert user_one.total_requests == 4
|
|
assert user_one.success_requests == 3
|
|
assert user_one.total_cost == 1.5
|
|
|
|
assert user_two.total_requests == 0
|
|
assert user_two.success_requests == 0
|
|
assert user_two.error_requests == 0
|
|
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(
|
|
start_date=date(2026, 3, 1),
|
|
end_date=date(2026, 3, 3),
|
|
timezone="Asia/Singapore",
|
|
)
|
|
|
|
result = StatsAggregatorService.compute_percentiles_by_local_day(cast(Any, db), time_range)
|
|
|
|
assert [row["date"] for row in result] == ["2026-03-01", "2026-03-02", "2026-03-03"]
|
|
assert all(row["p50_response_time_ms"] is None for row in result)
|
|
assert all(row["p50_first_byte_time_ms"] is None for row in result)
|