mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-01 17:00:21 +08:00
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 生命周期行为
This commit is contained in:
@@ -231,7 +231,8 @@ class Config:
|
||||
# 最终写入 DB 前仍会按 SystemConfigService 做脱敏与截断。
|
||||
self.usage_queue_stream_key = os.getenv("USAGE_QUEUE_STREAM_KEY", "usage:events")
|
||||
self.usage_queue_stream_group = os.getenv("USAGE_QUEUE_STREAM_GROUP", "usage_consumers")
|
||||
self.usage_queue_stream_maxlen = int(os.getenv("USAGE_QUEUE_STREAM_MAXLEN", "200000"))
|
||||
# 主队列只做短暂缓冲;成功消费后会立即删除,不应把 Redis 当历史存储。
|
||||
self.usage_queue_stream_maxlen = int(os.getenv("USAGE_QUEUE_STREAM_MAXLEN", "2000"))
|
||||
self.usage_queue_dlq_key = os.getenv("USAGE_QUEUE_DLQ_KEY", "usage:events:dlq")
|
||||
self.usage_queue_dlq_maxlen = int(os.getenv("USAGE_QUEUE_DLQ_MAXLEN", "5000"))
|
||||
self.usage_queue_consumer_batch = int(os.getenv("USAGE_QUEUE_CONSUMER_BATCH", "200"))
|
||||
|
||||
@@ -1,161 +0,0 @@
|
||||
"""Helpers for stable outbound request cache fingerprints."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
from collections.abc import Mapping
|
||||
from enum import Enum
|
||||
from typing import Any
|
||||
|
||||
# model 和 prompt_cache_key 统一包含在所有格式中,无需运行时动态添加
|
||||
_CACHE_RELEVANT_FIELDS_BY_FORMAT: dict[str, frozenset[str]] = {
|
||||
"openai:chat": frozenset({"model", "messages", "tools", "tool_choice", "prompt_cache_key"}),
|
||||
"openai:cli": frozenset(
|
||||
{"model", "input", "instructions", "tools", "tool_choice", "prompt_cache_key"}
|
||||
),
|
||||
"openai:compact": frozenset(
|
||||
{"model", "input", "instructions", "tools", "tool_choice", "prompt_cache_key"}
|
||||
),
|
||||
"claude:chat": frozenset(
|
||||
{"model", "system", "messages", "tools", "tool_choice", "prompt_cache_key"}
|
||||
),
|
||||
"gemini:chat": frozenset(
|
||||
{
|
||||
"model",
|
||||
"contents",
|
||||
"system_instruction",
|
||||
"systemInstruction",
|
||||
"tools",
|
||||
"tool_config",
|
||||
"toolConfig",
|
||||
"generation_config",
|
||||
"generationConfig",
|
||||
"prompt_cache_key",
|
||||
}
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
def _normalize_for_hash(value: Any) -> Any:
|
||||
if value is None or isinstance(value, (str, int, float, bool)):
|
||||
return value
|
||||
if isinstance(value, bytes):
|
||||
return value.decode("utf-8", errors="replace")
|
||||
if isinstance(value, Enum):
|
||||
return _normalize_for_hash(value.value)
|
||||
if isinstance(value, Mapping):
|
||||
return {str(key): _normalize_for_hash(item) for key, item in value.items()}
|
||||
if isinstance(value, (list, tuple)):
|
||||
return [_normalize_for_hash(item) for item in value]
|
||||
if isinstance(value, (set, frozenset)):
|
||||
normalized = [_normalize_for_hash(item) for item in value]
|
||||
return sorted(normalized, key=_stable_json_dumps)
|
||||
return str(value)
|
||||
|
||||
|
||||
def _stable_json_dumps(value: Any) -> str:
|
||||
return json.dumps(
|
||||
_normalize_for_hash(value),
|
||||
ensure_ascii=False,
|
||||
sort_keys=True,
|
||||
separators=(",", ":"),
|
||||
)
|
||||
|
||||
|
||||
def _hash_json_payload(value: Any) -> tuple[str, int]:
|
||||
"""Return (sha256_hex, json_byte_length) for the canonicalized JSON."""
|
||||
payload = _stable_json_dumps(value).encode("utf-8")
|
||||
return hashlib.sha256(payload).hexdigest(), len(payload)
|
||||
|
||||
|
||||
def _normalize_provider_api_format(provider_api_format: str | None) -> str | None:
|
||||
normalized = str(provider_api_format or "").strip().lower()
|
||||
return normalized or None
|
||||
|
||||
|
||||
def _get_prompt_cache_key(payload: Any) -> str | None:
|
||||
if not isinstance(payload, Mapping):
|
||||
return None
|
||||
prompt_cache_key = str(payload.get("prompt_cache_key") or "").strip()
|
||||
return prompt_cache_key or None
|
||||
|
||||
|
||||
def _extract_cache_relevant_payload(
|
||||
payload: Any, provider_api_format: str | None
|
||||
) -> tuple[Any, list[str]]:
|
||||
if not isinstance(payload, Mapping):
|
||||
return payload, []
|
||||
|
||||
fields = _CACHE_RELEVANT_FIELDS_BY_FORMAT.get(provider_api_format or "")
|
||||
if not fields:
|
||||
# 未知格式:整个 payload 参与哈希
|
||||
top_level_keys = sorted(str(key) for key in payload.keys())
|
||||
return dict(payload), top_level_keys
|
||||
|
||||
subset = {field: payload[field] for field in fields if field in payload}
|
||||
if not subset:
|
||||
top_level_keys = sorted(str(key) for key in payload.keys())
|
||||
return dict(payload), top_level_keys
|
||||
|
||||
return subset, sorted(subset.keys())
|
||||
|
||||
|
||||
def _build_field_fingerprints(payload: Any, field_names: list[str]) -> dict[str, dict[str, Any]]:
|
||||
if not isinstance(payload, Mapping) or not field_names:
|
||||
return {}
|
||||
|
||||
field_fingerprints: dict[str, dict[str, Any]] = {}
|
||||
for field_name in field_names:
|
||||
if field_name not in payload:
|
||||
continue
|
||||
field_sha256, field_bytes = _hash_json_payload(payload[field_name])
|
||||
field_fingerprints[field_name] = {
|
||||
"sha256": field_sha256,
|
||||
"bytes": field_bytes,
|
||||
}
|
||||
return field_fingerprints
|
||||
|
||||
|
||||
def build_request_cache_fingerprint(
|
||||
provider_request_body: Any,
|
||||
*,
|
||||
provider_api_format: str | None = None,
|
||||
) -> dict[str, Any] | None:
|
||||
"""Build stable hashes for the final outbound payload and its cache-relevant subset."""
|
||||
if provider_request_body is None:
|
||||
return None
|
||||
|
||||
normalized_format = _normalize_provider_api_format(provider_api_format)
|
||||
payload_sha256, payload_bytes = _hash_json_payload(provider_request_body)
|
||||
cache_relevant_payload, cache_relevant_keys = _extract_cache_relevant_payload(
|
||||
provider_request_body,
|
||||
normalized_format,
|
||||
)
|
||||
cache_relevant_sha256, cache_relevant_bytes = _hash_json_payload(cache_relevant_payload)
|
||||
field_fingerprints = _build_field_fingerprints(cache_relevant_payload, cache_relevant_keys)
|
||||
|
||||
top_level_keys = []
|
||||
if isinstance(provider_request_body, Mapping):
|
||||
top_level_keys = sorted(str(key) for key in provider_request_body.keys())
|
||||
|
||||
fingerprint: dict[str, Any] = {
|
||||
"version": 2,
|
||||
"provider_api_format": normalized_format,
|
||||
"payload_sha256": payload_sha256,
|
||||
"payload_bytes": payload_bytes,
|
||||
"cache_relevant_sha256": cache_relevant_sha256,
|
||||
"cache_relevant_bytes": cache_relevant_bytes,
|
||||
"top_level_keys": top_level_keys,
|
||||
"cache_relevant_keys": cache_relevant_keys,
|
||||
"field_fingerprints": field_fingerprints,
|
||||
}
|
||||
|
||||
prompt_cache_key = _get_prompt_cache_key(provider_request_body)
|
||||
if prompt_cache_key:
|
||||
fingerprint["prompt_cache_key"] = prompt_cache_key
|
||||
|
||||
return fingerprint
|
||||
|
||||
|
||||
__all__ = ["build_request_cache_fingerprint"]
|
||||
@@ -38,7 +38,6 @@ METADATA_KEEP_KEYS: frozenset[str] = frozenset(
|
||||
{
|
||||
"billing_snapshot",
|
||||
"billing_updated_at",
|
||||
"cache_fingerprint",
|
||||
"perf",
|
||||
"pool_summary",
|
||||
"scheduling_audit",
|
||||
|
||||
@@ -134,6 +134,9 @@ class UsageQueueConsumer:
|
||||
self._dlq_key = config.usage_queue_dlq_key
|
||||
self._dlq_maxlen = config.usage_queue_dlq_maxlen
|
||||
self._metrics_interval = config.usage_queue_metrics_interval_seconds
|
||||
# 清理长期闲置的旧 consumer,避免 Redis consumer group 元数据持续累积。
|
||||
# 仅清理 pending=0 且空闲时间足够长的 consumer,不影响正常重投递。
|
||||
self._stale_consumer_idle_ms = max(self._claim_idle_ms * 10, 60 * 60 * 1000)
|
||||
|
||||
@staticmethod
|
||||
def _is_duplicate_key_error(exc: IntegrityError) -> bool:
|
||||
@@ -158,6 +161,9 @@ class UsageQueueConsumer:
|
||||
async def start(self) -> None:
|
||||
if self._running:
|
||||
return
|
||||
redis_client = await get_redis_client(require_redis=False)
|
||||
if redis_client:
|
||||
await self._cleanup_stale_consumers(redis_client)
|
||||
self._running = True
|
||||
self._task = asyncio.create_task(self._run(), name="usage-queue-consumer")
|
||||
logger.info("[usage-queue] Consumer started: {}", self._consumer)
|
||||
@@ -172,8 +178,81 @@ class UsageQueueConsumer:
|
||||
await self._task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
redis_client = await get_redis_client(require_redis=False)
|
||||
if redis_client:
|
||||
await self._delete_consumer(redis_client, self._consumer)
|
||||
logger.info("[usage-queue] Consumer stopped: {}", self._consumer)
|
||||
|
||||
async def _delete_consumer(self, redis_client: Any, consumer_name: str) -> None:
|
||||
try:
|
||||
await redis_client.xgroup_delconsumer(
|
||||
self._stream_key,
|
||||
self._stream_group,
|
||||
consumer_name,
|
||||
)
|
||||
except ResponseError as exc:
|
||||
# Group/stream may already be gone during shutdown; ignore in that case.
|
||||
if "NOGROUP" in str(exc) or "ERR no such key" in str(exc):
|
||||
return
|
||||
logger.debug("[usage-queue] DELCONSUMER failed for {}: {}", consumer_name, exc)
|
||||
except Exception as exc:
|
||||
logger.debug("[usage-queue] DELCONSUMER failed for {}: {}", consumer_name, exc)
|
||||
|
||||
async def _cleanup_stale_consumers(self, redis_client: Any) -> None:
|
||||
try:
|
||||
consumers = await redis_client.xinfo_consumers(
|
||||
self._stream_key,
|
||||
self._stream_group,
|
||||
)
|
||||
except ResponseError as exc:
|
||||
if "NOGROUP" in str(exc):
|
||||
return
|
||||
logger.debug("[usage-queue] XINFO CONSUMERS failed: {}", exc)
|
||||
return
|
||||
except Exception as exc:
|
||||
logger.debug("[usage-queue] XINFO CONSUMERS failed: {}", exc)
|
||||
return
|
||||
|
||||
deleted = 0
|
||||
for consumer in consumers or []:
|
||||
if not isinstance(consumer, dict):
|
||||
continue
|
||||
consumer_name = str(consumer.get("name") or "").strip()
|
||||
if not consumer_name or consumer_name == self._consumer:
|
||||
continue
|
||||
|
||||
pending = int(consumer.get("pending", 0) or 0)
|
||||
idle_ms = int(consumer.get("idle", 0) or 0)
|
||||
if pending > 0 or idle_ms < self._stale_consumer_idle_ms:
|
||||
continue
|
||||
|
||||
await self._delete_consumer(redis_client, consumer_name)
|
||||
deleted += 1
|
||||
|
||||
if deleted:
|
||||
logger.info(
|
||||
"[usage-queue] Cleaned up {} stale consumers from group {}",
|
||||
deleted,
|
||||
self._stream_group,
|
||||
)
|
||||
|
||||
async def _ack_and_delete_messages(self, redis_client: Any, message_ids: list[str]) -> None:
|
||||
"""ACK messages and immediately delete them from the main stream.
|
||||
|
||||
usage:events is only meant to be a short-lived buffer. Once an event is
|
||||
successfully persisted (or moved to DLQ), keeping it in Redis only
|
||||
retains duplicate history and inflates memory.
|
||||
"""
|
||||
if not message_ids:
|
||||
return
|
||||
|
||||
pipe = redis_client.pipeline()
|
||||
for message_id in message_ids:
|
||||
pipe.xack(self._stream_key, self._stream_group, message_id)
|
||||
for message_id in message_ids:
|
||||
pipe.xdel(self._stream_key, message_id)
|
||||
await pipe.execute()
|
||||
|
||||
async def _run(self) -> None:
|
||||
while self._running:
|
||||
try:
|
||||
@@ -284,10 +363,7 @@ class UsageQueueConsumer:
|
||||
|
||||
# 使用 pipeline 批量 ACK 成功处理的消息
|
||||
if success_ids:
|
||||
pipe = redis_client.pipeline()
|
||||
for message_id in success_ids:
|
||||
pipe.xack(self._stream_key, self._stream_group, message_id)
|
||||
await pipe.execute()
|
||||
await self._ack_and_delete_messages(redis_client, success_ids)
|
||||
|
||||
async def _process_record_batch(
|
||||
self,
|
||||
@@ -307,11 +383,8 @@ class UsageQueueConsumer:
|
||||
# 批量写入
|
||||
await self._record_usage_batch(records)
|
||||
|
||||
# 使用 pipeline 批量 ACK 提升性能
|
||||
pipe = redis_client.pipeline()
|
||||
for message_id in message_ids:
|
||||
pipe.xack(self._stream_key, self._stream_group, message_id)
|
||||
await pipe.execute()
|
||||
# 写库成功后立即从主队列删除,避免 Redis 保留已入库历史。
|
||||
await self._ack_and_delete_messages(redis_client, message_ids)
|
||||
|
||||
logger.debug("[usage-queue] Batch processed {} records", len(records))
|
||||
|
||||
@@ -340,10 +413,7 @@ class UsageQueueConsumer:
|
||||
)
|
||||
# 批量 ACK 成功处理的消息
|
||||
if success_ids:
|
||||
pipe = redis_client.pipeline()
|
||||
for message_id in success_ids:
|
||||
pipe.xack(self._stream_key, self._stream_group, message_id)
|
||||
await pipe.execute()
|
||||
await self._ack_and_delete_messages(redis_client, success_ids)
|
||||
|
||||
async def _handle_processing_error(
|
||||
self,
|
||||
@@ -367,7 +437,7 @@ class UsageQueueConsumer:
|
||||
)
|
||||
else:
|
||||
await redis_client.xadd(self._dlq_key, dlq_fields)
|
||||
await redis_client.xack(self._stream_key, self._stream_group, message_id)
|
||||
await self._ack_and_delete_messages(redis_client, [message_id])
|
||||
logger.error(
|
||||
"[usage-queue] Message moved to DLQ after {} attempts: {}", retries, message_id
|
||||
)
|
||||
|
||||
@@ -12,7 +12,6 @@ from typing import Any
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.core.logger import logger
|
||||
from src.services.provider.cache_fingerprint import build_request_cache_fingerprint
|
||||
from src.services.system.audit import audit_service
|
||||
from src.services.usage.service import UsageService
|
||||
|
||||
@@ -36,8 +35,6 @@ class MessageTelemetry:
|
||||
*,
|
||||
request_metadata: dict[str, Any] | None = None,
|
||||
response_metadata: dict[str, Any] | None = None,
|
||||
provider_request_body: Any | None = None,
|
||||
provider_api_format: str | None = None,
|
||||
) -> dict[str, Any] | None:
|
||||
metadata: dict[str, Any] | None = None
|
||||
|
||||
@@ -48,23 +45,6 @@ class MessageTelemetry:
|
||||
elif response_metadata:
|
||||
metadata = dict(response_metadata)
|
||||
|
||||
fingerprint = build_request_cache_fingerprint(
|
||||
provider_request_body,
|
||||
provider_api_format=provider_api_format,
|
||||
)
|
||||
if fingerprint:
|
||||
if metadata is None:
|
||||
metadata = {}
|
||||
metadata["cache_fingerprint"] = fingerprint
|
||||
logger.debug(
|
||||
"[Telemetry] cache fingerprint: request_id={}, format={}, payload_sha256={}, cache_sha256={}, prompt_cache_key_present={}",
|
||||
self.request_id,
|
||||
fingerprint.get("provider_api_format"),
|
||||
str(fingerprint.get("payload_sha256") or "")[:12],
|
||||
str(fingerprint.get("cache_relevant_sha256") or "")[:12],
|
||||
bool(fingerprint.get("prompt_cache_key")),
|
||||
)
|
||||
|
||||
return metadata
|
||||
|
||||
async def calculate_cost(
|
||||
@@ -136,8 +116,6 @@ class MessageTelemetry:
|
||||
metadata = self._build_usage_metadata(
|
||||
request_metadata=request_metadata,
|
||||
response_metadata=response_metadata,
|
||||
provider_request_body=provider_request_body,
|
||||
provider_api_format=endpoint_api_format or api_format,
|
||||
)
|
||||
|
||||
usage = await UsageService.record_usage(
|
||||
@@ -262,8 +240,6 @@ class MessageTelemetry:
|
||||
|
||||
metadata = self._build_usage_metadata(
|
||||
request_metadata=request_metadata,
|
||||
provider_request_body=provider_request_body,
|
||||
provider_api_format=endpoint_api_format or api_format,
|
||||
)
|
||||
|
||||
await UsageService.record_usage(
|
||||
@@ -352,8 +328,6 @@ class MessageTelemetry:
|
||||
provider_name = provider or "unknown"
|
||||
metadata = self._build_usage_metadata(
|
||||
request_metadata=request_metadata,
|
||||
provider_request_body=provider_request_body,
|
||||
provider_api_format=endpoint_api_format or api_format,
|
||||
)
|
||||
|
||||
await UsageService.record_usage(
|
||||
|
||||
@@ -1,229 +0,0 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from src.config.settings import config
|
||||
from src.services.provider.cache_fingerprint import build_request_cache_fingerprint
|
||||
from src.services.usage._recording_helpers import sanitize_request_metadata
|
||||
from src.services.usage.service import UsageService
|
||||
from src.services.usage.telemetry import MessageTelemetry
|
||||
|
||||
|
||||
def _build_openai_cli_body() -> dict[str, Any]:
|
||||
return {
|
||||
"model": "gpt-5.4",
|
||||
"instructions": "You are precise.",
|
||||
"input": [{"role": "user", "content": [{"type": "input_text", "text": "hello"}]}],
|
||||
"tools": [
|
||||
{
|
||||
"type": "function",
|
||||
"name": "lookup_weather",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"required": ["city", "country"],
|
||||
"properties": {
|
||||
"country": {"type": "string"},
|
||||
"city": {"type": "string"},
|
||||
},
|
||||
},
|
||||
}
|
||||
],
|
||||
"temperature": 0.2,
|
||||
"prompt_cache_key": "pcache-123",
|
||||
}
|
||||
|
||||
|
||||
def test_build_request_cache_fingerprint_is_stable_for_dict_key_reordering() -> None:
|
||||
body_a = _build_openai_cli_body()
|
||||
body_b = {
|
||||
"prompt_cache_key": "pcache-123",
|
||||
"temperature": 0.2,
|
||||
"tools": [
|
||||
{
|
||||
"parameters": {
|
||||
"properties": {
|
||||
"city": {"type": "string"},
|
||||
"country": {"type": "string"},
|
||||
},
|
||||
"required": ["city", "country"],
|
||||
"type": "object",
|
||||
},
|
||||
"name": "lookup_weather",
|
||||
"type": "function",
|
||||
}
|
||||
],
|
||||
"input": [{"content": [{"text": "hello", "type": "input_text"}], "role": "user"}],
|
||||
"instructions": "You are precise.",
|
||||
"model": "gpt-5.4",
|
||||
}
|
||||
|
||||
fingerprint_a = build_request_cache_fingerprint(body_a, provider_api_format="openai:cli")
|
||||
fingerprint_b = build_request_cache_fingerprint(body_b, provider_api_format="openai:cli")
|
||||
|
||||
assert fingerprint_a is not None
|
||||
assert fingerprint_b is not None
|
||||
assert fingerprint_a["version"] == 2
|
||||
assert fingerprint_a["payload_sha256"] == fingerprint_b["payload_sha256"]
|
||||
assert fingerprint_a["cache_relevant_sha256"] == fingerprint_b["cache_relevant_sha256"]
|
||||
assert fingerprint_a["field_fingerprints"] == fingerprint_b["field_fingerprints"]
|
||||
assert fingerprint_a["prompt_cache_key"] == "pcache-123"
|
||||
assert fingerprint_a["cache_relevant_keys"] == [
|
||||
"input",
|
||||
"instructions",
|
||||
"model",
|
||||
"prompt_cache_key",
|
||||
"tools",
|
||||
]
|
||||
assert fingerprint_a["field_fingerprints"]["instructions"]["bytes"] > 0
|
||||
assert fingerprint_a["field_fingerprints"]["tools"]["bytes"] > 0
|
||||
|
||||
|
||||
def test_build_request_cache_fingerprint_ignores_non_prompt_fields_in_cache_hash() -> None:
|
||||
body_a = _build_openai_cli_body()
|
||||
body_b = _build_openai_cli_body()
|
||||
body_b["temperature"] = 0.9
|
||||
|
||||
fingerprint_a = build_request_cache_fingerprint(body_a, provider_api_format="openai:cli")
|
||||
fingerprint_b = build_request_cache_fingerprint(body_b, provider_api_format="openai:cli")
|
||||
|
||||
assert fingerprint_a is not None
|
||||
assert fingerprint_b is not None
|
||||
assert fingerprint_a["payload_sha256"] != fingerprint_b["payload_sha256"]
|
||||
assert fingerprint_a["cache_relevant_sha256"] == fingerprint_b["cache_relevant_sha256"]
|
||||
assert fingerprint_a["field_fingerprints"] == fingerprint_b["field_fingerprints"]
|
||||
assert "temperature" not in fingerprint_a["field_fingerprints"]
|
||||
|
||||
|
||||
def test_build_request_cache_fingerprint_tracks_prompt_changes() -> None:
|
||||
body_a = _build_openai_cli_body()
|
||||
body_b = _build_openai_cli_body()
|
||||
body_b["instructions"] = "You are terse."
|
||||
|
||||
fingerprint_a = build_request_cache_fingerprint(body_a, provider_api_format="openai:cli")
|
||||
fingerprint_b = build_request_cache_fingerprint(body_b, provider_api_format="openai:cli")
|
||||
|
||||
assert fingerprint_a is not None
|
||||
assert fingerprint_b is not None
|
||||
assert fingerprint_a["cache_relevant_sha256"] != fingerprint_b["cache_relevant_sha256"]
|
||||
assert (
|
||||
fingerprint_a["field_fingerprints"]["instructions"]["sha256"]
|
||||
!= fingerprint_b["field_fingerprints"]["instructions"]["sha256"]
|
||||
)
|
||||
assert (
|
||||
fingerprint_a["field_fingerprints"]["input"]["sha256"]
|
||||
== fingerprint_b["field_fingerprints"]["input"]["sha256"]
|
||||
)
|
||||
assert (
|
||||
fingerprint_a["field_fingerprints"]["tools"]["sha256"]
|
||||
== fingerprint_b["field_fingerprints"]["tools"]["sha256"]
|
||||
)
|
||||
|
||||
|
||||
def test_sanitize_request_metadata_preserves_cache_fingerprint(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setattr(config, "usage_metadata_max_bytes", 120, raising=False)
|
||||
|
||||
metadata = {
|
||||
"trace": {"payload": "x" * 400},
|
||||
"debug": {"payload": "y" * 400},
|
||||
"cache_fingerprint": {
|
||||
"payload_sha256": "a" * 64,
|
||||
"cache_relevant_sha256": "b" * 64,
|
||||
},
|
||||
}
|
||||
|
||||
sanitized = sanitize_request_metadata(metadata)
|
||||
|
||||
assert sanitized["_metadata_truncated"] is True
|
||||
assert sanitized["cache_fingerprint"]["payload_sha256"] == "a" * 64
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_message_telemetry_record_success_keeps_response_shape_and_adds_fingerprint(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
captured: dict[str, Any] = {}
|
||||
|
||||
async def _fake_record_usage(**kwargs: Any) -> Any:
|
||||
captured.update(kwargs)
|
||||
return SimpleNamespace(total_cost_usd=0.0, input_tokens=1, output_tokens=2)
|
||||
|
||||
monkeypatch.setattr(UsageService, "record_usage", _fake_record_usage)
|
||||
|
||||
telemetry = MessageTelemetry(
|
||||
db=SimpleNamespace(), # type: ignore[arg-type]
|
||||
user=None,
|
||||
api_key=None,
|
||||
request_id="req-cache-fingerprint",
|
||||
client_ip="127.0.0.1",
|
||||
)
|
||||
|
||||
await telemetry.record_success(
|
||||
provider="openai",
|
||||
model="gpt-5.4",
|
||||
input_tokens=1,
|
||||
output_tokens=2,
|
||||
response_time_ms=10,
|
||||
status_code=200,
|
||||
request_body={"messages": [{"role": "user", "content": "hello"}]},
|
||||
request_headers={"user-agent": "codex desktop"},
|
||||
response_body={"id": "resp-1"},
|
||||
response_headers={"x-test": "1"},
|
||||
provider_request_body=_build_openai_cli_body(),
|
||||
response_metadata={"model_version": "gpt-5.4-2026-03-01"},
|
||||
endpoint_api_format="openai:cli",
|
||||
)
|
||||
|
||||
metadata = captured["metadata"]
|
||||
assert metadata["model_version"] == "gpt-5.4-2026-03-01"
|
||||
assert "response" not in metadata
|
||||
assert metadata["cache_fingerprint"]["version"] == 2
|
||||
assert metadata["cache_fingerprint"]["provider_api_format"] == "openai:cli"
|
||||
assert metadata["cache_fingerprint"]["prompt_cache_key"] == "pcache-123"
|
||||
assert "instructions" in metadata["cache_fingerprint"]["field_fingerprints"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_message_telemetry_record_failure_keeps_request_metadata_and_adds_fingerprint(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
captured: dict[str, Any] = {}
|
||||
|
||||
async def _fake_record_usage(**kwargs: Any) -> Any:
|
||||
captured.update(kwargs)
|
||||
return SimpleNamespace()
|
||||
|
||||
monkeypatch.setattr(UsageService, "record_usage", _fake_record_usage)
|
||||
|
||||
telemetry = MessageTelemetry(
|
||||
db=SimpleNamespace(), # type: ignore[arg-type]
|
||||
user=None,
|
||||
api_key=None,
|
||||
request_id="req-cache-fingerprint-fail",
|
||||
client_ip="127.0.0.1",
|
||||
)
|
||||
|
||||
await telemetry.record_failure(
|
||||
provider="openai",
|
||||
model="gpt-5.4",
|
||||
response_time_ms=10,
|
||||
status_code=502,
|
||||
error_message="upstream failed",
|
||||
request_body={"messages": [{"role": "user", "content": "hello"}]},
|
||||
request_headers={"user-agent": "codex desktop"},
|
||||
is_stream=False,
|
||||
provider_request_body=_build_openai_cli_body(),
|
||||
request_metadata={"perf": {"ttfb_ms": 12}},
|
||||
endpoint_api_format="openai:cli",
|
||||
)
|
||||
|
||||
metadata = captured["metadata"]
|
||||
assert metadata["perf"]["ttfb_ms"] == 12
|
||||
assert metadata["cache_fingerprint"]["version"] == 2
|
||||
assert metadata["cache_fingerprint"]["provider_api_format"] == "openai:cli"
|
||||
assert metadata["cache_fingerprint"]["prompt_cache_key"] == "pcache-123"
|
||||
assert "input" in metadata["cache_fingerprint"]["field_fingerprints"]
|
||||
@@ -527,6 +527,10 @@ class MockRedisPipeline:
|
||||
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:
|
||||
@@ -534,6 +538,10 @@ class MockRedisPipeline:
|
||||
_, 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
|
||||
|
||||
|
||||
@@ -542,9 +550,12 @@ class MockRedisForConsumer:
|
||||
|
||||
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
|
||||
@@ -556,6 +567,13 @@ class MockRedisForConsumer:
|
||||
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,
|
||||
@@ -578,6 +596,9 @@ class MockRedisForConsumer:
|
||||
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,
|
||||
@@ -700,6 +721,57 @@ async def test_consumer_start_stop() -> None:
|
||||
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:
|
||||
"""测试成功处理消息"""
|
||||
@@ -735,6 +807,27 @@ async def test_consumer_process_messages_success(monkeypatch: Any) -> None:
|
||||
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 事件的重试行为"""
|
||||
@@ -804,6 +897,7 @@ async def test_consumer_process_messages_move_to_dlq(monkeypatch: Any) -> None:
|
||||
|
||||
# 消息应该被 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
|
||||
@@ -1203,8 +1297,14 @@ async def test_consumer_process_record_batch_success(monkeypatch: Any) -> None:
|
||||
records = mock_record_batch.call_args[0][1]
|
||||
assert len(records) == 3
|
||||
|
||||
# 验证所有消息被 ACK
|
||||
# 验证所有消息被 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
|
||||
@@ -1243,8 +1343,9 @@ async def test_consumer_process_record_batch_fallback(monkeypatch: Any) -> None:
|
||||
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
|
||||
# 消息被 ACK 并从主队列删除
|
||||
assert len(mock_redis.xack_calls) == 1
|
||||
assert mock_redis.xdel_calls == [(config.usage_queue_stream_key, "msg-1")]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
||||
Reference in New Issue
Block a user