mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-01 17:00:21 +08:00
refactor: 共享请求管道、按需懒加载、流式内存护栏与连接池治理
- 抽取 ApiRequestPipeline 单例,44 个路由文件共享同一实例 - Handler/Adapter 模块级 __getattr__ 延迟导入,减少启动时间 - 新增 ensure_stream_buffer_limit() 流式内存护栏(16MB 单行 / 32MB 总量) - HTTP 空闲连接清理与 curl_cffi LRU 会话池 - ensure_providers_bootstrapped 按需引导指定 provider_types - Usage 事件序列化迁移至 msgpack,Redis codec 隔离 - 启动预热任务(/readyz 就绪门控)与优雅关闭 - 通知邮件模块独立开关与 SMTP 配置校验 - CryptoService DCL 线程安全修复 - 通知模块开关 DB 查询 30s 内存缓存 - /readyz 对 unknown 状态返回 503 - 预热关闭 5s 超时保护 - 预热适配器逐个 try-except 容错 - FormatConversionRegistry 哨兵模式防并发重复物化 - 流式缓冲检查无条件执行 Closes #230 Co-authored-by: AAEE86 <ppk0227@hotmail.com>
This commit is contained in:
@@ -8,8 +8,10 @@ import pytest
|
||||
from src.api.handlers.base.upstream_stream_bridge import (
|
||||
aggregate_upstream_stream_to_internal_response,
|
||||
)
|
||||
from src.config.constants import StreamDefaults
|
||||
from src.core.api_format.conversion import register_default_normalizers
|
||||
from src.core.api_format.conversion.internal import TextBlock
|
||||
from src.core.exceptions import ProviderNotAvailableException
|
||||
|
||||
|
||||
async def _iter_stream_lines(lines: list[str]) -> AsyncIterator[bytes]:
|
||||
@@ -82,3 +84,62 @@ async def test_aggregate_claude_stream_uses_message_start_usage_when_message_del
|
||||
assert len(internal.content) == 1
|
||||
assert isinstance(internal.content[0], TextBlock)
|
||||
assert internal.content[0].text == "hello"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aggregate_stream_raises_when_buffer_exceeds_limit() -> None:
|
||||
register_default_normalizers()
|
||||
|
||||
async def _iter_overflow_bytes() -> AsyncIterator[bytes]:
|
||||
yield b"x" * (StreamDefaults.MAX_STREAM_BUFFER_BYTES + 1)
|
||||
|
||||
with pytest.raises(ProviderNotAvailableException):
|
||||
await aggregate_upstream_stream_to_internal_response(
|
||||
_iter_overflow_bytes(),
|
||||
provider_api_format="claude:cli",
|
||||
provider_name="claude_code",
|
||||
model="claude-sonnet-4-5",
|
||||
request_id="req_bridge_overflow",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aggregate_stream_raises_when_total_buffer_exceeds_hard_limit(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
register_default_normalizers()
|
||||
monkeypatch.setattr(StreamDefaults, "MAX_STREAM_BUFFER_BYTES", 64)
|
||||
monkeypatch.setattr(StreamDefaults, "MAX_STREAM_BUFFER_TOTAL_BYTES", 80)
|
||||
|
||||
async def _iter_total_overflow_bytes() -> AsyncIterator[bytes]:
|
||||
yield b":" + (b"a" * 30) + b"\n" + b":" + (b"b" * 30) + b"\n" + b":" + (b"c" * 30) + b"\n"
|
||||
|
||||
with pytest.raises(ProviderNotAvailableException):
|
||||
await aggregate_upstream_stream_to_internal_response(
|
||||
_iter_total_overflow_bytes(),
|
||||
provider_api_format="claude:cli",
|
||||
provider_name="claude_code",
|
||||
model="claude-sonnet-4-5",
|
||||
request_id="req_bridge_total_overflow",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aggregate_stream_allows_large_chunk_with_multiple_complete_lines(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
register_default_normalizers()
|
||||
monkeypatch.setattr(StreamDefaults, "MAX_STREAM_BUFFER_BYTES", 64)
|
||||
|
||||
async def _iter_multiline_bytes() -> AsyncIterator[bytes]:
|
||||
yield b":" + (b"a" * 30) + b"\n" + b":" + (b"b" * 30) + b"\n" + b":" + (b"c" * 30) + b"\n"
|
||||
|
||||
internal = await aggregate_upstream_stream_to_internal_response(
|
||||
_iter_multiline_bytes(),
|
||||
provider_api_format="claude:cli",
|
||||
provider_name="claude_code",
|
||||
model="claude-sonnet-4-5",
|
||||
request_id="req_bridge_multiline",
|
||||
)
|
||||
|
||||
assert internal is not None
|
||||
|
||||
41
tests/clients/test_redis_client.py
Normal file
41
tests/clients/test_redis_client.py
Normal file
@@ -0,0 +1,41 @@
|
||||
from src.clients import redis_client as redis_client_module
|
||||
from src.clients.redis_client import RedisClientManager, RedisState
|
||||
|
||||
|
||||
def _seed_open_circuit(manager: RedisClientManager) -> None:
|
||||
manager._circuit_open_until = 9999999999.0
|
||||
manager._consecutive_failures = 5
|
||||
manager._last_error = "boom"
|
||||
|
||||
|
||||
def test_reset_redis_circuit_breaker_resets_both_clients() -> None:
|
||||
old_global = redis_client_module._redis_manager
|
||||
old_usage = redis_client_module._usage_queue_redis_manager
|
||||
try:
|
||||
global_manager = RedisClientManager(client_name="global")
|
||||
usage_manager = RedisClientManager(client_name="usage")
|
||||
_seed_open_circuit(global_manager)
|
||||
_seed_open_circuit(usage_manager)
|
||||
|
||||
redis_client_module._redis_manager = global_manager
|
||||
redis_client_module._usage_queue_redis_manager = usage_manager
|
||||
|
||||
assert redis_client_module.reset_redis_circuit_breaker() is True
|
||||
assert global_manager.get_state() == RedisState.NOT_INITIALIZED
|
||||
assert usage_manager.get_state() == RedisState.NOT_INITIALIZED
|
||||
finally:
|
||||
redis_client_module._redis_manager = old_global
|
||||
redis_client_module._usage_queue_redis_manager = old_usage
|
||||
|
||||
|
||||
def test_reset_redis_circuit_breaker_returns_false_when_uninitialized() -> None:
|
||||
old_global = redis_client_module._redis_manager
|
||||
old_usage = redis_client_module._usage_queue_redis_manager
|
||||
try:
|
||||
redis_client_module._redis_manager = None
|
||||
redis_client_module._usage_queue_redis_manager = None
|
||||
|
||||
assert redis_client_module.reset_redis_circuit_breaker() is False
|
||||
finally:
|
||||
redis_client_module._redis_manager = old_global
|
||||
redis_client_module._usage_queue_redis_manager = old_usage
|
||||
@@ -19,11 +19,9 @@ def _make_key(
|
||||
return key
|
||||
|
||||
|
||||
@patch(
|
||||
"src.services.scheduling.candidate_builder.health_monitor.get_circuit_breaker_status",
|
||||
return_value=(True, None),
|
||||
)
|
||||
def test_kiro_quota_remaining_zero_skips(_mock_cb: MagicMock) -> None:
|
||||
@patch("src.services.scheduling.candidate_builder.get_health_monitor")
|
||||
def test_kiro_quota_remaining_zero_skips(mock_get_health_monitor: MagicMock) -> None:
|
||||
mock_get_health_monitor.return_value.get_circuit_breaker_status.return_value = (True, None)
|
||||
scheduler = CacheAwareScheduler()
|
||||
key = _make_key(upstream_metadata={"kiro": {"remaining": 0.0}})
|
||||
|
||||
@@ -38,11 +36,9 @@ def test_kiro_quota_remaining_zero_skips(_mock_cb: MagicMock) -> None:
|
||||
assert reason == "Kiro 账号配额剩余 0"
|
||||
|
||||
|
||||
@patch(
|
||||
"src.services.scheduling.candidate_builder.health_monitor.get_circuit_breaker_status",
|
||||
return_value=(True, None),
|
||||
)
|
||||
def test_kiro_quota_remaining_positive_allows(_mock_cb: MagicMock) -> None:
|
||||
@patch("src.services.scheduling.candidate_builder.get_health_monitor")
|
||||
def test_kiro_quota_remaining_positive_allows(mock_get_health_monitor: MagicMock) -> None:
|
||||
mock_get_health_monitor.return_value.get_circuit_breaker_status.return_value = (True, None)
|
||||
scheduler = CacheAwareScheduler()
|
||||
key = _make_key(upstream_metadata={"kiro": {"remaining": 1.0}})
|
||||
|
||||
@@ -57,11 +53,9 @@ def test_kiro_quota_remaining_positive_allows(_mock_cb: MagicMock) -> None:
|
||||
assert reason is None
|
||||
|
||||
|
||||
@patch(
|
||||
"src.services.scheduling.candidate_builder.health_monitor.get_circuit_breaker_status",
|
||||
return_value=(True, None),
|
||||
)
|
||||
def test_codex_weekly_quota_exhausted_skips(_mock_cb: MagicMock) -> None:
|
||||
@patch("src.services.scheduling.candidate_builder.get_health_monitor")
|
||||
def test_codex_weekly_quota_exhausted_skips(mock_get_health_monitor: MagicMock) -> None:
|
||||
mock_get_health_monitor.return_value.get_circuit_breaker_status.return_value = (True, None)
|
||||
scheduler = CacheAwareScheduler()
|
||||
key = _make_key(
|
||||
upstream_metadata={
|
||||
@@ -83,11 +77,9 @@ def test_codex_weekly_quota_exhausted_skips(_mock_cb: MagicMock) -> None:
|
||||
assert reason == "Codex 周限额剩余 0%"
|
||||
|
||||
|
||||
@patch(
|
||||
"src.services.scheduling.candidate_builder.health_monitor.get_circuit_breaker_status",
|
||||
return_value=(True, None),
|
||||
)
|
||||
def test_codex_5h_quota_exhausted_skips(_mock_cb: MagicMock) -> None:
|
||||
@patch("src.services.scheduling.candidate_builder.get_health_monitor")
|
||||
def test_codex_5h_quota_exhausted_skips(mock_get_health_monitor: MagicMock) -> None:
|
||||
mock_get_health_monitor.return_value.get_circuit_breaker_status.return_value = (True, None)
|
||||
scheduler = CacheAwareScheduler()
|
||||
key = _make_key(
|
||||
upstream_metadata={
|
||||
@@ -109,11 +101,9 @@ def test_codex_5h_quota_exhausted_skips(_mock_cb: MagicMock) -> None:
|
||||
assert reason == "Codex 5H 限额剩余 0%"
|
||||
|
||||
|
||||
@patch(
|
||||
"src.services.scheduling.candidate_builder.health_monitor.get_circuit_breaker_status",
|
||||
return_value=(True, None),
|
||||
)
|
||||
def test_codex_ignores_unrelated_metadata_fields(_mock_cb: MagicMock) -> None:
|
||||
@patch("src.services.scheduling.candidate_builder.get_health_monitor")
|
||||
def test_codex_ignores_unrelated_metadata_fields(mock_get_health_monitor: MagicMock) -> None:
|
||||
mock_get_health_monitor.return_value.get_circuit_breaker_status.return_value = (True, None)
|
||||
scheduler = CacheAwareScheduler()
|
||||
key = _make_key(
|
||||
upstream_metadata={
|
||||
@@ -136,11 +126,9 @@ def test_codex_ignores_unrelated_metadata_fields(_mock_cb: MagicMock) -> None:
|
||||
assert reason is None
|
||||
|
||||
|
||||
@patch(
|
||||
"src.services.scheduling.candidate_builder.health_monitor.get_circuit_breaker_status",
|
||||
return_value=(True, None),
|
||||
)
|
||||
def test_antigravity_model_quota_exhausted_skips(_mock_cb: MagicMock) -> None:
|
||||
@patch("src.services.scheduling.candidate_builder.get_health_monitor")
|
||||
def test_antigravity_model_quota_exhausted_skips(mock_get_health_monitor: MagicMock) -> None:
|
||||
mock_get_health_monitor.return_value.get_circuit_breaker_status.return_value = (True, None)
|
||||
scheduler = CacheAwareScheduler()
|
||||
key = _make_key(
|
||||
upstream_metadata={
|
||||
@@ -164,11 +152,9 @@ def test_antigravity_model_quota_exhausted_skips(_mock_cb: MagicMock) -> None:
|
||||
assert reason == "Antigravity 模型 ag-model 配额剩余 0%"
|
||||
|
||||
|
||||
@patch(
|
||||
"src.services.scheduling.candidate_builder.health_monitor.get_circuit_breaker_status",
|
||||
return_value=(True, None),
|
||||
)
|
||||
def test_antigravity_other_model_not_exhausted_allows(_mock_cb: MagicMock) -> None:
|
||||
@patch("src.services.scheduling.candidate_builder.get_health_monitor")
|
||||
def test_antigravity_other_model_not_exhausted_allows(mock_get_health_monitor: MagicMock) -> None:
|
||||
mock_get_health_monitor.return_value.get_circuit_breaker_status.return_value = (True, None)
|
||||
scheduler = CacheAwareScheduler()
|
||||
key = _make_key(
|
||||
upstream_metadata={
|
||||
@@ -192,11 +178,9 @@ def test_antigravity_other_model_not_exhausted_allows(_mock_cb: MagicMock) -> No
|
||||
assert reason is None
|
||||
|
||||
|
||||
@patch(
|
||||
"src.services.scheduling.candidate_builder.health_monitor.get_circuit_breaker_status",
|
||||
return_value=(True, None),
|
||||
)
|
||||
def test_antigravity_quota_uses_mapping_matched_model(_mock_cb: MagicMock) -> None:
|
||||
@patch("src.services.scheduling.candidate_builder.get_health_monitor")
|
||||
def test_antigravity_quota_uses_mapping_matched_model(mock_get_health_monitor: MagicMock) -> None:
|
||||
mock_get_health_monitor.return_value.get_circuit_breaker_status.return_value = (True, None)
|
||||
scheduler = CacheAwareScheduler()
|
||||
|
||||
# Request uses GlobalModel.name, but allowed_models only contains provider-side model id.
|
||||
|
||||
@@ -57,11 +57,12 @@ async def test_executor_records_health_by_provider_format() -> None:
|
||||
patch("src.services.request.executor.RequestCandidateService.mark_candidate_started"),
|
||||
patch("src.services.request.executor.RequestCandidateService.mark_candidate_success"),
|
||||
patch("src.services.request.executor.get_adaptive_reservation_manager") as mock_res_mgr,
|
||||
patch("src.services.request.executor.health_monitor.record_success") as record_success,
|
||||
patch("src.services.request.executor.get_health_monitor") as mock_get_health_monitor,
|
||||
):
|
||||
mock_res_mgr.return_value.calculate_reservation.return_value = MagicMock(
|
||||
ratio=0.0, phase="stable", confidence=1.0
|
||||
)
|
||||
record_success = mock_get_health_monitor.return_value.record_success
|
||||
|
||||
executor = RequestExecutor(
|
||||
db=db, concurrency_manager=concurrency_manager, adaptive_manager=adaptive_manager
|
||||
@@ -100,8 +101,9 @@ async def test_error_classifier_records_failure_by_provider_format() -> None:
|
||||
key.id = "k1"
|
||||
|
||||
with patch(
|
||||
"src.services.orchestration.error_handler.health_monitor.record_failure"
|
||||
) as record_failure:
|
||||
"src.services.orchestration.error_handler.get_health_monitor"
|
||||
) as mock_get_health_monitor:
|
||||
record_failure = mock_get_health_monitor.return_value.record_failure
|
||||
await classifier.handle_retriable_error(
|
||||
error=RuntimeError("boom"),
|
||||
provider=provider,
|
||||
|
||||
@@ -6,7 +6,7 @@ from types import SimpleNamespace
|
||||
|
||||
from src.services.provider.pool.config import PoolConfig, SchedulingPreset, ScoringWeights
|
||||
from src.services.provider.pool.strategies.multi_score import MultiScoreStrategy
|
||||
from src.services.provider.pool.strategy import get_pool_strategy, register_pool_strategy
|
||||
from src.services.provider.pool.strategy import get_pool_strategy
|
||||
|
||||
|
||||
def _context() -> dict:
|
||||
@@ -61,9 +61,8 @@ def test_multi_score_combines_health_and_cost() -> None:
|
||||
|
||||
|
||||
def test_multi_score_strategy_is_registered() -> None:
|
||||
register_pool_strategy("multi_score", MultiScoreStrategy())
|
||||
registered = get_pool_strategy("multi_score")
|
||||
assert registered is not None
|
||||
assert isinstance(registered, MultiScoreStrategy)
|
||||
|
||||
|
||||
def test_multi_score_preset_free_team_first_prefers_free_or_team() -> None:
|
||||
@@ -196,6 +195,51 @@ def test_multi_score_preset_recent_refresh_uses_codex_weekly_reset() -> None:
|
||||
assert s2 < s1
|
||||
|
||||
|
||||
def test_multi_score_codex_default_enables_recent_refresh_when_missing() -> None:
|
||||
strategy = MultiScoreStrategy()
|
||||
cfg = PoolConfig(
|
||||
scheduling_mode="multi_score",
|
||||
scheduling_presets=(SchedulingPreset(preset="cache_affinity", enabled=True),),
|
||||
)
|
||||
ctx = {
|
||||
"provider_type": "codex",
|
||||
"all_key_ids": ["k1", "k2"],
|
||||
"lru_scores": {},
|
||||
"keys_by_id": {
|
||||
"k1": _key_with_metadata({"codex": {"primary_reset_seconds": 600}}),
|
||||
"k2": _key_with_metadata({"codex": {"primary_reset_seconds": 120}}),
|
||||
},
|
||||
}
|
||||
s1 = strategy.compute_score(key_id="k1", config=cfg, context=ctx)
|
||||
s2 = strategy.compute_score(key_id="k2", config=cfg, context=ctx)
|
||||
assert s1 is not None and s2 is not None
|
||||
assert s2 < s1
|
||||
|
||||
|
||||
def test_multi_score_codex_recent_refresh_can_be_explicitly_disabled() -> None:
|
||||
strategy = MultiScoreStrategy()
|
||||
cfg = PoolConfig(
|
||||
scheduling_mode="multi_score",
|
||||
scheduling_presets=(
|
||||
SchedulingPreset(preset="cache_affinity", enabled=True),
|
||||
SchedulingPreset(preset="recent_refresh", enabled=False),
|
||||
),
|
||||
)
|
||||
ctx = {
|
||||
"provider_type": "codex",
|
||||
"all_key_ids": ["k1", "k2"],
|
||||
"lru_scores": {},
|
||||
"keys_by_id": {
|
||||
"k1": _key_with_metadata({"codex": {"primary_reset_seconds": 600}}),
|
||||
"k2": _key_with_metadata({"codex": {"primary_reset_seconds": 120}}),
|
||||
},
|
||||
}
|
||||
s1 = strategy.compute_score(key_id="k1", config=cfg, context=ctx)
|
||||
s2 = strategy.compute_score(key_id="k2", config=cfg, context=ctx)
|
||||
assert s1 is not None and s2 is not None
|
||||
assert s1 < s2
|
||||
|
||||
|
||||
def test_multi_score_preset_single_account_prefers_internal_priority_then_reverse_lru() -> None:
|
||||
strategy = MultiScoreStrategy()
|
||||
cfg = PoolConfig(
|
||||
|
||||
@@ -88,6 +88,72 @@ 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_prefers_codex_reset_at(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr("src.services.provider.pool.dimensions._helpers.time.time", lambda: 1000.0)
|
||||
key_obj = SimpleNamespace(
|
||||
provider_type="codex",
|
||||
upstream_metadata={
|
||||
"codex": {
|
||||
"primary_reset_at": 1200.0,
|
||||
"primary_reset_seconds": 9999.0,
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
assert extract_reset_seconds(key_obj) == pytest.approx(200.0)
|
||||
|
||||
|
||||
def test_extract_reset_seconds_codex_fallback_corrects_elapsed(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setattr("src.services.provider.pool.dimensions._helpers.time.time", lambda: 2000.0)
|
||||
key_obj = SimpleNamespace(
|
||||
provider_type="codex",
|
||||
upstream_metadata={
|
||||
"codex": {
|
||||
"primary_reset_seconds": 180.0,
|
||||
"updated_at": 1900.0,
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
assert extract_reset_seconds(key_obj) == pytest.approx(80.0)
|
||||
|
||||
|
||||
def test_extract_reset_seconds_codex_fallback_clamps_to_zero(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setattr("src.services.provider.pool.dimensions._helpers.time.time", lambda: 2000.0)
|
||||
key_obj = SimpleNamespace(
|
||||
provider_type="codex",
|
||||
upstream_metadata={
|
||||
"codex": {
|
||||
"primary_reset_seconds": 120.0,
|
||||
"updated_at": 1500.0,
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
assert extract_reset_seconds(key_obj) == pytest.approx(0.0)
|
||||
|
||||
|
||||
def test_extract_reset_seconds_codex_fallback_future_updated_at_no_inflation(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setattr("src.services.provider.pool.dimensions._helpers.time.time", lambda: 2000.0)
|
||||
key_obj = SimpleNamespace(
|
||||
provider_type="codex",
|
||||
upstream_metadata={
|
||||
"codex": {
|
||||
"primary_reset_seconds": 120.0,
|
||||
"updated_at": 2600.0,
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
assert extract_reset_seconds(key_obj) == pytest.approx(120.0)
|
||||
|
||||
|
||||
def test_resolve_pool_account_state_keeps_codex_metadata_block() -> None:
|
||||
state = resolve_pool_account_state(
|
||||
provider_type="codex",
|
||||
|
||||
@@ -44,6 +44,13 @@ async def test_maintenance_scheduler_start_skips_startup_task_when_disabled(
|
||||
assert created is False
|
||||
|
||||
|
||||
def test_http_client_idle_cleanup_interval_env_invalid(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setenv("HTTP_CLIENT_IDLE_CLEANUP_INTERVAL_MINUTES", "bad")
|
||||
assert MaintenanceScheduler._get_http_client_idle_cleanup_interval_minutes() == 5
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_candidate_cleanup_uses_dedicated_retention_and_batch_settings(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
|
||||
@@ -180,3 +180,73 @@ async def test_prepare_usage_record_infers_ttl_from_1h_cache_split(
|
||||
|
||||
assert _DummyBillingService.last_dimensions is not None
|
||||
assert _DummyBillingService.last_dimensions.get("cache_ttl_minutes") == 60
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_prepare_usage_record_deserializes_body_json_before_build_params(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
db = MagicMock()
|
||||
|
||||
monkeypatch.setattr("src.services.billing.service.BillingService", _DummyBillingService)
|
||||
monkeypatch.setattr(
|
||||
"src.services.usage._billing_integration.sanitize_request_metadata",
|
||||
lambda metadata: metadata,
|
||||
)
|
||||
|
||||
captured: dict[str, Any] = {}
|
||||
|
||||
def _capture_build_usage_params(**kwargs: Any) -> dict[str, Any]:
|
||||
captured.update(kwargs)
|
||||
return {"total_cost_usd": 0.0, "actual_total_cost_usd": 0.0}
|
||||
|
||||
monkeypatch.setattr(
|
||||
"src.services.usage._billing_integration.build_usage_params",
|
||||
_capture_build_usage_params,
|
||||
)
|
||||
|
||||
params = _build_params(db)
|
||||
params.request_body = '{"messages":[{"role":"user","content":"hello"}]}'
|
||||
params.provider_request_body = '{"tools":[{"name":"calc"}]}'
|
||||
params.response_body = '{"choices":[{"index":0}]}'
|
||||
params.client_response_body = '{"output":[{"type":"text"}]}'
|
||||
|
||||
await _TestUsageBillingIntegration._prepare_usage_record(params)
|
||||
|
||||
assert isinstance(captured["request_body"], dict)
|
||||
assert captured["request_body"]["messages"][0]["content"] == "hello"
|
||||
assert isinstance(captured["provider_request_body"], dict)
|
||||
assert isinstance(captured["response_body"], dict)
|
||||
assert isinstance(captured["client_response_body"], dict)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_prepare_usage_record_keeps_invalid_json_body_string(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
db = MagicMock()
|
||||
|
||||
monkeypatch.setattr("src.services.billing.service.BillingService", _DummyBillingService)
|
||||
monkeypatch.setattr(
|
||||
"src.services.usage._billing_integration.sanitize_request_metadata",
|
||||
lambda metadata: metadata,
|
||||
)
|
||||
|
||||
captured: dict[str, Any] = {}
|
||||
|
||||
def _capture_build_usage_params(**kwargs: Any) -> dict[str, Any]:
|
||||
captured.update(kwargs)
|
||||
return {"total_cost_usd": 0.0, "actual_total_cost_usd": 0.0}
|
||||
|
||||
monkeypatch.setattr(
|
||||
"src.services.usage._billing_integration.build_usage_params",
|
||||
_capture_build_usage_params,
|
||||
)
|
||||
|
||||
invalid_json = '{"content":"x...[truncated]'
|
||||
params = _build_params(db)
|
||||
params.request_body = invalid_json
|
||||
|
||||
await _TestUsageBillingIntegration._prepare_usage_record(params)
|
||||
|
||||
assert captured["request_body"] == invalid_json
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
import asyncio
|
||||
import json
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
from redis.exceptions import ResponseError
|
||||
@@ -26,13 +25,13 @@ from src.services.usage.telemetry_writer import (
|
||||
|
||||
class DummyRedis:
|
||||
def __init__(self) -> None:
|
||||
self.calls: list[tuple[str, dict[str, str], int | None, bool | None]] = []
|
||||
self.calls: list[tuple[str, dict[str, bytes], int | None, bool | None]] = []
|
||||
self.xadd_error: Exception | None = None
|
||||
|
||||
async def xadd(
|
||||
self,
|
||||
key: str,
|
||||
fields: dict[str, str],
|
||||
fields: dict[str, bytes],
|
||||
maxlen: int | None = None,
|
||||
approximate: bool | None = None,
|
||||
) -> str:
|
||||
@@ -118,20 +117,67 @@ async def test_usage_event_all_types() -> None:
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_usage_event_bytes_payload() -> None:
|
||||
"""测试 bytes 类型 payload 的反序列化"""
|
||||
async def test_usage_event_to_stream_fields_sanitizes_non_json_value() -> None:
|
||||
event = build_usage_event(
|
||||
event_type=UsageEventType.COMPLETED,
|
||||
request_id="req-non-json",
|
||||
data={"meta": {"obj": object()}},
|
||||
)
|
||||
fields = event.to_stream_fields()
|
||||
restored = UsageEvent.from_stream_fields(fields)
|
||||
assert isinstance(restored.data["meta"]["obj"], str)
|
||||
|
||||
|
||||
def test_usage_event_bytes_payload() -> None:
|
||||
"""测试 surrogateescape 字符串 payload 的反序列化(msgpack)"""
|
||||
event = build_usage_event(
|
||||
event_type=UsageEventType.COMPLETED,
|
||||
request_id="req-bytes",
|
||||
data={"key": "value"},
|
||||
)
|
||||
fields = event.to_stream_fields()
|
||||
# 模拟 Redis 返回 bytes
|
||||
fields["payload"] = fields["payload"].encode("utf-8") # type: ignore[assignment]
|
||||
# 模拟 decode_responses=True + surrogateescape 返回的 str
|
||||
fields["payload"] = fields["payload"].decode("utf-8", errors="surrogateescape") # type: ignore[assignment]
|
||||
restored = UsageEvent.from_stream_fields(fields)
|
||||
assert restored.request_id == "req-bytes"
|
||||
|
||||
|
||||
def test_usage_event_legacy_json_payload_string() -> None:
|
||||
"""兼容旧格式:payload 为 JSON 字符串"""
|
||||
fields = {
|
||||
"payload": json.dumps(
|
||||
{
|
||||
"v": 1,
|
||||
"type": UsageEventType.COMPLETED.value,
|
||||
"request_id": "req-legacy-json",
|
||||
"timestamp_ms": 123,
|
||||
"data": {"foo": "bar"},
|
||||
}
|
||||
)
|
||||
}
|
||||
restored = UsageEvent.from_stream_fields(fields)
|
||||
assert restored.request_id == "req-legacy-json"
|
||||
assert restored.data["foo"] == "bar"
|
||||
|
||||
|
||||
def test_usage_event_legacy_json_payload_bytes() -> None:
|
||||
"""兼容旧格式:payload 为 bytes(JSON)"""
|
||||
fields = {
|
||||
"payload": json.dumps(
|
||||
{
|
||||
"v": 1,
|
||||
"type": UsageEventType.COMPLETED.value,
|
||||
"request_id": "req-legacy-json-bytes",
|
||||
"timestamp_ms": 123,
|
||||
"data": {"foo": "bar"},
|
||||
}
|
||||
).encode("utf-8")
|
||||
}
|
||||
restored = UsageEvent.from_stream_fields(fields)
|
||||
assert restored.request_id == "req-legacy-json-bytes"
|
||||
assert restored.data["foo"] == "bar"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_usage_event_missing_payload() -> None:
|
||||
"""测试缺少 payload 字段时抛出异常"""
|
||||
@@ -159,14 +205,13 @@ def test_sanitize_payload_nested() -> None:
|
||||
|
||||
|
||||
def test_parse_body_json_string() -> None:
|
||||
"""测试 _parse_body 正确反序列化 JSON 字符串"""
|
||||
"""测试 _parse_body 在消费阶段不反序列化 JSON 字符串"""
|
||||
from src.services.usage.consumer_streams import _parse_body
|
||||
|
||||
# JSON 字符串应被解析为 dict
|
||||
# JSON 字符串应保持原样,反序列化延迟到写库前
|
||||
json_str = '{"messages": [{"role": "user", "content": "hello"}]}'
|
||||
result = _parse_body(json_str)
|
||||
assert isinstance(result, dict)
|
||||
assert result["messages"][0]["content"] == "hello"
|
||||
assert result == json_str
|
||||
|
||||
|
||||
def test_parse_body_dict_passthrough() -> None:
|
||||
@@ -370,8 +415,8 @@ async def test_queue_writer_failure_preserves_empty_request_headers(monkeypatch:
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_event_to_record_body_deserialization(monkeypatch: Any) -> None:
|
||||
"""测试 _event_to_record 正确反序列化 body 字符串为 dict"""
|
||||
async def test_event_to_record_body_passthrough() -> None:
|
||||
"""测试 _event_to_record 保留 body 字符串,交由写库阶段再解析"""
|
||||
from src.services.usage.consumer_streams import _event_to_record
|
||||
|
||||
# 模拟 QueueTelemetryWriter 产生的事件(body 被序列化为 JSON 字符串)
|
||||
@@ -392,11 +437,9 @@ async def test_event_to_record_body_deserialization(monkeypatch: Any) -> None:
|
||||
|
||||
record = _event_to_record(event)
|
||||
|
||||
# 验证 body 被正确反序列化为 dict
|
||||
assert isinstance(record["request_body"], dict)
|
||||
assert record["request_body"]["messages"][0]["content"] == "hello"
|
||||
assert isinstance(record["response_body"], dict)
|
||||
assert record["response_body"]["choices"][0]["message"]["content"] == "hi"
|
||||
# 消费阶段不做 json.loads,保留原字符串
|
||||
assert record["request_body"] == event.data["request_body"]
|
||||
assert record["response_body"] == event.data["response_body"]
|
||||
assert record["finalized_at"] is not None
|
||||
|
||||
|
||||
|
||||
40
tests/unit/test_curl_cffi_transport_pool.py
Normal file
40
tests/unit/test_curl_cffi_transport_pool.py
Normal file
@@ -0,0 +1,40 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections import OrderedDict
|
||||
|
||||
import pytest
|
||||
|
||||
import src.clients.curl_cffi_transport as transport_module
|
||||
|
||||
|
||||
class _DummySession:
|
||||
def __init__(self, **kwargs: object) -> None:
|
||||
self.kwargs = kwargs
|
||||
self.close_calls = 0
|
||||
|
||||
async def close(self) -> None:
|
||||
self.close_calls += 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_or_create_session_uses_lru_eviction(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(transport_module, "AsyncSession", _DummySession, raising=False)
|
||||
monkeypatch.setattr(transport_module, "_MAX_SESSIONS", 2, raising=False)
|
||||
monkeypatch.setattr(transport_module, "_session_pool", OrderedDict(), raising=False)
|
||||
|
||||
s1 = await transport_module._get_or_create_session("chrome120", "http://p1")
|
||||
s2 = await transport_module._get_or_create_session("chrome120", "http://p2")
|
||||
|
||||
# 命中 s1 使其变成最近使用,随后新增 s3 应淘汰 s2。
|
||||
s1_hit = await transport_module._get_or_create_session("chrome120", "http://p1")
|
||||
assert s1_hit is s1
|
||||
|
||||
s3 = await transport_module._get_or_create_session("chrome120", "http://p3")
|
||||
assert isinstance(s3, _DummySession)
|
||||
|
||||
assert len(transport_module._session_pool) == 2
|
||||
assert "chrome120::http://p1" in transport_module._session_pool
|
||||
assert "chrome120::http://p3" in transport_module._session_pool
|
||||
assert "chrome120::http://p2" not in transport_module._session_pool
|
||||
assert s2.close_calls == 1
|
||||
assert s1.close_calls == 0
|
||||
59
tests/unit/test_http_client_pool_cleanup.py
Normal file
59
tests/unit/test_http_client_pool_cleanup.py
Normal file
@@ -0,0 +1,59 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
|
||||
import pytest
|
||||
|
||||
from src.clients.http_client import HTTPClientPool
|
||||
|
||||
|
||||
class _DummyClient:
|
||||
def __init__(self, *, is_closed: bool = False) -> None:
|
||||
self.is_closed = is_closed
|
||||
self.close_calls = 0
|
||||
|
||||
async def aclose(self) -> None:
|
||||
self.close_calls += 1
|
||||
self.is_closed = True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cleanup_idle_clients_closes_stale_entries(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
now = time.time()
|
||||
stale_proxy = _DummyClient()
|
||||
active_proxy = _DummyClient()
|
||||
already_closed_proxy = _DummyClient(is_closed=True)
|
||||
stale_tunnel = _DummyClient()
|
||||
|
||||
monkeypatch.setattr(
|
||||
HTTPClientPool,
|
||||
"_proxy_clients",
|
||||
{
|
||||
"stale": (stale_proxy, now - 1200),
|
||||
"active": (active_proxy, now - 10),
|
||||
"closed": (already_closed_proxy, now - 1200),
|
||||
},
|
||||
raising=False,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
HTTPClientPool,
|
||||
"_tunnel_clients",
|
||||
{"tunnel-stale": (stale_tunnel, now - 1200)},
|
||||
raising=False,
|
||||
)
|
||||
|
||||
stats = await HTTPClientPool.cleanup_idle_clients(max_idle_seconds=600)
|
||||
|
||||
assert stats["proxy_closed"] == 1
|
||||
assert stats["tunnel_closed"] == 1
|
||||
assert stats["proxy_already_closed"] == 1
|
||||
assert stats["tunnel_already_closed"] == 0
|
||||
|
||||
assert stale_proxy.close_calls == 1
|
||||
assert stale_tunnel.close_calls == 1
|
||||
assert active_proxy.close_calls == 0
|
||||
|
||||
assert "active" in HTTPClientPool._proxy_clients
|
||||
assert "stale" not in HTTPClientPool._proxy_clients
|
||||
assert "closed" not in HTTPClientPool._proxy_clients
|
||||
assert "tunnel-stale" not in HTTPClientPool._tunnel_clients
|
||||
164
tests/unit/test_notification_email_module_toggle.py
Normal file
164
tests/unit/test_notification_email_module_toggle.py
Normal file
@@ -0,0 +1,164 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from types import SimpleNamespace
|
||||
from typing import Any, cast
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from starlette.requests import Request
|
||||
|
||||
import src.database as database_module
|
||||
from src.middleware.plugin_middleware import PluginMiddleware
|
||||
from src.modules.notification_email import notification_email_module
|
||||
from src.services.system.config import SystemConfigService
|
||||
|
||||
|
||||
def _build_request(path: str = "/boom") -> Request:
|
||||
scope = {
|
||||
"type": "http",
|
||||
"asgi": {"version": "3.0"},
|
||||
"http_version": "1.1",
|
||||
"method": "GET",
|
||||
"scheme": "http",
|
||||
"path": path,
|
||||
"raw_path": path.encode("utf-8"),
|
||||
"query_string": b"",
|
||||
"headers": [],
|
||||
"client": ("127.0.0.1", 12345),
|
||||
"server": ("testserver", 80),
|
||||
}
|
||||
return Request(scope)
|
||||
|
||||
|
||||
async def _dummy_app(scope, receive, send): # type: ignore[no-untyped-def]
|
||||
return
|
||||
|
||||
|
||||
def test_notification_email_module_validate_config_requires_smtp(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setattr(
|
||||
"src.services.email.email_sender.EmailSenderService.is_smtp_configured",
|
||||
lambda _db: False,
|
||||
)
|
||||
assert notification_email_module.validate_config is not None
|
||||
ok, message = notification_email_module.validate_config(object()) # type: ignore[arg-type]
|
||||
assert ok is False
|
||||
assert "SMTP" in message
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_error_plugins_skips_when_module_disabled(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
middleware = PluginMiddleware(_dummy_app)
|
||||
send_error = AsyncMock()
|
||||
middleware.plugin_manager = cast(
|
||||
Any,
|
||||
SimpleNamespace(
|
||||
get_plugin=lambda *_args, **_kwargs: SimpleNamespace(
|
||||
enabled=True, send_error=send_error
|
||||
)
|
||||
),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
middleware,
|
||||
"_is_notification_email_module_enabled",
|
||||
AsyncMock(return_value=False),
|
||||
)
|
||||
|
||||
await middleware._call_error_plugins(
|
||||
_build_request(),
|
||||
RuntimeError("boom"),
|
||||
start_time=time.time(),
|
||||
)
|
||||
|
||||
send_error.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_error_plugins_sends_when_module_enabled(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
middleware = PluginMiddleware(_dummy_app)
|
||||
send_error = AsyncMock()
|
||||
middleware.plugin_manager = cast(
|
||||
Any,
|
||||
SimpleNamespace(
|
||||
get_plugin=lambda *_args, **_kwargs: SimpleNamespace(
|
||||
enabled=True, send_error=send_error
|
||||
)
|
||||
),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
middleware,
|
||||
"_is_notification_email_module_enabled",
|
||||
AsyncMock(return_value=True),
|
||||
)
|
||||
request = _build_request("/internal-error")
|
||||
request.state.request_id = "req-1"
|
||||
|
||||
await middleware._call_error_plugins(
|
||||
request,
|
||||
RuntimeError("boom"),
|
||||
start_time=time.time(),
|
||||
)
|
||||
|
||||
send_error.assert_awaited_once()
|
||||
assert send_error.await_args is not None
|
||||
kwargs = send_error.await_args.kwargs
|
||||
assert kwargs["context"]["endpoint"] == "GET /internal-error"
|
||||
assert kwargs["context"]["request_id"] == "req-1"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_notification_email_switch_reads_from_system_config(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
middleware = PluginMiddleware(_dummy_app)
|
||||
request = _build_request("/x")
|
||||
|
||||
class _FakeDb:
|
||||
closed = False
|
||||
|
||||
def close(self) -> None:
|
||||
self.closed = True
|
||||
|
||||
fake_db = _FakeDb()
|
||||
monkeypatch.setattr(database_module, "create_session", lambda: fake_db)
|
||||
monkeypatch.setattr(SystemConfigService, "get_config", lambda *_args, **_kwargs: True)
|
||||
|
||||
enabled = await middleware._is_notification_email_module_enabled(request)
|
||||
assert enabled is True
|
||||
assert fake_db.closed is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_error_plugins_ignores_http_4xx(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
middleware = PluginMiddleware(_dummy_app)
|
||||
send_error = AsyncMock()
|
||||
middleware.plugin_manager = cast(
|
||||
Any,
|
||||
SimpleNamespace(
|
||||
get_plugin=lambda *_args, **_kwargs: SimpleNamespace(
|
||||
enabled=True, send_error=send_error
|
||||
)
|
||||
),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
middleware,
|
||||
"_is_notification_email_module_enabled",
|
||||
AsyncMock(return_value=True),
|
||||
)
|
||||
|
||||
await middleware._call_error_plugins(
|
||||
_build_request(),
|
||||
HTTPException(status_code=400, detail="bad request"),
|
||||
start_time=time.time(),
|
||||
)
|
||||
|
||||
send_error.assert_not_awaited()
|
||||
20
tests/unit/test_stream_usage_limits.py
Normal file
20
tests/unit/test_stream_usage_limits.py
Normal file
@@ -0,0 +1,20 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from src.services.usage.stream import _get_response_chunks_max_size
|
||||
|
||||
|
||||
def test_response_chunks_max_size_default(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.delenv("RESPONSE_CHUNKS_MAX_SIZE_MB", raising=False)
|
||||
assert _get_response_chunks_max_size() == 2 * 1024 * 1024
|
||||
|
||||
|
||||
def test_response_chunks_max_size_respects_env(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setenv("RESPONSE_CHUNKS_MAX_SIZE_MB", "2")
|
||||
assert _get_response_chunks_max_size() == 2 * 1024 * 1024
|
||||
|
||||
|
||||
def test_response_chunks_max_size_invalid_env_fallback(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setenv("RESPONSE_CHUNKS_MAX_SIZE_MB", "bad")
|
||||
assert _get_response_chunks_max_size() == 2 * 1024 * 1024
|
||||
Reference in New Issue
Block a user