feat: 流式空闲超时、健康监控查询优化、限流桶内存上限与维护清理修复

Close #233

Co-authored-by: AAEE86 <ppk0227@hotmail.com>

- cli_monitor_mixin: 引入 STREAM_IDLE_TIMEOUT_SECONDS(可通过环境变量配置),
  流传输开始后若超出空闲窗口无新 chunk 则提前取消并返回 504,避免长时间挂起
- stream_context: 新增 managed_recorded_bodies 上下文管理器,确保 chunks 在
  telemetry 完成后及时释放;stream_telemetry 使用该接口统一管理 response body 构建
- health endpoint: 将状态聚合改为 GROUP BY 直接统计,事件列表按 api_format
  单独查询,避免单次 limit 拉取大量记录导致的遗漏与性能问题;同时过滤不活跃
  provider/endpoint,与公开健康接口保持一致
- endpoint health service: 修正时间线数据按 endpoint_id 而非 key_id 聚合
- token_bucket: 引入 max_buckets/bucket_expiry 上限与定时清理,防止内存无限增长;
  修复 refill_rate=0 时 get_reset_time 除零异常;新增 _is_unlimited_rate_limit 判断
- maintenance_scheduler: 调整清理顺序(先删整行再按窗口清理),新增 newer_than
  边界参数,避免同一行在同一轮中被重复改写
- sync_execute: 新增 create_pending_usage 开关,允许已预创建记录的调用方跳过重复创建
- quota_reader / provider_ops balance: 小幅修复与健壮性提升
- Dockerfile: 添加 MALLOC_ARENA_MAX=2 环境变量以降低 gunicorn worker RSS
- 补充相关测试覆盖
This commit is contained in:
fawney19
2026-03-18 23:38:26 +08:00
parent 3d5b6141a5
commit 1d72a8f9c1
37 changed files with 1787 additions and 607 deletions

View File

@@ -49,6 +49,10 @@ ADMIN_PASSWORD=admin123456
# max-requests-jitter 会自动设置为 MAX_REQUESTS/20 (5%) # max-requests-jitter 会自动设置为 MAX_REQUESTS/20 (5%)
# MAX_REQUESTS=4000 # MAX_REQUESTS=4000
# glibc malloc arena 上限(默认 2
# 降低 malloc 内存碎片,减少 gunicorn worker RSS
# MALLOC_ARENA_MAX=2
# HTTP 连接池上限(默认总预算约 200按 worker 平分) # HTTP 连接池上限(默认总预算约 200按 worker 平分)
# 如果容器内存偏高,可继续下调;例如 2 worker 时设为 80-100 # 如果容器内存偏高,可继续下调;例如 2 worker 时设为 80-100
# HTTP_MAX_CONNECTIONS=100 # HTTP_MAX_CONNECTIONS=100
@@ -74,6 +78,11 @@ ADMIN_PASSWORD=admin123456
# - 需要更多调试上下文4 # - 需要更多调试上下文4
# RESPONSE_CHUNKS_MAX_SIZE_MB=2 # RESPONSE_CHUNKS_MAX_SIZE_MB=2
# 流式空闲超时(单位秒,默认 30
# 当流已经开始但连续一段时间没有任何新 chunk 时,提前中断并返回 504
# 避免一直等到 worker 超时(如 300s
# STREAM_IDLE_TIMEOUT_SECONDS=30
# API Key 前缀(默认 sk # API Key 前缀(默认 sk
# API_KEY_PREFIX=sk # API_KEY_PREFIX=sk

View File

@@ -25,7 +25,12 @@ RUN --mount=type=cache,target=/var/cache/apt,sharing=locked \
nginx \ nginx \
supervisor \ supervisor \
libpq5 \ libpq5 \
curl curl \
libjemalloc2
RUN set -eux; \
jemalloc_path="$(find /usr/lib -type f -name 'libjemalloc.so.2' | head -n1)"; \
[ -n "$jemalloc_path" ]; \
ln -sf "$jemalloc_path" /usr/local/lib/libjemalloc.so.2
# 从 base 镜像复制 Python 包 # 从 base 镜像复制 Python 包
COPY --from=builder /usr/local/lib/python3.13/site-packages /usr/local/lib/python3.13/site-packages COPY --from=builder /usr/local/lib/python3.13/site-packages /usr/local/lib/python3.13/site-packages
# 只复制需要的 Python 可执行文件 # 只复制需要的 Python 可执行文件
@@ -247,7 +252,7 @@ RUN printf '%s\n' \
'stdout_logfile_maxbytes=0' \ 'stdout_logfile_maxbytes=0' \
'stderr_logfile=/dev/stderr' \ 'stderr_logfile=/dev/stderr' \
'stderr_logfile_maxbytes=0' \ 'stderr_logfile_maxbytes=0' \
'environment=PYTHONUNBUFFERED=1,PYTHONIOENCODING=utf-8,LANG=C.UTF-8,LC_ALL=C.UTF-8,DOCKER_CONTAINER=true' \ 'environment=PYTHONUNBUFFERED=1,PYTHONIOENCODING=utf-8,LANG=C.UTF-8,LC_ALL=C.UTF-8,DOCKER_CONTAINER=true,LD_PRELOAD=/usr/local/lib/libjemalloc.so.2,MALLOC_CONF="background_thread:true,dirty_decay_ms:5000,muzzy_decay_ms:5000"' \
'' \ '' \
'[program:tunnel-hub]' \ '[program:tunnel-hub]' \
'command=/usr/local/bin/aether-hub --bind 0.0.0.0:8085' \ 'command=/usr/local/bin/aether-hub --bind 0.0.0.0:8085' \
@@ -268,6 +273,8 @@ ENV PYTHONUNBUFFERED=1 \
PYTHONIOENCODING=utf-8 \ PYTHONIOENCODING=utf-8 \
LANG=C.UTF-8 \ LANG=C.UTF-8 \
LC_ALL=C.UTF-8 \ LC_ALL=C.UTF-8 \
LD_PRELOAD=/usr/local/lib/libjemalloc.so.2 \
MALLOC_CONF=background_thread:true,dirty_decay_ms:5000,muzzy_decay_ms:5000 \
PORT=8084 \ PORT=8084 \
GUNICORN_WORKERS=2 \ GUNICORN_WORKERS=2 \
MAX_REQUESTS=4000 MAX_REQUESTS=4000

View File

@@ -34,7 +34,12 @@ RUN --mount=type=cache,target=/var/cache/apt,sharing=locked \
nginx \ nginx \
supervisor \ supervisor \
libpq5 \ libpq5 \
curl curl \
libjemalloc2
RUN set -eux; \
jemalloc_path="$(find /usr/lib -type f -name 'libjemalloc.so.2' | head -n1)"; \
[ -n "$jemalloc_path" ]; \
ln -sf "$jemalloc_path" /usr/local/lib/libjemalloc.so.2
# 从 base 镜像复制 Python 包 # 从 base 镜像复制 Python 包
COPY --from=builder /usr/local/lib/python3.13/site-packages /usr/local/lib/python3.13/site-packages COPY --from=builder /usr/local/lib/python3.13/site-packages /usr/local/lib/python3.13/site-packages
@@ -270,7 +275,7 @@ RUN printf '%s\n' \
'stdout_logfile_maxbytes=0' \ 'stdout_logfile_maxbytes=0' \
'stderr_logfile=/dev/stderr' \ 'stderr_logfile=/dev/stderr' \
'stderr_logfile_maxbytes=0' \ 'stderr_logfile_maxbytes=0' \
'environment=PYTHONUNBUFFERED=1,PYTHONIOENCODING=utf-8,LANG=C.UTF-8,LC_ALL=C.UTF-8,DOCKER_CONTAINER=true' \ 'environment=PYTHONUNBUFFERED=1,PYTHONIOENCODING=utf-8,LANG=C.UTF-8,LC_ALL=C.UTF-8,DOCKER_CONTAINER=true,LD_PRELOAD=/usr/local/lib/libjemalloc.so.2,MALLOC_CONF="background_thread:true,dirty_decay_ms:5000,muzzy_decay_ms:5000"' \
'' \ '' \
'[program:tunnel-hub]' \ '[program:tunnel-hub]' \
'command=/usr/local/bin/aether-hub --bind 0.0.0.0:8085' \ 'command=/usr/local/bin/aether-hub --bind 0.0.0.0:8085' \
@@ -294,6 +299,8 @@ ENV PYTHONUNBUFFERED=1 \
PYTHONIOENCODING=utf-8 \ PYTHONIOENCODING=utf-8 \
LANG=C.UTF-8 \ LANG=C.UTF-8 \
LC_ALL=C.UTF-8 \ LC_ALL=C.UTF-8 \
LD_PRELOAD=/usr/local/lib/libjemalloc.so.2 \
MALLOC_CONF=background_thread:true,dirty_decay_ms:5000,muzzy_decay_ms:5000 \
PORT=8084 \ PORT=8084 \
GUNICORN_WORKERS=2 \ GUNICORN_WORKERS=2 \
MAX_REQUESTS=4000 MAX_REQUESTS=4000

View File

@@ -84,15 +84,24 @@ export interface CodexResetStatus {
* @param resetSecs 相对剩余秒数(用于 fallback * @param resetSecs 相对剩余秒数(用于 fallback
* @param updatedAt 元数据更新时间Unix 秒) * @param updatedAt 元数据更新时间Unix 秒)
* @param _tick 响应式触发器(传入 tick.value 以触发响应式更新) * @param _tick 响应式触发器(传入 tick.value 以触发响应式更新)
* @param remainingPercent 当前窗口剩余额度百分比0-100100 表示满额不启动倒计时)
*/ */
export function getCodexResetCountdown( export function getCodexResetCountdown(
resetAt: number | null | undefined, resetAt: number | null | undefined,
resetSecs: number | null | undefined, resetSecs: number | null | undefined,
updatedAt: number | null | undefined, updatedAt: number | null | undefined,
_tick: number _tick: number,
remainingPercent?: number | null
): CodexResetStatus | null { ): CodexResetStatus | null {
void _tick void _tick
if (remainingPercent != null) {
const normalizedRemaining = Number(remainingPercent)
if (Number.isFinite(normalizedRemaining) && normalizedRemaining >= 100) {
return null
}
}
const nowSec = Math.floor(Date.now() / 1000) const nowSec = Math.floor(Date.now() / 1000)
let remaining: number let remaining: number

View File

@@ -573,18 +573,20 @@
/> />
</div> </div>
<div <div
v-if="key.upstream_metadata.codex.primary_reset_at || key.upstream_metadata.codex.primary_reset_seconds" v-if="(key.upstream_metadata.codex.primary_reset_at || key.upstream_metadata.codex.primary_reset_seconds) && shouldStartCodexResetCountdown(key.upstream_metadata.codex.primary_used_percent)"
class="text-[9px] mt-0.5 tabular-nums" class="text-[9px] mt-0.5 tabular-nums"
:class="getResetCountdownClass( :class="getResetCountdownClass(
key.upstream_metadata.codex.primary_reset_at, key.upstream_metadata.codex.primary_reset_at,
key.upstream_metadata.codex.primary_reset_seconds, key.upstream_metadata.codex.primary_reset_seconds,
key.upstream_metadata.codex.updated_at key.upstream_metadata.codex.updated_at,
key.upstream_metadata.codex.primary_used_percent
)" )"
> >
{{ getResetCountdownText( {{ getResetCountdownText(
key.upstream_metadata.codex.primary_reset_at, key.upstream_metadata.codex.primary_reset_at,
key.upstream_metadata.codex.primary_reset_seconds, key.upstream_metadata.codex.primary_reset_seconds,
key.upstream_metadata.codex.updated_at key.upstream_metadata.codex.updated_at,
key.upstream_metadata.codex.primary_used_percent
) }} ) }}
</div> </div>
</div> </div>
@@ -604,18 +606,21 @@
/> />
</div> </div>
<div <div
v-if="shouldStartCodexResetCountdown(key.upstream_metadata.codex.secondary_used_percent)"
class="text-[9px] mt-0.5 tabular-nums" class="text-[9px] mt-0.5 tabular-nums"
:class="getResetCountdownClass( :class="getResetCountdownClass(
key.upstream_metadata.codex.secondary_reset_at, key.upstream_metadata.codex.secondary_reset_at,
key.upstream_metadata.codex.secondary_reset_seconds, key.upstream_metadata.codex.secondary_reset_seconds,
key.upstream_metadata.codex.updated_at key.upstream_metadata.codex.updated_at,
key.upstream_metadata.codex.secondary_used_percent
)" )"
> >
<template v-if="key.upstream_metadata.codex.secondary_reset_at || key.upstream_metadata.codex.secondary_reset_seconds"> <template v-if="key.upstream_metadata.codex.secondary_reset_at || key.upstream_metadata.codex.secondary_reset_seconds">
{{ getResetCountdownText( {{ getResetCountdownText(
key.upstream_metadata.codex.secondary_reset_at, key.upstream_metadata.codex.secondary_reset_at,
key.upstream_metadata.codex.secondary_reset_seconds, key.upstream_metadata.codex.secondary_reset_seconds,
key.upstream_metadata.codex.updated_at key.upstream_metadata.codex.updated_at,
key.upstream_metadata.codex.secondary_used_percent
) }} ) }}
</template> </template>
<template v-else> <template v-else>
@@ -2549,9 +2554,16 @@ function getAntigravityQuotaSummary(metadata: UpstreamMetadata | null | undefine
function getResetCountdownText( function getResetCountdownText(
resetAt: number | null | undefined, resetAt: number | null | undefined,
resetSecs: number | null | undefined, resetSecs: number | null | undefined,
updatedAt: number | null | undefined updatedAt: number | null | undefined,
usedPercent: number | null | undefined
): string { ): string {
const status = getCodexResetCountdown(resetAt, resetSecs, updatedAt, countdownTick.value) const status = getCodexResetCountdown(
resetAt,
resetSecs,
updatedAt,
countdownTick.value,
toCodexRemainingPercent(usedPercent)
)
if (!status) return '' if (!status) return ''
return status.isExpired ? status.text : `${status.text} 后重置` return status.isExpired ? status.text : `${status.text} 后重置`
} }
@@ -2559,15 +2571,35 @@ function getResetCountdownText(
function getResetCountdownClass( function getResetCountdownClass(
resetAt: number | null | undefined, resetAt: number | null | undefined,
resetSecs: number | null | undefined, resetSecs: number | null | undefined,
updatedAt: number | null | undefined updatedAt: number | null | undefined,
usedPercent: number | null | undefined
): string { ): string {
const status = getCodexResetCountdown(resetAt, resetSecs, updatedAt, countdownTick.value) const status = getCodexResetCountdown(
resetAt,
resetSecs,
updatedAt,
countdownTick.value,
toCodexRemainingPercent(usedPercent)
)
if (!status || status.isExpired) return 'text-muted-foreground/70' if (!status || status.isExpired) return 'text-muted-foreground/70'
if (status.isCritical) return 'text-destructive font-medium animate-pulse' if (status.isCritical) return 'text-destructive font-medium animate-pulse'
if (status.isUrgent) return 'text-amber-500 dark:text-amber-400' if (status.isUrgent) return 'text-amber-500 dark:text-amber-400'
return 'text-muted-foreground/70' return 'text-muted-foreground/70'
} }
function toCodexRemainingPercent(usedPercent: number | null | undefined): number | null {
const normalizedUsed = Number(usedPercent)
if (!Number.isFinite(normalizedUsed)) return null
const clampedUsed = Math.min(Math.max(normalizedUsed, 0), 100)
return Math.max(100 - clampedUsed, 0)
}
function shouldStartCodexResetCountdown(usedPercent: number | null | undefined): boolean {
const remainingPercent = toCodexRemainingPercent(usedPercent)
if (remainingPercent == null) return true
return remainingPercent < 100
}
// 格式化重置时间 // 格式化重置时间
function formatResetTime(seconds: number): string { function formatResetTime(seconds: number): string {
const days = Math.floor(seconds / 86400) const days = Math.floor(seconds / 86400)

View File

@@ -2428,7 +2428,7 @@ function getQuotaProgressLabel(label: string): string {
function getQuotaProgressCountdown(item: QuotaProgressItem) { function getQuotaProgressCountdown(item: QuotaProgressItem) {
if ((item.label !== '5H' && item.label !== '周') || item.resetAtSeconds == null) return null if ((item.label !== '5H' && item.label !== '周') || item.resetAtSeconds == null) return null
return getCodexResetCountdown(item.resetAtSeconds, null, null, countdownTick.value) return getCodexResetCountdown(item.resetAtSeconds, null, null, countdownTick.value, item.remainingPercent)
} }
function getQuotaProgressCountdownText(item: QuotaProgressItem): string { function getQuotaProgressCountdownText(item: QuotaProgressItem): string {
@@ -2438,6 +2438,9 @@ function getQuotaProgressCountdownText(item: QuotaProgressItem): string {
} }
function getQuotaProgressTooltip(item: QuotaProgressItem): string { function getQuotaProgressTooltip(item: QuotaProgressItem): string {
if ((item.label === '5H' || item.label === '周') && item.remainingPercent >= 100) {
return ''
}
const detail = item.detail?.trim() || '' const detail = item.detail?.trim() || ''
const countdownText = getQuotaProgressCountdownText(item) const countdownText = getQuotaProgressCountdownText(item)
if (countdownText) { if (countdownText) {

View File

@@ -93,6 +93,32 @@ def _format_str(api_format_enum: Any) -> str:
return api_format_enum.value if hasattr(api_format_enum, "value") else str(api_format_enum) return api_format_enum.value if hasattr(api_format_enum, "value") else str(api_format_enum)
def _fetch_recent_attempts_for_api_format(
db: Session,
*,
api_format: str,
since: datetime,
per_format_limit: int,
) -> list[RequestCandidate]:
"""获取单个 API 格式最近的最终态请求,用于事件展示。"""
final_statuses = ["success", "failed", "skipped"]
return (
db.query(RequestCandidate)
.join(ProviderEndpoint, RequestCandidate.endpoint_id == ProviderEndpoint.id)
.join(Provider, ProviderEndpoint.provider_id == Provider.id)
.filter(
ProviderEndpoint.is_active.is_(True),
Provider.is_active.is_(True),
ProviderEndpoint.api_format == api_format,
RequestCandidate.created_at >= since,
RequestCandidate.status.in_(final_statuses),
)
.order_by(RequestCandidate.created_at.desc())
.limit(per_format_limit)
.all()
)
pipeline = get_pipeline() pipeline = get_pipeline()
@@ -379,7 +405,10 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
func.count(RequestCandidate.id).label("count"), func.count(RequestCandidate.id).label("count"),
) )
.join(RequestCandidate, ProviderEndpoint.id == RequestCandidate.endpoint_id) .join(RequestCandidate, ProviderEndpoint.id == RequestCandidate.endpoint_id)
.join(Provider, ProviderEndpoint.provider_id == Provider.id)
.filter( .filter(
ProviderEndpoint.is_active.is_(True),
Provider.is_active.is_(True),
RequestCandidate.created_at >= since, RequestCandidate.created_at >= since,
RequestCandidate.status.in_(final_statuses), RequestCandidate.status.in_(final_statuses),
) )
@@ -395,40 +424,15 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
status_counts[fmt] = {"success": 0, "failed": 0, "skipped": 0} status_counts[fmt] = {"success": 0, "failed": 0, "skipped": 0}
status_counts[fmt][status] = count status_counts[fmt][status] = count
# 3. 获取最近一段时间的 RequestCandidate限制数量 # 3. 为所有活跃格式生成监控数据(包括没有请求记录的
# 使用上面定义的 final_statuses排除中间状态
limit_rows = max(500, self.per_format_limit * 10)
rows = (
db.query(
RequestCandidate,
ProviderEndpoint.api_format,
ProviderEndpoint.provider_id,
)
.join(ProviderEndpoint, RequestCandidate.endpoint_id == ProviderEndpoint.id)
.filter(
RequestCandidate.created_at >= since,
RequestCandidate.status.in_(final_statuses),
)
.order_by(RequestCandidate.created_at.desc())
.limit(limit_rows)
.all()
)
grouped_attempts: dict[str, list[RequestCandidate]] = {}
for attempt, api_format_enum, provider_id in rows:
fmt = _format_str(api_format_enum)
if fmt not in grouped_attempts:
grouped_attempts[fmt] = []
# 只保留每个 API 格式最近 per_format_limit 条记录
if len(grouped_attempts[fmt]) < self.per_format_limit:
grouped_attempts[fmt].append(attempt)
# 4. 为所有活跃格式生成监控数据(包括没有请求记录的)
monitors: list[ApiFormatHealthMonitor] = [] monitors: list[ApiFormatHealthMonitor] = []
for api_format in all_formats: for api_format in all_formats:
attempts = grouped_attempts.get(api_format, []) attempts = _fetch_recent_attempts_for_api_format(
db=db,
api_format=api_format,
since=since,
per_format_limit=self.per_format_limit,
)
# 获取窗口内的真实统计数据 # 获取窗口内的真实统计数据
# 只统计最终状态success, failed, skipped # 只统计最终状态success, failed, skipped
# 中间状态available, pending, used, started不计入统计 # 中间状态available, pending, used, started不计入统计

View File

@@ -485,7 +485,7 @@ class BaseMessageHandler:
api_format: str | None = None, api_format: str | None = None,
request_headers: dict[str, Any] | None = None, request_headers: dict[str, Any] | None = None,
request_body: dict[str, Any] | None = None, request_body: dict[str, Any] | None = None,
) -> None: ) -> bool:
"""在请求开始时创建 pending 状态的 Usage 记录 """在请求开始时创建 pending 状态的 Usage 记录
让前端可以立即看到"处理中"的请求,提升用户体验。 让前端可以立即看到"处理中"的请求,提升用户体验。
@@ -498,6 +498,8 @@ class BaseMessageHandler:
api_format: API 格式 api_format: API 格式
request_headers: 原始请求头 request_headers: 原始请求头
request_body: 原始请求体 request_body: 原始请求体
Returns:
bool: True 表示已成功创建False 表示创建失败(调用方可按需回退处理)。
""" """
try: try:
UsageService.create_pending_usage( UsageService.create_pending_usage(
@@ -512,9 +514,11 @@ class BaseMessageHandler:
request_headers=request_headers, request_headers=request_headers,
request_body=request_body, request_body=request_body,
) )
return True
except Exception as exc: except Exception as exc:
# 创建失败不影响主流程 # 创建失败不影响主流程
logger.warning(f"[{self.request_id}] Failed to create pending usage: {exc}") logger.warning(f"[{self.request_id}] Failed to create pending usage: {exc}")
return False
def _update_usage_to_streaming(self, request_id: str | None = None) -> None: def _update_usage_to_streaming(self, request_id: str | None = None) -> None:
"""更新 Usage 状态为 streaming流式传输开始时调用 """更新 Usage 状态为 streaming流式传输开始时调用

View File

@@ -487,7 +487,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
model = getattr(converted_request, "model", original_request_body.get("model", "unknown")) model = getattr(converted_request, "model", original_request_body.get("model", "unknown"))
# 提前创建 pending 记录,让前端可以立即看到"处理中" # 提前创建 pending 记录,让前端可以立即看到"处理中"
self._create_pending_usage( pending_usage_created = self._create_pending_usage(
model=model, model=model,
is_stream=True, is_stream=True,
request_type="chat", request_type="chat",
@@ -579,6 +579,8 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
request_body_state=request_state, request_body_state=request_state,
request_headers=original_headers, request_headers=original_headers,
request_body=original_request_body, request_body=original_request_body,
# 预创建失败时,回退到 TaskService 侧创建,避免丢失 pending 状态。
create_pending_usage=not pending_usage_created,
) )
stream_generator = exec_result.response stream_generator = exec_result.response
provider_name = exec_result.provider_name or "unknown" provider_name = exec_result.provider_name or "unknown"

View File

@@ -113,7 +113,7 @@ class ChatSyncExecutor:
api_format = handler.allowed_api_formats[0] api_format = handler.allowed_api_formats[0]
# 提前创建 pending 记录,让前端可以立即看到"处理中" # 提前创建 pending 记录,让前端可以立即看到"处理中"
handler._create_pending_usage( pending_usage_created = handler._create_pending_usage(
model=model, model=model,
is_stream=False, is_stream=False,
request_type="chat", request_type="chat",
@@ -176,6 +176,8 @@ class ChatSyncExecutor:
request_body_state=request_state, request_body_state=request_state,
request_headers=original_headers, request_headers=original_headers,
request_body=original_request_body, request_body=original_request_body,
# 预创建失败时,回退到 TaskService 侧创建,避免丢失 pending 状态。
create_pending_usage=not pending_usage_created,
) )
actual_provider_name = exec_result.provider_name or "unknown" actual_provider_name = exec_result.provider_name or "unknown"
ctx.provider_id = exec_result.provider_id ctx.provider_id = exec_result.provider_id

View File

@@ -3,6 +3,7 @@
from __future__ import annotations from __future__ import annotations
import asyncio import asyncio
import os
import time import time
from collections.abc import AsyncGenerator from collections.abc import AsyncGenerator
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any
@@ -29,12 +30,23 @@ if TYPE_CHECKING:
from src.api.handlers.base.cli_protocol import CliHandlerProtocol from src.api.handlers.base.cli_protocol import CliHandlerProtocol
def _read_stream_idle_timeout_seconds() -> float:
raw_value = os.getenv("STREAM_IDLE_TIMEOUT_SECONDS", "30")
try:
parsed = float(raw_value)
except (TypeError, ValueError):
return 30.0
return parsed if parsed > 0 else 30.0
class CliMonitorMixin: class CliMonitorMixin:
"""监控和统计相关方法的 Mixin""" """监控和统计相关方法的 Mixin"""
# CancelledError 归因时,断连检查参数(秒) # CancelledError 归因时,断连检查参数(秒)
CANCEL_DISCONNECT_CHECK_TIMEOUT_SECONDS = 0.5 CANCEL_DISCONNECT_CHECK_TIMEOUT_SECONDS = 0.5
CANCEL_DISCONNECT_RETRY_DELAYS_SECONDS = (0.1, 0.2) CANCEL_DISCONNECT_RETRY_DELAYS_SECONDS = (0.1, 0.2)
# 流式传输过程中若长时间没有任何 chunk提前判定为 idle timeout避免一直等到 worker 超时
STREAM_IDLE_TIMEOUT_SECONDS = _read_stream_idle_timeout_seconds()
async def _probe_client_disconnect( async def _probe_client_disconnect(
self, self,
@@ -115,8 +127,29 @@ class CliMonitorMixin:
last_chunk_time = time_module.time() last_chunk_time = time_module.time()
chunk_count = 0 chunk_count = 0
stream_started = False
idle_timeout_triggered = False
idle_timeout = self.STREAM_IDLE_TIMEOUT_SECONDS
parent_task = asyncio.current_task()
idle_watch_task: asyncio.Task[None] | None = None
async def watch_stream_idle_timeout() -> None:
nonlocal idle_timeout_triggered
if parent_task is None:
return
poll_interval = min(1.0, max(0.1, idle_timeout / 5))
while not ctx.has_completion:
await asyncio.sleep(poll_interval)
if not stream_started:
continue
if (time_module.time() - last_chunk_time) <= idle_timeout:
continue
idle_timeout_triggered = True
parent_task.cancel()
return
try: try:
idle_watch_task = asyncio.create_task(watch_stream_idle_timeout())
if http_request is not None: if http_request is not None:
# 使用后台任务检测断连,完全不阻塞流式传输 # 使用后台任务检测断连,完全不阻塞流式传输
disconnected = False disconnected = False
@@ -149,6 +182,7 @@ class CliMonitorMixin:
ctx.status_code = 499 ctx.status_code = 499
ctx.error_message = "client_disconnected" ctx.error_message = "client_disconnected"
break break
stream_started = True
last_chunk_time = time_module.time() last_chunk_time = time_module.time()
chunk_count += 1 chunk_count += 1
yield chunk yield chunk
@@ -158,14 +192,33 @@ class CliMonitorMixin:
await check_task await check_task
except asyncio.CancelledError: except asyncio.CancelledError:
pass pass
if idle_watch_task is not None:
idle_watch_task.cancel()
try:
await idle_watch_task
except asyncio.CancelledError:
pass
else: else:
# 无 http_request仅被动监控 # 无 http_request仅被动监控
async for chunk in stream_generator: try:
last_chunk_time = time_module.time() async for chunk in stream_generator:
chunk_count += 1 stream_started = True
yield chunk last_chunk_time = time_module.time()
chunk_count += 1
yield chunk
finally:
if idle_watch_task is not None:
idle_watch_task.cancel()
try:
await idle_watch_task
except asyncio.CancelledError:
pass
except asyncio.CancelledError: except asyncio.CancelledError:
# 防御性清理:正常路径中 idle_watch_task 已由内部 finally 取消,
# 但若 CancelledError 在异常路径传播,确保不留孤儿 task。
if idle_watch_task is not None and not idle_watch_task.done():
idle_watch_task.cancel()
# 注意CancelledError 不等于"用户手动取消",它既可能是客户端断连触发, # 注意CancelledError 不等于"用户手动取消",它既可能是客户端断连触发,
# 也可能是服务端(重载/关停/内部取消)导致的协程取消。 # 也可能是服务端(重载/关停/内部取消)导致的协程取消。
# 这里尽量做一次"断连归因":仅当能确认客户端已断开时才记为 499 cancelled。 # 这里尽量做一次"断连归因":仅当能确认客户端已断开时才记为 499 cancelled。
@@ -173,6 +226,28 @@ class CliMonitorMixin:
if not ctx.has_completion: if not ctx.has_completion:
ctx.ensure_estimated_output_tokens() ctx.ensure_estimated_output_tokens()
if not ctx.has_completion and idle_timeout_triggered:
ctx.status_code = 504
ctx.error_message = "stream_idle_timeout"
cancel_origin = "stream_idle_timeout"
logger.warning(
f"ID:{ctx.request_id} | Stream idle timeout: "
f"idle_timeout={idle_timeout:g}s, "
f"chunks={chunk_count}, "
f"has_completion={ctx.has_completion}, "
f"time_since_last_chunk={time_since_last_chunk:.2f}s, "
f"output_tokens={ctx.output_tokens}"
)
ctx.upstream_response = (
f"cancel_origin={cancel_origin}, "
f"chunks={chunk_count}, "
f"has_completion={ctx.has_completion}, "
f"time_since_last_chunk={time_since_last_chunk:.2f}s, "
f"output_tokens={ctx.output_tokens}, "
f"idle_timeout={idle_timeout:g}s"
)
raise
is_client_disconnected = False is_client_disconnected = False
disconnect_check_uncertain = False disconnect_check_uncertain = False
if http_request is not None: if http_request is not None:
@@ -229,10 +304,14 @@ class CliMonitorMixin:
) )
raise raise
except httpx.TimeoutException as e: except httpx.TimeoutException as e:
if idle_watch_task is not None and not idle_watch_task.done():
idle_watch_task.cancel()
ctx.status_code = 504 ctx.status_code = 504
ctx.error_message = str(e) ctx.error_message = str(e)
raise raise
except Exception as e: except Exception as e:
if idle_watch_task is not None and not idle_watch_task.done():
idle_watch_task.cancel()
ctx.status_code = 500 ctx.status_code = 500
ctx.error_message = str(e) ctx.error_message = str(e)
raise raise
@@ -294,59 +373,149 @@ class CliMonitorMixin:
ctx, ctx.provider_request_body or original_request_body ctx, ctx.provider_request_body or original_request_body
) )
response_body = ctx.build_response_body(response_time_ms) with ctx.managed_recorded_bodies(response_time_ms) as recorded_bodies:
client_response_body = ctx.build_client_response_body(response_time_ms) # 根据状态码决定记录成功还是失败
# 499 = 客户端取消(不算系统失败);其他 4xx/5xx 视为失败
if ctx.status_code and ctx.status_code >= 400:
client_response_headers = ctx.client_response_headers or {
"content-type": "application/json"
}
# 根据状态码决定记录成功还是失败 if ctx.is_client_disconnected():
# 499 = 客户端取消(不算系统失败);其他 4xx/5xx 视为失败 # 客户端取消:记录为 cancelled不算系统失败
if ctx.status_code and ctx.status_code >= 400: request_metadata = self._merge_scheduling_metadata(
client_response_headers = ctx.client_response_headers or { {"perf": ctx.perf_metrics} if ctx.perf_metrics else None,
"content-type": "application/json" selected_key_id=ctx.key_id,
} candidate_keys=ctx.candidate_keys,
pool_summary=ctx.pool_summary,
if ctx.is_client_disconnected(): fallback_from_request=True,
# 客户端取消:记录为 cancelled不算系统失败 )
request_metadata = self._merge_scheduling_metadata( await bg_telemetry.record_cancelled(
{"perf": ctx.perf_metrics} if ctx.perf_metrics else None, provider=ctx.provider_name or "unknown",
selected_key_id=ctx.key_id, model=ctx.model,
candidate_keys=ctx.candidate_keys, response_time_ms=response_time_ms,
pool_summary=ctx.pool_summary, first_byte_time_ms=ctx.first_byte_time_ms,
fallback_from_request=True, status_code=ctx.status_code,
) request_headers=original_headers,
await bg_telemetry.record_cancelled( request_body=original_request_body,
provider=ctx.provider_name or "unknown", is_stream=True,
model=ctx.model, api_format=ctx.api_format,
response_time_ms=response_time_ms, api_family=self.api_family,
first_byte_time_ms=ctx.first_byte_time_ms, endpoint_kind=self.endpoint_kind,
status_code=ctx.status_code, provider_request_headers=ctx.provider_request_headers,
request_headers=original_headers, provider_request_body=ctx.provider_request_body,
request_body=original_request_body, input_tokens=ctx.input_tokens,
is_stream=True, output_tokens=ctx.output_tokens,
api_format=ctx.api_format, cache_creation_tokens=ctx.cache_creation_tokens,
api_family=self.api_family, cache_read_tokens=ctx.cached_tokens,
endpoint_kind=self.endpoint_kind, response_body=recorded_bodies.response_body,
provider_request_headers=ctx.provider_request_headers, client_response_body=recorded_bodies.client_response_body,
provider_request_body=ctx.provider_request_body, response_headers=ctx.response_headers,
input_tokens=ctx.input_tokens, client_response_headers=client_response_headers,
output_tokens=ctx.output_tokens, endpoint_api_format=ctx.provider_api_format or None,
cache_creation_tokens=ctx.cache_creation_tokens, has_format_conversion=ctx.has_format_conversion,
cache_read_tokens=ctx.cached_tokens, target_model=ctx.mapped_model,
response_body=response_body, request_metadata=request_metadata,
client_response_body=client_response_body, )
response_headers=ctx.response_headers, logger.debug("{} 流式响应被客户端取消", self.FORMAT_ID)
client_response_headers=client_response_headers, logger.info(
endpoint_api_format=ctx.provider_api_format or None, "[CANCEL] {} | {} | {} | {}ms | {} | in:{} out:{} cache:{}",
has_format_conversion=ctx.has_format_conversion, self.request_id[:8],
target_model=ctx.mapped_model, ctx.model,
request_metadata=request_metadata, ctx.provider_name,
) response_time_ms,
logger.debug("{} 流式响应被客户端取消", self.FORMAT_ID) ctx.status_code,
logger.info( ctx.input_tokens,
f"[CANCEL] {self.request_id[:8]} | {ctx.model} | {ctx.provider_name} | {response_time_ms}ms | " ctx.output_tokens,
f"{ctx.status_code} | in:{ctx.input_tokens} out:{ctx.output_tokens} cache:{ctx.cached_tokens}" ctx.cached_tokens,
) )
else:
# 服务端/上游异常:记录为失败
request_metadata = self._merge_scheduling_metadata(
{"perf": ctx.perf_metrics} if ctx.perf_metrics else None,
selected_key_id=ctx.key_id,
candidate_keys=ctx.candidate_keys,
pool_summary=ctx.pool_summary,
fallback_from_request=True,
)
await bg_telemetry.record_failure(
provider=ctx.provider_name or "unknown",
model=ctx.model,
response_time_ms=response_time_ms,
status_code=ctx.status_code,
error_message=ctx.error_message or f"HTTP {ctx.status_code}",
request_headers=original_headers,
request_body=original_request_body,
is_stream=True,
api_format=ctx.api_format,
api_family=self.api_family,
endpoint_kind=self.endpoint_kind,
provider_request_headers=ctx.provider_request_headers,
provider_request_body=ctx.provider_request_body,
# 预估 token 信息(来自 message_start 事件)
input_tokens=ctx.input_tokens,
output_tokens=ctx.output_tokens,
cache_creation_tokens=ctx.cache_creation_tokens,
cache_read_tokens=ctx.cached_tokens,
response_body=recorded_bodies.response_body,
client_response_body=recorded_bodies.client_response_body,
response_headers=ctx.response_headers,
client_response_headers=client_response_headers,
# 格式转换追踪
endpoint_api_format=ctx.provider_api_format or None,
has_format_conversion=ctx.has_format_conversion,
# 模型映射信息
target_model=ctx.mapped_model,
request_metadata=request_metadata,
)
logger.debug("{} 流式响应中断", self.FORMAT_ID)
logger.info(
"[FAIL] {} | {} | {} | {}ms | {} | in:{} out:{} cache:{}",
self.request_id[:8],
ctx.model,
ctx.provider_name,
response_time_ms,
ctx.status_code,
ctx.input_tokens,
ctx.output_tokens,
ctx.cached_tokens,
)
else: else:
# 服务端/上游异常:记录为失败 # 在记录统计前,允许子类从 parsed_chunks 中提取额外的元数据
self._finalize_stream_metadata(ctx)
# 流式格式转换汇总日志
if ctx.stream_conversion_event_count > 0:
logger.debug(
"[{}] 流式转换完成: {}->{}, total_events={}",
self.request_id[:8],
ctx.provider_api_format,
ctx.client_api_format,
ctx.stream_conversion_event_count,
)
# 流未正常完成(如上游截断/连接中断)且无 token 数据时,
# 从已收集的文本和请求体估算 tokens避免 usage 记录为 0
# 流式成功时,返回给客户端的是提供商响应头 + SSE 必需头
client_response_headers = filter_proxy_response_headers(
ctx.response_headers
)
client_response_headers.update(
{
"Cache-Control": "no-cache, no-transform",
"X-Accel-Buffering": "no",
"content-type": "text/event-stream",
}
)
logger.debug(
"[{}] 开始记录 Usage: provider={}, model={}, in={}, out={}",
ctx.request_id,
ctx.provider_name,
ctx.model,
ctx.input_tokens,
ctx.output_tokens,
)
request_metadata = self._merge_scheduling_metadata( request_metadata = self._merge_scheduling_metadata(
{"perf": ctx.perf_metrics} if ctx.perf_metrics else None, {"perf": ctx.perf_metrics} if ctx.perf_metrics else None,
selected_key_id=ctx.key_id, selected_key_id=ctx.key_id,
@@ -354,129 +523,60 @@ class CliMonitorMixin:
pool_summary=ctx.pool_summary, pool_summary=ctx.pool_summary,
fallback_from_request=True, fallback_from_request=True,
) )
await bg_telemetry.record_failure( total_cost = await bg_telemetry.record_success(
provider=ctx.provider_name or "unknown", provider=ctx.provider_name,
model=ctx.model, model=ctx.model,
input_tokens=ctx.input_tokens,
output_tokens=ctx.output_tokens,
response_time_ms=response_time_ms, response_time_ms=response_time_ms,
first_byte_time_ms=ctx.first_byte_time_ms, # 传递首字时间
status_code=ctx.status_code, status_code=ctx.status_code,
error_message=ctx.error_message or f"HTTP {ctx.status_code}",
request_headers=original_headers, request_headers=original_headers,
request_body=original_request_body, request_body=original_request_body,
response_headers=ctx.response_headers,
client_response_headers=client_response_headers,
response_body=recorded_bodies.response_body,
client_response_body=recorded_bodies.client_response_body,
provider_request_body=ctx.provider_request_body,
cache_creation_tokens=ctx.cache_creation_tokens,
cache_read_tokens=ctx.cached_tokens,
is_stream=True, is_stream=True,
provider_request_headers=ctx.provider_request_headers,
api_format=ctx.api_format, api_format=ctx.api_format,
api_family=self.api_family, api_family=self.api_family,
endpoint_kind=self.endpoint_kind, endpoint_kind=self.endpoint_kind,
provider_request_headers=ctx.provider_request_headers,
provider_request_body=ctx.provider_request_body,
# 预估 token 信息(来自 message_start 事件)
input_tokens=ctx.input_tokens,
output_tokens=ctx.output_tokens,
cache_creation_tokens=ctx.cache_creation_tokens,
cache_read_tokens=ctx.cached_tokens,
response_body=response_body,
client_response_body=client_response_body,
response_headers=ctx.response_headers,
client_response_headers=client_response_headers,
# 格式转换追踪 # 格式转换追踪
endpoint_api_format=ctx.provider_api_format or None, endpoint_api_format=ctx.provider_api_format or None,
has_format_conversion=ctx.has_format_conversion, has_format_conversion=ctx.has_format_conversion,
# Provider 侧追踪信息(用于记录真实成本)
provider_id=ctx.provider_id,
provider_endpoint_id=ctx.endpoint_id,
provider_api_key_id=ctx.key_id,
# 模型映射信息 # 模型映射信息
target_model=ctx.mapped_model, target_model=ctx.mapped_model,
# Provider 响应元数据(如 Gemini 的 modelVersion
response_metadata=(
ctx.response_metadata if ctx.response_metadata else None
),
request_metadata=request_metadata, request_metadata=request_metadata,
) )
logger.debug("{} 流式响应中断", self.FORMAT_ID)
logger.info(
f"[FAIL] {self.request_id[:8]} | {ctx.model} | {ctx.provider_name} | {response_time_ms}ms | "
f"{ctx.status_code} | in:{ctx.input_tokens} out:{ctx.output_tokens} cache:{ctx.cached_tokens}"
)
else:
# 在记录统计前,允许子类从 parsed_chunks 中提取额外的元数据
self._finalize_stream_metadata(ctx)
# 流式格式转换汇总日志
if ctx.stream_conversion_event_count > 0:
logger.debug( logger.debug(
"[{}] 流式转换完成: {}->{}, total_events={}", "[{}] Usage 记录完成: cost=${:.6f}", ctx.request_id, total_cost
self.request_id[:8], )
ctx.provider_api_format, # 简洁的请求完成摘要(两行格式)
ctx.client_api_format, ttfb_part = (
ctx.stream_conversion_event_count, f" | TTFB: {ctx.first_byte_time_ms}ms" if ctx.first_byte_time_ms else ""
)
logger.info(
"[OK] {} | {} | {}{}\n Total: {}ms | in:{} out:{}",
self.request_id[:8],
ctx.model,
ctx.provider_name,
ttfb_part,
response_time_ms,
ctx.input_tokens or 0,
ctx.output_tokens or 0,
) )
# 流未正常完成(如上游截断/连接中断)且无 token 数据时,
# 从已收集的文本和请求体估算 tokens避免 usage 记录为 0
# 流式成功时,返回给客户端的是提供商响应头 + SSE 必需头
client_response_headers = filter_proxy_response_headers(ctx.response_headers)
client_response_headers.update(
{
"Cache-Control": "no-cache, no-transform",
"X-Accel-Buffering": "no",
"content-type": "text/event-stream",
}
)
logger.debug(
f"[{ctx.request_id}] 开始记录 Usage: "
f"provider={ctx.provider_name}, model={ctx.model}, "
f"in={ctx.input_tokens}, out={ctx.output_tokens}"
)
request_metadata = self._merge_scheduling_metadata(
{"perf": ctx.perf_metrics} if ctx.perf_metrics else None,
selected_key_id=ctx.key_id,
candidate_keys=ctx.candidate_keys,
pool_summary=ctx.pool_summary,
fallback_from_request=True,
)
total_cost = await bg_telemetry.record_success(
provider=ctx.provider_name,
model=ctx.model,
input_tokens=ctx.input_tokens,
output_tokens=ctx.output_tokens,
response_time_ms=response_time_ms,
first_byte_time_ms=ctx.first_byte_time_ms, # 传递首字时间
status_code=ctx.status_code,
request_headers=original_headers,
request_body=original_request_body,
response_headers=ctx.response_headers,
client_response_headers=client_response_headers,
response_body=response_body,
client_response_body=client_response_body,
provider_request_body=ctx.provider_request_body,
cache_creation_tokens=ctx.cache_creation_tokens,
cache_read_tokens=ctx.cached_tokens,
is_stream=True,
provider_request_headers=ctx.provider_request_headers,
api_format=ctx.api_format,
api_family=self.api_family,
endpoint_kind=self.endpoint_kind,
# 格式转换追踪
endpoint_api_format=ctx.provider_api_format or None,
has_format_conversion=ctx.has_format_conversion,
# Provider 侧追踪信息(用于记录真实成本)
provider_id=ctx.provider_id,
provider_endpoint_id=ctx.endpoint_id,
provider_api_key_id=ctx.key_id,
# 模型映射信息
target_model=ctx.mapped_model,
# Provider 响应元数据(如 Gemini 的 modelVersion
response_metadata=ctx.response_metadata if ctx.response_metadata else None,
request_metadata=request_metadata,
)
logger.debug("[{}] Usage 记录完成: cost=${:.6f}", ctx.request_id, total_cost)
# 简洁的请求完成摘要(两行格式)
ttfb_part = (
f" | TTFB: {ctx.first_byte_time_ms}ms" if ctx.first_byte_time_ms else ""
)
logger.info(
"[OK] {} | {} | {}{}\n Total: {}ms | in:{} out:{}",
self.request_id[:8],
ctx.model,
ctx.provider_name,
ttfb_part,
response_time_ms,
ctx.input_tokens or 0,
ctx.output_tokens or 0,
)
# 更新候选记录的最终状态和延迟时间 # 更新候选记录的最终状态和延迟时间
# 注意RequestExecutor 会在流开始时过早地标记成功(只记录了连接建立的时间) # 注意RequestExecutor 会在流开始时过早地标记成功(只记录了连接建立的时间)
@@ -556,6 +656,9 @@ class CliMonitorMixin:
except Exception as e: except Exception as e:
logger.exception("记录流式统计信息时出错") logger.exception("记录流式统计信息时出错")
finally:
# 遥测写入完成后主动释放大对象列表,降低高并发长流的内存滞留。
ctx.release_recorded_chunks()
async def _record_stream_failure( async def _record_stream_failure(
self, self,
@@ -591,26 +694,30 @@ class CliMonitorMixin:
pool_summary=ctx.pool_summary, pool_summary=ctx.pool_summary,
fallback_from_request=True, fallback_from_request=True,
) )
await self.telemetry.record_failure( try:
provider=ctx.provider_name or "unknown", await self.telemetry.record_failure(
model=ctx.model, provider=ctx.provider_name or "unknown",
response_time_ms=response_time_ms, model=ctx.model,
status_code=status_code, response_time_ms=response_time_ms,
error_message=extract_client_error_message(error), status_code=status_code,
request_headers=original_headers, error_message=extract_client_error_message(error),
request_body=original_request_body, request_headers=original_headers,
is_stream=True, request_body=original_request_body,
api_format=ctx.api_format, is_stream=True,
api_family=self.api_family, api_format=ctx.api_format,
endpoint_kind=self.endpoint_kind, api_family=self.api_family,
provider_request_headers=ctx.provider_request_headers, endpoint_kind=self.endpoint_kind,
provider_request_body=ctx.provider_request_body, provider_request_headers=ctx.provider_request_headers,
response_headers=ctx.response_headers, provider_request_body=ctx.provider_request_body,
client_response_headers=client_response_headers, response_headers=ctx.response_headers,
# 格式转换追踪 client_response_headers=client_response_headers,
endpoint_api_format=ctx.provider_api_format or None, # 格式转换追踪
has_format_conversion=ctx.has_format_conversion, endpoint_api_format=ctx.provider_api_format or None,
# 模型映射信息 has_format_conversion=ctx.has_format_conversion,
target_model=ctx.mapped_model, # 模型映射信息
request_metadata=request_metadata, target_model=ctx.mapped_model,
) request_metadata=request_metadata,
)
finally:
# 失败路径同样可能持有 chunk 审计数据,及时释放。
ctx.release_recorded_chunks()

View File

@@ -81,7 +81,7 @@ class CliHandlerProtocol(Protocol):
api_format: str | None = ..., api_format: str | None = ...,
request_headers: dict[str, Any] | None = ..., request_headers: dict[str, Any] | None = ...,
request_body: dict[str, Any] | None = ..., request_body: dict[str, Any] | None = ...,
) -> None: ... ) -> bool: ...
def _build_request_metadata( def _build_request_metadata(
self, self,

View File

@@ -92,7 +92,7 @@ class CliStreamMixin:
client_api_format = self.primary_api_format client_api_format = self.primary_api_format
# 提前创建 pending 记录,让前端可以立即看到"处理中" # 提前创建 pending 记录,让前端可以立即看到"处理中"
self._create_pending_usage( pending_usage_created = self._create_pending_usage(
model=model, model=model,
is_stream=True, is_stream=True,
request_type="chat", request_type="chat",
@@ -168,6 +168,8 @@ class CliStreamMixin:
request_body_state=request_state, request_body_state=request_state,
request_headers=original_headers, request_headers=original_headers,
request_body=original_request_body, request_body=original_request_body,
# 预创建失败时,回退到 TaskService 侧创建,避免丢失 pending 状态。
create_pending_usage=not pending_usage_created,
) )
stream_generator = exec_result.response stream_generator = exec_result.response
provider_name = exec_result.provider_name or "unknown" provider_name = exec_result.provider_name or "unknown"
@@ -260,8 +262,7 @@ class CliStreamMixin:
) -> AsyncGenerator[bytes]: ) -> AsyncGenerator[bytes]:
"""执行流式请求并返回流生成器""" """执行流式请求并返回流生成器"""
# 重置上下文状态(重试时清除之前的数据,避免累积) # 重置上下文状态(重试时清除之前的数据,避免累积)
ctx.parsed_chunks = [] ctx.release_recorded_chunks()
ctx.provider_parsed_chunks = []
ctx.chunk_count = 0 ctx.chunk_count = 0
ctx.data_count = 0 ctx.data_count = 0
ctx.has_completion = False ctx.has_completion = False

View File

@@ -75,7 +75,7 @@ class CliSyncMixin:
sync_start_time = time.time() sync_start_time = time.time()
# 提前创建 pending 记录,让前端可以立即看到"处理中" # 提前创建 pending 记录,让前端可以立即看到"处理中"
self._create_pending_usage( pending_usage_created = self._create_pending_usage(
model=model, model=model,
is_stream=False, is_stream=False,
request_type="chat", request_type="chat",
@@ -390,6 +390,8 @@ class CliSyncMixin:
request_body_state=request_state, request_body_state=request_state,
request_headers=original_headers, request_headers=original_headers,
request_body=original_request_body, request_body=original_request_body,
# 预创建失败时,回退到 TaskService 侧创建,避免丢失 pending 状态。
create_pending_usage=not pending_usage_created,
) )
result = exec_result.response result = exec_result.response
actual_provider_name = exec_result.provider_name or "unknown" actual_provider_name = exec_result.provider_name or "unknown"

View File

@@ -12,8 +12,9 @@ from __future__ import annotations
import json import json
import time import time
from contextlib import contextmanager
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any, Iterator
if TYPE_CHECKING: if TYPE_CHECKING:
from src.core.api_format.conversion.stream_state import StreamState from src.core.api_format.conversion.stream_state import StreamState
@@ -49,6 +50,21 @@ def is_format_converted(
_MAX_COLLECTED_TEXT_CHARS = 16 * 1024 _MAX_COLLECTED_TEXT_CHARS = 16 * 1024
@dataclass
class RecordedStreamBodies:
"""统一封装 telemetry/usage 使用的流式响应体引用。"""
response_body: dict[str, Any] | None
client_response_body: dict[str, Any] | None
def ensure_populated(self, ctx: StreamContext, response_time_ms: int) -> None:
"""在 fallback 到需要 body 的路径时按需补建响应体。"""
if self.response_body is None:
self.response_body = ctx.build_response_body(response_time_ms)
if self.client_response_body is None:
self.client_response_body = ctx.build_client_response_body(response_time_ms)
@dataclass @dataclass
class StreamContext: class StreamContext:
""" """
@@ -160,8 +176,7 @@ class StreamContext:
在故障转移重试时调用,清除之前的数据避免累积。 在故障转移重试时调用,清除之前的数据避免累积。
保留 model 和 api_format重置其他所有状态。 保留 model 和 api_format重置其他所有状态。
""" """
self.parsed_chunks = [] self.release_recorded_chunks()
self.provider_parsed_chunks = []
self.chunk_count = 0 self.chunk_count = 0
self.data_count = 0 self.data_count = 0
self.has_completion = False self.has_completion = False
@@ -194,6 +209,32 @@ class StreamContext:
self.needs_conversion = False self.needs_conversion = False
self.selected_base_url = None self.selected_base_url = None
def release_recorded_chunks(self) -> None:
"""释放 telemetry/usage 已消费完的 chunk 列表,避免后台任务继续持有大对象。"""
self.parsed_chunks = []
self.provider_parsed_chunks = []
@contextmanager
def managed_recorded_bodies(
self,
response_time_ms: int,
*,
include_bodies: bool = True,
) -> Iterator[RecordedStreamBodies]:
"""统一管理响应体构建与 chunk 释放,避免 telemetry 路径重复写 finally。"""
recorded_bodies = RecordedStreamBodies(
response_body=self.build_response_body(response_time_ms) if include_bodies else None,
client_response_body=(
self.build_client_response_body(response_time_ms) if include_bodies else None
),
)
try:
yield recorded_bodies
finally:
self.release_recorded_chunks()
recorded_bodies.response_body = None
recorded_bodies.client_response_body = None
@property @property
def collected_text(self) -> str: def collected_text(self) -> str:
"""已收集的文本内容(按需拼接,避免在流式过程中频繁做字符串拷贝)""" """已收集的文本内容(按需拼接,避免在流式过程中频繁做字符串拷贝)"""

View File

@@ -99,6 +99,7 @@ class StreamTelemetryRecorder:
try: try:
writer = await self._get_telemetry_writer(bg_db, ctx, response_time_ms) writer = await self._get_telemetry_writer(bg_db, ctx, response_time_ms)
if writer is None: if writer is None:
ctx.release_recorded_chunks()
return return
# 兜底估算:流未正常完成且 token 均为 0 时,从请求体粗略估算。 # 兜底估算:流未正常完成且 token 均为 0 时,从请求体粗略估算。
# 覆盖成功但缺少 completion以及已传出部分数据后被中断的场景。 # 覆盖成功但缺少 completion以及已传出部分数据后被中断的场景。
@@ -114,53 +115,47 @@ class StreamTelemetryRecorder:
if isinstance(writer, QueueTelemetryWriter) if isinstance(writer, QueueTelemetryWriter)
else should_log_body else should_log_body
) )
response_body = ( with ctx.managed_recorded_bodies(
ctx.build_response_body(response_time_ms) if include_bodies else None response_time_ms, include_bodies=include_bodies
) ) as recorded_bodies:
client_response_body = ( try:
ctx.build_client_response_body(response_time_ms) if include_bodies else None await self._dispatch_record(
)
try:
await self._dispatch_record(
bg_db,
writer,
ctx,
original_headers,
original_request_body,
response_body,
response_time_ms,
client_response_body=client_response_body,
)
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}"
)
db_writer = self._build_db_writer(bg_db)
if db_writer is None:
await self._update_usage_status_directly(
bg_db, bg_db,
status=self._get_status_from_ctx(ctx), writer,
response_time_ms=response_time_ms, ctx,
status_code=ctx.status_code, original_headers,
original_request_body,
recorded_bodies.response_body,
response_time_ms,
client_response_body=recorded_bodies.client_response_body,
)
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}"
)
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 should_log_body:
recorded_bodies.ensure_populated(ctx, response_time_ms)
await self._dispatch_record(
bg_db,
db_writer,
ctx,
original_headers,
original_request_body,
recorded_bodies.response_body,
response_time_ms,
client_response_body=recorded_bodies.client_response_body,
) )
return
if response_body is None and should_log_body:
response_body = ctx.build_response_body(response_time_ms)
if client_response_body is None and should_log_body:
client_response_body = ctx.build_client_response_body(response_time_ms)
await self._dispatch_record(
bg_db,
db_writer,
ctx,
original_headers,
original_request_body,
response_body,
response_time_ms,
client_response_body=client_response_body,
)
# 更新候选记录状态 # 更新候选记录状态
await self._update_candidate_status(bg_db, ctx, response_time_ms, start_time) await self._update_candidate_status(bg_db, ctx, response_time_ms, start_time)
@@ -175,6 +170,9 @@ class StreamTelemetryRecorder:
response_time_ms=response_time_ms, response_time_ms=response_time_ms,
error_message=f"记录统计信息失败: {str(e)[:200]}", error_message=f"记录统计信息失败: {str(e)[:200]}",
) )
finally:
# 遥测写入后主动释放大列表,避免长流式请求对象滞留在 worker 堆中。
ctx.release_recorded_chunks()
async def _record_success( async def _record_success(
self, self,

View File

@@ -47,6 +47,32 @@ router = APIRouter(prefix="/api/public", tags=["System Catalog"])
pipeline = get_pipeline() pipeline = get_pipeline()
def _fetch_recent_public_health_attempts_for_api_format(
db: Session,
*,
api_format: str,
since: datetime,
per_format_limit: int,
) -> list[RequestCandidate]:
"""获取单个 API 格式最近的最终态请求,用于公开监控事件列表。"""
final_statuses = ["success", "failed", "skipped"]
return (
db.query(RequestCandidate)
.join(ProviderEndpoint, RequestCandidate.endpoint_id == ProviderEndpoint.id)
.join(Provider, ProviderEndpoint.provider_id == Provider.id)
.filter(
ProviderEndpoint.is_active.is_(True),
Provider.is_active.is_(True),
ProviderEndpoint.api_format == api_format,
RequestCandidate.created_at >= since,
RequestCandidate.status.in_(final_statuses),
)
.order_by(RequestCandidate.created_at.desc())
.limit(per_format_limit)
.all()
)
@router.get("/site-info") @router.get("/site-info")
def get_site_info(db: Session = Depends(get_db)) -> dict[str, str]: def get_site_info(db: Session = Depends(get_db)) -> dict[str, str]:
"""获取站点基本信息(公开接口,无需认证)""" """获取站点基本信息(公开接口,无需认证)"""
@@ -640,47 +666,51 @@ class PublicApiFormatHealthMonitorAdapter(PublicApiAdapter):
) )
endpoint_map[api_format].append(endpoint_id) endpoint_map[api_format].append(endpoint_id)
# 2. 获取最近一段时间的 RequestCandidate限制数量 # 2. 统计窗口内每个 API 格式的真实状态分布
# 只查询最终状态的记录success, failed, skipped
final_statuses = ["success", "failed", "skipped"] final_statuses = ["success", "failed", "skipped"]
limit_rows = max(500, self.per_format_limit * 10) status_counts_query = (
rows = (
db.query( db.query(
RequestCandidate,
ProviderEndpoint.api_format, ProviderEndpoint.api_format,
RequestCandidate.status,
func.count(RequestCandidate.id).label("count"),
) )
.join(ProviderEndpoint, RequestCandidate.endpoint_id == ProviderEndpoint.id) .join(RequestCandidate, ProviderEndpoint.id == RequestCandidate.endpoint_id)
.join(Provider, ProviderEndpoint.provider_id == Provider.id)
.filter( .filter(
ProviderEndpoint.is_active.is_(True),
Provider.is_active.is_(True),
RequestCandidate.created_at >= since, RequestCandidate.created_at >= since,
RequestCandidate.status.in_(final_statuses), RequestCandidate.status.in_(final_statuses),
) )
.order_by(RequestCandidate.created_at.desc()) .group_by(ProviderEndpoint.api_format, RequestCandidate.status)
.limit(limit_rows)
.all() .all()
) )
grouped_candidates: dict[str, list[RequestCandidate]] = {} status_counts: dict[str, dict[str, int]] = {}
for api_format_enum, status, count in status_counts_query:
for candidate, api_format_enum in rows:
api_format = ( api_format = (
api_format_enum.value if hasattr(api_format_enum, "value") else str(api_format_enum) api_format_enum.value if hasattr(api_format_enum, "value") else str(api_format_enum)
) )
if api_format not in grouped_candidates: if api_format not in status_counts:
grouped_candidates[api_format] = [] status_counts[api_format] = {"success": 0, "failed": 0, "skipped": 0}
status_counts[api_format][status] = count
if len(grouped_candidates[api_format]) < self.per_format_limit:
grouped_candidates[api_format].append(candidate)
# 3. 为所有活跃格式生成监控数据 # 3. 为所有活跃格式生成监控数据
monitors: list[PublicApiFormatHealthMonitor] = [] monitors: list[PublicApiFormatHealthMonitor] = []
for api_format in all_formats: for api_format in all_formats:
candidates = grouped_candidates.get(api_format, []) candidates = _fetch_recent_public_health_attempts_for_api_format(
db=db,
api_format=api_format,
since=since,
per_format_limit=self.per_format_limit,
)
# 统计 # 统计使用窗口内真实总数events 仅保留最近样本用于展示。
success_count = sum(1 for c in candidates if c.status == "success") format_stats = status_counts.get(api_format, {"success": 0, "failed": 0, "skipped": 0})
failed_count = sum(1 for c in candidates if c.status == "failed") success_count = format_stats.get("success", 0)
skipped_count = sum(1 for c in candidates if c.status == "skipped") failed_count = format_stats.get("failed", 0)
total_attempts = len(candidates) skipped_count = format_stats.get("skipped", 0)
total_attempts = success_count + failed_count + skipped_count
# 计算成功率 = success / (success + failed) # 计算成功率 = success / (success + failed)
actual_completed = success_count + failed_count actual_completed = success_count + failed_count

View File

@@ -49,7 +49,8 @@ class PluginManager:
# notification 默认不加载,避免未配置插件(如 email在启动时初始化失败并占用内存。 # notification 默认不加载,避免未配置插件(如 email在启动时初始化失败并占用内存。
DEFAULT_ENABLED_PLUGIN_MODULES: dict[str, tuple[str, ...]] = { DEFAULT_ENABLED_PLUGIN_MODULES: dict[str, tuple[str, ...]] = {
"auth": ("api_key",), "auth": ("api_key",),
"rate_limit": ("sliding_window",), # 默认切到 token_bucket全局 Redis 就绪时会自动切到分布式后端。
"rate_limit": ("token_bucket",),
"cache": ("memory",), "cache": ("memory",),
"monitor": ("prometheus",), "monitor": ("prometheus",),
"token": ("claude",), "token": ("claude",),

View File

@@ -29,6 +29,7 @@ class TokenBucket:
self.refill_rate = refill_rate self.refill_rate = refill_rate
self.tokens = capacity self.tokens = capacity
self.last_refill = time.time() self.last_refill = time.time()
self.last_access_time = self.last_refill
def _refill(self) -> None: def _refill(self) -> None:
"""补充令牌""" """补充令牌"""
@@ -36,6 +37,7 @@ class TokenBucket:
time_passed = now - self.last_refill time_passed = now - self.last_refill
tokens_to_add = time_passed * self.refill_rate tokens_to_add = time_passed * self.refill_rate
self.last_access_time = now
if tokens_to_add > 0: if tokens_to_add > 0:
self.tokens = min(self.capacity, self.tokens + tokens_to_add) self.tokens = min(self.capacity, self.tokens + tokens_to_add)
self.last_refill = now self.last_refill = now
@@ -64,7 +66,7 @@ class TokenBucket:
def get_reset_time(self) -> datetime: def get_reset_time(self) -> datetime:
"""获取下次完全恢复的时间""" """获取下次完全恢复的时间"""
if self.tokens >= self.capacity: if self.tokens >= self.capacity or self.refill_rate <= 0:
return datetime.now(timezone.utc) return datetime.now(timezone.utc)
tokens_needed = self.capacity - self.tokens tokens_needed = self.capacity - self.tokens
@@ -82,6 +84,10 @@ class TokenBucketStrategy(RateLimitStrategy):
- 适合处理不均匀的流量模式 - 适合处理不均匀的流量模式
""" """
DEFAULT_MAX_BUCKETS = 10000
DEFAULT_BUCKET_EXPIRY = 3600
DEFAULT_REDIS_RETRY_INTERVAL = 30.0
def __init__(self) -> None: def __init__(self) -> None:
super().__init__("token_bucket") super().__init__("token_bucket")
self.buckets: dict[str, TokenBucket] = {} self.buckets: dict[str, TokenBucket] = {}
@@ -90,11 +96,46 @@ class TokenBucketStrategy(RateLimitStrategy):
# 默认配置 # 默认配置
self.default_capacity = 100 # 默认桶容量 self.default_capacity = 100 # 默认桶容量
self.default_refill_rate = 10 # 默认每秒补充10个令牌 self.default_refill_rate = 10 # 默认每秒补充10个令牌
self.max_buckets = self.DEFAULT_MAX_BUCKETS
self.bucket_expiry = self.DEFAULT_BUCKET_EXPIRY
self._last_cleanup_time: float = time.time()
self._cleanup_interval = 300 # 每 5 分钟检查一次清理
# 可选的 Redis 后端 # 可选的 Redis 后端
self._redis_backend: RedisTokenBucketBackend | None = None self._redis_backend: RedisTokenBucketBackend | None = None
self._redis_checked = False self._redis_checked = False
self._backend_mode = os.getenv("RATE_LIMIT_BACKEND", "auto").lower() self._backend_mode = os.getenv("RATE_LIMIT_BACKEND", "auto").lower()
self._redis_retry_interval = self.DEFAULT_REDIS_RETRY_INTERVAL
self._next_redis_probe_time = 0.0
@staticmethod
def _is_unlimited_rate_limit(rate_limit: Any) -> bool:
"""显式传入 0/负数时,按“不限流”处理。"""
if rate_limit is None:
return False
try:
return int(rate_limit) <= 0
except (TypeError, ValueError):
return False
def _resolve_bucket_config(self, key: str, rate_limit: int | None = None) -> tuple[int, float]:
"""解析指定 key 当前应使用的桶容量和补充速率。"""
if rate_limit is not None:
normalized_rate_limit = int(rate_limit)
if normalized_rate_limit <= 0:
return 0, 0.0
return normalized_rate_limit, normalized_rate_limit / 60.0
if key.startswith("api_key:"):
return (
self.config.get("api_key_capacity", self.default_capacity),
self.config.get("api_key_refill_rate", self.default_refill_rate),
)
if key.startswith("user:"):
return (
self.config.get("user_capacity", self.default_capacity * 2),
self.config.get("user_refill_rate", self.default_refill_rate * 2),
)
return self.default_capacity, self.default_refill_rate
def _get_bucket(self, key: str, rate_limit: int | None = None) -> TokenBucket: def _get_bucket(self, key: str, rate_limit: int | None = None) -> TokenBucket:
""" """
@@ -107,41 +148,88 @@ class TokenBucketStrategy(RateLimitStrategy):
Returns: Returns:
令牌桶实例 令牌桶实例
""" """
if key not in self.buckets: capacity, refill_rate = self._resolve_bucket_config(key, rate_limit)
# 如果提供了rate_limit参数来自数据库优先使用 bucket = self.buckets.get(key)
if rate_limit is not None: if bucket is None:
# rate_limit 是每分钟请求数,转换为令牌桶参数 bucket = TokenBucket(capacity, refill_rate)
capacity = rate_limit # 桶容量等于每分钟限制 self.buckets[key] = bucket
refill_rate = rate_limit / 60.0 # 每秒补充的令牌数 return bucket
# 否则根据key的不同前缀使用不同的配置
elif key.startswith("api_key:"):
capacity = self.config.get("api_key_capacity", self.default_capacity)
refill_rate = self.config.get("api_key_refill_rate", self.default_refill_rate)
elif key.startswith("user:"):
capacity = self.config.get("user_capacity", self.default_capacity * 2)
refill_rate = self.config.get("user_refill_rate", self.default_refill_rate * 2)
else:
capacity = self.default_capacity
refill_rate = self.default_refill_rate
self.buckets[key] = TokenBucket(capacity, refill_rate) if bucket.capacity != capacity or bucket.refill_rate != refill_rate:
bucket._refill()
bucket.capacity = capacity
bucket.refill_rate = refill_rate
bucket.tokens = min(bucket.tokens, capacity)
return self.buckets[key] return bucket
def _cleanup_expired_buckets(self) -> int:
"""清理长时间未访问的桶,避免 key 集合无限增长。"""
current_time = time.time()
expired_keys = [
key
for key, bucket in self.buckets.items()
if current_time - bucket.last_access_time > self.bucket_expiry
]
for key in expired_keys:
del self.buckets[key]
if expired_keys:
logger.info("清理了 {} 个过期的令牌桶", len(expired_keys))
return len(expired_keys)
def _evict_lru_buckets(self, count: int) -> int:
"""达到容量上限时淘汰最久未使用的桶。"""
if not self.buckets or count <= 0:
return 0
sorted_keys = sorted(self.buckets, key=lambda key: self.buckets[key].last_access_time)
evicted = 0
for key in sorted_keys[:count]:
del self.buckets[key]
evicted += 1
if evicted:
logger.warning("LRU 淘汰了 {} 个令牌桶(达到容量上限)", evicted)
return evicted
async def _maybe_cleanup(self) -> None:
"""定期清理或淘汰桶,控制进程内桶数量。"""
current_time = time.time()
if current_time - self._last_cleanup_time > self._cleanup_interval:
self._cleanup_expired_buckets()
self._last_cleanup_time = current_time
if len(self.buckets) >= self.max_buckets:
evict_count = max(1, self.max_buckets // 10)
self._evict_lru_buckets(evict_count)
def _want_redis_backend(self) -> bool: def _want_redis_backend(self) -> bool:
return self._backend_mode in {"auto", "redis"} return self._backend_mode in {"auto", "redis"}
async def _ensure_backend(self) -> None: async def _ensure_backend(self) -> None:
if self._redis_checked: if self._redis_backend is not None:
return return
self._redis_checked = True
if not self._want_redis_backend(): if not self._want_redis_backend():
self._redis_checked = True
return
current_time = time.time()
if self._redis_checked and current_time < self._next_redis_probe_time:
return return
redis_client = get_redis_client_sync() redis_client = get_redis_client_sync()
if redis_client: if redis_client:
self._redis_backend = RedisTokenBucketBackend(redis_client) self._redis_backend = RedisTokenBucketBackend(redis_client)
self._redis_checked = True
self._next_redis_probe_time = 0.0
logger.info("速率限制改用 Redis 令牌桶后端") logger.info("速率限制改用 Redis 令牌桶后端")
elif self._backend_mode == "redis": return
self._redis_checked = True
self._next_redis_probe_time = current_time + self._redis_retry_interval
if self._backend_mode == "redis":
logger.warning("RATE_LIMIT_BACKEND=redis 但 Redis 客户端不可用,回退到内存桶") logger.warning("RATE_LIMIT_BACKEND=redis 但 Redis 客户端不可用,回退到内存桶")
async def check_limit(self, key: str, **kwargs: Any) -> RateLimitResult: async def check_limit(self, key: str, **kwargs: Any) -> RateLimitResult:
@@ -155,11 +243,14 @@ class TokenBucketStrategy(RateLimitStrategy):
Returns: Returns:
速率限制检查结果 速率限制检查结果
""" """
await self._ensure_backend()
rate_limit = kwargs.get("rate_limit") rate_limit = kwargs.get("rate_limit")
amount = kwargs.get("amount", 1) amount = kwargs.get("amount", 1)
if self._is_unlimited_rate_limit(rate_limit):
return RateLimitResult(allowed=True, remaining=0)
await self._ensure_backend()
if self._redis_backend: if self._redis_backend:
return await self._redis_backend.peek( return await self._redis_backend.peek(
key=key, key=key,
@@ -169,6 +260,7 @@ class TokenBucketStrategy(RateLimitStrategy):
) )
async with self._lock: async with self._lock:
await self._maybe_cleanup()
bucket = self._get_bucket(key, rate_limit) bucket = self._get_bucket(key, rate_limit)
remaining = bucket.get_remaining() remaining = bucket.get_remaining()
reset_at = bucket.get_reset_time() reset_at = bucket.get_reset_time()
@@ -203,13 +295,17 @@ class TokenBucketStrategy(RateLimitStrategy):
Returns: Returns:
是否成功消费 是否成功消费
""" """
rate_limit = kwargs.get("rate_limit")
if self._is_unlimited_rate_limit(rate_limit):
return True
await self._ensure_backend() await self._ensure_backend()
if self._redis_backend: if self._redis_backend:
success, remaining = await self._redis_backend.consume( success, remaining = await self._redis_backend.consume(
key=key, key=key,
capacity=self._resolve_capacity(key, kwargs.get("rate_limit")), capacity=self._resolve_capacity(key, rate_limit),
refill_rate=self._resolve_refill_rate(key, kwargs.get("rate_limit")), refill_rate=self._resolve_refill_rate(key, rate_limit),
amount=amount, amount=amount,
) )
if success: if success:
@@ -219,7 +315,8 @@ class TokenBucketStrategy(RateLimitStrategy):
return success return success
async with self._lock: async with self._lock:
bucket = self._get_bucket(key) await self._maybe_cleanup()
bucket = self._get_bucket(key, rate_limit)
success = bucket.consume(amount) success = bucket.consume(amount)
if success: if success:
@@ -270,6 +367,7 @@ class TokenBucketStrategy(RateLimitStrategy):
) )
async with self._lock: async with self._lock:
await self._maybe_cleanup()
bucket = self._get_bucket(key) bucket = self._get_bucket(key)
return { return {
"strategy": "token_bucket", "strategy": "token_bucket",
@@ -293,24 +391,17 @@ class TokenBucketStrategy(RateLimitStrategy):
super().configure(config) super().configure(config)
self.default_capacity = config.get("default_capacity", self.default_capacity) self.default_capacity = config.get("default_capacity", self.default_capacity)
self.default_refill_rate = config.get("default_refill_rate", self.default_refill_rate) self.default_refill_rate = config.get("default_refill_rate", self.default_refill_rate)
self.max_buckets = int(config.get("max_buckets", self.max_buckets))
self.bucket_expiry = int(config.get("bucket_expiry", self.bucket_expiry))
self._cleanup_interval = int(config.get("cleanup_interval", self._cleanup_interval))
def _resolve_capacity(self, key: str, rate_limit: int | None = None) -> int: def _resolve_capacity(self, key: str, rate_limit: int | None = None) -> int:
if rate_limit is not None: capacity, _ = self._resolve_bucket_config(key, rate_limit)
return rate_limit return capacity
if key.startswith("api_key:"):
return self.config.get("api_key_capacity", self.default_capacity)
if key.startswith("user:"):
return self.config.get("user_capacity", self.default_capacity * 2)
return self.default_capacity
def _resolve_refill_rate(self, key: str, rate_limit: int | None = None) -> float: def _resolve_refill_rate(self, key: str, rate_limit: int | None = None) -> float:
if rate_limit is not None: _, refill_rate = self._resolve_bucket_config(key, rate_limit)
return rate_limit / 60.0 return refill_rate
if key.startswith("api_key:"):
return self.config.get("api_key_refill_rate", self.default_refill_rate)
if key.startswith("user:"):
return self.config.get("user_refill_rate", self.default_refill_rate * 2)
return self.default_refill_rate
class RedisTokenBucketBackend: class RedisTokenBucketBackend:
@@ -365,6 +456,9 @@ class RedisTokenBucketBackend:
refill_rate: float, refill_rate: float,
amount: int, amount: int,
) -> RateLimitResult: ) -> RateLimitResult:
if capacity <= 0 or refill_rate <= 0:
return RateLimitResult(allowed=True, remaining=0)
bucket_key = self._redis_key(key) bucket_key = self._redis_key(key)
data = await self.redis.hmget(bucket_key, "tokens", "timestamp") data = await self.redis.hmget(bucket_key, "tokens", "timestamp")
tokens = data[0] tokens = data[0]
@@ -372,7 +466,7 @@ class RedisTokenBucketBackend:
if tokens is None or last_refill is None: if tokens is None or last_refill is None:
remaining = capacity remaining = capacity
reset_at = datetime.now(timezone.utc) + timedelta(seconds=capacity / refill_rate) reset_at = datetime.now(timezone.utc)
else: else:
tokens_value = float(tokens) tokens_value = float(tokens)
last_refill_value = float(last_refill) last_refill_value = float(last_refill)
@@ -407,6 +501,9 @@ class RedisTokenBucketBackend:
refill_rate: float, refill_rate: float,
amount: int, amount: int,
) -> tuple[bool, int]: ) -> tuple[bool, int]:
if capacity <= 0 or refill_rate <= 0:
return True, 0
result = await self._consume_script( result = await self._consume_script(
keys=[self._redis_key(key)], keys=[self._redis_key(key)],
args=[time.time(), capacity, refill_rate, amount], args=[time.time(), capacity, refill_rate, amount],

View File

@@ -69,7 +69,13 @@ class EndpointHealthService:
# 查询所有活跃的端点(一次性获取所有需要的数据) # 查询所有活跃的端点(一次性获取所有需要的数据)
endpoints = ( endpoints = (
db.query(ProviderEndpoint).join(Provider).filter(Provider.is_active.is_(True)).all() db.query(ProviderEndpoint)
.join(Provider)
.filter(
ProviderEndpoint.is_active.is_(True),
Provider.is_active.is_(True),
)
.all()
) )
# 收集所有 provider_ids # 收集所有 provider_ids
@@ -155,16 +161,13 @@ class EndpointHealthService:
format_stats[api_format]["health_scores"].append(health_score) format_stats[api_format]["health_scores"].append(health_score)
# 批量生成所有格式的时间线数据 # 批量生成所有格式的时间线数据
all_key_ids = [] format_endpoint_mapping: dict[str, list[str]] = {}
format_key_mapping: dict[str, list[str]] = {}
for api_format, stats in format_stats.items(): for api_format, stats in format_stats.items():
key_ids = stats["key_ids"] format_endpoint_mapping[api_format] = stats["endpoint_ids"]
format_key_mapping[api_format] = key_ids
all_key_ids.extend(key_ids)
# 一次性查询所有时间线数据 # 一次性查询所有时间线数据
timeline_data_map = EndpointHealthService._generate_timeline_batch( timeline_data_map = EndpointHealthService._generate_timeline_batch(
db, format_key_mapping, now, lookback_hours db, format_endpoint_mapping, now, lookback_hours
) )
# 生成结果 # 生成结果
@@ -234,7 +237,7 @@ class EndpointHealthService:
@staticmethod @staticmethod
def _generate_timeline_batch( def _generate_timeline_batch(
db: Session, db: Session,
format_key_mapping: dict[str, list[str]], format_endpoint_mapping: dict[str, list[str]],
now: datetime, now: datetime,
lookback_hours: int, lookback_hours: int,
segments: int = 100, segments: int = 100,
@@ -249,7 +252,7 @@ class EndpointHealthService:
Args: Args:
db: 数据库会话 db: 数据库会话
format_key_mapping: API格式 -> key_ids 的映射 format_endpoint_mapping: API格式 -> endpoint_ids 的映射
now: 当前时间 now: 当前时间
lookback_hours: 回溯小时数 lookback_hours: 回溯小时数
segments: 时间段数量 segments: 时间段数量
@@ -257,19 +260,19 @@ class EndpointHealthService:
Returns: Returns:
API格式 -> 时间线数据的映射 API格式 -> 时间线数据的映射
""" """
# 收集所有 key_ids # 收集所有 endpoint_ids
all_key_ids = [] all_endpoint_ids = []
for key_ids in format_key_mapping.values(): for endpoint_ids in format_endpoint_mapping.values():
all_key_ids.extend(key_ids) all_endpoint_ids.extend(endpoint_ids)
if not all_key_ids: if not all_endpoint_ids:
return { return {
api_format: { api_format: {
"timeline": ["unknown"] * 100, "timeline": ["unknown"] * 100,
"time_range_start": None, "time_range_start": None,
"time_range_end": None, "time_range_end": None,
} }
for api_format in format_key_mapping.keys() for api_format in format_endpoint_mapping.keys()
} }
# 参数校验API 层已通过 Query(ge=1) 保证,这里做防御性检查) # 参数校验API 层已通过 Query(ge=1) 保证,这里做防御性检查)
@@ -293,7 +296,7 @@ class EndpointHealthService:
candidate_stats = ( candidate_stats = (
db.query( db.query(
RequestCandidate.key_id, RequestCandidate.endpoint_id,
segment_expr, segment_expr,
func.count(RequestCandidate.id).label("total_count"), func.count(RequestCandidate.id).label("total_count"),
func.sum(case((RequestCandidate.status == "success", 1), else_=0)).label( func.sum(case((RequestCandidate.status == "success", 1), else_=0)).label(
@@ -306,20 +309,20 @@ class EndpointHealthService:
func.max(RequestCandidate.created_at).label("max_time"), func.max(RequestCandidate.created_at).label("max_time"),
) )
.filter( .filter(
RequestCandidate.key_id.in_(all_key_ids), RequestCandidate.endpoint_id.in_(all_endpoint_ids),
RequestCandidate.created_at >= start_time, RequestCandidate.created_at >= start_time,
RequestCandidate.created_at <= now, RequestCandidate.created_at <= now,
RequestCandidate.status.in_(final_statuses), RequestCandidate.status.in_(final_statuses),
) )
.group_by(RequestCandidate.key_id, segment_expr) .group_by(RequestCandidate.endpoint_id, segment_expr)
.all() .all()
) )
# 构建 key_id -> api_format 的反向映射 # 构建 endpoint_id -> api_format 的反向映射
key_to_format: dict[str, str] = {} endpoint_to_format: dict[str, str] = {}
for api_format, key_ids in format_key_mapping.items(): for api_format, endpoint_ids in format_endpoint_mapping.items():
for key_id in key_ids: for endpoint_id in endpoint_ids:
key_to_format[key_id] = api_format endpoint_to_format[endpoint_id] = api_format
# 按 api_format 和 segment 聚合数据 # 按 api_format 和 segment 聚合数据
format_segment_data: dict[str, dict[int, dict]] = defaultdict( format_segment_data: dict[str, dict[int, dict]] = defaultdict(
@@ -335,9 +338,9 @@ class EndpointHealthService:
) )
for row in candidate_stats: for row in candidate_stats:
key_id = row.key_id endpoint_id = row.endpoint_id
segment_idx = int(row.segment_idx) if row.segment_idx is not None else 0 segment_idx = int(row.segment_idx) if row.segment_idx is not None else 0
api_format = key_to_format.get(key_id) api_format = endpoint_to_format.get(endpoint_id)
if api_format and 0 <= segment_idx < segments: if api_format and 0 <= segment_idx < segments:
seg_data = format_segment_data[api_format][segment_idx] seg_data = format_segment_data[api_format][segment_idx]
@@ -355,7 +358,7 @@ class EndpointHealthService:
# 生成各格式的时间线 # 生成各格式的时间线
result: dict[str, dict[str, Any]] = {} result: dict[str, dict[str, Any]] = {}
for api_format in format_key_mapping.keys(): for api_format in format_endpoint_mapping.keys():
timeline = [] timeline = []
earliest_time = None earliest_time = None
latest_time = None latest_time = None
@@ -427,46 +430,10 @@ class EndpointHealthService:
"time_range_end": None, "time_range_end": None,
} }
# 基于 endpoint_ids 反推 provider_ids 与 api_format再选出支持该格式的 keys # 直接按 endpoint_id 聚合,避免共享 Key 时把不同 API 格式串桶。
endpoint_rows = ( format_endpoint_mapping = {"_single": endpoint_ids}
db.query(ProviderEndpoint.provider_id, ProviderEndpoint.api_format)
.filter(ProviderEndpoint.id.in_(endpoint_ids))
.all()
)
if not endpoint_rows:
return {
"timeline": ["unknown"] * 100,
"time_range_start": None,
"time_range_end": None,
}
provider_ids = {str(pid) for pid, _fmt in endpoint_rows}
# 同一调用中 endpoint_ids 来自同一 api_format上层已按格式分组
api_format = (
endpoint_rows[0][1].value
if hasattr(endpoint_rows[0][1], "value")
else str(endpoint_rows[0][1])
)
keys = (
db.query(ProviderAPIKey.id, ProviderAPIKey.api_formats)
.filter(ProviderAPIKey.provider_id.in_(provider_ids))
.all()
)
key_ids = [str(key_id) for key_id, formats in keys if api_format in (formats or [])]
if not key_ids:
return {
"timeline": ["unknown"] * 100,
"time_range_start": None,
"time_range_end": None,
}
# 使用批量查询
format_key_mapping = {"_single": key_ids}
result = EndpointHealthService._generate_timeline_batch( result = EndpointHealthService._generate_timeline_batch(
db, format_key_mapping, now, lookback_hours, segments db, format_endpoint_mapping, now, lookback_hours, segments
) )
return result.get( return result.get(

View File

@@ -140,6 +140,13 @@ def _extract_codex_weekly_reset_seconds(metadata: dict[str, Any]) -> float | Non
if not isinstance(codex, dict): if not isinstance(codex, dict):
return None return None
weekly_used_percent = safe_float(codex.get("primary_used_percent"))
if weekly_used_percent is not None:
clamped_used = max(0.0, min(weekly_used_percent, 100.0))
if clamped_used <= 1e-6:
# 周额度仍为满额时,不启用周窗口重置倒计时。
return None
now = time.time() now = time.time()
# 优先绝对时间戳,避免 reset_seconds 快照随时间漂移。 # 优先绝对时间戳,避免 reset_seconds 快照随时间漂移。

View File

@@ -69,6 +69,14 @@ def _format_quota_value(value: float) -> str:
return f"{value:.1f}" return f"{value:.1f}"
def _has_quota_consumption(used_percent_raw: Any) -> bool:
used = _to_float(used_percent_raw)
if used is None:
return False
clamped_used = max(0.0, min(used, 100.0))
return clamped_used > 1e-6
def _format_reset_after(seconds_raw: Any) -> str | None: def _format_reset_after(seconds_raw: Any) -> str | None:
seconds = _to_float(seconds_raw) seconds = _to_float(seconds_raw)
if seconds is None: if seconds is None:
@@ -223,7 +231,11 @@ class CodexQuotaReader(PoolQuotaReader):
primary_used = _to_float(self._data.get("primary_used_percent")) primary_used = _to_float(self._data.get("primary_used_percent"))
if primary_used is not None: if primary_used is not None:
part = f"周剩余 {_format_percent(100.0 - primary_used)}" part = f"周剩余 {_format_percent(100.0 - primary_used)}"
reset_text = _format_reset_after(self._data.get("primary_reset_seconds")) reset_text = (
_format_reset_after(self._data.get("primary_reset_seconds"))
if _has_quota_consumption(primary_used)
else None
)
if reset_text: if reset_text:
part = f"{part} ({reset_text})" part = f"{part} ({reset_text})"
parts.append(part) parts.append(part)
@@ -231,7 +243,11 @@ class CodexQuotaReader(PoolQuotaReader):
secondary_used = _to_float(self._data.get("secondary_used_percent")) secondary_used = _to_float(self._data.get("secondary_used_percent"))
if secondary_used is not None: if secondary_used is not None:
part = f"5H剩余 {_format_percent(100.0 - secondary_used)}" part = f"5H剩余 {_format_percent(100.0 - secondary_used)}"
reset_text = _format_reset_after(self._data.get("secondary_reset_seconds")) reset_text = (
_format_reset_after(self._data.get("secondary_reset_seconds"))
if _has_quota_consumption(secondary_used)
else None
)
if reset_text: if reset_text:
part = f"{part} ({reset_text})" part = f"{part} ({reset_text})"
parts.append(part) parts.append(part)

View File

@@ -4,7 +4,7 @@ Sub2API 余额查询操作
import asyncio import asyncio
import time import time
from typing import Any from typing import Any, cast
import httpx import httpx
@@ -40,10 +40,13 @@ class Sub2ApiBalanceAction(BalanceAction):
sub_endpoint = self.config.get("subscription_endpoint", "/api/v1/subscriptions/summary") sub_endpoint = self.config.get("subscription_endpoint", "/api/v1/subscriptions/summary")
try: try:
me_resp, sub_resp = await asyncio.gather( me_resp, sub_resp = cast(
client.get(me_endpoint), tuple[httpx.Response | BaseException, httpx.Response | BaseException],
client.get(sub_endpoint), await asyncio.gather(
return_exceptions=True, client.get(me_endpoint),
client.get(sub_endpoint),
return_exceptions=True,
),
) )
response_time_ms = int((time.time() - start_time) * 1000) response_time_ms = int((time.time() - start_time) * 1000)

View File

@@ -4,7 +4,7 @@ YesCode 余额查询操作
import asyncio import asyncio
from datetime import datetime, timedelta from datetime import datetime, timedelta
from typing import Any from typing import Any, cast
import httpx import httpx
@@ -37,8 +37,9 @@ async def fetch_yescode_combined_data(
balance_task = client.get(f"{base_url}/api/v1/user/balance") balance_task = client.get(f"{base_url}/api/v1/user/balance")
profile_task = client.get(f"{base_url}/api/v1/auth/profile") profile_task = client.get(f"{base_url}/api/v1/auth/profile")
balance_resp, profile_resp = await asyncio.gather( balance_resp, profile_resp = cast(
balance_task, profile_task, return_exceptions=True tuple[httpx.Response | BaseException, httpx.Response | BaseException],
await asyncio.gather(balance_task, profile_task, return_exceptions=True),
) )
# 解析 balance 接口 # 解析 balance 接口

View File

@@ -975,21 +975,23 @@ class MaintenanceScheduler:
try: try:
now = datetime.now(timezone.utc) now = datetime.now(timezone.utc)
# 1. 压缩详细日志 (body 字段 -> 压缩字段)
detail_cutoff = now - timedelta(days=detail_retention) detail_cutoff = now - timedelta(days=detail_retention)
body_compressed = self._cleanup_body_fields(detail_cutoff, batch_size)
# 2. 清理压缩字段
compressed_cutoff = now - timedelta(days=compressed_retention) compressed_cutoff = now - timedelta(days=compressed_retention)
compressed_cleaned = self._cleanup_compressed_fields(compressed_cutoff, batch_size)
# 3. 清理请求头
header_cutoff = now - timedelta(days=header_retention) header_cutoff = now - timedelta(days=header_retention)
header_cleaned = self._cleanup_header_fields(header_cutoff, batch_size)
# 4. 删除过期记录
log_cutoff = now - timedelta(days=log_retention) log_cutoff = now - timedelta(days=log_retention)
# 先删最老的整行,再按窗口处理剩余记录,避免同一行在一轮里被重复改写。
records_deleted = self._delete_old_records(log_cutoff, batch_size) records_deleted = self._delete_old_records(log_cutoff, batch_size)
header_cleaned = self._cleanup_header_fields(
header_cutoff, batch_size, newer_than=log_cutoff
)
body_cleaned = self._cleanup_stale_body_fields(
compressed_cutoff, batch_size, newer_than=log_cutoff
)
# 仅压缩 7-30 天窗口内的 body更老记录直接清空 body不再先压缩再清理。
body_compressed = self._cleanup_body_fields(
detail_cutoff, batch_size, newer_than=compressed_cutoff
)
# 5. 清理过期的API Keys # 5. 清理过期的API Keys
keys_db = create_session() keys_db = create_session()
@@ -1005,7 +1007,7 @@ class MaintenanceScheduler:
logger.info( logger.info(
f"清理完成: 压缩 {body_compressed} 条, " f"清理完成: 压缩 {body_compressed} 条, "
f"清理压缩字段 {compressed_cleaned} 条, " f"清理body {body_cleaned} 条, "
f"清理header {header_cleaned} 条, " f"清理header {header_cleaned} 条, "
f"删除记录 {records_deleted} 条, " f"删除记录 {records_deleted} 条, "
f"清理过期Keys {keys_cleaned}" f"清理过期Keys {keys_cleaned}"
@@ -1016,10 +1018,16 @@ class MaintenanceScheduler:
loop = asyncio.get_running_loop() loop = asyncio.get_running_loop()
await loop.run_in_executor(None, _do_cleanup) await loop.run_in_executor(None, _do_cleanup)
def _cleanup_body_fields(self, cutoff_time: datetime, batch_size: int) -> int: def _cleanup_body_fields(
self,
cutoff_time: datetime,
batch_size: int,
*,
newer_than: datetime | None = None,
) -> int:
"""压缩 request_body 和 response_body 字段到压缩字段 """压缩 request_body 和 response_body 字段到压缩字段
逐条处理,确保每条记录都正确更新(同步方法,在线程池中调用) 仅处理指定时间窗口内仍保留原始 body 的记录,避免对更老记录重复写放大。
""" """
from sqlalchemy import null, update from sqlalchemy import null, update
@@ -1027,84 +1035,45 @@ class MaintenanceScheduler:
no_progress_count = 0 no_progress_count = 0
memory_safe_batch_size = max(1, min(batch_size, 25)) memory_safe_batch_size = max(1, min(batch_size, 25))
if newer_than is not None and newer_than >= cutoff_time:
logger.warning(
"压缩 body 字段跳过: 无效时间窗口 newer_than={} cutoff_time={}",
newer_than,
cutoff_time,
)
return 0
while True: while True:
batch_db = create_session() batch_db = create_session()
try: try:
record_ids = [ query = batch_db.query(
row.id Usage.id,
for row in ( Usage.request_body,
batch_db.query(Usage.id) Usage.response_body,
.filter(Usage.created_at < cutoff_time) Usage.provider_request_body,
.filter( Usage.client_response_body,
(Usage.request_body.isnot(None)) ).filter(Usage.created_at < cutoff_time)
| (Usage.response_body.isnot(None)) if newer_than is not None:
| (Usage.provider_request_body.isnot(None)) query = query.filter(Usage.created_at >= newer_than)
| (Usage.client_response_body.isnot(None)) records = (
) query.filter(
.order_by(Usage.created_at.asc(), Usage.id.asc()) (Usage.request_body.isnot(None))
.limit(memory_safe_batch_size) | (Usage.response_body.isnot(None))
.all() | (Usage.provider_request_body.isnot(None))
| (Usage.client_response_body.isnot(None))
) )
] .order_by(Usage.created_at.asc(), Usage.id.asc())
except Exception as e: .limit(memory_safe_batch_size)
logger.exception("加载待压缩 body 记录 ID 失败: {}", e) .all()
try: )
batch_db.rollback()
except Exception:
pass
break
finally:
batch_db.close()
if not record_ids: if not records:
break break
batch_success = 0 batch_success = 0
batch_progress = False batch_progress = False
for record in records:
for record_id in record_ids: result = batch_db.execute(
record_db = create_session()
try:
record = (
record_db.query(
Usage.id,
Usage.request_body,
Usage.response_body,
Usage.provider_request_body,
Usage.client_response_body,
)
.filter(Usage.id == record_id)
.first()
)
if record is None:
continue
has_body_payload = (
record.request_body is not None
or record.response_body is not None
or record.provider_request_body is not None
or record.client_response_body is not None
)
if not has_body_payload:
result = record_db.execute(
update(Usage)
.where(Usage.id == record.id)
.values(
request_body=null(),
response_body=null(),
provider_request_body=null(),
client_response_body=null(),
)
.execution_options(synchronize_session=False)
)
record_db.commit()
if result.rowcount > 0:
batch_progress = True
continue
result = record_db.execute(
update(Usage) update(Usage)
.where(Usage.id == record.id) .where(Usage.id == record.id)
.values( .values(
@@ -1133,20 +1102,20 @@ class MaintenanceScheduler:
) )
.execution_options(synchronize_session=False) .execution_options(synchronize_session=False)
) )
record_db.commit()
if result.rowcount > 0: if result.rowcount > 0:
batch_success += 1 batch_success += 1
batch_progress = True batch_progress = True
except Exception as e:
logger.warning("压缩记录 {} 失败: {}", record_id, e) batch_db.commit()
try: except Exception as e:
record_db.rollback() logger.warning("压缩 body 批次失败: {}", e)
except Exception: try:
pass batch_db.rollback()
continue except Exception:
finally: pass
record_db.close() break
finally:
batch_db.close()
if not batch_progress: if not batch_progress:
no_progress_count += 1 no_progress_count += 1
@@ -1166,27 +1135,47 @@ class MaintenanceScheduler:
return total_compressed return total_compressed
def _cleanup_compressed_fields(self, cutoff_time: datetime, batch_size: int) -> int: def _cleanup_stale_body_fields(
"""清理压缩字段删除压缩的body self,
cutoff_time: datetime,
batch_size: int,
*,
newer_than: datetime | None = None,
) -> int:
"""清理已超过压缩保留期的 body 字段
每批使用短生命周期 session同步方法在线程池中调用 直接清空 raw/compressed body避免更老记录先压缩再马上被清掉。
""" """
from sqlalchemy import null, update from sqlalchemy import null, update
total_cleaned = 0 total_cleaned = 0
if newer_than is not None and newer_than >= cutoff_time:
logger.warning(
"清理 body 字段跳过: 无效时间窗口 newer_than={} cutoff_time={}",
newer_than,
cutoff_time,
)
return 0
while True: while True:
batch_db = create_session() batch_db = create_session()
try: try:
query = batch_db.query(Usage.id).filter(Usage.created_at < cutoff_time)
if newer_than is not None:
query = query.filter(Usage.created_at >= newer_than)
records_to_clean = ( records_to_clean = (
batch_db.query(Usage.id) query.filter(
.filter(Usage.created_at < cutoff_time) (Usage.request_body.isnot(None))
.filter( | (Usage.response_body.isnot(None))
(Usage.request_body_compressed.isnot(None)) | (Usage.provider_request_body.isnot(None))
| (Usage.client_response_body.isnot(None))
| (Usage.request_body_compressed.isnot(None))
| (Usage.response_body_compressed.isnot(None)) | (Usage.response_body_compressed.isnot(None))
| (Usage.provider_request_body_compressed.isnot(None)) | (Usage.provider_request_body_compressed.isnot(None))
| (Usage.client_response_body_compressed.isnot(None)) | (Usage.client_response_body_compressed.isnot(None))
) )
.order_by(Usage.created_at.asc(), Usage.id.asc())
.limit(batch_size) .limit(batch_size)
.all() .all()
) )
@@ -1200,6 +1189,10 @@ class MaintenanceScheduler:
update(Usage) update(Usage)
.where(Usage.id.in_(record_ids)) .where(Usage.id.in_(record_ids))
.values( .values(
request_body=null(),
response_body=null(),
provider_request_body=null(),
client_response_body=null(),
request_body_compressed=null(), request_body_compressed=null(),
response_body_compressed=null(), response_body_compressed=null(),
provider_request_body_compressed=null(), provider_request_body_compressed=null(),
@@ -1211,14 +1204,14 @@ class MaintenanceScheduler:
batch_db.commit() batch_db.commit()
if rows_updated == 0: if rows_updated == 0:
logger.warning("清理压缩字段: rowcount=0可能存在问题") logger.warning("清理 body 字段: rowcount=0可能存在问题")
break break
total_cleaned += rows_updated total_cleaned += rows_updated
logger.debug(f"已清理 {rows_updated} 条记录的压缩字段,累计 {total_cleaned}") logger.debug(f"已清理 {rows_updated} 条记录的 body 字段,累计 {total_cleaned}")
except Exception as e: except Exception as e:
logger.exception(f"清理压缩字段失败: {e}") logger.exception(f"清理 body 字段失败: {e}")
try: try:
batch_db.rollback() batch_db.rollback()
except Exception: except Exception:
@@ -1229,7 +1222,13 @@ class MaintenanceScheduler:
return total_cleaned return total_cleaned
def _cleanup_header_fields(self, cutoff_time: datetime, batch_size: int) -> int: def _cleanup_header_fields(
self,
cutoff_time: datetime,
batch_size: int,
*,
newer_than: datetime | None = None,
) -> int:
"""清理 request_headers, response_headers 和 provider_request_headers 字段 """清理 request_headers, response_headers 和 provider_request_headers 字段
每批使用短生命周期 session同步方法在线程池中调用 每批使用短生命周期 session同步方法在线程池中调用
@@ -1238,17 +1237,28 @@ class MaintenanceScheduler:
total_cleaned = 0 total_cleaned = 0
if newer_than is not None and newer_than >= cutoff_time:
logger.warning(
"清理 header 字段跳过: 无效时间窗口 newer_than={} cutoff_time={}",
newer_than,
cutoff_time,
)
return 0
while True: while True:
batch_db = create_session() batch_db = create_session()
try: try:
query = batch_db.query(Usage.id).filter(Usage.created_at < cutoff_time)
if newer_than is not None:
query = query.filter(Usage.created_at >= newer_than)
records_to_clean = ( records_to_clean = (
batch_db.query(Usage.id) query.filter(
.filter(Usage.created_at < cutoff_time)
.filter(
(Usage.request_headers.isnot(None)) (Usage.request_headers.isnot(None))
| (Usage.response_headers.isnot(None)) | (Usage.response_headers.isnot(None))
| (Usage.provider_request_headers.isnot(None)) | (Usage.provider_request_headers.isnot(None))
| (Usage.client_response_headers.isnot(None))
) )
.order_by(Usage.created_at.asc(), Usage.id.asc())
.limit(batch_size) .limit(batch_size)
.all() .all()
) )
@@ -1265,6 +1275,7 @@ class MaintenanceScheduler:
request_headers=null(), request_headers=null(),
response_headers=null(), response_headers=null(),
provider_request_headers=null(), provider_request_headers=null(),
client_response_headers=null(),
) )
) )
@@ -1300,6 +1311,7 @@ class MaintenanceScheduler:
records_to_delete = ( records_to_delete = (
batch_db.query(Usage.id) batch_db.query(Usage.id)
.filter(Usage.created_at < cutoff_time) .filter(Usage.created_at < cutoff_time)
.order_by(Usage.created_at.asc(), Usage.id.asc())
.limit(batch_size) .limit(batch_size)
.all() .all()
) )

View File

@@ -69,6 +69,7 @@ class SyncTaskExecutionService:
request_body_state: RequestBodyState | None, request_body_state: RequestBodyState | None,
request_headers: dict[str, Any] | None, request_headers: dict[str, Any] | None,
request_body: dict[str, Any] | None, request_body: dict[str, Any] | None,
create_pending_usage: bool = True,
) -> ExecutionResult: ) -> ExecutionResult:
""" """
Unified candidate traversal loop for SYNC. Unified candidate traversal loop for SYNC.
@@ -143,26 +144,32 @@ class SyncTaskExecutionService:
affinity_key = str(user_api_key.id) affinity_key = str(user_api_key.id)
user_id = str(user_api_key.user_id) user_id = str(user_api_key.user_id)
api_format_norm = normalize_endpoint_signature(api_format) api_format_norm = normalize_endpoint_signature(api_format)
user: User | None = None
username_snapshot = None username_snapshot = None
api_key_name_snapshot = getattr(user_api_key, "name", None) api_key_name_snapshot = getattr(user_api_key, "name", None)
# Keep pending usage creation behavior consistent with previous behavior
try: try:
user = self.db.query(User).filter(User.id == user_api_key.user_id).first() user = self.db.query(User).filter(User.id == user_api_key.user_id).first()
username_snapshot = getattr(user, "username", None) if user else None username_snapshot = getattr(user, "username", None) if user else None
UsageService.create_pending_usage(
db=self.db,
request_id=request_id,
user=user,
api_key=user_api_key,
model=model_name,
is_stream=is_stream,
api_format=api_format_norm,
request_headers=request_headers,
request_body=request_body,
)
except Exception as exc: except Exception as exc:
logger.warning("创建 pending 使用记录失败: {}", str(exc)) # username 仅用于审计快照,不应阻塞主请求链路。
logger.warning("查询用户快照失败: {}", str(exc))
# 默认由 TaskService 创建 pending 使用记录;已预创建的调用方可关闭。
if create_pending_usage:
try:
UsageService.create_pending_usage(
db=self.db,
request_id=request_id,
user=user,
api_key=user_api_key,
model=model_name,
is_stream=is_stream,
api_format=api_format_norm,
request_headers=request_headers,
request_body=request_body,
)
except Exception as exc:
logger.warning("创建 pending 使用记录失败: {}", str(exc))
all_candidates, global_model_id = await candidate_resolver.fetch_candidates( all_candidates, global_model_id = await candidate_resolver.fetch_candidates(
api_format=api_format_norm, api_format=api_format_norm,

View File

@@ -126,6 +126,7 @@ class TaskService:
supported_auth_types: set[str] | None = None, supported_auth_types: set[str] | None = None,
allow_format_conversion: bool = False, allow_format_conversion: bool = False,
max_candidates: int | None = None, max_candidates: int | None = None,
create_pending_usage: bool = True,
) -> ExecutionResult: ) -> ExecutionResult:
"""兼容入口:默认绑定到 TaskService 内部执行路由。""" """兼容入口:默认绑定到 TaskService 内部执行路由。"""
return await self._execute_facade_ops.execute( return await self._execute_facade_ops.execute(
@@ -146,6 +147,7 @@ class TaskService:
supported_auth_types=supported_auth_types, supported_auth_types=supported_auth_types,
allow_format_conversion=allow_format_conversion, allow_format_conversion=allow_format_conversion,
max_candidates=max_candidates, max_candidates=max_candidates,
create_pending_usage=create_pending_usage,
) )
async def _execute_internal( async def _execute_internal(
@@ -168,6 +170,7 @@ class TaskService:
supported_auth_types: set[str] | None = None, supported_auth_types: set[str] | None = None,
allow_format_conversion: bool = False, allow_format_conversion: bool = False,
max_candidates: int | None = None, max_candidates: int | None = None,
create_pending_usage: bool = True,
) -> ExecutionResult: ) -> ExecutionResult:
if task_mode == TaskMode.ASYNC: if task_mode == TaskMode.ASYNC:
if extract_external_task_id is None: if extract_external_task_id is None:
@@ -250,6 +253,7 @@ class TaskService:
request_body_state=request_body_state, request_body_state=request_body_state,
request_headers=request_headers, request_headers=request_headers,
request_body=request_body, request_body=request_body,
create_pending_usage=create_pending_usage,
) )
async def execute_sync_candidates( async def execute_sync_candidates(

View File

@@ -39,6 +39,17 @@ async def _yield_once_then_cancel(ctx: StreamContext) -> AsyncGenerator[bytes, N
raise asyncio.CancelledError() raise asyncio.CancelledError()
async def _yield_once_then_hang(ctx: StreamContext) -> AsyncGenerator[bytes, None]:
ctx.append_text("partial output")
yield b"data: chunk\n\n"
await asyncio.sleep(3600)
async def _yield_after_delay_then_complete() -> AsyncGenerator[bytes, None]:
await asyncio.sleep(0.25)
yield b"data: first\n\n"
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_create_monitored_stream_marks_client_disconnected_when_confirmed() -> None: async def test_create_monitored_stream_marks_client_disconnected_when_confirmed() -> None:
monitor = _DummyMonitor() monitor = _DummyMonitor()
@@ -113,3 +124,38 @@ async def test_create_monitored_stream_estimates_output_tokens_before_unknown_ca
assert ctx.error_message == "cancelled_unknown" assert ctx.error_message == "cancelled_unknown"
assert ctx.output_tokens == expected_output_tokens assert ctx.output_tokens == expected_output_tokens
assert f"output_tokens={expected_output_tokens}" in (ctx.upstream_response or "") assert f"output_tokens={expected_output_tokens}" in (ctx.upstream_response or "")
@pytest.mark.asyncio
async def test_create_monitored_stream_marks_idle_timeout_before_worker_timeout() -> None:
monitor = _DummyMonitor()
monitor.CANCEL_DISCONNECT_RETRY_DELAYS_SECONDS = ()
monitor.STREAM_IDLE_TIMEOUT_SECONDS = 1.0
ctx = StreamContext(model="test-model", api_format="openai:cli", request_id="req-idle-timeout")
monitored = monitor._create_monitored_stream(ctx, _yield_once_then_hang(ctx), None)
with pytest.raises(asyncio.CancelledError):
async for _ in monitored:
pass
expected_output_tokens = max(1, len("partial output") // 4)
assert ctx.status_code == 504
assert ctx.error_message == "stream_idle_timeout"
assert ctx.output_tokens == expected_output_tokens
assert "cancel_origin=stream_idle_timeout" in (ctx.upstream_response or "")
@pytest.mark.asyncio
async def test_create_monitored_stream_does_not_idle_timeout_before_first_chunk() -> None:
monitor = _DummyMonitor()
monitor.CANCEL_DISCONNECT_RETRY_DELAYS_SECONDS = ()
monitor.STREAM_IDLE_TIMEOUT_SECONDS = 0.05
ctx = StreamContext(model="test-model", api_format="openai:cli", request_id="req-first-chunk")
monitored = monitor._create_monitored_stream(ctx, _yield_after_delay_then_complete(), None)
chunks = [chunk async for chunk in monitored]
assert chunks == [b"data: first\n\n"]
assert ctx.status_code == 200
assert ctx.error_message is None

View File

@@ -23,7 +23,7 @@ class _DummySyncHandler(CliSyncMixin):
) -> str: ) -> str:
return str(request_body.get("model") or "unknown") return str(request_body.get("model") or "unknown")
def _create_pending_usage(self, **kwargs: object) -> None: def _create_pending_usage(self, **kwargs: object) -> bool:
self.pending_calls.append(kwargs) self.pending_calls.append(kwargs)
raise _StopExecution() raise _StopExecution()
@@ -41,7 +41,7 @@ class _DummyStreamHandler(CliStreamMixin):
) -> str: ) -> str:
return str(request_body.get("model") or "unknown") return str(request_body.get("model") or "unknown")
def _create_pending_usage(self, **kwargs: object) -> None: def _create_pending_usage(self, **kwargs: object) -> bool:
self.pending_calls.append(kwargs) self.pending_calls.append(kwargs)
raise _StopExecution() raise _StopExecution()

View File

@@ -1,3 +1,9 @@
import os
import pytest
os.environ.setdefault("JWT_SECRET_KEY", "test-secret-key")
from src.api.handlers.base import stream_context from src.api.handlers.base import stream_context
from src.api.handlers.base.stream_context import StreamContext from src.api.handlers.base.stream_context import StreamContext
@@ -50,7 +56,36 @@ def test_reset_for_retry_clears_state() -> None:
assert ctx.error_message is None assert ctx.error_message is None
def test_record_first_byte_time(monkeypatch) -> None: def test_release_recorded_chunks_clears_both_chunk_lists() -> None:
ctx = StreamContext(model="test-model", api_format="openai:chat")
ctx.parsed_chunks.append({"type": "client"})
ctx.provider_parsed_chunks.append({"type": "provider"})
ctx.release_recorded_chunks()
assert ctx.parsed_chunks == []
assert ctx.provider_parsed_chunks == []
def test_managed_recorded_bodies_builds_then_releases_chunks() -> None:
ctx = StreamContext(model="test-model", api_format="openai:chat")
ctx.parsed_chunks.append({"type": "client"})
ctx.provider_parsed_chunks.append({"type": "provider"})
ctx.data_count = 1
with ctx.managed_recorded_bodies(123) as recorded_bodies:
assert recorded_bodies.response_body is not None
assert recorded_bodies.response_body["chunks"] == [{"type": "provider"}]
assert recorded_bodies.client_response_body is not None
assert recorded_bodies.client_response_body["chunks"] == [{"type": "client"}]
assert ctx.parsed_chunks == []
assert ctx.provider_parsed_chunks == []
assert recorded_bodies.response_body is None
assert recorded_bodies.client_response_body is None
def test_record_first_byte_time(monkeypatch: pytest.MonkeyPatch) -> None:
"""测试记录首字时间""" """测试记录首字时间"""
ctx = StreamContext(model="claude-3", api_format="claude_messages") ctx = StreamContext(model="claude-3", api_format="claude_messages")
start_time = 100.0 start_time = 100.0
@@ -63,7 +98,7 @@ def test_record_first_byte_time(monkeypatch) -> None:
assert ctx.first_byte_time_ms == 12 assert ctx.first_byte_time_ms == 12
def test_record_first_byte_time_idempotent(monkeypatch) -> None: def test_record_first_byte_time_idempotent(monkeypatch: pytest.MonkeyPatch) -> None:
"""测试首字时间只记录一次""" """测试首字时间只记录一次"""
ctx = StreamContext(model="claude-3", api_format="claude_messages") ctx = StreamContext(model="claude-3", api_format="claude_messages")
start_time = 100.0 start_time = 100.0
@@ -82,7 +117,7 @@ def test_record_first_byte_time_idempotent(monkeypatch) -> None:
assert first_value == second_value assert first_value == second_value
def test_reset_for_retry_clears_first_byte_time(monkeypatch) -> None: def test_reset_for_retry_clears_first_byte_time(monkeypatch: pytest.MonkeyPatch) -> None:
"""测试重试时清除首字时间""" """测试重试时清除首字时间"""
ctx = StreamContext(model="claude-3", api_format="claude_messages") ctx = StreamContext(model="claude-3", api_format="claude_messages")
start_time = 100.0 start_time = 100.0

View File

@@ -1,11 +1,15 @@
from __future__ import annotations from __future__ import annotations
import os
import time import time
from types import SimpleNamespace from types import SimpleNamespace
from typing import Any
from unittest.mock import AsyncMock from unittest.mock import AsyncMock
import pytest import pytest
os.environ.setdefault("JWT_SECRET_KEY", "test-secret-key")
from src.api.handlers.base import stream_telemetry as stream_telemetry_module from src.api.handlers.base import stream_telemetry as stream_telemetry_module
from src.api.handlers.base.stream_context import StreamContext from src.api.handlers.base.stream_context import StreamContext
from src.api.handlers.base.stream_telemetry import StreamTelemetryRecorder from src.api.handlers.base.stream_telemetry import StreamTelemetryRecorder
@@ -64,3 +68,69 @@ async def test_record_stream_stats_estimates_tokens_for_failed_partial_stream(
assert ctx.input_tokens > 0 assert ctx.input_tokens > 0
assert ctx.output_tokens == max(1, len("partial output") // 4) assert ctx.output_tokens == max(1, len("partial output") // 4)
recorder._dispatch_record.assert_awaited_once() # type: ignore[attr-defined] recorder._dispatch_record.assert_awaited_once() # type: ignore[attr-defined]
@pytest.mark.asyncio
async def test_record_stream_stats_releases_parsed_chunks_after_dispatch(
monkeypatch: pytest.MonkeyPatch,
) -> None:
recorder = StreamTelemetryRecorder(
request_id="req-release",
user_id="1",
api_key_id="2",
client_ip="127.0.0.1",
format_id="openai:chat",
)
recorder._get_telemetry_writer = AsyncMock( # type: ignore[method-assign]
return_value=SimpleNamespace(include_bodies=True)
)
recorder._update_candidate_status = AsyncMock() # type: ignore[method-assign]
dispatch_payloads: list[dict[str, Any]] = []
async def _capture_dispatch(*args: Any, **kwargs: Any) -> None:
dispatch_payloads.append(
{
"response_body": args[5],
"client_response_body": kwargs.get("client_response_body"),
}
)
recorder._dispatch_record = _capture_dispatch # type: ignore[method-assign]
ctx = StreamContext(
model="test-model",
api_format="openai:chat",
request_id="req-release",
user_id=1,
api_key_id=2,
)
ctx.provider_name = "test-provider"
ctx.parsed_chunks.extend([{"type": "chunk-1"}, {"type": "chunk-2"}])
ctx.data_count = 2
ctx.chunk_count = 2
monkeypatch.setattr(stream_telemetry_module, "get_db", lambda: iter([_DummyDb()]))
monkeypatch.setattr(
stream_telemetry_module.SystemConfigService,
"should_log_body",
lambda _db: True,
)
monkeypatch.setattr(stream_telemetry_module.config, "stream_stats_delay", 0)
await recorder.record_stream_stats(
ctx,
original_headers={},
original_request_body={"input": [{"content": "hello world"}]},
start_time=time.time(),
)
assert len(dispatch_payloads) == 1
response_body = dispatch_payloads[0]["response_body"]
assert response_body["chunks"] == [{"type": "chunk-1"}, {"type": "chunk-2"}]
assert response_body["metadata"]["stream"] is True
assert response_body["metadata"]["total_chunks"] == 2
assert response_body["metadata"]["data_count"] == 2
assert dispatch_payloads[0]["client_response_body"] is None
assert ctx.parsed_chunks == []
assert ctx.provider_parsed_chunks == []

View File

@@ -0,0 +1,146 @@
from __future__ import annotations
from datetime import datetime, timezone
from types import SimpleNamespace
from typing import Any, cast
from unittest.mock import MagicMock
import pytest
from src.api.admin.endpoints.health import AdminApiFormatHealthMonitorAdapter
from src.api.public.catalog import PublicApiFormatHealthMonitorAdapter
def _build_query(result: object) -> MagicMock:
query = MagicMock()
query.join.return_value = query
query.distinct.return_value = query
query.filter.return_value = query
query.group_by.return_value = query
query.order_by.return_value = query
query.limit.return_value = query
query.all.return_value = result
return query
def _expr_texts(query: MagicMock) -> list[str]:
return [str(arg) for arg in query.filter.call_args.args]
@pytest.mark.asyncio
async def test_admin_api_format_health_monitor_filters_inactive_sources(
monkeypatch: pytest.MonkeyPatch,
) -> None:
endpoint_query = _build_query([("openai:compact", "ep-active", "provider-active")])
key_query = _build_query([("provider-active", ["openai:compact"])])
status_query = _build_query([("openai:compact", "success", 3)])
rows_query = _build_query([])
db = MagicMock()
db.query.side_effect = [endpoint_query, key_query, status_query, rows_query]
monkeypatch.setattr(
"src.api.admin.endpoints.health.EndpointHealthService._generate_timeline_from_usage",
lambda **_: {
"timeline": ["healthy"] * 100,
"time_range_start": None,
"time_range_end": None,
},
)
context = SimpleNamespace(
db=db,
request=SimpleNamespace(state=SimpleNamespace()),
add_audit_metadata=lambda **kwargs: None,
)
adapter = AdminApiFormatHealthMonitorAdapter(lookback_hours=6, per_format_limit=20)
await adapter.handle(cast(Any, context))
status_filters = _expr_texts(status_query)
rows_filters = _expr_texts(rows_query)
assert any("provider_endpoints.is_active" in expr for expr in status_filters)
assert any("providers.is_active" in expr for expr in status_filters)
assert any("provider_endpoints.is_active" in expr for expr in rows_filters)
assert any("providers.is_active" in expr for expr in rows_filters)
@pytest.mark.asyncio
async def test_public_api_format_health_monitor_uses_real_counts_not_sampled_events(
monkeypatch: pytest.MonkeyPatch,
) -> None:
now = datetime.now(timezone.utc)
active_formats_query = _build_query([("openai:compact",)])
endpoint_rows_query = _build_query([("openai:compact", "ep-active")])
status_query = _build_query(
[
("openai:compact", "success", 7),
("openai:compact", "failed", 3),
("openai:compact", "skipped", 5),
]
)
rows_query = _build_query(
[
SimpleNamespace(
status="failed",
status_code=500,
latency_ms=321,
error_type="provider_error",
finished_at=now,
started_at=None,
created_at=now,
),
SimpleNamespace(
status="success",
status_code=200,
latency_ms=123,
error_type=None,
finished_at=now,
started_at=None,
created_at=now,
),
]
)
db = MagicMock()
db.query.side_effect = [
active_formats_query,
endpoint_rows_query,
status_query,
rows_query,
]
monkeypatch.setattr(
"src.api.public.catalog.EndpointHealthService._generate_timeline_from_usage",
lambda **_: {
"timeline": ["healthy"] * 100,
"time_range_start": None,
"time_range_end": now,
},
)
monkeypatch.setattr(
"src.core.api_format.get_local_path_for_endpoint",
lambda api_format: f"/{api_format}",
)
context = SimpleNamespace(
db=db,
request=SimpleNamespace(state=SimpleNamespace()),
)
adapter = PublicApiFormatHealthMonitorAdapter(lookback_hours=6, per_format_limit=20)
result = await adapter.handle(cast(Any, context))
monitor = result["formats"][0]
assert monitor["api_format"] == "openai:compact"
assert monitor["total_attempts"] == 15
assert monitor["success_count"] == 7
assert monitor["failed_count"] == 3
assert monitor["skipped_count"] == 5
assert monitor["success_rate"] == pytest.approx(0.7)
assert len(monitor["events"]) == 2
rows_filters = _expr_texts(rows_query)
assert any("provider_endpoints.is_active" in expr for expr in rows_filters)
assert any("providers.is_active" in expr for expr in rows_filters)

View File

@@ -0,0 +1,50 @@
import os
import pytest
os.environ.setdefault("JWT_SECRET_KEY", "test-secret-key")
from src.plugins.manager import PluginManager
def _build_manager() -> PluginManager:
return PluginManager(
config={
"auth": {"api_key": False},
"cache": {"memory": False},
"monitor": {"prometheus": False},
"token": {"claude": False},
"load_balancer": {"sticky_priority": False},
}
)
def test_rate_limit_defaults_to_token_bucket_when_unconfigured() -> None:
manager = _build_manager()
plugin = manager.get_plugin("rate_limit")
assert plugin is not None
assert plugin.name == "token_bucket"
@pytest.mark.asyncio
async def test_default_rate_limit_plugin_honors_dynamic_rate_limit() -> None:
manager = _build_manager()
plugin = manager.get_plugin("rate_limit")
assert plugin is not None
first = await plugin.check_limit("public_ip:test", rate_limit=2)
assert first.allowed is True
await plugin.consume("public_ip:test", amount=1, rate_limit=2)
second = await plugin.check_limit("public_ip:test", rate_limit=2)
assert second.allowed is True
await plugin.consume("public_ip:test", amount=1, rate_limit=2)
third = await plugin.check_limit("public_ip:test", rate_limit=2)
assert third.allowed is False
assert third.remaining == 0
assert third.retry_after is not None

View File

@@ -0,0 +1,111 @@
from __future__ import annotations
from datetime import datetime, timezone
import pytest
import src.plugins.rate_limit.token_bucket as token_bucket_module
from src.plugins.rate_limit.token_bucket import RedisTokenBucketBackend, TokenBucketStrategy
@pytest.mark.asyncio
async def test_token_bucket_cleans_up_expired_buckets(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("RATE_LIMIT_BACKEND", "memory")
strategy = TokenBucketStrategy()
strategy.configure({"bucket_expiry": 1, "cleanup_interval": 0})
await strategy.check_limit("api_key:stale")
strategy.buckets["api_key:stale"].last_access_time -= 3600
await strategy.check_limit("api_key:fresh")
assert "api_key:stale" not in strategy.buckets
assert "api_key:fresh" in strategy.buckets
@pytest.mark.asyncio
async def test_token_bucket_reconfigures_existing_bucket_when_rate_limit_changes(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setenv("RATE_LIMIT_BACKEND", "memory")
strategy = TokenBucketStrategy()
await strategy.check_limit("user:42", rate_limit=120)
bucket = strategy.buckets["user:42"]
bucket.tokens = 90
await strategy.check_limit("user:42", rate_limit=30)
updated_bucket = strategy.buckets["user:42"]
assert updated_bucket.capacity == 30
assert updated_bucket.refill_rate == 0.5
assert updated_bucket.tokens <= 30
@pytest.mark.asyncio
async def test_token_bucket_treats_non_positive_dynamic_rate_limit_as_unlimited(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setenv("RATE_LIMIT_BACKEND", "memory")
strategy = TokenBucketStrategy()
result = await strategy.check_limit("public_ip:test", rate_limit=0)
consumed = await strategy.consume("public_ip:test", amount=1, rate_limit=0)
assert result.allowed is True
assert consumed is True
assert "public_ip:test" not in strategy.buckets
class _FakeRedisClient:
async def hmget(self, _key: str, *_fields: str) -> list[None]:
return [None, None]
def register_script(self, _script: str): # type: ignore[no-untyped-def]
async def _runner(*args, **kwargs): # type: ignore[no-untyped-def]
return [1, 0, 0]
return _runner
@pytest.mark.asyncio
async def test_token_bucket_retries_redis_backend_probe_after_initial_miss(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setenv("RATE_LIMIT_BACKEND", "auto")
strategy = TokenBucketStrategy()
strategy._redis_retry_interval = 0
fake_redis = _FakeRedisClient()
calls = {"count": 0}
def _fake_get_redis_client_sync(): # type: ignore[no-untyped-def]
calls["count"] += 1
if calls["count"] == 1:
return None
return fake_redis
monkeypatch.setattr(
token_bucket_module,
"get_redis_client_sync",
_fake_get_redis_client_sync,
)
await strategy.check_limit("public_ip:first")
assert strategy._redis_backend is None
await strategy.check_limit("public_ip:second")
assert strategy._redis_backend is not None
assert calls["count"] == 2
@pytest.mark.asyncio
async def test_redis_token_bucket_missing_bucket_reports_reset_now() -> None:
backend = RedisTokenBucketBackend(_FakeRedisClient())
result = await backend.peek("public_ip:test", capacity=60, refill_rate=1.0, amount=1)
assert result.allowed is True
assert result.remaining == 60
assert result.reset_at is not None
assert abs((result.reset_at - datetime.now(timezone.utc)).total_seconds()) < 2

View File

@@ -36,6 +36,22 @@ def test_codex_reader_preserves_summary_formats() -> None:
assert credits_reader.display_summary() == "积分 12.35" assert credits_reader.display_summary() == "积分 12.35"
def test_codex_reader_hides_reset_countdown_when_remaining_is_full() -> None:
reader = get_quota_reader(
"codex",
{
"codex": {
"primary_used_percent": 0.0,
"primary_reset_seconds": 266400,
"secondary_used_percent": 0.0,
"secondary_reset_seconds": 3600,
}
},
)
assert reader.display_summary() == "周剩余 100.0% | 5H剩余 100.0%"
def test_antigravity_reader_keeps_used_percent_fallbacks() -> None: def test_antigravity_reader_keeps_used_percent_fallbacks() -> None:
reader = get_quota_reader( reader = get_quota_reader(
"antigravity", "antigravity",
@@ -88,6 +104,20 @@ def test_extract_reset_seconds_uses_codex_weekly_reset_for_codex_provider() -> N
assert extract_reset_seconds(key_obj) == pytest.approx(1800.0) assert extract_reset_seconds(key_obj) == pytest.approx(1800.0)
def test_extract_reset_seconds_codex_weekly_full_quota_returns_none() -> None:
key_obj = SimpleNamespace(
provider_type="codex",
upstream_metadata={
"codex": {
"primary_used_percent": 0.0,
"primary_reset_seconds": 1800.0,
}
},
)
assert extract_reset_seconds(key_obj) is None
def test_extract_reset_seconds_prefers_codex_reset_at(monkeypatch: pytest.MonkeyPatch) -> None: def test_extract_reset_seconds_prefers_codex_reset_at(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr("src.services.provider.pool.dimensions._helpers.time.time", lambda: 1000.0) monkeypatch.setattr("src.services.provider.pool.dimensions._helpers.time.time", lambda: 1000.0)
key_obj = SimpleNamespace( key_obj = SimpleNamespace(

View File

@@ -1,6 +1,7 @@
from __future__ import annotations from __future__ import annotations
import inspect import inspect
from datetime import datetime, timedelta, timezone
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import MagicMock from unittest.mock import MagicMock
@@ -124,44 +125,29 @@ async def test_candidate_cleanup_uses_dedicated_retention_and_batch_settings(
assert batch_two.closed is True assert batch_two.closed is True
def test_cleanup_body_fields_loads_ids_then_processes_records_individually( def test_cleanup_body_fields_batches_records_with_single_commit(
monkeypatch: pytest.MonkeyPatch, monkeypatch: pytest.MonkeyPatch,
) -> None: ) -> None:
scheduler = MaintenanceScheduler() scheduler = MaintenanceScheduler()
class _IdBatchSession: class _BatchSession:
def __init__(self, ids: list[str]) -> None: def __init__(self, records: list[SimpleNamespace]) -> None:
self.ids = ids self.records = records
self.closed = False
self.query_obj = MagicMock()
filtered = self.query_obj.filter.return_value
filtered.filter.return_value.order_by.return_value.limit.return_value.all.return_value = [
SimpleNamespace(id=value) for value in ids
]
def query(self, *args): # type: ignore[no-untyped-def]
self.query_args = args
return self.query_obj
def rollback(self) -> None:
raise AssertionError("rollback should not be called")
def close(self) -> None:
self.closed = True
class _RecordSession:
def __init__(self, record: SimpleNamespace) -> None:
self.record = record
self.closed = False self.closed = False
self.committed = False self.committed = False
self.executed = 0
self.query_obj = MagicMock() self.query_obj = MagicMock()
self.query_obj.filter.return_value.first.return_value = record filtered = self.query_obj.filter.return_value
filtered.filter.return_value.order_by.return_value.limit.return_value.all.return_value = (
records
)
def query(self, *args): # type: ignore[no-untyped-def] def query(self, *args): # type: ignore[no-untyped-def]
self.query_args = args self.query_args = args
return self.query_obj return self.query_obj
def execute(self, _statement): # type: ignore[no-untyped-def] def execute(self, _statement): # type: ignore[no-untyped-def]
self.executed += 1
return SimpleNamespace(rowcount=1) return SimpleNamespace(rowcount=1)
def commit(self) -> None: def commit(self) -> None:
@@ -173,27 +159,26 @@ def test_cleanup_body_fields_loads_ids_then_processes_records_individually(
def close(self) -> None: def close(self) -> None:
self.closed = True self.closed = True
batch_one = _IdBatchSession(["usage-1", "usage-2"]) batch_one = _BatchSession(
record_one = _RecordSession( [
SimpleNamespace( SimpleNamespace(
id="usage-1", id="usage-1",
request_body={"hello": "world"}, request_body={"hello": "world"},
response_body=None, response_body=None,
provider_request_body=None, provider_request_body=None,
client_response_body=None, client_response_body=None,
) ),
SimpleNamespace(
id="usage-2",
request_body=None,
response_body={"ok": True},
provider_request_body=None,
client_response_body=None,
),
]
) )
record_two = _RecordSession( batch_two = _BatchSession([])
SimpleNamespace( sessions = iter([batch_one, batch_two])
id="usage-2",
request_body=None,
response_body={"ok": True},
provider_request_body=None,
client_response_body=None,
)
)
batch_two = _IdBatchSession([])
sessions = iter([batch_one, record_one, record_two, batch_two])
monkeypatch.setattr( monkeypatch.setattr(
maintenance_scheduler_module, maintenance_scheduler_module,
@@ -212,12 +197,173 @@ def test_cleanup_body_fields_loads_ids_then_processes_records_individually(
) )
assert compressed == 2 assert compressed == 2
assert batch_one.query_args == (maintenance_scheduler_module.Usage.id,) assert len(batch_one.query_args) == 5
assert len(record_one.query_args) == 5 assert batch_one.executed == 2
assert len(record_two.query_args) == 5 assert batch_one.committed is True
assert batch_one.closed is True assert batch_one.closed is True
assert batch_two.closed is True assert batch_two.closed is True
assert record_one.committed is True
assert record_two.committed is True
assert record_one.closed is True def test_cleanup_header_fields_clears_client_response_headers(
assert record_two.closed is True monkeypatch: pytest.MonkeyPatch,
) -> None:
scheduler = MaintenanceScheduler()
class _BatchSession:
def __init__(self, ids: list[str]) -> None:
self.ids = ids
self.closed = False
self.committed = False
self.query_obj = MagicMock()
self.filtered_by_time = MagicMock()
self.filtered_by_headers = MagicMock()
self.query_obj.filter.return_value = self.filtered_by_time
self.filtered_by_time.filter.return_value = self.filtered_by_headers
self.filtered_by_headers.order_by.return_value.limit.return_value.all.return_value = [
SimpleNamespace(id=value) for value in ids
]
self.executed_statements: list[str] = []
def query(self, *args): # type: ignore[no-untyped-def]
self.query_args = args
return self.query_obj
def execute(self, statement): # type: ignore[no-untyped-def]
self.executed_statements.append(str(statement))
return SimpleNamespace(rowcount=len(self.ids))
def commit(self) -> None:
self.committed = True
def rollback(self) -> None:
raise AssertionError("rollback should not be called")
def close(self) -> None:
self.closed = True
batch_one = _BatchSession(["usage-1"])
batch_two = _BatchSession([])
sessions = iter([batch_one, batch_two])
monkeypatch.setattr(
maintenance_scheduler_module,
"create_session",
lambda: next(sessions),
)
cleaned = scheduler._cleanup_header_fields(
cutoff_time=SimpleNamespace(), # type: ignore[arg-type]
batch_size=1000,
)
header_filter = str(batch_one.filtered_by_time.filter.call_args.args[0])
assert cleaned == 1
assert batch_one.query_args == (maintenance_scheduler_module.Usage.id,)
assert "client_response_headers" in header_filter
assert "client_response_headers" in batch_one.executed_statements[0]
assert batch_one.committed is True
assert batch_one.closed is True
assert batch_two.closed is True
@pytest.mark.asyncio
async def test_perform_cleanup_deletes_first_and_uses_non_overlapping_windows(
monkeypatch: pytest.MonkeyPatch,
) -> None:
scheduler = MaintenanceScheduler()
fixed_now = datetime(2026, 3, 18, 3, 0, 0, tzinfo=timezone.utc)
class _FakeDateTime(datetime):
@classmethod
def now(cls, tz=None): # type: ignore[override]
if tz is None:
return fixed_now.replace(tzinfo=None)
return fixed_now.astimezone(tz)
class _FakeLoop:
async def run_in_executor(self, _executor, func): # type: ignore[no-untyped-def]
return func()
class _ConfigSession:
def close(self) -> None:
return None
calls: list[tuple[str, datetime, int, datetime | None]] = []
def _record(name: str, count: int):
def _inner(
cutoff_time: datetime,
batch_size: int,
*,
newer_than: datetime | None = None,
) -> int:
calls.append((name, cutoff_time, batch_size, newer_than))
return count
return _inner
config_values = {
"enable_auto_cleanup": True,
"detail_log_retention_days": 7,
"compressed_log_retention_days": 30,
"header_retention_days": 90,
"log_retention_days": 365,
"cleanup_batch_size": 123,
"auto_delete_expired_keys": False,
}
monkeypatch.setattr(maintenance_scheduler_module, "datetime", _FakeDateTime)
monkeypatch.setattr(
maintenance_scheduler_module.asyncio, "get_running_loop", lambda: _FakeLoop()
)
monkeypatch.setattr(
maintenance_scheduler_module,
"create_session",
lambda: _ConfigSession(),
)
monkeypatch.setattr(
maintenance_scheduler_module.SystemConfigService,
"get_config",
lambda _db, key, default=None: config_values.get(key, default),
)
monkeypatch.setattr(
scheduler,
"_delete_old_records",
lambda cutoff_time, batch_size: calls.append(("delete", cutoff_time, batch_size, None))
or 5,
)
monkeypatch.setattr(
scheduler,
"_cleanup_header_fields",
_record("header", 4),
)
monkeypatch.setattr(
scheduler,
"_cleanup_stale_body_fields",
_record("body", 3),
)
monkeypatch.setattr(
scheduler,
"_cleanup_body_fields",
_record("compress", 2),
)
monkeypatch.setattr(
maintenance_scheduler_module.ApiKeyService,
"cleanup_expired_keys",
lambda _db, auto_delete=False: 0,
)
await scheduler._perform_cleanup()
detail_cutoff = fixed_now - timedelta(days=7)
compressed_cutoff = fixed_now - timedelta(days=30)
header_cutoff = fixed_now - timedelta(days=90)
log_cutoff = fixed_now - timedelta(days=365)
assert calls == [
("delete", log_cutoff, 123, None),
("header", header_cutoff, 123, log_cutoff),
("body", compressed_cutoff, 123, log_cutoff),
("compress", detail_cutoff, 123, compressed_cutoff),
]

View File

@@ -0,0 +1,173 @@
from __future__ import annotations
from datetime import datetime, timedelta, timezone
from types import SimpleNamespace
from typing import Any, cast
from unittest.mock import MagicMock
import pytest
from src.services.health.endpoint import EndpointHealthService
class _FakeQuery:
def __init__(self, rows: list[SimpleNamespace]) -> None:
self._rows = rows
def filter(self, *args: Any, **kwargs: Any) -> _FakeQuery:
return self
def group_by(self, *args: Any, **kwargs: Any) -> _FakeQuery:
return self
def all(self) -> list[SimpleNamespace]:
return self._rows
class _FakeDb:
def __init__(self, rows: list[SimpleNamespace]) -> None:
self._rows = rows
def query(self, *args: Any, **kwargs: Any) -> _FakeQuery:
return _FakeQuery(self._rows)
def _expr_texts(query: MagicMock) -> list[str]:
return [str(arg) for arg in query.filter.call_args.args]
def test_generate_timeline_batch_keeps_compact_and_cli_isolated() -> None:
now = datetime(2026, 3, 18, 12, 0, tzinfo=timezone.utc)
db = _FakeDb(
[
SimpleNamespace(
endpoint_id="endpoint-compact",
segment_idx=0,
total_count=2,
success_count=2,
failed_count=0,
min_time=now - timedelta(minutes=55),
max_time=now - timedelta(minutes=40),
),
SimpleNamespace(
endpoint_id="endpoint-cli",
segment_idx=0,
total_count=3,
success_count=0,
failed_count=3,
min_time=now - timedelta(minutes=54),
max_time=now - timedelta(minutes=39),
),
]
)
result = EndpointHealthService._generate_timeline_batch(
db=cast(Any, db),
format_endpoint_mapping={
"openai:compact": ["endpoint-compact"],
"openai:cli": ["endpoint-cli"],
},
now=now,
lookback_hours=1,
segments=4,
)
assert result["openai:compact"]["timeline"][0] == "healthy"
assert result["openai:cli"]["timeline"][0] == "unhealthy"
def test_generate_timeline_from_usage_uses_endpoint_ids_directly(
monkeypatch: pytest.MonkeyPatch,
) -> None:
expected = {
"timeline": ["healthy", "warning"],
"time_range_start": "start",
"time_range_end": "end",
}
captured: dict[str, object] = {}
def _fake_generate_timeline_batch(
db: Any,
format_endpoint_mapping: dict[str, list[str]],
now: datetime,
lookback_hours: int,
segments: int,
) -> dict[str, dict[str, Any]]:
captured["db"] = db
captured["mapping"] = format_endpoint_mapping
captured["lookback_hours"] = lookback_hours
captured["segments"] = segments
return {"_single": expected}
monkeypatch.setattr(
EndpointHealthService,
"_generate_timeline_batch",
staticmethod(_fake_generate_timeline_batch),
)
db = cast(Any, object())
now = datetime(2026, 3, 18, 12, 0, tzinfo=timezone.utc)
result = EndpointHealthService._generate_timeline_from_usage(
db=db,
endpoint_ids=["endpoint-compact"],
now=now,
lookback_hours=6,
segments=2,
)
assert result == expected
assert captured["db"] is db
assert captured["mapping"] == {"_single": ["endpoint-compact"]}
assert captured["lookback_hours"] == 6
assert captured["segments"] == 2
def test_get_endpoint_health_by_format_filters_inactive_endpoints(
monkeypatch: pytest.MonkeyPatch,
) -> None:
endpoint_query = MagicMock()
endpoint_query.join.return_value = endpoint_query
endpoint_query.filter.return_value = endpoint_query
endpoint_query.all.return_value = [
SimpleNamespace(
id="endpoint-compact",
provider_id="provider-1",
api_format="openai:compact",
is_active=True,
)
]
key_query = MagicMock()
key_query.filter.return_value = key_query
key_query.options.return_value = key_query
key_query.all.return_value = []
db = MagicMock()
db.query.side_effect = [endpoint_query, key_query]
monkeypatch.setattr(
EndpointHealthService,
"_generate_timeline_batch",
staticmethod(
lambda db, format_endpoint_mapping, now, lookback_hours: {
"openai:compact": {
"timeline": ["unknown"] * 100,
"time_range_start": None,
"time_range_end": None,
}
}
),
)
EndpointHealthService.get_endpoint_health_by_format(
db=cast(Any, db),
lookback_hours=6,
include_admin_fields=False,
use_cache=False,
)
filters = _expr_texts(endpoint_query)
assert any("provider_endpoints.is_active" in expr for expr in filters)
assert any("providers.is_active" in expr for expr in filters)