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

@@ -41,5 +41,4 @@ ADMIN_PASSWORD=admin123456
# CORS 配置(允许跨域的源,多个源用逗号分隔)
# 示例: http://localhost:3000,https://example.com
# 默认: * (允许所有源)
# CORS_ORIGINS=*
# CORS_ORIGINS=*

View File

@@ -20,6 +20,7 @@ from src.config.settings import config
from src.core.logger import logger
from src.database import get_db
from src.models.database import ApiKey, User
from src.services.usage.telemetry_writer import DbTelemetryWriter, QueueTelemetryWriter, TelemetryWriter
class StreamTelemetryRecorder:
@@ -90,60 +91,39 @@ class StreamTelemetryRecorder:
bg_db = next(db_gen)
try:
user = bg_db.query(User).filter(User.id == self.user_id).first()
api_key_obj = bg_db.query(ApiKey).filter(ApiKey.id == self.api_key_id).first()
if not user or not api_key_obj:
logger.warning(
f"[{self.request_id}] User or ApiKey not found, updating status directly"
)
if ctx.is_success():
status = "completed"
elif ctx.is_client_disconnected():
status = "cancelled"
else:
status = "failed"
await self._update_usage_status_directly(
bg_db,
status=status,
response_time_ms=response_time_ms,
status_code=ctx.status_code,
)
writer = await self._get_telemetry_writer(bg_db, ctx, response_time_ms)
if writer is None:
return
bg_telemetry = MessageTelemetry(
bg_db, user, api_key_obj, self.request_id, self.client_ip
)
actual_request_body = ctx.provider_request_body or original_request_body
response_body = ctx.build_response_body(response_time_ms)
response_body = None
if not isinstance(writer, QueueTelemetryWriter) or config.usage_queue_include_bodies:
response_body = ctx.build_response_body(response_time_ms)
if ctx.is_success():
await self._record_success(
bg_telemetry,
ctx,
original_headers,
actual_request_body,
response_body,
response_time_ms,
try:
await self._dispatch_record(
writer, ctx, original_headers, actual_request_body,
response_body, response_time_ms,
)
elif ctx.is_client_disconnected():
await self._record_cancelled(
bg_telemetry,
ctx,
original_headers,
actual_request_body,
response_body,
response_time_ms,
except Exception as writer_error:
if not isinstance(writer, QueueTelemetryWriter):
raise
logger.warning(
f"[{self.request_id}] Queue writer failed, falling back to DB: {writer_error}"
)
else:
await self._record_failure(
bg_telemetry,
ctx,
original_headers,
actual_request_body,
response_body,
response_time_ms,
db_writer = self._build_db_writer(bg_db)
if db_writer is None:
await self._update_usage_status_directly(
bg_db,
status=self._get_status_from_ctx(ctx),
response_time_ms=response_time_ms,
status_code=ctx.status_code,
)
return
if response_body is None:
response_body = ctx.build_response_body(response_time_ms)
await self._dispatch_record(
db_writer, ctx, original_headers, actual_request_body,
response_body, response_time_ms,
)
# 更新候选记录状态
@@ -162,11 +142,11 @@ class StreamTelemetryRecorder:
async def _record_success(
self,
telemetry: MessageTelemetry,
writer: TelemetryWriter,
ctx: StreamContext,
original_headers: Dict[str, str],
actual_request_body: Dict[str, Any],
response_body: Dict[str, Any],
response_body: Optional[Dict[str, Any]],
response_time_ms: int,
) -> None:
"""记录成功的请求"""
@@ -178,7 +158,7 @@ class StreamTelemetryRecorder:
"content-type": "text/event-stream",
})
await telemetry.record_success(
await writer.record_success(
provider=ctx.provider_name or "unknown",
model=ctx.model,
input_tokens=ctx.input_tokens,
@@ -200,6 +180,10 @@ class StreamTelemetryRecorder:
provider_endpoint_id=ctx.endpoint_id,
provider_api_key_id=ctx.key_id,
target_model=ctx.mapped_model,
request_type="chat",
metadata={"stream": True, "content_length": ctx.data_count},
endpoint_api_format=ctx.provider_api_format,
has_format_conversion=ctx.needs_conversion,
)
logger.debug(f"{self.format_id} 流式响应完成")
@@ -207,18 +191,18 @@ class StreamTelemetryRecorder:
async def _record_failure(
self,
telemetry: MessageTelemetry,
writer: TelemetryWriter,
ctx: StreamContext,
original_headers: Dict[str, str],
actual_request_body: Dict[str, Any],
response_body: Dict[str, Any],
response_body: Optional[Dict[str, Any]],
response_time_ms: int,
) -> None:
"""记录失败的请求"""
# 失败时返回给客户端的是 JSON 错误响应,如果没有设置则使用默认值
client_response_headers = ctx.client_response_headers or {"content-type": "application/json"}
await telemetry.record_failure(
await writer.record_failure(
provider=ctx.provider_name or "unknown",
model=ctx.model,
response_time_ms=response_time_ms,
@@ -237,6 +221,10 @@ class StreamTelemetryRecorder:
response_headers=ctx.response_headers,
client_response_headers=client_response_headers,
target_model=ctx.mapped_model,
request_type="chat",
metadata={"stream": True, "content_length": ctx.data_count},
endpoint_api_format=ctx.provider_api_format,
has_format_conversion=ctx.needs_conversion,
)
logger.debug(f"{self.format_id} 流式响应中断")
@@ -246,17 +234,17 @@ class StreamTelemetryRecorder:
async def _record_cancelled(
self,
telemetry: MessageTelemetry,
writer: TelemetryWriter,
ctx: StreamContext,
original_headers: Dict[str, str],
actual_request_body: Dict[str, Any],
response_body: Dict[str, Any],
response_body: Optional[Dict[str, Any]],
response_time_ms: int,
) -> None:
"""记录客户端取消的请求"""
client_response_headers = ctx.client_response_headers or {"content-type": "application/json"}
await telemetry.record_cancelled(
await writer.record_cancelled(
provider=ctx.provider_name or "unknown",
model=ctx.model,
response_time_ms=response_time_ms,
@@ -275,6 +263,10 @@ class StreamTelemetryRecorder:
response_headers=ctx.response_headers,
client_response_headers=client_response_headers,
target_model=ctx.mapped_model,
request_type="chat",
metadata={"stream": True, "content_length": ctx.data_count},
endpoint_api_format=ctx.provider_api_format,
has_format_conversion=ctx.needs_conversion,
)
logger.debug(f"{self.format_id} 流式响应被客户端取消")
@@ -383,3 +375,72 @@ class StreamTelemetryRecorder:
logger.debug(f"[{self.request_id}] Usage 状态已更新: {status}")
except Exception as e:
logger.error(f"[{self.request_id}] 直接更新 Usage 状态失败: {e}")
async def _get_telemetry_writer(
self, bg_db: Session, ctx: StreamContext, response_time_ms: int
) -> Optional[TelemetryWriter]:
if config.usage_queue_enabled and self.user_id and self.api_key_id:
return QueueTelemetryWriter(
request_id=self.request_id,
user_id=self.user_id,
api_key_id=self.api_key_id,
)
db_writer = self._build_db_writer(bg_db)
if db_writer is None:
await self._update_usage_status_directly(
bg_db,
status=self._get_status_from_ctx(ctx),
response_time_ms=response_time_ms,
status_code=ctx.status_code,
)
return None
return db_writer
async def _dispatch_record(
self,
writer: TelemetryWriter,
ctx: StreamContext,
original_headers: Dict[str, str],
actual_request_body: Dict[str, Any],
response_body: Optional[Dict[str, Any]],
response_time_ms: int,
) -> None:
"""根据上下文状态分发到对应的记录方法"""
if ctx.is_success():
await self._record_success(
writer, ctx, original_headers, actual_request_body,
response_body, response_time_ms,
)
elif ctx.is_client_disconnected():
await self._record_cancelled(
writer, ctx, original_headers, actual_request_body,
response_body, response_time_ms,
)
else:
await self._record_failure(
writer, ctx, original_headers, actual_request_body,
response_body, response_time_ms,
)
def _get_status_from_ctx(self, ctx: StreamContext) -> str:
"""根据上下文获取状态字符串"""
if ctx.is_success():
return "completed"
if ctx.is_client_disconnected():
return "cancelled"
return "failed"
def _build_db_writer(self, bg_db: Session) -> Optional[DbTelemetryWriter]:
user = bg_db.query(User).filter(User.id == self.user_id).first()
api_key_obj = bg_db.query(ApiKey).filter(ApiKey.id == self.api_key_id).first()
if not user or not api_key_obj:
logger.warning(
f"[{self.request_id}] User or ApiKey not found, updating status directly"
)
return None
bg_telemetry = MessageTelemetry(
bg_db, user, api_key_obj, self.request_id, self.client_ip
)
return DbTelemetryWriter(bg_telemetry)

View File

@@ -180,6 +180,44 @@ class Config:
self.stream_stats_delay = float(os.getenv("STREAM_STATS_DELAY", "0.1"))
self.stream_first_byte_timeout = float(os.getenv("STREAM_FIRST_BYTE_TIMEOUT", "30.0"))
# Usage 队列配置Redis Streams
# 默认启用队列模式,通过 Redis Streams 异步写入 DB提升响应性能
self.usage_queue_enabled = os.getenv("USAGE_QUEUE_ENABLED", "true").lower() == "true"
# 默认传输 headers/bodies由系统设置request_log_level决定最终存储内容
self.usage_queue_include_headers = (
os.getenv("USAGE_QUEUE_INCLUDE_HEADERS", "true").lower() == "true"
)
self.usage_queue_include_bodies = (
os.getenv("USAGE_QUEUE_INCLUDE_BODIES", "true").lower() == "true"
)
# 0 表示不截断由系统设置max_request/response_body_size统一控制
self.usage_queue_body_max_bytes = int(
os.getenv("USAGE_QUEUE_BODY_MAX_BYTES", "0")
)
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")
)
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")
)
self.usage_queue_consumer_block_ms = int(
os.getenv("USAGE_QUEUE_CONSUMER_BLOCK_MS", "500")
)
self.usage_queue_claim_idle_ms = int(os.getenv("USAGE_QUEUE_CLAIM_IDLE_MS", "30000"))
self.usage_queue_claim_interval_seconds = float(
os.getenv("USAGE_QUEUE_CLAIM_INTERVAL_SECONDS", "5")
)
self.usage_queue_max_retries = int(os.getenv("USAGE_QUEUE_MAX_RETRIES", "2"))
self.usage_queue_metrics_interval_seconds = float(
os.getenv("USAGE_QUEUE_METRICS_INTERVAL_SECONDS", "30")
)
# Thinking 整流器配置
# THINKING_RECTIFIER_ENABLED: 是否启用 Thinking 整流器
# 当遇到跨 Provider 的 thinking 签名错误时,自动整流请求体后重试

View File

@@ -147,6 +147,13 @@ async def lifespan(app: FastAPI):
await init_batch_committer()
logger.info("[OK] 批量提交器已启动,数据库写入性能优化已启用")
# 初始化 Usage 队列消费者(可选)
if config.usage_queue_enabled:
logger.info("初始化 Usage 队列消费者...")
from src.services.usage.consumer_streams import start_usage_queue_consumer
await start_usage_queue_consumer()
# 初始化插件系统
logger.info("初始化插件系统...")
plugin_manager = get_plugin_manager()
@@ -248,6 +255,13 @@ async def lifespan(app: FastAPI):
await shutdown_batch_committer()
logger.info("[OK] 批量提交器已停止,所有待提交数据已保存")
# 停止 Usage 队列消费者
if config.usage_queue_enabled:
logger.info("停止 Usage 队列消费者...")
from src.services.usage.consumer_streams import stop_usage_queue_consumer
await stop_usage_queue_consumer()
# 停止维护调度器
if maintenance_scheduler:
logger.info("停止系统维护调度器...")

View File

@@ -8,9 +8,12 @@
import hashlib
import time
from typing import Optional
from typing import TYPE_CHECKING, Optional
from starlette.requests import Request
if TYPE_CHECKING:
from sqlalchemy.orm import Session
from starlette.types import ASGIApp, Message, Receive, Scope, Send
from src.config import config
@@ -105,6 +108,7 @@ class PluginMiddleware:
if message["type"] == "http.response.start":
response_status_code = message.get("status", 0)
await self._maybe_release_streaming_db_session(request, message)
await send(message)
@@ -154,6 +158,43 @@ class PluginMiddleware:
"body": body,
})
def _finalize_db_session(
self,
db: "Session",
*,
should_commit: bool,
should_rollback: bool,
log_prefix: str = "",
) -> None:
"""统一的数据库会话清理逻辑
Args:
db: SQLAlchemy 会话
should_commit: 是否需要提交(仅当 should_rollback=False 时生效)
should_rollback: 是否需要回滚(优先于 commit
log_prefix: 日志前缀,用于区分调用场景
"""
try:
if should_rollback:
try:
db.rollback()
except Exception as rollback_error:
logger.debug(f"{log_prefix}回滚事务时出错(可忽略): {rollback_error}")
elif should_commit:
try:
db.commit()
except Exception as commit_error:
logger.error(f"{log_prefix}提交失败: {commit_error}")
try:
db.rollback()
except Exception:
pass
finally:
try:
db.close()
except Exception as close_error:
logger.debug(f"{log_prefix}关闭数据库连接时出错(可忽略): {close_error}")
async def _cleanup_db_session(
self, request: Request, exception: Optional[Exception]
) -> None:
@@ -167,37 +208,56 @@ class PluginMiddleware:
"""
from sqlalchemy.orm import Session
if getattr(request.state, "db_released_early", False):
return
if not getattr(request.state, "db_managed_by_middleware", False):
return
db = getattr(request.state, "db", None)
if not isinstance(db, Session):
return
# 检查是否由路由层已经提交
tx_committed_by_route = getattr(request.state, "tx_committed_by_route", False)
self._finalize_db_session(
db,
should_commit=not tx_committed_by_route and exception is None,
should_rollback=exception is not None,
)
try:
if exception is not None:
# 发生异常,回滚事务(无论谁负责提交)
try:
db.rollback()
except Exception as rollback_error:
logger.debug(f"回滚事务时出错(可忽略): {rollback_error}")
elif not tx_committed_by_route:
# 正常完成且路由未自行提交,由中间件提交事务
try:
db.commit()
except Exception as commit_error:
logger.error(f"关键事务提交失败: {commit_error}")
try:
db.rollback()
except Exception:
pass
# 如果 tx_committed_by_route 为 True跳过 commit路由已提交
finally:
# 关闭会话,归还连接到连接池
try:
db.close()
except Exception as close_error:
logger.debug(f"关闭数据库连接时出错(可忽略): {close_error}")
async def _maybe_release_streaming_db_session(
self, request: Request, message: Message
) -> None:
"""在 SSE 响应开始时提前释放请求级 DB session。"""
if getattr(request.state, "db_released_early", False):
return
if not getattr(request.state, "db_managed_by_middleware", False):
return
headers = message.get("headers") or []
content_type = None
for key, value in headers:
if key.lower() == b"content-type":
content_type = value.decode("utf-8", errors="ignore").lower()
break
if not content_type or "text/event-stream" not in content_type:
return
from sqlalchemy.orm import Session
db = getattr(request.state, "db", None)
if not isinstance(db, Session):
return
tx_committed_by_route = getattr(request.state, "tx_committed_by_route", False)
self._finalize_db_session(
db,
should_commit=not tx_committed_by_route,
should_rollback=False,
log_prefix="流式响应提前",
)
request.state.db = None
request.state.db_released_early = True
def _get_client_ip(self, request: Request) -> str:
"""

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

File diff suppressed because it is too large Load Diff