feat: 添加异步使用量记录队列系统

- 新增 QueueTelemetryWriter 支持通过内存队列异步记录使用量
- 新增 UsageQueueConsumer 消费者,支持批量处理和流式消费
- 重构 StreamTelemetryRecorder 支持队列写入模式
- UsageService 新增 record_usage_batch 批量记录方法
- 新增 USAGE_QUEUE_ENABLED 和 USAGE_QUEUE_INCLUDE_BODIES 配置项
- 状态值新增 cancelled 类型支持
This commit is contained in:
fawney19
2026-01-28 09:54:18 +08:00
parent b97e029c82
commit 865be3a6bf
10 changed files with 2536 additions and 87 deletions

View 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

View 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,
)

View File

@@ -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,

View 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
# 缓存 tokencache_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