mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 09:20:22 +08:00
refactor: Antigravity/Codex 服务重构为插件化适配器架构
- 将 Antigravity 和 Codex 从独立模块迁移至 src/services/provider/adapters/ 插件体系 - 新增 provider_types 和 oauth_token 模块,移除 maintenance_scheduler 中的 OAuth 定时刷新 - 增强 admin API:扩展 keys 和 provider_query 端点,新增 dashboard 路由 - 大幅增强 ProviderDetailDrawer 组件,新增 AntigravityQuotaDialog - 改进 handler 基类(chat/cli)和错误分类器 - 优化 fetch_scheduler 和 upstream_fetcher - 前端 UI 组件清理和优化 - 更新测试以匹配新模块结构
This commit is contained in:
@@ -1,5 +1,6 @@
|
||||
import asyncio
|
||||
import json
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
@@ -25,10 +26,16 @@ from src.services.usage.telemetry_writer import (
|
||||
|
||||
class DummyRedis:
|
||||
def __init__(self) -> None:
|
||||
self.calls = []
|
||||
self.xadd_error = None
|
||||
self.calls: list[tuple[str, dict[str, str], int | None, bool | None]] = []
|
||||
self.xadd_error: Exception | None = None
|
||||
|
||||
async def xadd(self, key, fields, maxlen=None, approximate=None):
|
||||
async def xadd(
|
||||
self,
|
||||
key: str,
|
||||
fields: dict[str, str],
|
||||
maxlen: int | None = None,
|
||||
approximate: bool | None = None,
|
||||
) -> str:
|
||||
if self.xadd_error:
|
||||
raise self.xadd_error
|
||||
self.calls.append((key, fields, maxlen, approximate))
|
||||
@@ -36,7 +43,7 @@ class DummyRedis:
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_usage_event_roundtrip():
|
||||
async def test_usage_event_roundtrip() -> None:
|
||||
event = build_usage_event(
|
||||
event_type=UsageEventType.COMPLETED,
|
||||
request_id="req-1",
|
||||
@@ -52,10 +59,10 @@ async def test_usage_event_roundtrip():
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_queue_writer_publishes_event(monkeypatch):
|
||||
async def test_queue_writer_publishes_event(monkeypatch: Any) -> None:
|
||||
dummy = DummyRedis()
|
||||
|
||||
async def _get_redis_client(require_redis=False):
|
||||
async def _get_redis_client(require_redis: bool = False) -> Any:
|
||||
return dummy
|
||||
|
||||
monkeypatch.setattr("src.services.usage.telemetry_writer.get_redis_client", _get_redis_client)
|
||||
@@ -96,7 +103,7 @@ async def test_queue_writer_publishes_event(monkeypatch):
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_usage_event_all_types():
|
||||
async def test_usage_event_all_types() -> None:
|
||||
"""测试所有事件类型的序列化/反序列化"""
|
||||
for event_type in UsageEventType:
|
||||
event = build_usage_event(
|
||||
@@ -111,7 +118,7 @@ async def test_usage_event_all_types():
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_usage_event_bytes_payload():
|
||||
async def test_usage_event_bytes_payload() -> None:
|
||||
"""测试 bytes 类型 payload 的反序列化"""
|
||||
event = build_usage_event(
|
||||
event_type=UsageEventType.COMPLETED,
|
||||
@@ -120,19 +127,19 @@ async def test_usage_event_bytes_payload():
|
||||
)
|
||||
fields = event.to_stream_fields()
|
||||
# 模拟 Redis 返回 bytes
|
||||
fields["payload"] = fields["payload"].encode("utf-8")
|
||||
fields["payload"] = fields["payload"].encode("utf-8") # type: ignore[assignment]
|
||||
restored = UsageEvent.from_stream_fields(fields)
|
||||
assert restored.request_id == "req-bytes"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_usage_event_missing_payload():
|
||||
async def test_usage_event_missing_payload() -> None:
|
||||
"""测试缺少 payload 字段时抛出异常"""
|
||||
with pytest.raises(ValueError, match="Missing payload"):
|
||||
UsageEvent.from_stream_fields({})
|
||||
|
||||
|
||||
def test_sanitize_payload_nested():
|
||||
def test_sanitize_payload_nested() -> None:
|
||||
"""测试 sanitize_payload 处理嵌套结构"""
|
||||
data = {
|
||||
"str": "hello",
|
||||
@@ -151,7 +158,7 @@ def test_sanitize_payload_nested():
|
||||
assert isinstance(result["custom"], str)
|
||||
|
||||
|
||||
def test_parse_body_json_string():
|
||||
def test_parse_body_json_string() -> None:
|
||||
"""测试 _parse_body 正确反序列化 JSON 字符串"""
|
||||
from src.services.usage.consumer_streams import _parse_body
|
||||
|
||||
@@ -162,7 +169,7 @@ def test_parse_body_json_string():
|
||||
assert result["messages"][0]["content"] == "hello"
|
||||
|
||||
|
||||
def test_parse_body_dict_passthrough():
|
||||
def test_parse_body_dict_passthrough() -> None:
|
||||
"""测试 _parse_body 直接返回 dict"""
|
||||
from src.services.usage.consumer_streams import _parse_body
|
||||
|
||||
@@ -171,14 +178,14 @@ def test_parse_body_dict_passthrough():
|
||||
assert result is data # 应该是同一个对象
|
||||
|
||||
|
||||
def test_parse_body_none():
|
||||
def test_parse_body_none() -> None:
|
||||
"""测试 _parse_body 处理 None"""
|
||||
from src.services.usage.consumer_streams import _parse_body
|
||||
|
||||
assert _parse_body(None) is None
|
||||
|
||||
|
||||
def test_parse_body_truncated_string():
|
||||
def test_parse_body_truncated_string() -> None:
|
||||
"""测试 _parse_body 处理被截断的 JSON 字符串"""
|
||||
from src.services.usage.consumer_streams import _parse_body
|
||||
|
||||
@@ -192,7 +199,7 @@ def test_parse_body_truncated_string():
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_db_telemetry_writer_filters_kwargs():
|
||||
async def test_db_telemetry_writer_filters_kwargs() -> None:
|
||||
"""测试 DbTelemetryWriter 过滤不支持的参数"""
|
||||
mock_telemetry = MagicMock()
|
||||
mock_telemetry.record_success = AsyncMock()
|
||||
@@ -219,7 +226,7 @@ async def test_db_telemetry_writer_filters_kwargs():
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_db_telemetry_writer_all_methods():
|
||||
async def test_db_telemetry_writer_all_methods() -> None:
|
||||
"""测试 DbTelemetryWriter 的所有方法"""
|
||||
mock_telemetry = MagicMock()
|
||||
mock_telemetry.record_success = AsyncMock()
|
||||
@@ -238,11 +245,11 @@ async def test_db_telemetry_writer_all_methods():
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_queue_writer_record_failure(monkeypatch):
|
||||
async def test_queue_writer_record_failure(monkeypatch: Any) -> None:
|
||||
"""测试 QueueTelemetryWriter.record_failure"""
|
||||
dummy = DummyRedis()
|
||||
|
||||
async def _get_redis_client(require_redis=False):
|
||||
async def _get_redis_client(require_redis: bool = False) -> Any:
|
||||
return dummy
|
||||
|
||||
monkeypatch.setattr("src.services.usage.telemetry_writer.get_redis_client", _get_redis_client)
|
||||
@@ -275,11 +282,11 @@ async def test_queue_writer_record_failure(monkeypatch):
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_queue_writer_record_cancelled(monkeypatch):
|
||||
async def test_queue_writer_record_cancelled(monkeypatch: Any) -> None:
|
||||
"""测试 QueueTelemetryWriter.record_cancelled"""
|
||||
dummy = DummyRedis()
|
||||
|
||||
async def _get_redis_client(require_redis=False):
|
||||
async def _get_redis_client(require_redis: bool = False) -> Any:
|
||||
return dummy
|
||||
|
||||
monkeypatch.setattr("src.services.usage.telemetry_writer.get_redis_client", _get_redis_client)
|
||||
@@ -296,11 +303,11 @@ async def test_queue_writer_record_cancelled(monkeypatch):
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_queue_writer_include_headers_bodies(monkeypatch):
|
||||
async def test_queue_writer_include_headers_bodies(monkeypatch: Any) -> None:
|
||||
"""测试 include_headers 和 include_bodies 配置"""
|
||||
dummy = DummyRedis()
|
||||
|
||||
async def _get_redis_client(require_redis=False):
|
||||
async def _get_redis_client(require_redis: bool = False) -> Any:
|
||||
return dummy
|
||||
|
||||
monkeypatch.setattr("src.services.usage.telemetry_writer.get_redis_client", _get_redis_client)
|
||||
@@ -332,7 +339,7 @@ async def test_queue_writer_include_headers_bodies(monkeypatch):
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_event_to_record_body_deserialization(monkeypatch):
|
||||
async def test_event_to_record_body_deserialization(monkeypatch: Any) -> None:
|
||||
"""测试 _event_to_record 正确反序列化 body 字符串为 dict"""
|
||||
from src.services.usage.consumer_streams import _event_to_record
|
||||
|
||||
@@ -361,11 +368,11 @@ async def test_event_to_record_body_deserialization(monkeypatch):
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_queue_writer_body_truncation(monkeypatch):
|
||||
async def test_queue_writer_body_truncation(monkeypatch: Any) -> None:
|
||||
"""测试 body 超长截断"""
|
||||
dummy = DummyRedis()
|
||||
|
||||
async def _get_redis_client(require_redis=False):
|
||||
async def _get_redis_client(require_redis: bool = False) -> Any:
|
||||
return dummy
|
||||
|
||||
monkeypatch.setattr("src.services.usage.telemetry_writer.get_redis_client", _get_redis_client)
|
||||
@@ -393,10 +400,10 @@ async def test_queue_writer_body_truncation(monkeypatch):
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_queue_writer_redis_unavailable(monkeypatch):
|
||||
async def test_queue_writer_redis_unavailable(monkeypatch: Any) -> None:
|
||||
"""测试 Redis 不可用时抛出异常"""
|
||||
|
||||
async def _get_redis_client(require_redis=False):
|
||||
async def _get_redis_client(require_redis: bool = False) -> Any:
|
||||
return None
|
||||
|
||||
monkeypatch.setattr("src.services.usage.telemetry_writer.get_redis_client", _get_redis_client)
|
||||
@@ -411,12 +418,12 @@ async def test_queue_writer_redis_unavailable(monkeypatch):
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_queue_writer_xadd_error(monkeypatch):
|
||||
async def test_queue_writer_xadd_error(monkeypatch: Any) -> None:
|
||||
"""测试 XADD 失败时抛出异常"""
|
||||
dummy = DummyRedis()
|
||||
dummy.xadd_error = Exception("Redis connection lost")
|
||||
|
||||
async def _get_redis_client(require_redis=False):
|
||||
async def _get_redis_client(require_redis: bool = False) -> Any:
|
||||
return dummy
|
||||
|
||||
monkeypatch.setattr("src.services.usage.telemetry_writer.get_redis_client", _get_redis_client)
|
||||
@@ -438,13 +445,13 @@ class MockRedisPipeline:
|
||||
|
||||
def __init__(self, parent: "MockRedisForConsumer") -> None:
|
||||
self._parent = parent
|
||||
self._commands = []
|
||||
self._commands: list[tuple[str, ...]] = []
|
||||
|
||||
def xack(self, key, group, message_id):
|
||||
def xack(self, key: str, group: str, message_id: str) -> Any:
|
||||
self._commands.append(("xack", key, group, message_id))
|
||||
return self
|
||||
|
||||
async def execute(self):
|
||||
async def execute(self) -> Any:
|
||||
results = []
|
||||
for cmd in self._commands:
|
||||
if cmd[0] == "xack":
|
||||
@@ -458,54 +465,69 @@ class MockRedisForConsumer:
|
||||
"""模拟 Redis 客户端用于消费者测试"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.xgroup_create_calls = []
|
||||
self.xreadgroup_results = []
|
||||
self.xautoclaim_results = []
|
||||
self.xack_calls = []
|
||||
self.xadd_calls = []
|
||||
self.xpending_range_results = []
|
||||
self.xlen_result = 0
|
||||
self.xpending_result = {"pending": 0}
|
||||
self.xgroup_create_error = None
|
||||
self.xgroup_create_calls: list[tuple[str, str, str, bool]] = []
|
||||
self.xreadgroup_results: list[Any] = []
|
||||
self.xautoclaim_results: list[Any] = []
|
||||
self.xack_calls: list[tuple[str, str, str]] = []
|
||||
self.xadd_calls: list[tuple[str, dict[str, str], int | None, bool | None]] = []
|
||||
self.xpending_range_results: list[Any] = []
|
||||
self.xlen_result: int = 0
|
||||
self.xpending_result: dict[str, int] | tuple[int, str, str, list[Any]] = {"pending": 0}
|
||||
self.xgroup_create_error: Exception | None = None
|
||||
|
||||
async def xgroup_create(self, key, group, id, mkstream=False):
|
||||
async def xgroup_create(self, key: str, group: str, id: str, mkstream: bool = False) -> Any:
|
||||
self.xgroup_create_calls.append((key, group, id, mkstream))
|
||||
if self.xgroup_create_error:
|
||||
raise self.xgroup_create_error
|
||||
|
||||
async def xreadgroup(self, groupname, consumername, streams, count, block):
|
||||
async def xreadgroup(
|
||||
self,
|
||||
groupname: str,
|
||||
consumername: str,
|
||||
streams: dict[str, str],
|
||||
count: int,
|
||||
block: int,
|
||||
) -> Any:
|
||||
if self.xreadgroup_results:
|
||||
return self.xreadgroup_results.pop(0)
|
||||
return None
|
||||
|
||||
async def xautoclaim(self, key, group, consumer, min_idle_time, start_id, count):
|
||||
async def xautoclaim(
|
||||
self, key: str, group: str, consumer: str, min_idle_time: int, start_id: str, count: int
|
||||
) -> Any:
|
||||
if self.xautoclaim_results:
|
||||
return self.xautoclaim_results.pop(0)
|
||||
return None
|
||||
|
||||
async def xack(self, key, group, message_id):
|
||||
async def xack(self, key: str, group: str, message_id: str) -> Any:
|
||||
self.xack_calls.append((key, group, message_id))
|
||||
|
||||
async def xadd(self, key, fields, maxlen=None, approximate=None):
|
||||
async def xadd(
|
||||
self,
|
||||
key: str,
|
||||
fields: dict[str, str],
|
||||
maxlen: int | None = None,
|
||||
approximate: bool | None = None,
|
||||
) -> Any:
|
||||
self.xadd_calls.append((key, fields, maxlen, approximate))
|
||||
return "dlq-1-0"
|
||||
|
||||
async def xpending_range(self, key, group, min, max, count):
|
||||
async def xpending_range(self, key: str, group: str, min: str, max: str, count: int) -> Any:
|
||||
if self.xpending_range_results:
|
||||
return self.xpending_range_results.pop(0)
|
||||
return []
|
||||
|
||||
async def xlen(self, key):
|
||||
async def xlen(self, key: str) -> Any:
|
||||
return self.xlen_result
|
||||
|
||||
async def xpending(self, key, group):
|
||||
async def xpending(self, key: str, group: str) -> Any:
|
||||
return self.xpending_result
|
||||
|
||||
def pipeline(self):
|
||||
def pipeline(self) -> Any:
|
||||
return MockRedisPipeline(self)
|
||||
|
||||
|
||||
def test_consumer_name():
|
||||
def test_consumer_name() -> None:
|
||||
"""测试消费者名称生成"""
|
||||
name = _consumer_name()
|
||||
assert ":" in name
|
||||
@@ -516,11 +538,11 @@ def test_consumer_name():
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ensure_stream_group_creates_group(monkeypatch):
|
||||
async def test_ensure_stream_group_creates_group(monkeypatch: Any) -> None:
|
||||
"""测试创建消费者组"""
|
||||
mock_redis = MockRedisForConsumer()
|
||||
|
||||
async def _get_redis_client(require_redis=False):
|
||||
async def _get_redis_client(require_redis: bool = False) -> Any:
|
||||
return mock_redis
|
||||
|
||||
monkeypatch.setattr("src.services.usage.consumer_streams.get_redis_client", _get_redis_client)
|
||||
@@ -535,12 +557,12 @@ async def test_ensure_stream_group_creates_group(monkeypatch):
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ensure_stream_group_handles_busygroup(monkeypatch):
|
||||
async def test_ensure_stream_group_handles_busygroup(monkeypatch: Any) -> None:
|
||||
"""测试消费者组已存在时不抛异常"""
|
||||
mock_redis = MockRedisForConsumer()
|
||||
mock_redis.xgroup_create_error = ResponseError("BUSYGROUP Consumer Group name already exists")
|
||||
|
||||
async def _get_redis_client(require_redis=False):
|
||||
async def _get_redis_client(require_redis: bool = False) -> Any:
|
||||
return mock_redis
|
||||
|
||||
monkeypatch.setattr("src.services.usage.consumer_streams.get_redis_client", _get_redis_client)
|
||||
@@ -550,12 +572,12 @@ async def test_ensure_stream_group_handles_busygroup(monkeypatch):
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ensure_stream_group_raises_other_errors(monkeypatch):
|
||||
async def test_ensure_stream_group_raises_other_errors(monkeypatch: Any) -> None:
|
||||
"""测试其他 Redis 错误时抛出异常"""
|
||||
mock_redis = MockRedisForConsumer()
|
||||
mock_redis.xgroup_create_error = ResponseError("SOME OTHER ERROR")
|
||||
|
||||
async def _get_redis_client(require_redis=False):
|
||||
async def _get_redis_client(require_redis: bool = False) -> Any:
|
||||
return mock_redis
|
||||
|
||||
monkeypatch.setattr("src.services.usage.consumer_streams.get_redis_client", _get_redis_client)
|
||||
@@ -565,10 +587,10 @@ async def test_ensure_stream_group_raises_other_errors(monkeypatch):
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ensure_stream_group_no_redis(monkeypatch):
|
||||
async def test_ensure_stream_group_no_redis(monkeypatch: Any) -> None:
|
||||
"""测试 Redis 不可用时直接返回"""
|
||||
|
||||
async def _get_redis_client(require_redis=False):
|
||||
async def _get_redis_client(require_redis: bool = False) -> Any:
|
||||
return None
|
||||
|
||||
monkeypatch.setattr("src.services.usage.consumer_streams.get_redis_client", _get_redis_client)
|
||||
@@ -578,13 +600,13 @@ async def test_ensure_stream_group_no_redis(monkeypatch):
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_consumer_start_stop():
|
||||
async def test_consumer_start_stop() -> None:
|
||||
"""测试消费者启动和停止"""
|
||||
consumer = UsageQueueConsumer()
|
||||
assert not consumer._running
|
||||
|
||||
# Mock _run 避免真正执行
|
||||
consumer._run = AsyncMock()
|
||||
consumer._run = AsyncMock() # type: ignore[method-assign]
|
||||
|
||||
await consumer.start()
|
||||
assert consumer._running
|
||||
@@ -603,7 +625,7 @@ async def test_consumer_start_stop():
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_consumer_process_messages_success(monkeypatch):
|
||||
async def test_consumer_process_messages_success(monkeypatch: Any) -> None:
|
||||
"""测试成功处理消息"""
|
||||
mock_redis = MockRedisForConsumer()
|
||||
|
||||
@@ -626,7 +648,7 @@ async def test_consumer_process_messages_success(monkeypatch):
|
||||
consumer = UsageQueueConsumer()
|
||||
|
||||
# Mock 批量处理方法
|
||||
consumer._process_record_batch = AsyncMock()
|
||||
consumer._process_record_batch = AsyncMock() # type: ignore[method-assign]
|
||||
|
||||
await consumer._process_messages(mock_redis, messages)
|
||||
|
||||
@@ -638,7 +660,7 @@ async def test_consumer_process_messages_success(monkeypatch):
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_consumer_process_messages_error_retry(monkeypatch):
|
||||
async def test_consumer_process_messages_error_retry(monkeypatch: Any) -> None:
|
||||
"""测试处理消息失败时 STREAMING 事件的重试行为"""
|
||||
mock_redis = MockRedisForConsumer()
|
||||
# 返回重试次数小于 max_retries
|
||||
@@ -660,7 +682,7 @@ async def test_consumer_process_messages_error_retry(monkeypatch):
|
||||
# 在设置 config 后创建 consumer,以便缓存正确的配置值
|
||||
consumer = UsageQueueConsumer()
|
||||
# 模拟 STREAMING 事件处理失败
|
||||
consumer._apply_streaming_event = AsyncMock(side_effect=ValueError("Processing error"))
|
||||
consumer._apply_streaming_event = AsyncMock(side_effect=ValueError("Processing error")) # type: ignore[method-assign]
|
||||
|
||||
await consumer._process_messages(mock_redis, messages)
|
||||
|
||||
@@ -673,7 +695,7 @@ async def test_consumer_process_messages_error_retry(monkeypatch):
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_consumer_process_messages_move_to_dlq(monkeypatch):
|
||||
async def test_consumer_process_messages_move_to_dlq(monkeypatch: Any) -> None:
|
||||
"""测试消息超过最大重试次数后移入 DLQ"""
|
||||
mock_redis = MockRedisForConsumer()
|
||||
# 返回重试次数 >= max_retries
|
||||
@@ -712,7 +734,7 @@ async def test_consumer_process_messages_move_to_dlq(monkeypatch):
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_consumer_get_delivery_count_dict_format():
|
||||
async def test_consumer_get_delivery_count_dict_format() -> None:
|
||||
"""测试获取消息投递次数(dict 格式)"""
|
||||
mock_redis = MockRedisForConsumer()
|
||||
mock_redis.xpending_range_results = [[{"times_delivered": 3}]]
|
||||
@@ -723,7 +745,7 @@ async def test_consumer_get_delivery_count_dict_format():
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_consumer_get_delivery_count_tuple_format():
|
||||
async def test_consumer_get_delivery_count_tuple_format() -> None:
|
||||
"""测试获取消息投递次数(tuple 格式)"""
|
||||
mock_redis = MockRedisForConsumer()
|
||||
# (message_id, consumer, idle_time, times_delivered)
|
||||
@@ -735,7 +757,7 @@ async def test_consumer_get_delivery_count_tuple_format():
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_consumer_get_delivery_count_empty():
|
||||
async def test_consumer_get_delivery_count_empty() -> None:
|
||||
"""测试消息不存在时返回 0"""
|
||||
mock_redis = MockRedisForConsumer()
|
||||
mock_redis.xpending_range_results = [[]]
|
||||
@@ -746,7 +768,7 @@ async def test_consumer_get_delivery_count_empty():
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_consumer_read_new_messages(monkeypatch):
|
||||
async def test_consumer_read_new_messages(monkeypatch: Any) -> None:
|
||||
"""测试读取新消息"""
|
||||
mock_redis = MockRedisForConsumer()
|
||||
|
||||
@@ -761,7 +783,7 @@ async def test_consumer_read_new_messages(monkeypatch):
|
||||
]
|
||||
|
||||
consumer = UsageQueueConsumer()
|
||||
consumer._process_messages = AsyncMock()
|
||||
consumer._process_messages = AsyncMock() # type: ignore[method-assign]
|
||||
|
||||
await consumer._read_new(mock_redis)
|
||||
|
||||
@@ -772,14 +794,14 @@ async def test_consumer_read_new_messages(monkeypatch):
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_consumer_maybe_claim_pending_respects_interval(monkeypatch):
|
||||
async def test_consumer_maybe_claim_pending_respects_interval(monkeypatch: Any) -> None:
|
||||
"""测试 claim 间隔限制"""
|
||||
mock_redis = MockRedisForConsumer()
|
||||
|
||||
consumer = UsageQueueConsumer()
|
||||
consumer._last_claim = 999999999999.0 # 未来时间
|
||||
|
||||
consumer._process_messages = AsyncMock()
|
||||
consumer._process_messages = AsyncMock() # type: ignore[method-assign]
|
||||
|
||||
await consumer._maybe_claim_pending(mock_redis)
|
||||
|
||||
@@ -788,14 +810,14 @@ async def test_consumer_maybe_claim_pending_respects_interval(monkeypatch):
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_consumer_maybe_claim_pending_xautoclaim_error(monkeypatch):
|
||||
async def test_consumer_maybe_claim_pending_xautoclaim_error(monkeypatch: Any) -> None:
|
||||
"""测试 XAUTOCLAIM 失败时优雅处理"""
|
||||
mock_redis = MockRedisForConsumer()
|
||||
|
||||
async def failing_xautoclaim(*args, **kwargs):
|
||||
async def failing_xautoclaim(*args: Any, **kwargs: Any) -> Any:
|
||||
raise ResponseError("XAUTOCLAIM error")
|
||||
|
||||
mock_redis.xautoclaim = failing_xautoclaim
|
||||
mock_redis.xautoclaim = failing_xautoclaim # type: ignore[method-assign]
|
||||
|
||||
consumer = UsageQueueConsumer()
|
||||
consumer._last_claim = 0 # 确保会尝试 claim
|
||||
@@ -805,7 +827,7 @@ async def test_consumer_maybe_claim_pending_xautoclaim_error(monkeypatch):
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_consumer_log_metrics_dict_pending():
|
||||
async def test_consumer_log_metrics_dict_pending() -> None:
|
||||
"""测试 metrics 日志(dict 格式 pending)"""
|
||||
mock_redis = MockRedisForConsumer()
|
||||
mock_redis.xlen_result = 100
|
||||
@@ -825,7 +847,7 @@ async def test_consumer_log_metrics_dict_pending():
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_consumer_log_metrics_tuple_pending():
|
||||
async def test_consumer_log_metrics_tuple_pending() -> None:
|
||||
"""测试 metrics 日志(tuple 格式 pending)"""
|
||||
mock_redis = MockRedisForConsumer()
|
||||
mock_redis.xlen_result = 50
|
||||
@@ -845,11 +867,11 @@ async def test_consumer_log_metrics_tuple_pending():
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_consumer_apply_event_streaming(monkeypatch):
|
||||
async def test_consumer_apply_event_streaming(monkeypatch: Any) -> None:
|
||||
"""测试处理 STREAMING 事件"""
|
||||
mock_db = MagicMock()
|
||||
|
||||
def mock_create_session():
|
||||
def mock_create_session() -> Any:
|
||||
return mock_db
|
||||
|
||||
monkeypatch.setattr("src.services.usage.consumer_streams.create_session", mock_create_session)
|
||||
@@ -884,12 +906,12 @@ async def test_consumer_apply_event_streaming(monkeypatch):
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_consumer_apply_event_completed(monkeypatch):
|
||||
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():
|
||||
def mock_create_session() -> Any:
|
||||
return mock_db
|
||||
|
||||
monkeypatch.setattr("src.services.usage.consumer_streams.create_session", mock_create_session)
|
||||
@@ -925,12 +947,12 @@ async def test_consumer_apply_event_completed(monkeypatch):
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_consumer_apply_event_failed(monkeypatch):
|
||||
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():
|
||||
def mock_create_session() -> Any:
|
||||
return mock_db
|
||||
|
||||
monkeypatch.setattr("src.services.usage.consumer_streams.create_session", mock_create_session)
|
||||
@@ -962,12 +984,12 @@ async def test_consumer_apply_event_failed(monkeypatch):
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_consumer_apply_event_cancelled(monkeypatch):
|
||||
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():
|
||||
def mock_create_session() -> Any:
|
||||
return mock_db
|
||||
|
||||
monkeypatch.setattr("src.services.usage.consumer_streams.create_session", mock_create_session)
|
||||
@@ -997,7 +1019,7 @@ async def test_consumer_apply_event_cancelled(monkeypatch):
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_consumer_process_messages_batch(monkeypatch):
|
||||
async def test_consumer_process_messages_batch(monkeypatch: Any) -> None:
|
||||
"""测试批量处理消息"""
|
||||
mock_redis = MockRedisForConsumer()
|
||||
|
||||
@@ -1023,7 +1045,7 @@ async def test_consumer_process_messages_batch(monkeypatch):
|
||||
consumer = UsageQueueConsumer()
|
||||
|
||||
# Mock 批量处理
|
||||
consumer._process_record_batch = AsyncMock()
|
||||
consumer._process_record_batch = AsyncMock() # type: ignore[method-assign]
|
||||
|
||||
await consumer._process_messages(mock_redis, messages)
|
||||
|
||||
@@ -1032,7 +1054,7 @@ async def test_consumer_process_messages_batch(monkeypatch):
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_consumer_process_messages_separates_streaming(monkeypatch):
|
||||
async def test_consumer_process_messages_separates_streaming(monkeypatch: Any) -> None:
|
||||
"""测试消息分类:STREAMING 和其他事件分开处理"""
|
||||
mock_redis = MockRedisForConsumer()
|
||||
|
||||
@@ -1054,8 +1076,8 @@ async def test_consumer_process_messages_separates_streaming(monkeypatch):
|
||||
]
|
||||
|
||||
consumer = UsageQueueConsumer()
|
||||
consumer._apply_streaming_event = AsyncMock()
|
||||
consumer._process_record_batch = AsyncMock()
|
||||
consumer._apply_streaming_event = AsyncMock() # type: ignore[method-assign]
|
||||
consumer._process_record_batch = AsyncMock() # type: ignore[method-assign]
|
||||
|
||||
await consumer._process_messages(mock_redis, messages)
|
||||
|
||||
@@ -1072,12 +1094,12 @@ async def test_consumer_process_messages_separates_streaming(monkeypatch):
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_consumer_process_record_batch_success(monkeypatch):
|
||||
async def test_consumer_process_record_batch_success(monkeypatch: Any) -> None:
|
||||
"""测试批量记录处理成功"""
|
||||
mock_redis = MockRedisForConsumer()
|
||||
mock_db = MagicMock()
|
||||
|
||||
def mock_create_session():
|
||||
def mock_create_session() -> Any:
|
||||
return mock_db
|
||||
|
||||
monkeypatch.setattr("src.services.usage.consumer_streams.create_session", mock_create_session)
|
||||
@@ -1112,13 +1134,13 @@ async def test_consumer_process_record_batch_success(monkeypatch):
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_consumer_process_record_batch_fallback(monkeypatch):
|
||||
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():
|
||||
def mock_create_session() -> Any:
|
||||
return mock_db
|
||||
|
||||
monkeypatch.setattr("src.services.usage.consumer_streams.create_session", mock_create_session)
|
||||
@@ -1157,11 +1179,11 @@ async def test_consumer_process_record_batch_fallback(monkeypatch):
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_consumer_apply_streaming_event(monkeypatch):
|
||||
async def test_consumer_apply_streaming_event(monkeypatch: Any) -> None:
|
||||
"""测试 STREAMING 事件处理"""
|
||||
mock_db = MagicMock()
|
||||
|
||||
def mock_create_session():
|
||||
def mock_create_session() -> Any:
|
||||
return mock_db
|
||||
|
||||
monkeypatch.setattr("src.services.usage.consumer_streams.create_session", mock_create_session)
|
||||
@@ -1195,12 +1217,12 @@ async def test_consumer_apply_streaming_event(monkeypatch):
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_consumer_apply_record_event(monkeypatch):
|
||||
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():
|
||||
def mock_create_session() -> Any:
|
||||
return mock_db
|
||||
|
||||
monkeypatch.setattr("src.services.usage.consumer_streams.create_session", mock_create_session)
|
||||
@@ -1226,3 +1248,80 @@ async def test_consumer_apply_record_event(monkeypatch):
|
||||
call_kwargs = mock_record_usage.call_args.kwargs
|
||||
assert call_kwargs["request_id"] == "req-record"
|
||||
assert call_kwargs["input_tokens"] == 100
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_record_usage_batch_updates_when_status_completed_billing_pending(
|
||||
monkeypatch: Any,
|
||||
) -> None:
|
||||
"""回归测试:
|
||||
usage-queue 模式下,handler 可能会先把 Usage.status 直接更新为 completed(为减少 UI 延迟),
|
||||
但 billing_status 仍为 pending。此时 completed 事件仍应更新详情字段(如 response_headers/body),
|
||||
并将 billing_status 结算为 settled。
|
||||
"""
|
||||
|
||||
from src.models.database import Usage
|
||||
from src.services.usage.service import UsageService
|
||||
|
||||
class DummyQuery:
|
||||
def __init__(self, all_result: list[Any]) -> None:
|
||||
self._all_result = all_result
|
||||
|
||||
def filter(self, *args: Any, **kwargs: Any) -> "DummyQuery":
|
||||
return self
|
||||
|
||||
def all(self) -> list[Any]:
|
||||
return self._all_result
|
||||
|
||||
existing = Usage(
|
||||
request_id="req-usage-batch-1",
|
||||
provider_name="pending",
|
||||
model="gemini-3-pro-image-preview",
|
||||
status="completed",
|
||||
billing_status="pending",
|
||||
)
|
||||
assert existing.status == "completed"
|
||||
assert existing.billing_status == "pending"
|
||||
assert existing.finalized_at is None
|
||||
|
||||
db = MagicMock()
|
||||
db.query.side_effect = lambda model: (
|
||||
DummyQuery([existing]) if model is Usage else DummyQuery([])
|
||||
)
|
||||
|
||||
usage_params = {
|
||||
"status": "completed",
|
||||
"response_headers": {"content-type": "text/event-stream"},
|
||||
"response_body": {"chunks": [{"foo": "bar"}], "metadata": {"stream": True}},
|
||||
}
|
||||
|
||||
monkeypatch.setattr(
|
||||
UsageService,
|
||||
"_prepare_usage_records_batch",
|
||||
AsyncMock(return_value=[(usage_params, 0.0, None)]),
|
||||
)
|
||||
|
||||
def _fake_update(existing_usage: Any, params: dict[str, Any], _target_model: Any) -> None:
|
||||
existing_usage.status = params.get("status", existing_usage.status)
|
||||
existing_usage.response_headers = params.get("response_headers")
|
||||
existing_usage.response_body = params.get("response_body")
|
||||
|
||||
monkeypatch.setattr(UsageService, "_update_existing_usage", _fake_update)
|
||||
|
||||
result = await UsageService.record_usage_batch(
|
||||
db,
|
||||
[
|
||||
{
|
||||
"request_id": "req-usage-batch-1",
|
||||
"provider": "Antigravity反代",
|
||||
"model": "gemini-3-pro-image-preview",
|
||||
"status": "completed",
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
assert result and result[0] is existing
|
||||
assert existing.response_headers == usage_params["response_headers"]
|
||||
assert existing.response_body == usage_params["response_body"]
|
||||
assert existing.billing_status == "settled"
|
||||
assert existing.finalized_at is not None
|
||||
|
||||
Reference in New Issue
Block a user