feat: 流式空闲超时、健康监控查询优化、限流桶内存上限与维护清理修复

Close #233

Co-authored-by: AAEE86 <ppk0227@hotmail.com>

- cli_monitor_mixin: 引入 STREAM_IDLE_TIMEOUT_SECONDS(可通过环境变量配置),
  流传输开始后若超出空闲窗口无新 chunk 则提前取消并返回 504,避免长时间挂起
- stream_context: 新增 managed_recorded_bodies 上下文管理器,确保 chunks 在
  telemetry 完成后及时释放;stream_telemetry 使用该接口统一管理 response body 构建
- health endpoint: 将状态聚合改为 GROUP BY 直接统计,事件列表按 api_format
  单独查询,避免单次 limit 拉取大量记录导致的遗漏与性能问题;同时过滤不活跃
  provider/endpoint,与公开健康接口保持一致
- endpoint health service: 修正时间线数据按 endpoint_id 而非 key_id 聚合
- token_bucket: 引入 max_buckets/bucket_expiry 上限与定时清理,防止内存无限增长;
  修复 refill_rate=0 时 get_reset_time 除零异常;新增 _is_unlimited_rate_limit 判断
- maintenance_scheduler: 调整清理顺序(先删整行再按窗口清理),新增 newer_than
  边界参数,避免同一行在同一轮中被重复改写
- sync_execute: 新增 create_pending_usage 开关,允许已预创建记录的调用方跳过重复创建
- quota_reader / provider_ops balance: 小幅修复与健壮性提升
- Dockerfile: 添加 MALLOC_ARENA_MAX=2 环境变量以降低 gunicorn worker RSS
- 补充相关测试覆盖
This commit is contained in:
fawney19
2026-03-18 23:38:26 +08:00
parent 3d5b6141a5
commit 1d72a8f9c1
37 changed files with 1787 additions and 607 deletions

View File

@@ -39,6 +39,17 @@ async def _yield_once_then_cancel(ctx: StreamContext) -> AsyncGenerator[bytes, N
raise asyncio.CancelledError()
async def _yield_once_then_hang(ctx: StreamContext) -> AsyncGenerator[bytes, None]:
ctx.append_text("partial output")
yield b"data: chunk\n\n"
await asyncio.sleep(3600)
async def _yield_after_delay_then_complete() -> AsyncGenerator[bytes, None]:
await asyncio.sleep(0.25)
yield b"data: first\n\n"
@pytest.mark.asyncio
async def test_create_monitored_stream_marks_client_disconnected_when_confirmed() -> None:
monitor = _DummyMonitor()
@@ -113,3 +124,38 @@ async def test_create_monitored_stream_estimates_output_tokens_before_unknown_ca
assert ctx.error_message == "cancelled_unknown"
assert ctx.output_tokens == expected_output_tokens
assert f"output_tokens={expected_output_tokens}" in (ctx.upstream_response or "")
@pytest.mark.asyncio
async def test_create_monitored_stream_marks_idle_timeout_before_worker_timeout() -> None:
monitor = _DummyMonitor()
monitor.CANCEL_DISCONNECT_RETRY_DELAYS_SECONDS = ()
monitor.STREAM_IDLE_TIMEOUT_SECONDS = 1.0
ctx = StreamContext(model="test-model", api_format="openai:cli", request_id="req-idle-timeout")
monitored = monitor._create_monitored_stream(ctx, _yield_once_then_hang(ctx), None)
with pytest.raises(asyncio.CancelledError):
async for _ in monitored:
pass
expected_output_tokens = max(1, len("partial output") // 4)
assert ctx.status_code == 504
assert ctx.error_message == "stream_idle_timeout"
assert ctx.output_tokens == expected_output_tokens
assert "cancel_origin=stream_idle_timeout" in (ctx.upstream_response or "")
@pytest.mark.asyncio
async def test_create_monitored_stream_does_not_idle_timeout_before_first_chunk() -> None:
monitor = _DummyMonitor()
monitor.CANCEL_DISCONNECT_RETRY_DELAYS_SECONDS = ()
monitor.STREAM_IDLE_TIMEOUT_SECONDS = 0.05
ctx = StreamContext(model="test-model", api_format="openai:cli", request_id="req-first-chunk")
monitored = monitor._create_monitored_stream(ctx, _yield_after_delay_then_complete(), None)
chunks = [chunk async for chunk in monitored]
assert chunks == [b"data: first\n\n"]
assert ctx.status_code == 200
assert ctx.error_message is None

View File

@@ -23,7 +23,7 @@ class _DummySyncHandler(CliSyncMixin):
) -> str:
return str(request_body.get("model") or "unknown")
def _create_pending_usage(self, **kwargs: object) -> None:
def _create_pending_usage(self, **kwargs: object) -> bool:
self.pending_calls.append(kwargs)
raise _StopExecution()
@@ -41,7 +41,7 @@ class _DummyStreamHandler(CliStreamMixin):
) -> str:
return str(request_body.get("model") or "unknown")
def _create_pending_usage(self, **kwargs: object) -> None:
def _create_pending_usage(self, **kwargs: object) -> bool:
self.pending_calls.append(kwargs)
raise _StopExecution()

View File

@@ -1,3 +1,9 @@
import os
import pytest
os.environ.setdefault("JWT_SECRET_KEY", "test-secret-key")
from src.api.handlers.base import stream_context
from src.api.handlers.base.stream_context import StreamContext
@@ -50,7 +56,36 @@ def test_reset_for_retry_clears_state() -> None:
assert ctx.error_message is None
def test_record_first_byte_time(monkeypatch) -> None:
def test_release_recorded_chunks_clears_both_chunk_lists() -> None:
ctx = StreamContext(model="test-model", api_format="openai:chat")
ctx.parsed_chunks.append({"type": "client"})
ctx.provider_parsed_chunks.append({"type": "provider"})
ctx.release_recorded_chunks()
assert ctx.parsed_chunks == []
assert ctx.provider_parsed_chunks == []
def test_managed_recorded_bodies_builds_then_releases_chunks() -> None:
ctx = StreamContext(model="test-model", api_format="openai:chat")
ctx.parsed_chunks.append({"type": "client"})
ctx.provider_parsed_chunks.append({"type": "provider"})
ctx.data_count = 1
with ctx.managed_recorded_bodies(123) as recorded_bodies:
assert recorded_bodies.response_body is not None
assert recorded_bodies.response_body["chunks"] == [{"type": "provider"}]
assert recorded_bodies.client_response_body is not None
assert recorded_bodies.client_response_body["chunks"] == [{"type": "client"}]
assert ctx.parsed_chunks == []
assert ctx.provider_parsed_chunks == []
assert recorded_bodies.response_body is None
assert recorded_bodies.client_response_body is None
def test_record_first_byte_time(monkeypatch: pytest.MonkeyPatch) -> None:
"""测试记录首字时间"""
ctx = StreamContext(model="claude-3", api_format="claude_messages")
start_time = 100.0
@@ -63,7 +98,7 @@ def test_record_first_byte_time(monkeypatch) -> None:
assert ctx.first_byte_time_ms == 12
def test_record_first_byte_time_idempotent(monkeypatch) -> None:
def test_record_first_byte_time_idempotent(monkeypatch: pytest.MonkeyPatch) -> None:
"""测试首字时间只记录一次"""
ctx = StreamContext(model="claude-3", api_format="claude_messages")
start_time = 100.0
@@ -82,7 +117,7 @@ def test_record_first_byte_time_idempotent(monkeypatch) -> None:
assert first_value == second_value
def test_reset_for_retry_clears_first_byte_time(monkeypatch) -> None:
def test_reset_for_retry_clears_first_byte_time(monkeypatch: pytest.MonkeyPatch) -> None:
"""测试重试时清除首字时间"""
ctx = StreamContext(model="claude-3", api_format="claude_messages")
start_time = 100.0

View File

@@ -1,11 +1,15 @@
from __future__ import annotations
import os
import time
from types import SimpleNamespace
from typing import Any
from unittest.mock import AsyncMock
import pytest
os.environ.setdefault("JWT_SECRET_KEY", "test-secret-key")
from src.api.handlers.base import stream_telemetry as stream_telemetry_module
from src.api.handlers.base.stream_context import StreamContext
from src.api.handlers.base.stream_telemetry import StreamTelemetryRecorder
@@ -64,3 +68,69 @@ async def test_record_stream_stats_estimates_tokens_for_failed_partial_stream(
assert ctx.input_tokens > 0
assert ctx.output_tokens == max(1, len("partial output") // 4)
recorder._dispatch_record.assert_awaited_once() # type: ignore[attr-defined]
@pytest.mark.asyncio
async def test_record_stream_stats_releases_parsed_chunks_after_dispatch(
monkeypatch: pytest.MonkeyPatch,
) -> None:
recorder = StreamTelemetryRecorder(
request_id="req-release",
user_id="1",
api_key_id="2",
client_ip="127.0.0.1",
format_id="openai:chat",
)
recorder._get_telemetry_writer = AsyncMock( # type: ignore[method-assign]
return_value=SimpleNamespace(include_bodies=True)
)
recorder._update_candidate_status = AsyncMock() # type: ignore[method-assign]
dispatch_payloads: list[dict[str, Any]] = []
async def _capture_dispatch(*args: Any, **kwargs: Any) -> None:
dispatch_payloads.append(
{
"response_body": args[5],
"client_response_body": kwargs.get("client_response_body"),
}
)
recorder._dispatch_record = _capture_dispatch # type: ignore[method-assign]
ctx = StreamContext(
model="test-model",
api_format="openai:chat",
request_id="req-release",
user_id=1,
api_key_id=2,
)
ctx.provider_name = "test-provider"
ctx.parsed_chunks.extend([{"type": "chunk-1"}, {"type": "chunk-2"}])
ctx.data_count = 2
ctx.chunk_count = 2
monkeypatch.setattr(stream_telemetry_module, "get_db", lambda: iter([_DummyDb()]))
monkeypatch.setattr(
stream_telemetry_module.SystemConfigService,
"should_log_body",
lambda _db: True,
)
monkeypatch.setattr(stream_telemetry_module.config, "stream_stats_delay", 0)
await recorder.record_stream_stats(
ctx,
original_headers={},
original_request_body={"input": [{"content": "hello world"}]},
start_time=time.time(),
)
assert len(dispatch_payloads) == 1
response_body = dispatch_payloads[0]["response_body"]
assert response_body["chunks"] == [{"type": "chunk-1"}, {"type": "chunk-2"}]
assert response_body["metadata"]["stream"] is True
assert response_body["metadata"]["total_chunks"] == 2
assert response_body["metadata"]["data_count"] == 2
assert dispatch_payloads[0]["client_response_body"] is None
assert ctx.parsed_chunks == []
assert ctx.provider_parsed_chunks == []

View File

@@ -0,0 +1,146 @@
from __future__ import annotations
from datetime import datetime, timezone
from types import SimpleNamespace
from typing import Any, cast
from unittest.mock import MagicMock
import pytest
from src.api.admin.endpoints.health import AdminApiFormatHealthMonitorAdapter
from src.api.public.catalog import PublicApiFormatHealthMonitorAdapter
def _build_query(result: object) -> MagicMock:
query = MagicMock()
query.join.return_value = query
query.distinct.return_value = query
query.filter.return_value = query
query.group_by.return_value = query
query.order_by.return_value = query
query.limit.return_value = query
query.all.return_value = result
return query
def _expr_texts(query: MagicMock) -> list[str]:
return [str(arg) for arg in query.filter.call_args.args]
@pytest.mark.asyncio
async def test_admin_api_format_health_monitor_filters_inactive_sources(
monkeypatch: pytest.MonkeyPatch,
) -> None:
endpoint_query = _build_query([("openai:compact", "ep-active", "provider-active")])
key_query = _build_query([("provider-active", ["openai:compact"])])
status_query = _build_query([("openai:compact", "success", 3)])
rows_query = _build_query([])
db = MagicMock()
db.query.side_effect = [endpoint_query, key_query, status_query, rows_query]
monkeypatch.setattr(
"src.api.admin.endpoints.health.EndpointHealthService._generate_timeline_from_usage",
lambda **_: {
"timeline": ["healthy"] * 100,
"time_range_start": None,
"time_range_end": None,
},
)
context = SimpleNamespace(
db=db,
request=SimpleNamespace(state=SimpleNamespace()),
add_audit_metadata=lambda **kwargs: None,
)
adapter = AdminApiFormatHealthMonitorAdapter(lookback_hours=6, per_format_limit=20)
await adapter.handle(cast(Any, context))
status_filters = _expr_texts(status_query)
rows_filters = _expr_texts(rows_query)
assert any("provider_endpoints.is_active" in expr for expr in status_filters)
assert any("providers.is_active" in expr for expr in status_filters)
assert any("provider_endpoints.is_active" in expr for expr in rows_filters)
assert any("providers.is_active" in expr for expr in rows_filters)
@pytest.mark.asyncio
async def test_public_api_format_health_monitor_uses_real_counts_not_sampled_events(
monkeypatch: pytest.MonkeyPatch,
) -> None:
now = datetime.now(timezone.utc)
active_formats_query = _build_query([("openai:compact",)])
endpoint_rows_query = _build_query([("openai:compact", "ep-active")])
status_query = _build_query(
[
("openai:compact", "success", 7),
("openai:compact", "failed", 3),
("openai:compact", "skipped", 5),
]
)
rows_query = _build_query(
[
SimpleNamespace(
status="failed",
status_code=500,
latency_ms=321,
error_type="provider_error",
finished_at=now,
started_at=None,
created_at=now,
),
SimpleNamespace(
status="success",
status_code=200,
latency_ms=123,
error_type=None,
finished_at=now,
started_at=None,
created_at=now,
),
]
)
db = MagicMock()
db.query.side_effect = [
active_formats_query,
endpoint_rows_query,
status_query,
rows_query,
]
monkeypatch.setattr(
"src.api.public.catalog.EndpointHealthService._generate_timeline_from_usage",
lambda **_: {
"timeline": ["healthy"] * 100,
"time_range_start": None,
"time_range_end": now,
},
)
monkeypatch.setattr(
"src.core.api_format.get_local_path_for_endpoint",
lambda api_format: f"/{api_format}",
)
context = SimpleNamespace(
db=db,
request=SimpleNamespace(state=SimpleNamespace()),
)
adapter = PublicApiFormatHealthMonitorAdapter(lookback_hours=6, per_format_limit=20)
result = await adapter.handle(cast(Any, context))
monitor = result["formats"][0]
assert monitor["api_format"] == "openai:compact"
assert monitor["total_attempts"] == 15
assert monitor["success_count"] == 7
assert monitor["failed_count"] == 3
assert monitor["skipped_count"] == 5
assert monitor["success_rate"] == pytest.approx(0.7)
assert len(monitor["events"]) == 2
rows_filters = _expr_texts(rows_query)
assert any("provider_endpoints.is_active" in expr for expr in rows_filters)
assert any("providers.is_active" in expr for expr in rows_filters)

View File

@@ -0,0 +1,50 @@
import os
import pytest
os.environ.setdefault("JWT_SECRET_KEY", "test-secret-key")
from src.plugins.manager import PluginManager
def _build_manager() -> PluginManager:
return PluginManager(
config={
"auth": {"api_key": False},
"cache": {"memory": False},
"monitor": {"prometheus": False},
"token": {"claude": False},
"load_balancer": {"sticky_priority": False},
}
)
def test_rate_limit_defaults_to_token_bucket_when_unconfigured() -> None:
manager = _build_manager()
plugin = manager.get_plugin("rate_limit")
assert plugin is not None
assert plugin.name == "token_bucket"
@pytest.mark.asyncio
async def test_default_rate_limit_plugin_honors_dynamic_rate_limit() -> None:
manager = _build_manager()
plugin = manager.get_plugin("rate_limit")
assert plugin is not None
first = await plugin.check_limit("public_ip:test", rate_limit=2)
assert first.allowed is True
await plugin.consume("public_ip:test", amount=1, rate_limit=2)
second = await plugin.check_limit("public_ip:test", rate_limit=2)
assert second.allowed is True
await plugin.consume("public_ip:test", amount=1, rate_limit=2)
third = await plugin.check_limit("public_ip:test", rate_limit=2)
assert third.allowed is False
assert third.remaining == 0
assert third.retry_after is not None

View File

@@ -0,0 +1,111 @@
from __future__ import annotations
from datetime import datetime, timezone
import pytest
import src.plugins.rate_limit.token_bucket as token_bucket_module
from src.plugins.rate_limit.token_bucket import RedisTokenBucketBackend, TokenBucketStrategy
@pytest.mark.asyncio
async def test_token_bucket_cleans_up_expired_buckets(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("RATE_LIMIT_BACKEND", "memory")
strategy = TokenBucketStrategy()
strategy.configure({"bucket_expiry": 1, "cleanup_interval": 0})
await strategy.check_limit("api_key:stale")
strategy.buckets["api_key:stale"].last_access_time -= 3600
await strategy.check_limit("api_key:fresh")
assert "api_key:stale" not in strategy.buckets
assert "api_key:fresh" in strategy.buckets
@pytest.mark.asyncio
async def test_token_bucket_reconfigures_existing_bucket_when_rate_limit_changes(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setenv("RATE_LIMIT_BACKEND", "memory")
strategy = TokenBucketStrategy()
await strategy.check_limit("user:42", rate_limit=120)
bucket = strategy.buckets["user:42"]
bucket.tokens = 90
await strategy.check_limit("user:42", rate_limit=30)
updated_bucket = strategy.buckets["user:42"]
assert updated_bucket.capacity == 30
assert updated_bucket.refill_rate == 0.5
assert updated_bucket.tokens <= 30
@pytest.mark.asyncio
async def test_token_bucket_treats_non_positive_dynamic_rate_limit_as_unlimited(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setenv("RATE_LIMIT_BACKEND", "memory")
strategy = TokenBucketStrategy()
result = await strategy.check_limit("public_ip:test", rate_limit=0)
consumed = await strategy.consume("public_ip:test", amount=1, rate_limit=0)
assert result.allowed is True
assert consumed is True
assert "public_ip:test" not in strategy.buckets
class _FakeRedisClient:
async def hmget(self, _key: str, *_fields: str) -> list[None]:
return [None, None]
def register_script(self, _script: str): # type: ignore[no-untyped-def]
async def _runner(*args, **kwargs): # type: ignore[no-untyped-def]
return [1, 0, 0]
return _runner
@pytest.mark.asyncio
async def test_token_bucket_retries_redis_backend_probe_after_initial_miss(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setenv("RATE_LIMIT_BACKEND", "auto")
strategy = TokenBucketStrategy()
strategy._redis_retry_interval = 0
fake_redis = _FakeRedisClient()
calls = {"count": 0}
def _fake_get_redis_client_sync(): # type: ignore[no-untyped-def]
calls["count"] += 1
if calls["count"] == 1:
return None
return fake_redis
monkeypatch.setattr(
token_bucket_module,
"get_redis_client_sync",
_fake_get_redis_client_sync,
)
await strategy.check_limit("public_ip:first")
assert strategy._redis_backend is None
await strategy.check_limit("public_ip:second")
assert strategy._redis_backend is not None
assert calls["count"] == 2
@pytest.mark.asyncio
async def test_redis_token_bucket_missing_bucket_reports_reset_now() -> None:
backend = RedisTokenBucketBackend(_FakeRedisClient())
result = await backend.peek("public_ip:test", capacity=60, refill_rate=1.0, amount=1)
assert result.allowed is True
assert result.remaining == 60
assert result.reset_at is not None
assert abs((result.reset_at - datetime.now(timezone.utc)).total_seconds()) < 2

View File

@@ -36,6 +36,22 @@ def test_codex_reader_preserves_summary_formats() -> None:
assert credits_reader.display_summary() == "积分 12.35"
def test_codex_reader_hides_reset_countdown_when_remaining_is_full() -> None:
reader = get_quota_reader(
"codex",
{
"codex": {
"primary_used_percent": 0.0,
"primary_reset_seconds": 266400,
"secondary_used_percent": 0.0,
"secondary_reset_seconds": 3600,
}
},
)
assert reader.display_summary() == "周剩余 100.0% | 5H剩余 100.0%"
def test_antigravity_reader_keeps_used_percent_fallbacks() -> None:
reader = get_quota_reader(
"antigravity",
@@ -88,6 +104,20 @@ def test_extract_reset_seconds_uses_codex_weekly_reset_for_codex_provider() -> N
assert extract_reset_seconds(key_obj) == pytest.approx(1800.0)
def test_extract_reset_seconds_codex_weekly_full_quota_returns_none() -> None:
key_obj = SimpleNamespace(
provider_type="codex",
upstream_metadata={
"codex": {
"primary_used_percent": 0.0,
"primary_reset_seconds": 1800.0,
}
},
)
assert extract_reset_seconds(key_obj) is None
def test_extract_reset_seconds_prefers_codex_reset_at(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr("src.services.provider.pool.dimensions._helpers.time.time", lambda: 1000.0)
key_obj = SimpleNamespace(

View File

@@ -1,6 +1,7 @@
from __future__ import annotations
import inspect
from datetime import datetime, timedelta, timezone
from types import SimpleNamespace
from unittest.mock import MagicMock
@@ -124,44 +125,29 @@ async def test_candidate_cleanup_uses_dedicated_retention_and_batch_settings(
assert batch_two.closed is True
def test_cleanup_body_fields_loads_ids_then_processes_records_individually(
def test_cleanup_body_fields_batches_records_with_single_commit(
monkeypatch: pytest.MonkeyPatch,
) -> None:
scheduler = MaintenanceScheduler()
class _IdBatchSession:
def __init__(self, ids: list[str]) -> None:
self.ids = ids
self.closed = False
self.query_obj = MagicMock()
filtered = self.query_obj.filter.return_value
filtered.filter.return_value.order_by.return_value.limit.return_value.all.return_value = [
SimpleNamespace(id=value) for value in ids
]
def query(self, *args): # type: ignore[no-untyped-def]
self.query_args = args
return self.query_obj
def rollback(self) -> None:
raise AssertionError("rollback should not be called")
def close(self) -> None:
self.closed = True
class _RecordSession:
def __init__(self, record: SimpleNamespace) -> None:
self.record = record
class _BatchSession:
def __init__(self, records: list[SimpleNamespace]) -> None:
self.records = records
self.closed = False
self.committed = False
self.executed = 0
self.query_obj = MagicMock()
self.query_obj.filter.return_value.first.return_value = record
filtered = self.query_obj.filter.return_value
filtered.filter.return_value.order_by.return_value.limit.return_value.all.return_value = (
records
)
def query(self, *args): # type: ignore[no-untyped-def]
self.query_args = args
return self.query_obj
def execute(self, _statement): # type: ignore[no-untyped-def]
self.executed += 1
return SimpleNamespace(rowcount=1)
def commit(self) -> None:
@@ -173,27 +159,26 @@ def test_cleanup_body_fields_loads_ids_then_processes_records_individually(
def close(self) -> None:
self.closed = True
batch_one = _IdBatchSession(["usage-1", "usage-2"])
record_one = _RecordSession(
SimpleNamespace(
id="usage-1",
request_body={"hello": "world"},
response_body=None,
provider_request_body=None,
client_response_body=None,
)
batch_one = _BatchSession(
[
SimpleNamespace(
id="usage-1",
request_body={"hello": "world"},
response_body=None,
provider_request_body=None,
client_response_body=None,
),
SimpleNamespace(
id="usage-2",
request_body=None,
response_body={"ok": True},
provider_request_body=None,
client_response_body=None,
),
]
)
record_two = _RecordSession(
SimpleNamespace(
id="usage-2",
request_body=None,
response_body={"ok": True},
provider_request_body=None,
client_response_body=None,
)
)
batch_two = _IdBatchSession([])
sessions = iter([batch_one, record_one, record_two, batch_two])
batch_two = _BatchSession([])
sessions = iter([batch_one, batch_two])
monkeypatch.setattr(
maintenance_scheduler_module,
@@ -212,12 +197,173 @@ def test_cleanup_body_fields_loads_ids_then_processes_records_individually(
)
assert compressed == 2
assert batch_one.query_args == (maintenance_scheduler_module.Usage.id,)
assert len(record_one.query_args) == 5
assert len(record_two.query_args) == 5
assert len(batch_one.query_args) == 5
assert batch_one.executed == 2
assert batch_one.committed is True
assert batch_one.closed is True
assert batch_two.closed is True
assert record_one.committed is True
assert record_two.committed is True
assert record_one.closed is True
assert record_two.closed is True
def test_cleanup_header_fields_clears_client_response_headers(
monkeypatch: pytest.MonkeyPatch,
) -> None:
scheduler = MaintenanceScheduler()
class _BatchSession:
def __init__(self, ids: list[str]) -> None:
self.ids = ids
self.closed = False
self.committed = False
self.query_obj = MagicMock()
self.filtered_by_time = MagicMock()
self.filtered_by_headers = MagicMock()
self.query_obj.filter.return_value = self.filtered_by_time
self.filtered_by_time.filter.return_value = self.filtered_by_headers
self.filtered_by_headers.order_by.return_value.limit.return_value.all.return_value = [
SimpleNamespace(id=value) for value in ids
]
self.executed_statements: list[str] = []
def query(self, *args): # type: ignore[no-untyped-def]
self.query_args = args
return self.query_obj
def execute(self, statement): # type: ignore[no-untyped-def]
self.executed_statements.append(str(statement))
return SimpleNamespace(rowcount=len(self.ids))
def commit(self) -> None:
self.committed = True
def rollback(self) -> None:
raise AssertionError("rollback should not be called")
def close(self) -> None:
self.closed = True
batch_one = _BatchSession(["usage-1"])
batch_two = _BatchSession([])
sessions = iter([batch_one, batch_two])
monkeypatch.setattr(
maintenance_scheduler_module,
"create_session",
lambda: next(sessions),
)
cleaned = scheduler._cleanup_header_fields(
cutoff_time=SimpleNamespace(), # type: ignore[arg-type]
batch_size=1000,
)
header_filter = str(batch_one.filtered_by_time.filter.call_args.args[0])
assert cleaned == 1
assert batch_one.query_args == (maintenance_scheduler_module.Usage.id,)
assert "client_response_headers" in header_filter
assert "client_response_headers" in batch_one.executed_statements[0]
assert batch_one.committed is True
assert batch_one.closed is True
assert batch_two.closed is True
@pytest.mark.asyncio
async def test_perform_cleanup_deletes_first_and_uses_non_overlapping_windows(
monkeypatch: pytest.MonkeyPatch,
) -> None:
scheduler = MaintenanceScheduler()
fixed_now = datetime(2026, 3, 18, 3, 0, 0, tzinfo=timezone.utc)
class _FakeDateTime(datetime):
@classmethod
def now(cls, tz=None): # type: ignore[override]
if tz is None:
return fixed_now.replace(tzinfo=None)
return fixed_now.astimezone(tz)
class _FakeLoop:
async def run_in_executor(self, _executor, func): # type: ignore[no-untyped-def]
return func()
class _ConfigSession:
def close(self) -> None:
return None
calls: list[tuple[str, datetime, int, datetime | None]] = []
def _record(name: str, count: int):
def _inner(
cutoff_time: datetime,
batch_size: int,
*,
newer_than: datetime | None = None,
) -> int:
calls.append((name, cutoff_time, batch_size, newer_than))
return count
return _inner
config_values = {
"enable_auto_cleanup": True,
"detail_log_retention_days": 7,
"compressed_log_retention_days": 30,
"header_retention_days": 90,
"log_retention_days": 365,
"cleanup_batch_size": 123,
"auto_delete_expired_keys": False,
}
monkeypatch.setattr(maintenance_scheduler_module, "datetime", _FakeDateTime)
monkeypatch.setattr(
maintenance_scheduler_module.asyncio, "get_running_loop", lambda: _FakeLoop()
)
monkeypatch.setattr(
maintenance_scheduler_module,
"create_session",
lambda: _ConfigSession(),
)
monkeypatch.setattr(
maintenance_scheduler_module.SystemConfigService,
"get_config",
lambda _db, key, default=None: config_values.get(key, default),
)
monkeypatch.setattr(
scheduler,
"_delete_old_records",
lambda cutoff_time, batch_size: calls.append(("delete", cutoff_time, batch_size, None))
or 5,
)
monkeypatch.setattr(
scheduler,
"_cleanup_header_fields",
_record("header", 4),
)
monkeypatch.setattr(
scheduler,
"_cleanup_stale_body_fields",
_record("body", 3),
)
monkeypatch.setattr(
scheduler,
"_cleanup_body_fields",
_record("compress", 2),
)
monkeypatch.setattr(
maintenance_scheduler_module.ApiKeyService,
"cleanup_expired_keys",
lambda _db, auto_delete=False: 0,
)
await scheduler._perform_cleanup()
detail_cutoff = fixed_now - timedelta(days=7)
compressed_cutoff = fixed_now - timedelta(days=30)
header_cutoff = fixed_now - timedelta(days=90)
log_cutoff = fixed_now - timedelta(days=365)
assert calls == [
("delete", log_cutoff, 123, None),
("header", header_cutoff, 123, log_cutoff),
("body", compressed_cutoff, 123, log_cutoff),
("compress", detail_cutoff, 123, compressed_cutoff),
]

View File

@@ -0,0 +1,173 @@
from __future__ import annotations
from datetime import datetime, timedelta, timezone
from types import SimpleNamespace
from typing import Any, cast
from unittest.mock import MagicMock
import pytest
from src.services.health.endpoint import EndpointHealthService
class _FakeQuery:
def __init__(self, rows: list[SimpleNamespace]) -> None:
self._rows = rows
def filter(self, *args: Any, **kwargs: Any) -> _FakeQuery:
return self
def group_by(self, *args: Any, **kwargs: Any) -> _FakeQuery:
return self
def all(self) -> list[SimpleNamespace]:
return self._rows
class _FakeDb:
def __init__(self, rows: list[SimpleNamespace]) -> None:
self._rows = rows
def query(self, *args: Any, **kwargs: Any) -> _FakeQuery:
return _FakeQuery(self._rows)
def _expr_texts(query: MagicMock) -> list[str]:
return [str(arg) for arg in query.filter.call_args.args]
def test_generate_timeline_batch_keeps_compact_and_cli_isolated() -> None:
now = datetime(2026, 3, 18, 12, 0, tzinfo=timezone.utc)
db = _FakeDb(
[
SimpleNamespace(
endpoint_id="endpoint-compact",
segment_idx=0,
total_count=2,
success_count=2,
failed_count=0,
min_time=now - timedelta(minutes=55),
max_time=now - timedelta(minutes=40),
),
SimpleNamespace(
endpoint_id="endpoint-cli",
segment_idx=0,
total_count=3,
success_count=0,
failed_count=3,
min_time=now - timedelta(minutes=54),
max_time=now - timedelta(minutes=39),
),
]
)
result = EndpointHealthService._generate_timeline_batch(
db=cast(Any, db),
format_endpoint_mapping={
"openai:compact": ["endpoint-compact"],
"openai:cli": ["endpoint-cli"],
},
now=now,
lookback_hours=1,
segments=4,
)
assert result["openai:compact"]["timeline"][0] == "healthy"
assert result["openai:cli"]["timeline"][0] == "unhealthy"
def test_generate_timeline_from_usage_uses_endpoint_ids_directly(
monkeypatch: pytest.MonkeyPatch,
) -> None:
expected = {
"timeline": ["healthy", "warning"],
"time_range_start": "start",
"time_range_end": "end",
}
captured: dict[str, object] = {}
def _fake_generate_timeline_batch(
db: Any,
format_endpoint_mapping: dict[str, list[str]],
now: datetime,
lookback_hours: int,
segments: int,
) -> dict[str, dict[str, Any]]:
captured["db"] = db
captured["mapping"] = format_endpoint_mapping
captured["lookback_hours"] = lookback_hours
captured["segments"] = segments
return {"_single": expected}
monkeypatch.setattr(
EndpointHealthService,
"_generate_timeline_batch",
staticmethod(_fake_generate_timeline_batch),
)
db = cast(Any, object())
now = datetime(2026, 3, 18, 12, 0, tzinfo=timezone.utc)
result = EndpointHealthService._generate_timeline_from_usage(
db=db,
endpoint_ids=["endpoint-compact"],
now=now,
lookback_hours=6,
segments=2,
)
assert result == expected
assert captured["db"] is db
assert captured["mapping"] == {"_single": ["endpoint-compact"]}
assert captured["lookback_hours"] == 6
assert captured["segments"] == 2
def test_get_endpoint_health_by_format_filters_inactive_endpoints(
monkeypatch: pytest.MonkeyPatch,
) -> None:
endpoint_query = MagicMock()
endpoint_query.join.return_value = endpoint_query
endpoint_query.filter.return_value = endpoint_query
endpoint_query.all.return_value = [
SimpleNamespace(
id="endpoint-compact",
provider_id="provider-1",
api_format="openai:compact",
is_active=True,
)
]
key_query = MagicMock()
key_query.filter.return_value = key_query
key_query.options.return_value = key_query
key_query.all.return_value = []
db = MagicMock()
db.query.side_effect = [endpoint_query, key_query]
monkeypatch.setattr(
EndpointHealthService,
"_generate_timeline_batch",
staticmethod(
lambda db, format_endpoint_mapping, now, lookback_hours: {
"openai:compact": {
"timeline": ["unknown"] * 100,
"time_range_start": None,
"time_range_end": None,
}
}
),
)
EndpointHealthService.get_endpoint_health_by_format(
db=cast(Any, db),
lookback_hours=6,
include_admin_fields=False,
use_cache=False,
)
filters = _expr_texts(endpoint_query)
assert any("provider_endpoints.is_active" in expr for expr in filters)
assert any("providers.is_active" in expr for expr in filters)