mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 17:30:23 +08:00
feat: 添加异步使用量记录队列系统
- 新增 QueueTelemetryWriter 支持通过内存队列异步记录使用量 - 新增 UsageQueueConsumer 消费者,支持批量处理和流式消费 - 重构 StreamTelemetryRecorder 支持队列写入模式 - UsageService 新增 record_usage_batch 批量记录方法 - 新增 USAGE_QUEUE_ENABLED 和 USAGE_QUEUE_INCLUDE_BODIES 配置项 - 状态值新增 cancelled 类型支持
This commit is contained in:
503
src/services/usage/consumer_streams.py
Normal file
503
src/services/usage/consumer_streams.py
Normal file
@@ -0,0 +1,503 @@
|
||||
"""
|
||||
Usage Redis Streams consumer.
|
||||
|
||||
高性能消费者实现,支持批量处理和单次提交多条记录。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import socket
|
||||
import time
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
from redis.exceptions import ResponseError
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.clients.redis_client import get_redis_client
|
||||
from src.config.settings import config
|
||||
from src.core.logger import logger
|
||||
from src.database.database import create_session
|
||||
from src.services.usage.events import UsageEvent, UsageEventType
|
||||
from src.services.usage.service import UsageService
|
||||
|
||||
|
||||
def _consumer_name() -> str:
|
||||
host = socket.gethostname() or "unknown"
|
||||
return f"{host}:{os.getpid()}"
|
||||
|
||||
|
||||
def _parse_body(value: Any) -> Any:
|
||||
"""将 JSON 字符串 body 反序列化为 dict,否则原样返回。
|
||||
|
||||
QueueTelemetryWriter 会将 body 序列化为 JSON 字符串以便传输,
|
||||
消费者需要将其反序列化回 dict 以正确存入 JSON 列。
|
||||
"""
|
||||
if value is None or isinstance(value, dict):
|
||||
return value
|
||||
if isinstance(value, str):
|
||||
try:
|
||||
return json.loads(value)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
# 解析失败,保留原字符串(可能已被截断)
|
||||
return value
|
||||
return value
|
||||
|
||||
|
||||
def _event_to_record(event: UsageEvent) -> Dict[str, Any]:
|
||||
"""将 UsageEvent 转换为 record_usage_batch 所需的字典格式"""
|
||||
data = event.data
|
||||
status = "completed"
|
||||
if event.event_type == UsageEventType.FAILED:
|
||||
status = "failed"
|
||||
elif event.event_type == UsageEventType.CANCELLED:
|
||||
status = "cancelled"
|
||||
|
||||
return {
|
||||
"request_id": event.request_id,
|
||||
"user_id": data.get("user_id"),
|
||||
"api_key_id": data.get("api_key_id"),
|
||||
"provider": data.get("provider") or "unknown",
|
||||
"model": data.get("model") or "unknown",
|
||||
"input_tokens": data.get("input_tokens") or 0,
|
||||
"output_tokens": data.get("output_tokens") or 0,
|
||||
"cache_creation_input_tokens": data.get("cache_creation_input_tokens") or 0,
|
||||
"cache_read_input_tokens": data.get("cache_read_input_tokens") or 0,
|
||||
"request_type": data.get("request_type") or "chat",
|
||||
"api_format": data.get("api_format"),
|
||||
"endpoint_api_format": data.get("endpoint_api_format"),
|
||||
"has_format_conversion": data.get("has_format_conversion"),
|
||||
"is_stream": data.get("is_stream", True),
|
||||
"response_time_ms": data.get("response_time_ms"),
|
||||
"first_byte_time_ms": data.get("first_byte_time_ms"),
|
||||
"status_code": data.get("status_code") or 200,
|
||||
"error_message": data.get("error_message"),
|
||||
"metadata": data.get("metadata"),
|
||||
"request_headers": data.get("request_headers"),
|
||||
"request_body": _parse_body(data.get("request_body")),
|
||||
"provider_request_headers": data.get("provider_request_headers"),
|
||||
"response_headers": data.get("response_headers"),
|
||||
"client_response_headers": data.get("client_response_headers"),
|
||||
"response_body": _parse_body(data.get("response_body")),
|
||||
"provider_id": data.get("provider_id"),
|
||||
"provider_endpoint_id": data.get("provider_endpoint_id"),
|
||||
"provider_api_key_id": data.get("provider_api_key_id"),
|
||||
"status": status,
|
||||
"target_model": data.get("target_model"),
|
||||
}
|
||||
|
||||
|
||||
async def ensure_usage_stream_group() -> None:
|
||||
redis_client = await get_redis_client(require_redis=False)
|
||||
if not redis_client:
|
||||
return
|
||||
try:
|
||||
await redis_client.xgroup_create(
|
||||
config.usage_queue_stream_key,
|
||||
config.usage_queue_stream_group,
|
||||
id="0-0",
|
||||
mkstream=True,
|
||||
)
|
||||
logger.info(
|
||||
f"[usage-queue] Created consumer group {config.usage_queue_stream_group}"
|
||||
)
|
||||
except ResponseError as exc:
|
||||
if "BUSYGROUP" in str(exc):
|
||||
return
|
||||
raise
|
||||
|
||||
|
||||
class UsageQueueConsumer:
|
||||
"""Usage 队列消费者
|
||||
|
||||
性能优化:
|
||||
- 缓存配置值避免重复属性访问
|
||||
- STREAMING 事件使用 pipeline 批量 ACK
|
||||
- 记录事件批量写入数据库
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._running = False
|
||||
self._task: Optional[asyncio.Task] = None
|
||||
self._consumer = _consumer_name()
|
||||
self._last_claim = 0.0
|
||||
self._last_metrics_log = 0.0
|
||||
# 缓存配置值,避免热路径上的属性访问开销
|
||||
self._stream_key = config.usage_queue_stream_key
|
||||
self._stream_group = config.usage_queue_stream_group
|
||||
self._batch_size = config.usage_queue_consumer_batch
|
||||
self._block_ms = config.usage_queue_consumer_block_ms
|
||||
self._claim_idle_ms = config.usage_queue_claim_idle_ms
|
||||
self._claim_interval = config.usage_queue_claim_interval_seconds
|
||||
self._max_retries = config.usage_queue_max_retries
|
||||
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
|
||||
|
||||
async def start(self) -> None:
|
||||
if self._running:
|
||||
return
|
||||
self._running = True
|
||||
self._task = asyncio.create_task(self._run(), name="usage-queue-consumer")
|
||||
logger.info(f"[usage-queue] Consumer started: {self._consumer}")
|
||||
|
||||
async def stop(self) -> None:
|
||||
if not self._running:
|
||||
return
|
||||
self._running = False
|
||||
if self._task:
|
||||
self._task.cancel()
|
||||
try:
|
||||
await self._task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
logger.info(f"[usage-queue] Consumer stopped: {self._consumer}")
|
||||
|
||||
async def _run(self) -> None:
|
||||
while self._running:
|
||||
try:
|
||||
redis_client = await get_redis_client(require_redis=False)
|
||||
if not redis_client:
|
||||
await asyncio.sleep(1)
|
||||
continue
|
||||
|
||||
await self._maybe_claim_pending(redis_client)
|
||||
await self._read_new(redis_client)
|
||||
await self._log_metrics(redis_client)
|
||||
except asyncio.CancelledError:
|
||||
break
|
||||
except Exception as exc:
|
||||
logger.exception(f"[usage-queue] Consumer loop error: {exc}")
|
||||
await asyncio.sleep(1)
|
||||
|
||||
async def _maybe_claim_pending(self, redis_client) -> None:
|
||||
now = time.time()
|
||||
if now - self._last_claim < self._claim_interval:
|
||||
return
|
||||
self._last_claim = now
|
||||
try:
|
||||
result = await redis_client.xautoclaim(
|
||||
self._stream_key,
|
||||
self._stream_group,
|
||||
self._consumer,
|
||||
min_idle_time=self._claim_idle_ms,
|
||||
start_id="0-0",
|
||||
count=self._batch_size,
|
||||
)
|
||||
except ResponseError as exc:
|
||||
logger.warning(f"[usage-queue] XAUTOCLAIM failed: {exc}")
|
||||
return
|
||||
if not result:
|
||||
return
|
||||
_, messages = result[:2]
|
||||
await self._process_messages(redis_client, messages)
|
||||
|
||||
async def _read_new(self, redis_client) -> None:
|
||||
result = await redis_client.xreadgroup(
|
||||
groupname=self._stream_group,
|
||||
consumername=self._consumer,
|
||||
streams={self._stream_key: ">"},
|
||||
count=self._batch_size,
|
||||
block=self._block_ms,
|
||||
)
|
||||
if not result:
|
||||
return
|
||||
for _stream, messages in result:
|
||||
await self._process_messages(redis_client, messages)
|
||||
|
||||
async def _process_messages(self, redis_client, messages: list) -> None:
|
||||
"""批量处理消息,区分 STREAMING(状态更新)和其他事件(记录写入)"""
|
||||
if not messages:
|
||||
return
|
||||
|
||||
# 分类消息
|
||||
streaming_messages: List[Tuple[str, UsageEvent]] = []
|
||||
record_messages: List[Tuple[str, Dict[str, Any], UsageEvent]] = []
|
||||
failed_messages: List[Tuple[str, Dict[str, Any], Exception]] = []
|
||||
|
||||
for message_id, fields in messages:
|
||||
try:
|
||||
event = UsageEvent.from_stream_fields(fields)
|
||||
if event.event_type == UsageEventType.STREAMING:
|
||||
streaming_messages.append((message_id, event))
|
||||
else:
|
||||
record_messages.append((message_id, fields, event))
|
||||
except Exception as exc:
|
||||
failed_messages.append((message_id, fields, exc))
|
||||
|
||||
# 批量处理 STREAMING 事件(状态更新)
|
||||
if streaming_messages:
|
||||
await self._process_streaming_batch(redis_client, streaming_messages)
|
||||
|
||||
# 批量处理记录事件
|
||||
if record_messages:
|
||||
await self._process_record_batch(redis_client, record_messages)
|
||||
|
||||
# 处理解析失败的消息
|
||||
for message_id, fields, exc in failed_messages:
|
||||
await self._handle_processing_error(redis_client, message_id, fields, exc)
|
||||
|
||||
async def _process_streaming_batch(
|
||||
self,
|
||||
redis_client,
|
||||
messages: List[Tuple[str, UsageEvent]],
|
||||
) -> None:
|
||||
"""批量处理 STREAMING 事件(状态更新)"""
|
||||
success_ids: List[str] = []
|
||||
|
||||
for message_id, event in messages:
|
||||
try:
|
||||
await self._apply_streaming_event(event)
|
||||
success_ids.append(message_id)
|
||||
except Exception as exc:
|
||||
await self._handle_processing_error(redis_client, message_id, {}, exc)
|
||||
|
||||
# 使用 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()
|
||||
|
||||
async def _process_record_batch(
|
||||
self,
|
||||
redis_client,
|
||||
messages: List[Tuple[str, Dict[str, Any], UsageEvent]],
|
||||
) -> None:
|
||||
"""批量处理记录类型的事件"""
|
||||
db = create_session()
|
||||
|
||||
try:
|
||||
# 准备批量记录数据
|
||||
records: List[Dict[str, Any]] = []
|
||||
message_ids: List[str] = []
|
||||
message_fields: List[Dict[str, Any]] = []
|
||||
|
||||
for message_id, fields, event in messages:
|
||||
records.append(_event_to_record(event))
|
||||
message_ids.append(message_id)
|
||||
message_fields.append(fields)
|
||||
|
||||
# 批量写入
|
||||
await UsageService.record_usage_batch(db, 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()
|
||||
|
||||
logger.debug(f"[usage-queue] Batch processed {len(records)} records")
|
||||
|
||||
except Exception as exc:
|
||||
# 批量处理失败,回退到逐条处理(复用已创建的 db session)
|
||||
logger.warning(f"[usage-queue] Batch processing failed, falling back to individual: {exc}")
|
||||
try:
|
||||
db.rollback() # 清理批量失败的事务状态
|
||||
except Exception:
|
||||
pass
|
||||
success_ids: List[str] = []
|
||||
for message_id, fields, event in messages:
|
||||
try:
|
||||
await self._apply_record_event(event, db=db)
|
||||
success_ids.append(message_id)
|
||||
except Exception as individual_exc:
|
||||
await self._handle_processing_error(redis_client, message_id, fields, individual_exc)
|
||||
# 批量 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()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
async def _handle_processing_error(
|
||||
self,
|
||||
redis_client,
|
||||
message_id: str,
|
||||
fields: Dict[str, Any],
|
||||
error: Exception,
|
||||
) -> None:
|
||||
retries = await self._get_delivery_count(redis_client, message_id)
|
||||
if retries >= self._max_retries:
|
||||
try:
|
||||
dlq_fields = dict(fields)
|
||||
dlq_fields["source_id"] = message_id
|
||||
dlq_fields["error"] = str(error)[:200]
|
||||
if self._dlq_maxlen > 0:
|
||||
await redis_client.xadd(
|
||||
self._dlq_key,
|
||||
dlq_fields,
|
||||
maxlen=self._dlq_maxlen,
|
||||
approximate=True,
|
||||
)
|
||||
else:
|
||||
await redis_client.xadd(self._dlq_key, dlq_fields)
|
||||
await redis_client.xack(self._stream_key, self._stream_group, message_id)
|
||||
logger.error(
|
||||
f"[usage-queue] Message moved to DLQ after {retries} attempts: {message_id}"
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.error(f"[usage-queue] Failed to move message to DLQ: {exc}")
|
||||
else:
|
||||
logger.warning(
|
||||
f"[usage-queue] Processing failed (attempt {retries}): {message_id} error={error}"
|
||||
)
|
||||
|
||||
async def _get_delivery_count(self, redis_client, message_id: str) -> int:
|
||||
try:
|
||||
pending = await redis_client.xpending_range(
|
||||
self._stream_key,
|
||||
self._stream_group,
|
||||
min=message_id,
|
||||
max=message_id,
|
||||
count=1,
|
||||
)
|
||||
if not pending:
|
||||
return 0
|
||||
info = pending[0]
|
||||
if isinstance(info, dict):
|
||||
return int(info.get("times_delivered", 0))
|
||||
if isinstance(info, (list, tuple)) and len(info) >= 4:
|
||||
return int(info[3])
|
||||
except Exception:
|
||||
pass
|
||||
return 0
|
||||
|
||||
async def _apply_streaming_event(self, event: UsageEvent) -> None:
|
||||
"""处理 STREAMING 事件(状态更新)"""
|
||||
data = event.data
|
||||
db = create_session()
|
||||
try:
|
||||
UsageService.update_usage_status(
|
||||
db=db,
|
||||
request_id=event.request_id,
|
||||
status="streaming",
|
||||
provider=data.get("provider"),
|
||||
target_model=data.get("target_model"),
|
||||
first_byte_time_ms=data.get("first_byte_time_ms"),
|
||||
provider_id=data.get("provider_id"),
|
||||
provider_endpoint_id=data.get("provider_endpoint_id"),
|
||||
provider_api_key_id=data.get("provider_api_key_id"),
|
||||
api_format=data.get("api_format"),
|
||||
endpoint_api_format=data.get("endpoint_api_format"),
|
||||
has_format_conversion=data.get("has_format_conversion"),
|
||||
)
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
async def _apply_record_event(
|
||||
self, event: UsageEvent, db: Optional[Session] = None
|
||||
) -> None:
|
||||
"""处理记录类型事件(逐条写入,用于 fallback)
|
||||
|
||||
Args:
|
||||
event: 使用事件
|
||||
db: 可选的数据库会话。如果提供,复用该会话;否则创建新会话
|
||||
"""
|
||||
from src.models.database import ApiKey, User
|
||||
|
||||
data = event.data
|
||||
own_session = db is None
|
||||
if own_session:
|
||||
db = create_session()
|
||||
try:
|
||||
status = "completed"
|
||||
if event.event_type == UsageEventType.FAILED:
|
||||
status = "failed"
|
||||
elif event.event_type == UsageEventType.CANCELLED:
|
||||
status = "cancelled"
|
||||
|
||||
user = None
|
||||
api_key = None
|
||||
if data.get("user_id"):
|
||||
user = db.query(User).filter(User.id == data["user_id"]).first()
|
||||
if data.get("api_key_id"):
|
||||
api_key = db.query(ApiKey).filter(ApiKey.id == data["api_key_id"]).first()
|
||||
|
||||
await UsageService.record_usage(
|
||||
db=db,
|
||||
user=user,
|
||||
api_key=api_key,
|
||||
provider=data.get("provider") or "unknown",
|
||||
model=data.get("model") or "unknown",
|
||||
input_tokens=int(data.get("input_tokens") or 0),
|
||||
output_tokens=int(data.get("output_tokens") or 0),
|
||||
cache_creation_input_tokens=int(data.get("cache_creation_input_tokens") or 0),
|
||||
cache_read_input_tokens=int(data.get("cache_read_input_tokens") or 0),
|
||||
request_type=data.get("request_type") or "chat",
|
||||
api_format=data.get("api_format"),
|
||||
endpoint_api_format=data.get("endpoint_api_format"),
|
||||
has_format_conversion=bool(data.get("has_format_conversion") or False),
|
||||
is_stream=bool(data.get("is_stream", True)),
|
||||
response_time_ms=data.get("response_time_ms"),
|
||||
first_byte_time_ms=data.get("first_byte_time_ms"),
|
||||
status_code=int(data.get("status_code") or 200),
|
||||
error_message=data.get("error_message"),
|
||||
metadata=data.get("metadata"),
|
||||
request_headers=data.get("request_headers"),
|
||||
request_body=_parse_body(data.get("request_body")),
|
||||
provider_request_headers=data.get("provider_request_headers"),
|
||||
response_headers=data.get("response_headers"),
|
||||
client_response_headers=data.get("client_response_headers"),
|
||||
response_body=_parse_body(data.get("response_body")),
|
||||
request_id=event.request_id,
|
||||
provider_id=data.get("provider_id"),
|
||||
provider_endpoint_id=data.get("provider_endpoint_id"),
|
||||
provider_api_key_id=data.get("provider_api_key_id"),
|
||||
status=status,
|
||||
target_model=data.get("target_model"),
|
||||
)
|
||||
finally:
|
||||
if own_session:
|
||||
db.close()
|
||||
|
||||
async def _apply_event(self, event: UsageEvent) -> None:
|
||||
"""处理单个事件(兼容旧接口,用于测试)"""
|
||||
if event.event_type == UsageEventType.STREAMING:
|
||||
await self._apply_streaming_event(event)
|
||||
else:
|
||||
await self._apply_record_event(event)
|
||||
|
||||
async def _log_metrics(self, redis_client) -> None:
|
||||
now = time.time()
|
||||
if now - self._last_metrics_log < self._metrics_interval:
|
||||
return
|
||||
self._last_metrics_log = now
|
||||
try:
|
||||
stream_len = await redis_client.xlen(self._stream_key)
|
||||
pending = await redis_client.xpending(self._stream_key, self._stream_group)
|
||||
pending_count = 0
|
||||
if isinstance(pending, dict):
|
||||
pending_count = int(pending.get("pending", 0))
|
||||
elif isinstance(pending, (list, tuple)) and pending:
|
||||
pending_count = int(pending[0])
|
||||
logger.info(
|
||||
f"[usage-queue] backlog={stream_len} pending={pending_count}"
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.debug(f"[usage-queue] metrics log failed: {exc}")
|
||||
|
||||
|
||||
_consumer_instance: Optional[UsageQueueConsumer] = None
|
||||
|
||||
|
||||
async def start_usage_queue_consumer() -> Optional[UsageQueueConsumer]:
|
||||
global _consumer_instance
|
||||
if not config.usage_queue_enabled:
|
||||
return None
|
||||
await ensure_usage_stream_group()
|
||||
if _consumer_instance is None:
|
||||
_consumer_instance = UsageQueueConsumer()
|
||||
await _consumer_instance.start()
|
||||
return _consumer_instance
|
||||
|
||||
|
||||
async def stop_usage_queue_consumer() -> None:
|
||||
global _consumer_instance
|
||||
if _consumer_instance:
|
||||
await _consumer_instance.stop()
|
||||
_consumer_instance = None
|
||||
87
src/services/usage/events.py
Normal file
87
src/services/usage/events.py
Normal file
@@ -0,0 +1,87 @@
|
||||
"""
|
||||
Usage 事件定义与序列化工具(用于 Redis Streams)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from enum import Enum
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
USAGE_EVENT_VERSION = 1
|
||||
|
||||
|
||||
class UsageEventType(str, Enum):
|
||||
STREAMING = "streaming"
|
||||
COMPLETED = "completed"
|
||||
FAILED = "failed"
|
||||
CANCELLED = "cancelled"
|
||||
|
||||
|
||||
def now_ms() -> int:
|
||||
return int(time.time() * 1000)
|
||||
|
||||
|
||||
def _sanitize_value(value: Any) -> Any:
|
||||
if value is None or isinstance(value, (str, int, float, bool)):
|
||||
return value
|
||||
if isinstance(value, dict):
|
||||
return {str(k): _sanitize_value(v) for k, v in value.items()}
|
||||
if isinstance(value, list):
|
||||
return [_sanitize_value(item) for item in value]
|
||||
return str(value)
|
||||
|
||||
|
||||
def sanitize_payload(data: Dict[str, Any]) -> Dict[str, Any]:
|
||||
return {str(k): _sanitize_value(v) for k, v in data.items()}
|
||||
|
||||
|
||||
@dataclass
|
||||
class UsageEvent:
|
||||
event_type: UsageEventType
|
||||
request_id: str
|
||||
timestamp_ms: int
|
||||
data: Dict[str, Any]
|
||||
|
||||
def to_stream_fields(self) -> Dict[str, str]:
|
||||
payload = {
|
||||
"v": USAGE_EVENT_VERSION,
|
||||
"type": self.event_type.value,
|
||||
"request_id": self.request_id,
|
||||
"timestamp_ms": self.timestamp_ms,
|
||||
"data": sanitize_payload(self.data),
|
||||
}
|
||||
return {"payload": json.dumps(payload, ensure_ascii=False)}
|
||||
|
||||
@classmethod
|
||||
def from_stream_fields(cls, fields: Dict[str, Any]) -> "UsageEvent":
|
||||
raw = fields.get("payload")
|
||||
if not raw:
|
||||
raise ValueError("Missing payload field in usage event")
|
||||
if isinstance(raw, bytes):
|
||||
raw = raw.decode("utf-8", errors="ignore")
|
||||
payload = json.loads(raw)
|
||||
event_type = UsageEventType(payload["type"])
|
||||
return cls(
|
||||
event_type=event_type,
|
||||
request_id=payload["request_id"],
|
||||
timestamp_ms=int(payload.get("timestamp_ms", 0)),
|
||||
data=payload.get("data", {}) or {},
|
||||
)
|
||||
|
||||
|
||||
def build_usage_event(
|
||||
*,
|
||||
event_type: UsageEventType,
|
||||
request_id: str,
|
||||
data: Dict[str, Any],
|
||||
timestamp_ms: Optional[int] = None,
|
||||
) -> UsageEvent:
|
||||
return UsageEvent(
|
||||
event_type=event_type,
|
||||
request_id=request_id,
|
||||
timestamp_ms=timestamp_ms or now_ms(),
|
||||
data=data,
|
||||
)
|
||||
@@ -80,7 +80,12 @@ class UsageRecordParams:
|
||||
raise ValueError(f"无效的 HTTP 状态码: {self.status_code}")
|
||||
|
||||
# 状态值校验
|
||||
valid_statuses = {"pending", "streaming", "completed", "failed"}
|
||||
# - pending: 请求已创建,等待处理
|
||||
# - streaming: 流式响应进行中
|
||||
# - completed: 请求成功完成
|
||||
# - failed: 请求失败(上游错误、超时等)
|
||||
# - cancelled: 客户端主动断开连接
|
||||
valid_statuses = {"pending", "streaming", "completed", "failed", "cancelled"}
|
||||
if self.status not in valid_statuses:
|
||||
raise ValueError(f"无效的状态值: {self.status},有效值: {valid_statuses}")
|
||||
|
||||
@@ -1007,6 +1012,216 @@ class UsageService:
|
||||
|
||||
return usage
|
||||
|
||||
@classmethod
|
||||
async def record_usage_batch(
|
||||
cls,
|
||||
db: Session,
|
||||
records: List[Dict[str, Any]],
|
||||
) -> List[Usage]:
|
||||
"""批量记录使用量(高性能版,单次提交多条记录)
|
||||
|
||||
此方法针对高并发场景优化,特点:
|
||||
- 批量插入 Usage 记录,减少 commit 次数
|
||||
- 聚合更新用户/API Key 统计(按 user_id/api_key_id 分组)
|
||||
- 聚合更新 GlobalModel 和 Provider 统计
|
||||
|
||||
Args:
|
||||
db: 数据库会话
|
||||
records: 记录列表,每条记录包含 record_usage 所需的参数
|
||||
|
||||
Returns:
|
||||
创建的 Usage 记录列表
|
||||
"""
|
||||
if not records:
|
||||
return []
|
||||
|
||||
from collections import defaultdict
|
||||
from sqlalchemy import update
|
||||
from src.models.database import ApiKey as ApiKeyModel, User as UserModel, GlobalModel
|
||||
|
||||
usages: List[Usage] = []
|
||||
user_costs: Dict[str, float] = defaultdict(float) # user_id -> total_cost
|
||||
apikey_stats: Dict[str, Dict[str, Any]] = defaultdict(
|
||||
lambda: {"requests": 0, "cost": 0.0, "is_standalone": False}
|
||||
)
|
||||
model_counts: Dict[str, int] = defaultdict(int) # model -> count
|
||||
provider_costs: Dict[str, float] = defaultdict(float) # provider_id -> cost
|
||||
|
||||
# 批量预取 User 和 ApiKey,避免 N+1 查询
|
||||
user_ids = {r.get("user_id") for r in records if r.get("user_id")}
|
||||
api_key_ids = {r.get("api_key_id") for r in records if r.get("api_key_id")}
|
||||
|
||||
users_map: Dict[str, User] = {}
|
||||
if user_ids:
|
||||
users = db.query(User).filter(User.id.in_(user_ids)).all()
|
||||
users_map = {str(u.id): u for u in users}
|
||||
|
||||
api_keys_map: Dict[str, ApiKey] = {}
|
||||
if api_key_ids:
|
||||
api_keys = db.query(ApiKey).filter(ApiKey.id.in_(api_key_ids)).all()
|
||||
api_keys_map = {str(k.id): k for k in api_keys}
|
||||
|
||||
skipped_count = 0
|
||||
total_count = len(records)
|
||||
|
||||
for record in records:
|
||||
try:
|
||||
# 从预取的 map 中获取 user 和 api_key 对象
|
||||
user_id = record.get("user_id")
|
||||
api_key_id = record.get("api_key_id")
|
||||
user = users_map.get(str(user_id)) if user_id else None
|
||||
api_key = api_keys_map.get(str(api_key_id)) if api_key_id else None
|
||||
|
||||
# 准备记录参数
|
||||
request_id = record.get("request_id") or str(uuid.uuid4())[:8]
|
||||
params = UsageRecordParams(
|
||||
db=db,
|
||||
user=user,
|
||||
api_key=api_key,
|
||||
provider=record.get("provider") or "unknown",
|
||||
model=record.get("model") or "unknown",
|
||||
input_tokens=int(record.get("input_tokens") or 0),
|
||||
output_tokens=int(record.get("output_tokens") or 0),
|
||||
cache_creation_input_tokens=int(record.get("cache_creation_input_tokens") or 0),
|
||||
cache_read_input_tokens=int(record.get("cache_read_input_tokens") or 0),
|
||||
request_type=record.get("request_type") or "chat",
|
||||
api_format=record.get("api_format"),
|
||||
endpoint_api_format=record.get("endpoint_api_format"),
|
||||
has_format_conversion=bool(record.get("has_format_conversion")),
|
||||
is_stream=bool(record.get("is_stream", True)),
|
||||
response_time_ms=record.get("response_time_ms"),
|
||||
first_byte_time_ms=record.get("first_byte_time_ms"),
|
||||
status_code=int(record.get("status_code") or 200),
|
||||
error_message=record.get("error_message"),
|
||||
metadata=record.get("metadata"),
|
||||
request_headers=record.get("request_headers"),
|
||||
request_body=record.get("request_body"),
|
||||
provider_request_headers=record.get("provider_request_headers"),
|
||||
response_headers=record.get("response_headers"),
|
||||
client_response_headers=record.get("client_response_headers"),
|
||||
response_body=record.get("response_body"),
|
||||
request_id=request_id,
|
||||
provider_id=record.get("provider_id"),
|
||||
provider_endpoint_id=record.get("provider_endpoint_id"),
|
||||
provider_api_key_id=record.get("provider_api_key_id"),
|
||||
status=record.get("status") or "completed",
|
||||
cache_ttl_minutes=record.get("cache_ttl_minutes"),
|
||||
use_tiered_pricing=record.get("use_tiered_pricing", True),
|
||||
target_model=record.get("target_model"),
|
||||
)
|
||||
|
||||
usage_params, total_cost = await cls._prepare_usage_record(params)
|
||||
|
||||
# 创建 Usage 记录
|
||||
usage = Usage(**usage_params)
|
||||
db.add(usage)
|
||||
usages.append(usage)
|
||||
|
||||
# 聚合统计
|
||||
model_name = record.get("model") or "unknown"
|
||||
model_counts[model_name] += 1
|
||||
|
||||
provider_id = record.get("provider_id")
|
||||
if provider_id:
|
||||
actual_cost = usage_params.get("actual_total_cost_usd", 0)
|
||||
provider_costs[provider_id] += actual_cost
|
||||
|
||||
# 用户统计(独立 Key 不计入创建者)
|
||||
if user and not (api_key and api_key.is_standalone):
|
||||
user_costs[str(user.id)] += total_cost
|
||||
|
||||
# API Key 统计
|
||||
if api_key:
|
||||
key_id = str(api_key.id)
|
||||
apikey_stats[key_id]["requests"] += 1
|
||||
apikey_stats[key_id]["cost"] += total_cost
|
||||
apikey_stats[key_id]["is_standalone"] = api_key.is_standalone
|
||||
|
||||
except Exception as e:
|
||||
skipped_count += 1
|
||||
logger.warning(f"批量记录中跳过无效记录: {e}, request_id={record.get('request_id')}")
|
||||
continue
|
||||
|
||||
# 统计跳过的记录,失败率超过 10% 时提升日志级别
|
||||
if skipped_count > 0:
|
||||
skip_ratio = skipped_count / total_count if total_count > 0 else 0
|
||||
if skip_ratio > 0.1:
|
||||
logger.error(
|
||||
f"批量记录失败率过高: {skipped_count}/{total_count} ({skip_ratio:.1%}) 条记录被跳过"
|
||||
)
|
||||
else:
|
||||
logger.warning(
|
||||
f"批量记录部分失败: {skipped_count}/{total_count} 条记录被跳过"
|
||||
)
|
||||
|
||||
# 批量更新 GlobalModel 使用计数
|
||||
for model_name, count in model_counts.items():
|
||||
db.execute(
|
||||
update(GlobalModel)
|
||||
.where(GlobalModel.name == model_name)
|
||||
.values(usage_count=GlobalModel.usage_count + count)
|
||||
)
|
||||
|
||||
# 批量更新 Provider 月度使用量
|
||||
for provider_id, cost in provider_costs.items():
|
||||
if cost > 0:
|
||||
db.execute(
|
||||
update(Provider)
|
||||
.where(Provider.id == provider_id)
|
||||
.values(monthly_used_usd=Provider.monthly_used_usd + cost)
|
||||
)
|
||||
|
||||
# 批量更新用户使用量
|
||||
from sqlalchemy import func as sql_func
|
||||
for user_id, cost in user_costs.items():
|
||||
if cost > 0:
|
||||
db.execute(
|
||||
update(UserModel)
|
||||
.where(UserModel.id == user_id)
|
||||
.values(
|
||||
used_usd=UserModel.used_usd + cost,
|
||||
total_usd=UserModel.total_usd + cost,
|
||||
updated_at=sql_func.now(),
|
||||
)
|
||||
)
|
||||
|
||||
# 批量更新 API Key 统计
|
||||
for key_id, stats in apikey_stats.items():
|
||||
if stats["is_standalone"]:
|
||||
db.execute(
|
||||
update(ApiKeyModel)
|
||||
.where(ApiKeyModel.id == key_id)
|
||||
.values(
|
||||
total_requests=ApiKeyModel.total_requests + stats["requests"],
|
||||
total_cost_usd=ApiKeyModel.total_cost_usd + stats["cost"],
|
||||
balance_used_usd=ApiKeyModel.balance_used_usd + stats["cost"],
|
||||
last_used_at=sql_func.now(),
|
||||
updated_at=sql_func.now(),
|
||||
)
|
||||
)
|
||||
else:
|
||||
db.execute(
|
||||
update(ApiKeyModel)
|
||||
.where(ApiKeyModel.id == key_id)
|
||||
.values(
|
||||
total_requests=ApiKeyModel.total_requests + stats["requests"],
|
||||
total_cost_usd=ApiKeyModel.total_cost_usd + stats["cost"],
|
||||
last_used_at=sql_func.now(),
|
||||
updated_at=sql_func.now(),
|
||||
)
|
||||
)
|
||||
|
||||
# 单次提交所有更改
|
||||
try:
|
||||
db.commit()
|
||||
logger.debug(f"批量记录 {len(usages)} 条使用记录成功")
|
||||
except Exception as e:
|
||||
logger.error(f"批量提交使用记录时出错: {e}")
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
return usages
|
||||
|
||||
@staticmethod
|
||||
def check_user_quota(
|
||||
db: Session,
|
||||
|
||||
212
src/services/usage/telemetry_writer.py
Normal file
212
src/services/usage/telemetry_writer.py
Normal file
@@ -0,0 +1,212 @@
|
||||
"""
|
||||
Telemetry writer abstraction for stream usage.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from src.api.handlers.base.base_handler import MessageTelemetry
|
||||
from src.clients.redis_client import get_redis_client
|
||||
from src.config.settings import config
|
||||
from src.core.logger import logger
|
||||
from src.services.usage.events import UsageEventType, build_usage_event
|
||||
|
||||
|
||||
class TelemetryWriter(ABC):
|
||||
@abstractmethod
|
||||
async def record_success(self, **kwargs: Any) -> None:
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
async def record_failure(self, **kwargs: Any) -> None:
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
async def record_cancelled(self, **kwargs: Any) -> None:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class DbTelemetryWriter(TelemetryWriter):
|
||||
"""通过 MessageTelemetry 写入数据库的 Writer"""
|
||||
|
||||
# MessageTelemetry 不支持的参数,需要过滤掉
|
||||
# - request_type: MessageTelemetry 内部固定为 "chat",无需外部传入
|
||||
# - metadata: MessageTelemetry 不支持额外元数据字段
|
||||
_IGNORED_KWARGS = frozenset({"request_type", "metadata"})
|
||||
|
||||
def __init__(self, telemetry: MessageTelemetry) -> None:
|
||||
self._telemetry = telemetry
|
||||
|
||||
def _filter_kwargs(self, kwargs: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""过滤掉 MessageTelemetry 不支持的参数"""
|
||||
return {k: v for k, v in kwargs.items() if k not in self._IGNORED_KWARGS}
|
||||
|
||||
async def record_success(self, **kwargs: Any) -> None:
|
||||
await self._telemetry.record_success(**self._filter_kwargs(kwargs))
|
||||
|
||||
async def record_failure(self, **kwargs: Any) -> None:
|
||||
await self._telemetry.record_failure(**self._filter_kwargs(kwargs))
|
||||
|
||||
async def record_cancelled(self, **kwargs: Any) -> None:
|
||||
await self._telemetry.record_cancelled(**self._filter_kwargs(kwargs))
|
||||
|
||||
|
||||
class QueueTelemetryWriter(TelemetryWriter):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
request_id: str,
|
||||
user_id: str,
|
||||
api_key_id: str,
|
||||
) -> None:
|
||||
self.request_id = request_id
|
||||
self.user_id = user_id
|
||||
self.api_key_id = api_key_id
|
||||
|
||||
async def record_success(self, **kwargs: Any) -> None:
|
||||
await self._publish_event(UsageEventType.COMPLETED, **kwargs)
|
||||
|
||||
async def record_failure(self, **kwargs: Any) -> None:
|
||||
await self._publish_event(UsageEventType.FAILED, **kwargs)
|
||||
|
||||
async def record_cancelled(self, **kwargs: Any) -> None:
|
||||
await self._publish_event(UsageEventType.CANCELLED, **kwargs)
|
||||
|
||||
async def _publish_event(self, event_type: UsageEventType, **kwargs: Any) -> None:
|
||||
redis_client = await get_redis_client(require_redis=False)
|
||||
if not redis_client:
|
||||
raise RuntimeError("Redis unavailable for usage queue")
|
||||
|
||||
data = self._build_event_data(**kwargs)
|
||||
event = build_usage_event(
|
||||
event_type=event_type,
|
||||
request_id=self.request_id,
|
||||
data=data,
|
||||
)
|
||||
maxlen = config.usage_queue_stream_maxlen
|
||||
try:
|
||||
if maxlen > 0:
|
||||
await redis_client.xadd(
|
||||
config.usage_queue_stream_key,
|
||||
event.to_stream_fields(),
|
||||
maxlen=maxlen,
|
||||
approximate=True,
|
||||
)
|
||||
else:
|
||||
await redis_client.xadd(config.usage_queue_stream_key, event.to_stream_fields())
|
||||
except Exception as exc:
|
||||
logger.error(f"[usage-queue] XADD failed: {exc}")
|
||||
raise
|
||||
|
||||
def _truncate_body(self, value: Any) -> Optional[str]:
|
||||
"""将 body 序列化为字符串,超长时截断并添加标记"""
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
raw = json.dumps(value, ensure_ascii=False)
|
||||
except TypeError:
|
||||
raw = str(value)
|
||||
max_bytes = config.usage_queue_body_max_bytes
|
||||
if max_bytes > 0 and len(raw) > max_bytes:
|
||||
# 截断并添加标记,预留 15 字符给标记
|
||||
truncate_at = max(0, max_bytes - 15)
|
||||
raw = raw[:truncate_at] + "...[truncated]"
|
||||
return raw
|
||||
|
||||
def _build_event_data(self, **kwargs: Any) -> Dict[str, Any]:
|
||||
# 必需字段
|
||||
data: Dict[str, Any] = {
|
||||
"request_id": self.request_id,
|
||||
"user_id": self.user_id,
|
||||
"api_key_id": self.api_key_id,
|
||||
}
|
||||
|
||||
# 可选字段 - 只添加非 None/非默认值,减少 payload 大小
|
||||
# 注意:消费者端需要处理缺失字段的默认值
|
||||
if kwargs.get("provider"):
|
||||
data["provider"] = kwargs["provider"]
|
||||
if kwargs.get("model"):
|
||||
data["model"] = kwargs["model"]
|
||||
if kwargs.get("target_model"):
|
||||
data["target_model"] = kwargs["target_model"]
|
||||
|
||||
# Token 计数 - 0 是常见值,但仍需传递
|
||||
input_tokens = kwargs.get("input_tokens", 0)
|
||||
output_tokens = kwargs.get("output_tokens", 0)
|
||||
if input_tokens:
|
||||
data["input_tokens"] = input_tokens
|
||||
if output_tokens:
|
||||
data["output_tokens"] = output_tokens
|
||||
|
||||
# 缓存 token(cache_creation_tokens -> cache_creation_input_tokens 映射)
|
||||
cache_creation = kwargs.get("cache_creation_tokens", 0)
|
||||
cache_read = kwargs.get("cache_read_tokens", 0)
|
||||
if cache_creation:
|
||||
data["cache_creation_input_tokens"] = cache_creation
|
||||
if cache_read:
|
||||
data["cache_read_input_tokens"] = cache_read
|
||||
|
||||
# 时间指标
|
||||
if kwargs.get("response_time_ms") is not None:
|
||||
data["response_time_ms"] = kwargs["response_time_ms"]
|
||||
if kwargs.get("first_byte_time_ms") is not None:
|
||||
data["first_byte_time_ms"] = kwargs["first_byte_time_ms"]
|
||||
|
||||
# 状态信息
|
||||
status_code = kwargs.get("status_code", 200)
|
||||
if status_code != 200:
|
||||
data["status_code"] = status_code
|
||||
if kwargs.get("error_message"):
|
||||
data["error_message"] = kwargs["error_message"]
|
||||
|
||||
# 格式信息
|
||||
request_type = kwargs.get("request_type", "chat")
|
||||
if request_type != "chat":
|
||||
data["request_type"] = request_type
|
||||
if kwargs.get("api_format"):
|
||||
data["api_format"] = kwargs["api_format"]
|
||||
if kwargs.get("endpoint_api_format"):
|
||||
data["endpoint_api_format"] = kwargs["endpoint_api_format"]
|
||||
if kwargs.get("has_format_conversion"):
|
||||
data["has_format_conversion"] = True
|
||||
|
||||
# 流式标记 - 默认 True,只记录 False
|
||||
if not kwargs.get("is_stream", True):
|
||||
data["is_stream"] = False
|
||||
|
||||
# Provider 追踪
|
||||
if kwargs.get("provider_id"):
|
||||
data["provider_id"] = kwargs["provider_id"]
|
||||
if kwargs.get("provider_endpoint_id"):
|
||||
data["provider_endpoint_id"] = kwargs["provider_endpoint_id"]
|
||||
if kwargs.get("provider_api_key_id"):
|
||||
data["provider_api_key_id"] = kwargs["provider_api_key_id"]
|
||||
|
||||
# 元数据
|
||||
if kwargs.get("metadata"):
|
||||
data["metadata"] = kwargs["metadata"]
|
||||
|
||||
# 可选:Headers
|
||||
if config.usage_queue_include_headers:
|
||||
if kwargs.get("request_headers"):
|
||||
data["request_headers"] = kwargs["request_headers"]
|
||||
if kwargs.get("provider_request_headers"):
|
||||
data["provider_request_headers"] = kwargs["provider_request_headers"]
|
||||
if kwargs.get("response_headers"):
|
||||
data["response_headers"] = kwargs["response_headers"]
|
||||
if kwargs.get("client_response_headers"):
|
||||
data["client_response_headers"] = kwargs["client_response_headers"]
|
||||
|
||||
# 可选:Bodies
|
||||
if config.usage_queue_include_bodies:
|
||||
request_body = self._truncate_body(kwargs.get("request_body"))
|
||||
response_body = self._truncate_body(kwargs.get("response_body"))
|
||||
if request_body:
|
||||
data["request_body"] = request_body
|
||||
if response_body:
|
||||
data["response_body"] = response_body
|
||||
|
||||
return data
|
||||
Reference in New Issue
Block a user