mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
fix(redis): 按事件循环隔离 Redis 连接,防止子线程 asyncio.run 导致连接泄漏
RedisClientManager 新增 _redis_by_loop 字典,按 event loop id 维护独立连接, 避免 usage consumer 在 asyncio.to_thread + asyncio.run 场景下复用主循环连接。 同步重构 consumer_streams 的写库路径:_apply_record_event 统一走 record_usage_batch,_apply_streaming_event 改为 to_thread 执行同步 DB 操作, 移除冗余的 session 传递和手动 rollback 逻辑。
This commit is contained in:
@@ -1,3 +1,7 @@
|
||||
import asyncio
|
||||
|
||||
import pytest
|
||||
|
||||
from src.clients import redis_client as redis_client_module
|
||||
from src.clients.redis_client import RedisClientManager, RedisState
|
||||
|
||||
@@ -39,3 +43,49 @@ def test_reset_redis_circuit_breaker_returns_false_when_uninitialized() -> None:
|
||||
finally:
|
||||
redis_client_module._redis_manager = old_global
|
||||
redis_client_module._usage_queue_redis_manager = old_usage
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_redis_client_isolated_per_event_loop(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
old_global = redis_client_module._redis_manager
|
||||
try:
|
||||
redis_client_module._redis_manager = RedisClientManager(client_name="global")
|
||||
|
||||
class DummyRedis:
|
||||
def __init__(self, name: str) -> None:
|
||||
self.name = name
|
||||
self.closed = False
|
||||
|
||||
async def ping(self) -> bool:
|
||||
return True
|
||||
|
||||
async def close(self) -> None:
|
||||
self.closed = True
|
||||
|
||||
created_clients: list[DummyRedis] = []
|
||||
|
||||
async def fake_from_url(*args: object, **kwargs: object) -> DummyRedis:
|
||||
client = DummyRedis(f"client-{len(created_clients)}")
|
||||
created_clients.append(client)
|
||||
return client
|
||||
|
||||
monkeypatch.setattr(redis_client_module.aioredis, "from_url", fake_from_url)
|
||||
|
||||
main_client = await redis_client_module.get_redis_client()
|
||||
|
||||
def _get_client_from_thread() -> object:
|
||||
return asyncio.run(redis_client_module.get_redis_client())
|
||||
|
||||
thread_client = await asyncio.to_thread(_get_client_from_thread)
|
||||
|
||||
assert main_client is created_clients[0]
|
||||
assert thread_client is created_clients[1]
|
||||
assert thread_client is not main_client
|
||||
assert redis_client_module.get_redis_client_sync() is main_client
|
||||
|
||||
await redis_client_module.close_redis_client()
|
||||
|
||||
assert created_clients[0].closed is True
|
||||
assert created_clients[1].closed is True
|
||||
finally:
|
||||
redis_client_module._redis_manager = old_global
|
||||
|
||||
@@ -985,17 +985,16 @@ async def test_consumer_apply_event_streaming(monkeypatch: Any) -> None:
|
||||
async def test_consumer_apply_event_completed(monkeypatch: Any) -> None:
|
||||
"""测试处理 COMPLETED 事件"""
|
||||
mock_db = MagicMock()
|
||||
mock_db.query.return_value.filter.return_value.first.return_value = None
|
||||
|
||||
def mock_create_session() -> Any:
|
||||
return mock_db
|
||||
|
||||
monkeypatch.setattr("src.services.usage.consumer_streams.create_session", mock_create_session)
|
||||
|
||||
mock_record_usage = AsyncMock()
|
||||
mock_record_batch = AsyncMock(return_value=[])
|
||||
monkeypatch.setattr(
|
||||
"src.services.usage.consumer_streams.UsageService.record_usage",
|
||||
mock_record_usage,
|
||||
"src.services.usage.consumer_streams.UsageService.record_usage_batch",
|
||||
mock_record_batch,
|
||||
)
|
||||
|
||||
consumer = UsageQueueConsumer()
|
||||
@@ -1014,29 +1013,29 @@ async def test_consumer_apply_event_completed(monkeypatch: Any) -> None:
|
||||
|
||||
await consumer._apply_event(event)
|
||||
|
||||
mock_record_usage.assert_called_once()
|
||||
call_kwargs = mock_record_usage.call_args.kwargs
|
||||
assert call_kwargs["request_id"] == "req-done"
|
||||
assert call_kwargs["status"] == "completed"
|
||||
assert call_kwargs["provider"] == "openai"
|
||||
assert call_kwargs["input_tokens"] == 100
|
||||
mock_record_batch.assert_awaited_once()
|
||||
records = mock_record_batch.call_args[0][1]
|
||||
assert len(records) == 1
|
||||
assert records[0]["request_id"] == "req-done"
|
||||
assert records[0]["status"] == "completed"
|
||||
assert records[0]["provider"] == "openai"
|
||||
assert records[0]["input_tokens"] == 100
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_consumer_apply_event_failed(monkeypatch: Any) -> None:
|
||||
"""测试处理 FAILED 事件"""
|
||||
mock_db = MagicMock()
|
||||
mock_db.query.return_value.filter.return_value.first.return_value = None
|
||||
|
||||
def mock_create_session() -> Any:
|
||||
return mock_db
|
||||
|
||||
monkeypatch.setattr("src.services.usage.consumer_streams.create_session", mock_create_session)
|
||||
|
||||
mock_record_usage = AsyncMock()
|
||||
mock_record_batch = AsyncMock(return_value=[])
|
||||
monkeypatch.setattr(
|
||||
"src.services.usage.consumer_streams.UsageService.record_usage",
|
||||
mock_record_usage,
|
||||
"src.services.usage.consumer_streams.UsageService.record_usage_batch",
|
||||
mock_record_batch,
|
||||
)
|
||||
|
||||
consumer = UsageQueueConsumer()
|
||||
@@ -1054,26 +1053,25 @@ async def test_consumer_apply_event_failed(monkeypatch: Any) -> None:
|
||||
|
||||
await consumer._apply_event(event)
|
||||
|
||||
call_kwargs = mock_record_usage.call_args.kwargs
|
||||
assert call_kwargs["status"] == "failed"
|
||||
assert call_kwargs["error_message"] == "Rate limited"
|
||||
records = mock_record_batch.call_args[0][1]
|
||||
assert records[0]["status"] == "failed"
|
||||
assert records[0]["error_message"] == "Rate limited"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_consumer_apply_event_cancelled(monkeypatch: Any) -> None:
|
||||
"""测试处理 CANCELLED 事件"""
|
||||
mock_db = MagicMock()
|
||||
mock_db.query.return_value.filter.return_value.first.return_value = None
|
||||
|
||||
def mock_create_session() -> Any:
|
||||
return mock_db
|
||||
|
||||
monkeypatch.setattr("src.services.usage.consumer_streams.create_session", mock_create_session)
|
||||
|
||||
mock_record_usage = AsyncMock()
|
||||
mock_record_batch = AsyncMock(return_value=[])
|
||||
monkeypatch.setattr(
|
||||
"src.services.usage.consumer_streams.UsageService.record_usage",
|
||||
mock_record_usage,
|
||||
"src.services.usage.consumer_streams.UsageService.record_usage_batch",
|
||||
mock_record_batch,
|
||||
)
|
||||
|
||||
consumer = UsageQueueConsumer()
|
||||
@@ -1087,8 +1085,8 @@ async def test_consumer_apply_event_cancelled(monkeypatch: Any) -> None:
|
||||
|
||||
await consumer._apply_event(event)
|
||||
|
||||
call_kwargs = mock_record_usage.call_args.kwargs
|
||||
assert call_kwargs["status"] == "cancelled"
|
||||
records = mock_record_batch.call_args[0][1]
|
||||
assert records[0]["status"] == "cancelled"
|
||||
|
||||
|
||||
# ============ 批量处理测试 ============
|
||||
@@ -1214,27 +1212,19 @@ async def test_consumer_process_record_batch_fallback(monkeypatch: Any) -> None:
|
||||
"""测试批量处理失败时回退到逐条处理"""
|
||||
mock_redis = MockRedisForConsumer()
|
||||
mock_db = MagicMock()
|
||||
mock_db.query.return_value.filter.return_value.first.return_value = None
|
||||
|
||||
def mock_create_session() -> Any:
|
||||
return mock_db
|
||||
|
||||
monkeypatch.setattr("src.services.usage.consumer_streams.create_session", mock_create_session)
|
||||
|
||||
# 批量处理失败
|
||||
mock_record_batch = AsyncMock(side_effect=Exception("Batch failed"))
|
||||
# 首次批量处理失败,回退到单条 batch 成功
|
||||
mock_record_batch = AsyncMock(side_effect=[Exception("Batch failed"), []])
|
||||
monkeypatch.setattr(
|
||||
"src.services.usage.consumer_streams.UsageService.record_usage_batch",
|
||||
mock_record_batch,
|
||||
)
|
||||
|
||||
# 单条处理成功
|
||||
mock_record_usage = AsyncMock()
|
||||
monkeypatch.setattr(
|
||||
"src.services.usage.consumer_streams.UsageService.record_usage",
|
||||
mock_record_usage,
|
||||
)
|
||||
|
||||
consumer = UsageQueueConsumer()
|
||||
|
||||
event = build_usage_event(
|
||||
@@ -1248,8 +1238,11 @@ async def test_consumer_process_record_batch_fallback(monkeypatch: Any) -> None:
|
||||
[("msg-1", event.to_stream_fields(), event)],
|
||||
)
|
||||
|
||||
# 验证回退到单条处理
|
||||
mock_record_usage.assert_called_once()
|
||||
# 验证回退到单条 batch 处理
|
||||
assert mock_record_batch.await_count == 2
|
||||
fallback_records = mock_record_batch.call_args_list[1][0][1]
|
||||
assert len(fallback_records) == 1
|
||||
assert fallback_records[0]["request_id"] == "req-fallback"
|
||||
# 消息被 ACK
|
||||
assert len(mock_redis.xack_calls) == 1
|
||||
|
||||
@@ -1296,17 +1289,16 @@ async def test_consumer_apply_streaming_event(monkeypatch: Any) -> None:
|
||||
async def test_consumer_apply_record_event(monkeypatch: Any) -> None:
|
||||
"""测试记录事件单独处理"""
|
||||
mock_db = MagicMock()
|
||||
mock_db.query.return_value.filter.return_value.first.return_value = None
|
||||
|
||||
def mock_create_session() -> Any:
|
||||
return mock_db
|
||||
|
||||
monkeypatch.setattr("src.services.usage.consumer_streams.create_session", mock_create_session)
|
||||
|
||||
mock_record_usage = AsyncMock()
|
||||
mock_record_batch = AsyncMock(return_value=[])
|
||||
monkeypatch.setattr(
|
||||
"src.services.usage.consumer_streams.UsageService.record_usage",
|
||||
mock_record_usage,
|
||||
"src.services.usage.consumer_streams.UsageService.record_usage_batch",
|
||||
mock_record_batch,
|
||||
)
|
||||
|
||||
consumer = UsageQueueConsumer()
|
||||
@@ -1320,10 +1312,11 @@ async def test_consumer_apply_record_event(monkeypatch: Any) -> None:
|
||||
|
||||
await consumer._apply_record_event(event)
|
||||
|
||||
mock_record_usage.assert_called_once()
|
||||
call_kwargs = mock_record_usage.call_args.kwargs
|
||||
assert call_kwargs["request_id"] == "req-record"
|
||||
assert call_kwargs["input_tokens"] == 100
|
||||
mock_record_batch.assert_awaited_once()
|
||||
records = mock_record_batch.call_args[0][1]
|
||||
assert len(records) == 1
|
||||
assert records[0]["request_id"] == "req-record"
|
||||
assert records[0]["input_tokens"] == 100
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
||||
Reference in New Issue
Block a user