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:
fawney19
2026-03-14 11:59:07 +08:00
parent 45985f1c04
commit e0286aebe3
111 changed files with 2775 additions and 1102 deletions

View File

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

View 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

View File

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

View File

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

View File

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

View File

@@ -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",

View File

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

View File

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

View File

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

View 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

View 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

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

View 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