mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
feat: 添加异步使用量记录队列系统
- 新增 QueueTelemetryWriter 支持通过内存队列异步记录使用量 - 新增 UsageQueueConsumer 消费者,支持批量处理和流式消费 - 重构 StreamTelemetryRecorder 支持队列写入模式 - UsageService 新增 record_usage_batch 批量记录方法 - 新增 USAGE_QUEUE_ENABLED 和 USAGE_QUEUE_INCLUDE_BODIES 配置项 - 状态值新增 cancelled 类型支持
This commit is contained in:
@@ -41,5 +41,4 @@ ADMIN_PASSWORD=admin123456
|
||||
# CORS 配置(允许跨域的源,多个源用逗号分隔)
|
||||
# 示例: http://localhost:3000,https://example.com
|
||||
# 默认: * (允许所有源)
|
||||
# CORS_ORIGINS=*
|
||||
|
||||
# CORS_ORIGINS=*
|
||||
@@ -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)
|
||||
|
||||
@@ -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 签名错误时,自动整流请求体后重试
|
||||
|
||||
14
src/main.py
14
src/main.py
@@ -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("停止系统维护调度器...")
|
||||
|
||||
@@ -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:
|
||||
"""
|
||||
|
||||
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
|
||||
1260
tests/services/test_usage_queue_events.py
Normal file
1260
tests/services/test_usage_queue_events.py
Normal file
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user