Files
Aether/tests/services/test_usage_queue_events.py
fawney19 086efe6efe refactor(usage-queue): ACK 后立即删除消息,清理 consumer 元数据
- 消费成功/移入 DLQ 后立即 XDEL,避免主流 Redis 保留已入库历史
- 缩小 usage_queue_stream_maxlen 默认值:200000 -> 2000(仅作短暂缓冲)
- 启动时清理 pending=0 且长期闲置的旧 consumer(防 consumer group 元数据累积)
- 停机时主动 XGROUP DELCONSUMER 移除自身
- 移除 cache_fingerprint 模块及其对 telemetry/recording_helpers 的引用
- 同步更新相关测试,验证 xdel 调用及 consumer 生命周期行为
2026-03-19 02:09:09 +08:00

1574 lines
50 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
import json
from typing import Any
from unittest.mock import AsyncMock, MagicMock
import pytest
from redis.exceptions import ResponseError
from src.config.settings import config
from src.services.usage.consumer_streams import (
UsageQueueConsumer,
_consumer_name,
ensure_usage_stream_group,
)
from src.services.usage.events import (
UsageEvent,
UsageEventType,
build_usage_event,
sanitize_payload,
)
from src.services.usage.telemetry_writer import (
DbTelemetryWriter,
QueueTelemetryWriter,
)
class DummyRedis:
def __init__(self) -> 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, bytes],
maxlen: int | None = None,
approximate: bool | None = None,
) -> str:
if self.xadd_error:
raise self.xadd_error
self.calls.append((key, fields, maxlen, approximate))
return "1-0"
@pytest.mark.asyncio
async def test_usage_event_roundtrip() -> None:
event = build_usage_event(
event_type=UsageEventType.COMPLETED,
request_id="req-1",
data={"foo": "bar"},
timestamp_ms=123,
)
fields = event.to_stream_fields()
restored = UsageEvent.from_stream_fields(fields)
assert restored.event_type == UsageEventType.COMPLETED
assert restored.request_id == "req-1"
assert restored.timestamp_ms == 123
assert restored.data["foo"] == "bar"
@pytest.mark.asyncio
async def test_queue_writer_publishes_event(monkeypatch: Any) -> None:
dummy = DummyRedis()
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)
old_stream_key = config.usage_queue_stream_key
old_maxlen = config.usage_queue_stream_maxlen
try:
config.usage_queue_stream_key = "usage:events:test"
config.usage_queue_stream_maxlen = 0
writer = QueueTelemetryWriter(
request_id="req-2",
user_id="user-1",
api_key_id="key-1",
log_level="basic",
)
await writer.record_success(
provider="test",
model="model",
input_tokens=1,
output_tokens=2,
response_time_ms=10,
status_code=200,
)
finally:
config.usage_queue_stream_key = old_stream_key
config.usage_queue_stream_maxlen = old_maxlen
assert dummy.calls
key, fields, _, _ = dummy.calls[0]
assert key == "usage:events:test"
event = UsageEvent.from_stream_fields(fields)
assert event.data["user_id"] == "user-1"
assert event.data["api_key_id"] == "key-1"
# ============ events.py 测试 ============
@pytest.mark.asyncio
async def test_usage_event_all_types() -> None:
"""测试所有事件类型的序列化/反序列化"""
for event_type in UsageEventType:
event = build_usage_event(
event_type=event_type,
request_id=f"req-{event_type.value}",
data={"type": event_type.value},
)
fields = event.to_stream_fields()
restored = UsageEvent.from_stream_fields(fields)
assert restored.event_type == event_type
assert restored.data["type"] == event_type.value
@pytest.mark.asyncio
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()
# 模拟 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 字段时抛出异常"""
with pytest.raises(ValueError, match="Missing payload"):
UsageEvent.from_stream_fields({})
def test_sanitize_payload_nested() -> None:
"""测试 sanitize_payload 处理嵌套结构"""
data = {
"str": "hello",
"int": 42,
"float": 3.14,
"bool": True,
"none": None,
"list": [1, "two", {"nested": True}],
"dict": {"a": 1, "b": [2, 3]},
"custom": object(), # 非基础类型,应转为 str
}
result = sanitize_payload(data)
assert result["str"] == "hello"
assert result["int"] == 42
assert result["list"] == [1, "two", {"nested": True}]
assert isinstance(result["custom"], str)
def test_parse_body_json_string() -> None:
"""测试 _parse_body 在消费阶段不反序列化 JSON 字符串"""
from src.services.usage.consumer_streams import _parse_body
# JSON 字符串应保持原样,反序列化延迟到写库前
json_str = '{"messages": [{"role": "user", "content": "hello"}]}'
result = _parse_body(json_str)
assert result == json_str
def test_parse_body_dict_passthrough() -> None:
"""测试 _parse_body 直接返回 dict"""
from src.services.usage.consumer_streams import _parse_body
data = {"key": "value"}
result = _parse_body(data)
assert result is data # 应该是同一个对象
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() -> None:
"""测试 _parse_body 处理被截断的 JSON 字符串"""
from src.services.usage.consumer_streams import _parse_body
# 被截断的字符串无法解析,应原样返回
truncated = '{"content": "x...[truncated]'
result = _parse_body(truncated)
assert result == truncated
# ============ telemetry_writer.py 测试 ============
@pytest.mark.asyncio
async def test_db_telemetry_writer_filters_kwargs() -> None:
"""测试 DbTelemetryWriter 过滤不支持的参数"""
mock_telemetry = MagicMock()
mock_telemetry.record_success = AsyncMock()
mock_telemetry.record_failure = AsyncMock()
mock_telemetry.record_cancelled = AsyncMock()
writer = DbTelemetryWriter(mock_telemetry)
# 调用 record_success包含应被过滤的参数
await writer.record_success(
provider="test",
model="gpt-4",
request_type="chat", # 应被过滤
metadata={"foo": "bar"}, # 应被过滤
input_tokens=100,
)
mock_telemetry.record_success.assert_called_once()
call_kwargs = mock_telemetry.record_success.call_args.kwargs
assert "request_type" not in call_kwargs
assert "metadata" not in call_kwargs
assert call_kwargs["provider"] == "test"
assert call_kwargs["input_tokens"] == 100
@pytest.mark.asyncio
async def test_db_telemetry_writer_all_methods() -> None:
"""测试 DbTelemetryWriter 的所有方法"""
mock_telemetry = MagicMock()
mock_telemetry.record_success = AsyncMock()
mock_telemetry.record_failure = AsyncMock()
mock_telemetry.record_cancelled = AsyncMock()
writer = DbTelemetryWriter(mock_telemetry)
await writer.record_success(provider="p1")
await writer.record_failure(provider="p2")
await writer.record_cancelled(provider="p3")
mock_telemetry.record_success.assert_called_once()
mock_telemetry.record_failure.assert_called_once()
mock_telemetry.record_cancelled.assert_called_once()
@pytest.mark.asyncio
async def test_queue_writer_record_failure(monkeypatch: Any) -> None:
"""测试 QueueTelemetryWriter.record_failure"""
dummy = DummyRedis()
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)
old_maxlen = config.usage_queue_stream_maxlen
try:
config.usage_queue_stream_maxlen = 100
writer = QueueTelemetryWriter(
request_id="req-fail",
user_id="user-1",
api_key_id="key-1",
)
await writer.record_failure(
provider="test",
model="model",
error_message="something went wrong",
status_code=500,
)
finally:
config.usage_queue_stream_maxlen = old_maxlen
assert len(dummy.calls) == 1
_, fields, maxlen, approx = dummy.calls[0]
assert maxlen == 100
assert approx is True
event = UsageEvent.from_stream_fields(fields)
assert event.event_type == UsageEventType.FAILED
assert event.data["error_message"] == "something went wrong"
@pytest.mark.asyncio
async def test_queue_writer_record_cancelled(monkeypatch: Any) -> None:
"""测试 QueueTelemetryWriter.record_cancelled"""
dummy = DummyRedis()
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)
writer = QueueTelemetryWriter(
request_id="req-cancel",
user_id="user-1",
api_key_id="key-1",
)
await writer.record_cancelled(provider="test", model="model")
event = UsageEvent.from_stream_fields(dummy.calls[0][1])
assert event.event_type == UsageEventType.CANCELLED
@pytest.mark.asyncio
async def test_queue_writer_include_headers_bodies(monkeypatch: Any) -> None:
"""测试 include_headers 和 include_bodies 配置"""
dummy = DummyRedis()
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)
writer = QueueTelemetryWriter(
request_id="req-full",
user_id="user-1",
api_key_id="key-1",
log_level="full",
sensitive_headers=["authorization"],
max_request_body_size=0,
max_response_body_size=0,
)
await writer.record_success(
provider="test",
model="model",
request_headers={"Authorization": "Bearer xxx"},
response_headers={"Content-Type": "application/json"},
request_body={"messages": [{"role": "user", "content": "hi"}]},
response_body={"choices": [{"message": {"content": "hello"}}]},
)
event = UsageEvent.from_stream_fields(dummy.calls[0][1])
# Sensitive header should be masked before going into Redis
assert event.data["request_headers"]["Authorization"].startswith("Bear")
assert "****" in event.data["request_headers"]["Authorization"]
assert "request_body" in event.data
assert "response_body" in event.data
@pytest.mark.asyncio
async def test_queue_writer_failure_preserves_empty_request_headers(monkeypatch: Any) -> None:
"""失败事件在传入空请求头时也应保留 request_headers 字段。"""
dummy = DummyRedis()
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)
writer = QueueTelemetryWriter(
request_id="req-empty-hdr",
user_id="user-1",
api_key_id="key-1",
log_level="full",
)
await writer.record_failure(
provider="test",
model="model",
status_code=500,
error_message="boom",
request_headers={},
request_body={"message": "x"},
)
event = UsageEvent.from_stream_fields(dummy.calls[0][1])
assert "request_headers" in event.data
assert event.data["request_headers"] == {}
assert event.data["request_body"] == {"message": "x"}
@pytest.mark.asyncio
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 字符串)
event = build_usage_event(
event_type=UsageEventType.COMPLETED,
request_id="req-body-test",
timestamp_ms=1_700_000_000_123,
data={
"user_id": "user-1",
"api_key_id": "key-1",
"provider": "test",
"model": "gpt-4",
# body 是 JSON 字符串QueueTelemetryWriter._truncate_body 的输出)
"request_body": '{"messages": [{"role": "user", "content": "hello"}]}',
"response_body": '{"choices": [{"message": {"content": "hi"}}]}',
},
)
record = _event_to_record(event)
# 消费阶段不做 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
@pytest.mark.asyncio
async def test_queue_writer_body_truncation(monkeypatch: Any) -> None:
"""测试 body 超长截断"""
dummy = DummyRedis()
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)
writer = QueueTelemetryWriter(
request_id="req-trunc",
user_id="user-1",
api_key_id="key-1",
log_level="full",
max_request_body_size=50,
max_response_body_size=0,
)
long_body = {"content": "x" * 1000}
await writer.record_success(
provider="test",
model="model",
request_body=long_body,
)
event = UsageEvent.from_stream_fields(dummy.calls[0][1])
body = event.data["request_body"]
assert isinstance(body, dict)
assert body.get("_truncated") is True
assert len(body.get("_content") or "") <= 50
@pytest.mark.asyncio
async def test_queue_writer_redis_unavailable(monkeypatch: Any) -> None:
"""测试 Redis 不可用时抛出异常"""
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)
writer = QueueTelemetryWriter(
request_id="req-no-redis",
user_id="user-1",
api_key_id="key-1",
)
with pytest.raises(RuntimeError, match="Redis unavailable"):
await writer.record_success(provider="test", model="model")
@pytest.mark.asyncio
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: bool = False) -> Any:
return dummy
monkeypatch.setattr("src.services.usage.telemetry_writer.get_redis_client", _get_redis_client)
writer = QueueTelemetryWriter(
request_id="req-xadd-fail",
user_id="user-1",
api_key_id="key-1",
)
with pytest.raises(Exception, match="Redis connection lost"):
await writer.record_success(provider="test", model="model")
# ============ consumer_streams.py 测试 ============
class MockRedisPipeline:
"""模拟 Redis Pipeline"""
def __init__(self, parent: "MockRedisForConsumer") -> None:
self._parent = parent
self._commands: list[tuple[str, ...]] = []
def xack(self, key: str, group: str, message_id: str) -> Any:
self._commands.append(("xack", key, group, message_id))
return self
def xdel(self, key: str, message_id: str) -> Any:
self._commands.append(("xdel", key, message_id))
return self
async def execute(self) -> Any:
results = []
for cmd in self._commands:
if cmd[0] == "xack":
_, key, group, message_id = cmd
self._parent.xack_calls.append((key, group, message_id))
results.append(1)
elif cmd[0] == "xdel":
_, key, message_id = cmd
self._parent.xdel_calls.append((key, message_id))
results.append(1)
return results
class MockRedisForConsumer:
"""模拟 Redis 客户端用于消费者测试"""
def __init__(self) -> None:
self.xgroup_create_calls: list[tuple[str, str, str, bool]] = []
self.xgroup_delconsumer_calls: list[tuple[str, str, str]] = []
self.xinfo_consumers_result: list[dict[str, Any]] = []
self.xreadgroup_results: list[Any] = []
self.xautoclaim_results: list[Any] = []
self.xack_calls: list[tuple[str, str, str]] = []
self.xdel_calls: list[tuple[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: 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 xgroup_delconsumer(self, key: str, group: str, consumer: str) -> int:
self.xgroup_delconsumer_calls.append((key, group, consumer))
return 1
async def xinfo_consumers(self, key: str, group: str) -> list[dict[str, Any]]:
return list(self.xinfo_consumers_result)
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: 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: str, group: str, message_id: str) -> Any:
self.xack_calls.append((key, group, message_id))
async def xdel(self, key: str, message_id: str) -> Any:
self.xdel_calls.append((key, message_id))
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: 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: str) -> Any:
return self.xlen_result
async def xpending(self, key: str, group: str) -> Any:
return self.xpending_result
def pipeline(self) -> Any:
return MockRedisPipeline(self)
def test_consumer_name() -> None:
"""测试消费者名称生成"""
name = _consumer_name()
assert ":" in name
# 应包含 PID
import os
assert str(os.getpid()) in name
@pytest.mark.asyncio
async def test_ensure_stream_group_creates_group(monkeypatch: Any) -> None:
"""测试创建消费者组"""
mock_redis = MockRedisForConsumer()
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)
await ensure_usage_stream_group()
assert len(mock_redis.xgroup_create_calls) == 1
key, group, id_, mkstream = mock_redis.xgroup_create_calls[0]
assert key == config.usage_queue_stream_key
assert group == config.usage_queue_stream_group
assert mkstream is True
@pytest.mark.asyncio
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: bool = False) -> Any:
return mock_redis
monkeypatch.setattr("src.services.usage.consumer_streams.get_redis_client", _get_redis_client)
# 不应抛出异常
await ensure_usage_stream_group()
@pytest.mark.asyncio
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: bool = False) -> Any:
return mock_redis
monkeypatch.setattr("src.services.usage.consumer_streams.get_redis_client", _get_redis_client)
with pytest.raises(ResponseError, match="SOME OTHER ERROR"):
await ensure_usage_stream_group()
@pytest.mark.asyncio
async def test_ensure_stream_group_no_redis(monkeypatch: Any) -> None:
"""测试 Redis 不可用时直接返回"""
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)
# 不应抛出异常
await ensure_usage_stream_group()
@pytest.mark.asyncio
async def test_consumer_start_stop() -> None:
"""测试消费者启动和停止"""
consumer = UsageQueueConsumer()
assert not consumer._running
# Mock _run 避免真正执行
consumer._run = AsyncMock() # type: ignore[method-assign]
await consumer.start()
assert consumer._running
assert consumer._task is not None
# 重复启动应该是幂等的
await consumer.start()
assert consumer._running
await consumer.stop()
assert not consumer._running
# 重复停止应该是幂等的
await consumer.stop()
assert not consumer._running
@pytest.mark.asyncio
async def test_consumer_start_cleans_stale_consumers(monkeypatch: Any) -> None:
"""启动时应清理 pending=0 且长期闲置的旧 consumer。"""
mock_redis = MockRedisForConsumer()
current_consumer = _consumer_name()
mock_redis.xinfo_consumers_result = [
{"name": current_consumer, "pending": 0, "idle": 999999999},
{"name": "stale:1", "pending": 0, "idle": 999999999},
{"name": "busy:1", "pending": 1, "idle": 999999999},
{"name": "fresh:1", "pending": 0, "idle": 1000},
]
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)
consumer = UsageQueueConsumer()
consumer._run = AsyncMock() # type: ignore[method-assign]
await consumer.start()
await consumer.stop()
assert ("usage:events", "usage_consumers", "stale:1") in mock_redis.xgroup_delconsumer_calls
assert ("usage:events", "usage_consumers", "busy:1") not in mock_redis.xgroup_delconsumer_calls
assert ("usage:events", "usage_consumers", "fresh:1") not in mock_redis.xgroup_delconsumer_calls
@pytest.mark.asyncio
async def test_consumer_stop_removes_self_from_group(monkeypatch: Any) -> None:
"""停机时应主动删除自身 consumer避免 group 元数据累积。"""
mock_redis = MockRedisForConsumer()
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)
consumer = UsageQueueConsumer()
consumer._run = AsyncMock() # type: ignore[method-assign]
await consumer.start()
await consumer.stop()
assert (
config.usage_queue_stream_key,
config.usage_queue_stream_group,
consumer._consumer,
) in mock_redis.xgroup_delconsumer_calls
@pytest.mark.asyncio
async def test_consumer_process_messages_success(monkeypatch: Any) -> None:
"""测试成功处理消息"""
mock_redis = MockRedisForConsumer()
# 创建测试事件
event = build_usage_event(
event_type=UsageEventType.COMPLETED,
request_id="req-test-1",
data={
"user_id": None,
"api_key_id": None,
"provider": "test-provider",
"model": "test-model",
"input_tokens": 10,
"output_tokens": 20,
},
)
message_id = "1-0"
messages = [(message_id, event.to_stream_fields())]
consumer = UsageQueueConsumer()
# Mock 批量处理方法
consumer._process_record_batch = AsyncMock() # type: ignore[method-assign]
await consumer._process_messages(mock_redis, messages)
# 验证批量处理被调用
consumer._process_record_batch.assert_called_once()
call_messages = consumer._process_record_batch.call_args[0][1]
assert len(call_messages) == 1
assert call_messages[0][2].request_id == "req-test-1"
@pytest.mark.asyncio
async def test_consumer_process_streaming_batch_deletes_messages() -> None:
"""STREAMING 事件处理成功后应立即从主队列删除。"""
mock_redis = MockRedisForConsumer()
consumer = UsageQueueConsumer()
consumer._apply_streaming_event = AsyncMock() # type: ignore[method-assign]
event = build_usage_event(
event_type=UsageEventType.STREAMING,
request_id="req-stream-ok",
data={"provider": "test"},
)
await consumer._process_streaming_batch(mock_redis, [("stream-1", event)])
assert mock_redis.xack_calls == [
(config.usage_queue_stream_key, config.usage_queue_stream_group, "stream-1")
]
assert mock_redis.xdel_calls == [(config.usage_queue_stream_key, "stream-1")]
@pytest.mark.asyncio
async def test_consumer_process_messages_error_retry(monkeypatch: Any) -> None:
"""测试处理消息失败时 STREAMING 事件的重试行为"""
mock_redis = MockRedisForConsumer()
# 返回重试次数小于 max_retries
mock_redis.xpending_range_results = [[{"times_delivered": 2}]]
# 使用 STREAMING 事件测试单条处理失败
event = build_usage_event(
event_type=UsageEventType.STREAMING,
request_id="req-err",
data={"provider": "test"},
)
message_id = "err-1-0"
messages = [(message_id, event.to_stream_fields())]
old_max_retries = config.usage_queue_max_retries
try:
config.usage_queue_max_retries = 5
# 在设置 config 后创建 consumer以便缓存正确的配置值
consumer = UsageQueueConsumer()
# 模拟 STREAMING 事件处理失败
consumer._apply_streaming_event = AsyncMock(side_effect=ValueError("Processing error")) # type: ignore[method-assign]
await consumer._process_messages(mock_redis, messages)
# 消息不应该被 ack还能重试
assert len(mock_redis.xack_calls) == 0
# 不应该移入 DLQ重试次数未达上限
assert len(mock_redis.xadd_calls) == 0
finally:
config.usage_queue_max_retries = old_max_retries
@pytest.mark.asyncio
async def test_consumer_process_messages_move_to_dlq(monkeypatch: Any) -> None:
"""测试消息超过最大重试次数后移入 DLQ"""
mock_redis = MockRedisForConsumer()
# 返回重试次数 >= max_retries
mock_redis.xpending_range_results = [[{"times_delivered": 10}]]
# 测试解析失败的消息会被移入 DLQ
invalid_fields = {"payload": "invalid json"}
message_id = "dlq-1-0"
messages = [(message_id, invalid_fields)]
old_max_retries = config.usage_queue_max_retries
old_dlq_maxlen = config.usage_queue_dlq_maxlen
try:
config.usage_queue_max_retries = 5
config.usage_queue_dlq_maxlen = 1000
# 在设置 config 后创建 consumer以便缓存正确的配置值
consumer = UsageQueueConsumer()
await consumer._process_messages(mock_redis, messages)
# 应该移入 DLQ解析失败的消息
assert len(mock_redis.xadd_calls) == 1
dlq_key, dlq_fields, maxlen, approx = mock_redis.xadd_calls[0]
assert dlq_key == config.usage_queue_dlq_key
assert dlq_fields["source_id"] == message_id
assert "error" in dlq_fields
assert maxlen == 1000
assert approx is True
# 消息应该被 ack
assert len(mock_redis.xack_calls) == 1
assert mock_redis.xdel_calls == [(config.usage_queue_stream_key, message_id)]
finally:
config.usage_queue_max_retries = old_max_retries
config.usage_queue_dlq_maxlen = old_dlq_maxlen
@pytest.mark.asyncio
async def test_consumer_get_delivery_count_dict_format() -> None:
"""测试获取消息投递次数dict 格式)"""
mock_redis = MockRedisForConsumer()
mock_redis.xpending_range_results = [[{"times_delivered": 3}]]
consumer = UsageQueueConsumer()
count = await consumer._get_delivery_count(mock_redis, "msg-1")
assert count == 3
@pytest.mark.asyncio
async def test_consumer_get_delivery_count_tuple_format() -> None:
"""测试获取消息投递次数tuple 格式)"""
mock_redis = MockRedisForConsumer()
# (message_id, consumer, idle_time, times_delivered)
mock_redis.xpending_range_results = [[("msg-1", "consumer-1", 1000, 5)]]
consumer = UsageQueueConsumer()
count = await consumer._get_delivery_count(mock_redis, "msg-1")
assert count == 5
@pytest.mark.asyncio
async def test_consumer_get_delivery_count_empty() -> None:
"""测试消息不存在时返回 0"""
mock_redis = MockRedisForConsumer()
mock_redis.xpending_range_results = [[]]
consumer = UsageQueueConsumer()
count = await consumer._get_delivery_count(mock_redis, "nonexistent")
assert count == 0
@pytest.mark.asyncio
async def test_consumer_read_new_messages(monkeypatch: Any) -> None:
"""测试读取新消息"""
mock_redis = MockRedisForConsumer()
event = build_usage_event(
event_type=UsageEventType.COMPLETED,
request_id="req-new",
data={"provider": "test"},
)
# xreadgroup 返回格式: [(stream_key, [(msg_id, fields), ...])]
mock_redis.xreadgroup_results = [
[(config.usage_queue_stream_key, [("new-1-0", event.to_stream_fields())])]
]
consumer = UsageQueueConsumer()
consumer._process_messages = AsyncMock() # type: ignore[method-assign]
await consumer._read_new(mock_redis)
consumer._process_messages.assert_called_once()
_, messages = consumer._process_messages.call_args[0]
assert len(messages) == 1
assert messages[0][0] == "new-1-0"
@pytest.mark.asyncio
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() # type: ignore[method-assign]
await consumer._maybe_claim_pending(mock_redis)
# 由于间隔未到,不应该处理消息
consumer._process_messages.assert_not_called()
@pytest.mark.asyncio
async def test_consumer_maybe_claim_pending_xautoclaim_error(monkeypatch: Any) -> None:
"""测试 XAUTOCLAIM 失败时优雅处理"""
mock_redis = MockRedisForConsumer()
async def failing_xautoclaim(*args: Any, **kwargs: Any) -> Any:
raise ResponseError("XAUTOCLAIM error")
mock_redis.xautoclaim = failing_xautoclaim # type: ignore[method-assign]
consumer = UsageQueueConsumer()
consumer._last_claim = 0 # 确保会尝试 claim
# 不应抛出异常
await consumer._maybe_claim_pending(mock_redis)
@pytest.mark.asyncio
async def test_consumer_log_metrics_dict_pending() -> None:
"""测试 metrics 日志dict 格式 pending"""
mock_redis = MockRedisForConsumer()
mock_redis.xlen_result = 100
mock_redis.xpending_result = {"pending": 10}
consumer = UsageQueueConsumer()
consumer._last_metrics_log = 0 # 确保会执行
old_interval = config.usage_queue_metrics_interval_seconds
try:
config.usage_queue_metrics_interval_seconds = 0
await consumer._log_metrics(mock_redis)
# 不抛出异常即可
finally:
config.usage_queue_metrics_interval_seconds = old_interval
@pytest.mark.asyncio
async def test_consumer_log_metrics_tuple_pending() -> None:
"""测试 metrics 日志tuple 格式 pending"""
mock_redis = MockRedisForConsumer()
mock_redis.xlen_result = 50
# tuple 格式: (pending_count, min_id, max_id, consumers)
mock_redis.xpending_result = (5, "1-0", "5-0", [])
consumer = UsageQueueConsumer()
consumer._last_metrics_log = 0
old_interval = config.usage_queue_metrics_interval_seconds
try:
config.usage_queue_metrics_interval_seconds = 0
await consumer._log_metrics(mock_redis)
finally:
config.usage_queue_metrics_interval_seconds = old_interval
@pytest.mark.asyncio
async def test_consumer_apply_event_streaming(monkeypatch: Any) -> None:
"""测试处理 STREAMING 事件"""
mock_db = MagicMock()
def mock_create_session() -> Any:
return mock_db
monkeypatch.setattr("src.services.usage.consumer_streams.create_session", mock_create_session)
mock_update_status = MagicMock()
monkeypatch.setattr(
"src.services.usage.consumer_streams.UsageService.update_usage_status",
mock_update_status,
)
consumer = UsageQueueConsumer()
event = UsageEvent(
event_type=UsageEventType.STREAMING,
request_id="req-stream",
timestamp_ms=123,
data={
"provider": "test",
"target_model": "gpt-4",
"first_byte_time_ms": 100,
},
)
await consumer._apply_event(event)
mock_update_status.assert_called_once()
call_kwargs = mock_update_status.call_args.kwargs
assert call_kwargs["request_id"] == "req-stream"
assert call_kwargs["status"] == "streaming"
assert call_kwargs["provider"] == "test"
mock_db.close.assert_called_once()
@pytest.mark.asyncio
async def test_consumer_apply_event_completed(monkeypatch: Any) -> None:
"""测试处理 COMPLETED 事件"""
mock_db = MagicMock()
def mock_create_session() -> Any:
return mock_db
monkeypatch.setattr("src.services.usage.consumer_streams.create_session", mock_create_session)
mock_record_batch = AsyncMock(return_value=[])
monkeypatch.setattr(
"src.services.usage.consumer_streams.UsageService.record_usage_batch",
mock_record_batch,
)
consumer = UsageQueueConsumer()
event = UsageEvent(
event_type=UsageEventType.COMPLETED,
request_id="req-done",
timestamp_ms=123,
data={
"provider": "openai",
"model": "gpt-4",
"input_tokens": 100,
"output_tokens": 200,
},
)
await consumer._apply_event(event)
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()
def mock_create_session() -> Any:
return mock_db
monkeypatch.setattr("src.services.usage.consumer_streams.create_session", mock_create_session)
mock_record_batch = AsyncMock(return_value=[])
monkeypatch.setattr(
"src.services.usage.consumer_streams.UsageService.record_usage_batch",
mock_record_batch,
)
consumer = UsageQueueConsumer()
event = UsageEvent(
event_type=UsageEventType.FAILED,
request_id="req-fail",
timestamp_ms=123,
data={
"provider": "openai",
"model": "gpt-4",
"error_message": "Rate limited",
},
)
await consumer._apply_event(event)
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()
def mock_create_session() -> Any:
return mock_db
monkeypatch.setattr("src.services.usage.consumer_streams.create_session", mock_create_session)
mock_record_batch = AsyncMock(return_value=[])
monkeypatch.setattr(
"src.services.usage.consumer_streams.UsageService.record_usage_batch",
mock_record_batch,
)
consumer = UsageQueueConsumer()
event = UsageEvent(
event_type=UsageEventType.CANCELLED,
request_id="req-cancel",
timestamp_ms=123,
data={"provider": "openai", "model": "gpt-4"},
)
await consumer._apply_event(event)
records = mock_record_batch.call_args[0][1]
assert records[0]["status"] == "cancelled"
# ============ 批量处理测试 ============
@pytest.mark.asyncio
async def test_consumer_process_messages_batch(monkeypatch: Any) -> None:
"""测试批量处理消息"""
mock_redis = MockRedisForConsumer()
# 创建多个测试事件
events = []
for i in range(3):
event = build_usage_event(
event_type=UsageEventType.COMPLETED,
request_id=f"req-batch-{i}",
data={
"user_id": None,
"api_key_id": None,
"provider": "test-provider",
"model": "test-model",
"input_tokens": 10 * (i + 1),
"output_tokens": 20 * (i + 1),
},
)
events.append((f"msg-{i}", event.to_stream_fields(), event))
messages = [(msg_id, fields) for msg_id, fields, _ in events]
consumer = UsageQueueConsumer()
# Mock 批量处理
consumer._process_record_batch = AsyncMock() # type: ignore[method-assign]
await consumer._process_messages(mock_redis, messages)
# 验证批量处理被调用
consumer._process_record_batch.assert_called_once()
@pytest.mark.asyncio
async def test_consumer_process_messages_separates_streaming(monkeypatch: Any) -> None:
"""测试消息分类STREAMING 和其他事件分开处理"""
mock_redis = MockRedisForConsumer()
# 创建混合事件
streaming_event = build_usage_event(
event_type=UsageEventType.STREAMING,
request_id="req-stream",
data={"provider": "test"},
)
completed_event = build_usage_event(
event_type=UsageEventType.COMPLETED,
request_id="req-done",
data={"provider": "test", "model": "gpt-4"},
)
messages = [
("msg-stream", streaming_event.to_stream_fields()),
("msg-done", completed_event.to_stream_fields()),
]
consumer = UsageQueueConsumer()
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)
# STREAMING 事件单独处理
consumer._apply_streaming_event.assert_called_once()
streaming_call_event = consumer._apply_streaming_event.call_args[0][0]
assert streaming_call_event.request_id == "req-stream"
# 记录事件批量处理
consumer._process_record_batch.assert_called_once()
batch_call_messages = consumer._process_record_batch.call_args[0][1]
assert len(batch_call_messages) == 1
assert batch_call_messages[0][2].request_id == "req-done"
@pytest.mark.asyncio
async def test_consumer_process_record_batch_success(monkeypatch: Any) -> None:
"""测试批量记录处理成功"""
mock_redis = MockRedisForConsumer()
mock_db = MagicMock()
def mock_create_session() -> Any:
return mock_db
monkeypatch.setattr("src.services.usage.consumer_streams.create_session", mock_create_session)
mock_record_batch = AsyncMock(return_value=[])
monkeypatch.setattr(
"src.services.usage.consumer_streams.UsageService.record_usage_batch",
mock_record_batch,
)
consumer = UsageQueueConsumer()
# 创建批量消息
events_data = []
for i in range(3):
event = build_usage_event(
event_type=UsageEventType.COMPLETED,
request_id=f"req-{i}",
data={"provider": "test", "model": "gpt-4", "input_tokens": 10, "output_tokens": 20},
)
events_data.append((f"msg-{i}", event.to_stream_fields(), event))
await consumer._process_record_batch(mock_redis, events_data)
# 验证批量写入被调用
mock_record_batch.assert_called_once()
records = mock_record_batch.call_args[0][1]
assert len(records) == 3
# 验证所有消息被 ACK 并从主队列删除
assert len(mock_redis.xack_calls) == 3
assert len(mock_redis.xdel_calls) == 3
assert mock_redis.xdel_calls == [
(config.usage_queue_stream_key, "msg-0"),
(config.usage_queue_stream_key, "msg-1"),
(config.usage_queue_stream_key, "msg-2"),
]
@pytest.mark.asyncio
async def test_consumer_process_record_batch_fallback(monkeypatch: Any) -> None:
"""测试批量处理失败时回退到逐条处理"""
mock_redis = MockRedisForConsumer()
mock_db = MagicMock()
def mock_create_session() -> Any:
return mock_db
monkeypatch.setattr("src.services.usage.consumer_streams.create_session", mock_create_session)
# 首次批量处理失败,回退到单条 batch 成功
mock_record_batch = AsyncMock(side_effect=[Exception("Batch failed"), []])
monkeypatch.setattr(
"src.services.usage.consumer_streams.UsageService.record_usage_batch",
mock_record_batch,
)
consumer = UsageQueueConsumer()
event = build_usage_event(
event_type=UsageEventType.COMPLETED,
request_id="req-fallback",
data={"provider": "test", "model": "gpt-4"},
)
await consumer._process_record_batch(
mock_redis,
[("msg-1", event.to_stream_fields(), event)],
)
# 验证回退到单条 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
assert mock_redis.xdel_calls == [(config.usage_queue_stream_key, "msg-1")]
@pytest.mark.asyncio
async def test_consumer_apply_streaming_event(monkeypatch: Any) -> None:
"""测试 STREAMING 事件处理"""
mock_db = MagicMock()
def mock_create_session() -> Any:
return mock_db
monkeypatch.setattr("src.services.usage.consumer_streams.create_session", mock_create_session)
mock_update_status = MagicMock()
monkeypatch.setattr(
"src.services.usage.consumer_streams.UsageService.update_usage_status",
mock_update_status,
)
consumer = UsageQueueConsumer()
event = UsageEvent(
event_type=UsageEventType.STREAMING,
request_id="req-streaming",
timestamp_ms=123,
data={
"provider": "openai",
"target_model": "gpt-4",
"first_byte_time_ms": 150,
},
)
await consumer._apply_streaming_event(event)
mock_update_status.assert_called_once()
call_kwargs = mock_update_status.call_args.kwargs
assert call_kwargs["request_id"] == "req-streaming"
assert call_kwargs["status"] == "streaming"
assert call_kwargs["first_byte_time_ms"] == 150
@pytest.mark.asyncio
async def test_consumer_apply_record_event(monkeypatch: Any) -> None:
"""测试记录事件单独处理"""
mock_db = MagicMock()
def mock_create_session() -> Any:
return mock_db
monkeypatch.setattr("src.services.usage.consumer_streams.create_session", mock_create_session)
mock_record_batch = AsyncMock(return_value=[])
monkeypatch.setattr(
"src.services.usage.consumer_streams.UsageService.record_usage_batch",
mock_record_batch,
)
consumer = UsageQueueConsumer()
event = UsageEvent(
event_type=UsageEventType.COMPLETED,
request_id="req-record",
timestamp_ms=123,
data={"provider": "openai", "model": "gpt-4", "input_tokens": 100},
)
await consumer._apply_record_event(event)
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
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 options(self, *args: Any, **kwargs: Any) -> "DummyQuery":
return self
def filter(self, *args: Any, **kwargs: Any) -> "DummyQuery":
return self
def with_for_update(self) -> "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
@pytest.mark.asyncio
async def test_record_usage_batch_uses_orm_insert_for_new_records(
monkeypatch: Any,
) -> None:
"""确保批量新建 Usage 走 ORM add 路径,以便复用统一结算逻辑。"""
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 options(self, *args: Any, **kwargs: Any) -> "DummyQuery":
return self
def filter(self, *args: Any, **kwargs: Any) -> "DummyQuery":
return self
def with_for_update(self) -> "DummyQuery":
return self
def all(self) -> list[Any]:
return self._all_result
usage_query_calls = {"count": 0}
def _query_side_effect(model: Any) -> Any:
if model is Usage:
usage_query_calls["count"] += 1
return DummyQuery([])
return DummyQuery([])
db = MagicMock()
db.query.side_effect = _query_side_effect
usage_params = {
"request_id": "req-usage-batch-new",
"provider_name": "openai",
"model": "gpt-4",
"status": "completed",
}
monkeypatch.setattr(
UsageService,
"_prepare_usage_records_batch",
AsyncMock(return_value=[(usage_params, 0.0, None)]),
)
result = await UsageService.record_usage_batch(
db,
[
{
"request_id": "req-usage-batch-new",
"provider": "openai",
"model": "gpt-4",
"status": "completed",
}
],
)
db.add.assert_called_once()
db.bulk_insert_mappings.assert_not_called()
assert result and isinstance(result[0], Usage)
assert result[0].request_id == "req-usage-batch-new"
assert result[0].billing_status == "settled"
assert result[0].finalized_at is not None