mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-01 17:00:21 +08:00
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:
@@ -49,6 +49,10 @@ ADMIN_PASSWORD=admin123456
|
||||
# max-requests-jitter 会自动设置为 MAX_REQUESTS/20 (5%)
|
||||
# MAX_REQUESTS=4000
|
||||
|
||||
# glibc malloc arena 上限(默认 2)
|
||||
# 降低 malloc 内存碎片,减少 gunicorn worker RSS
|
||||
# MALLOC_ARENA_MAX=2
|
||||
|
||||
# HTTP 连接池上限(默认总预算约 200,按 worker 平分)
|
||||
# 如果容器内存偏高,可继续下调;例如 2 worker 时设为 80-100
|
||||
# HTTP_MAX_CONNECTIONS=100
|
||||
@@ -74,6 +78,11 @@ ADMIN_PASSWORD=admin123456
|
||||
# - 需要更多调试上下文:4
|
||||
# RESPONSE_CHUNKS_MAX_SIZE_MB=2
|
||||
|
||||
# 流式空闲超时(单位秒,默认 30)
|
||||
# 当流已经开始但连续一段时间没有任何新 chunk 时,提前中断并返回 504,
|
||||
# 避免一直等到 worker 超时(如 300s)
|
||||
# STREAM_IDLE_TIMEOUT_SECONDS=30
|
||||
|
||||
# API Key 前缀(默认 sk)
|
||||
# API_KEY_PREFIX=sk
|
||||
|
||||
|
||||
@@ -25,7 +25,12 @@ RUN --mount=type=cache,target=/var/cache/apt,sharing=locked \
|
||||
nginx \
|
||||
supervisor \
|
||||
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 包
|
||||
COPY --from=builder /usr/local/lib/python3.13/site-packages /usr/local/lib/python3.13/site-packages
|
||||
# 只复制需要的 Python 可执行文件
|
||||
@@ -247,7 +252,7 @@ RUN printf '%s\n' \
|
||||
'stdout_logfile_maxbytes=0' \
|
||||
'stderr_logfile=/dev/stderr' \
|
||||
'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]' \
|
||||
'command=/usr/local/bin/aether-hub --bind 0.0.0.0:8085' \
|
||||
@@ -268,6 +273,8 @@ ENV PYTHONUNBUFFERED=1 \
|
||||
PYTHONIOENCODING=utf-8 \
|
||||
LANG=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 \
|
||||
GUNICORN_WORKERS=2 \
|
||||
MAX_REQUESTS=4000
|
||||
|
||||
@@ -34,7 +34,12 @@ RUN --mount=type=cache,target=/var/cache/apt,sharing=locked \
|
||||
nginx \
|
||||
supervisor \
|
||||
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 包
|
||||
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' \
|
||||
'stderr_logfile=/dev/stderr' \
|
||||
'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]' \
|
||||
'command=/usr/local/bin/aether-hub --bind 0.0.0.0:8085' \
|
||||
@@ -294,6 +299,8 @@ ENV PYTHONUNBUFFERED=1 \
|
||||
PYTHONIOENCODING=utf-8 \
|
||||
LANG=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 \
|
||||
GUNICORN_WORKERS=2 \
|
||||
MAX_REQUESTS=4000
|
||||
|
||||
@@ -84,15 +84,24 @@ export interface CodexResetStatus {
|
||||
* @param resetSecs 相对剩余秒数(用于 fallback)
|
||||
* @param updatedAt 元数据更新时间(Unix 秒)
|
||||
* @param _tick 响应式触发器(传入 tick.value 以触发响应式更新)
|
||||
* @param remainingPercent 当前窗口剩余额度百分比(0-100,100 表示满额不启动倒计时)
|
||||
*/
|
||||
export function getCodexResetCountdown(
|
||||
resetAt: number | null | undefined,
|
||||
resetSecs: number | null | undefined,
|
||||
updatedAt: number | null | undefined,
|
||||
_tick: number
|
||||
_tick: number,
|
||||
remainingPercent?: number | null
|
||||
): CodexResetStatus | null {
|
||||
void _tick
|
||||
|
||||
if (remainingPercent != null) {
|
||||
const normalizedRemaining = Number(remainingPercent)
|
||||
if (Number.isFinite(normalizedRemaining) && normalizedRemaining >= 100) {
|
||||
return null
|
||||
}
|
||||
}
|
||||
|
||||
const nowSec = Math.floor(Date.now() / 1000)
|
||||
let remaining: number
|
||||
|
||||
|
||||
@@ -573,18 +573,20 @@
|
||||
/>
|
||||
</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="getResetCountdownClass(
|
||||
key.upstream_metadata.codex.primary_reset_at,
|
||||
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(
|
||||
key.upstream_metadata.codex.primary_reset_at,
|
||||
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>
|
||||
@@ -604,18 +606,21 @@
|
||||
/>
|
||||
</div>
|
||||
<div
|
||||
v-if="shouldStartCodexResetCountdown(key.upstream_metadata.codex.secondary_used_percent)"
|
||||
class="text-[9px] mt-0.5 tabular-nums"
|
||||
:class="getResetCountdownClass(
|
||||
key.upstream_metadata.codex.secondary_reset_at,
|
||||
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">
|
||||
{{ getResetCountdownText(
|
||||
key.upstream_metadata.codex.secondary_reset_at,
|
||||
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 v-else>
|
||||
@@ -2549,9 +2554,16 @@ function getAntigravityQuotaSummary(metadata: UpstreamMetadata | null | undefine
|
||||
function getResetCountdownText(
|
||||
resetAt: number | null | undefined,
|
||||
resetSecs: number | null | undefined,
|
||||
updatedAt: number | null | undefined
|
||||
updatedAt: number | null | undefined,
|
||||
usedPercent: number | null | undefined
|
||||
): string {
|
||||
const status = getCodexResetCountdown(resetAt, resetSecs, updatedAt, countdownTick.value)
|
||||
const status = getCodexResetCountdown(
|
||||
resetAt,
|
||||
resetSecs,
|
||||
updatedAt,
|
||||
countdownTick.value,
|
||||
toCodexRemainingPercent(usedPercent)
|
||||
)
|
||||
if (!status) return ''
|
||||
return status.isExpired ? status.text : `${status.text} 后重置`
|
||||
}
|
||||
@@ -2559,15 +2571,35 @@ function getResetCountdownText(
|
||||
function getResetCountdownClass(
|
||||
resetAt: number | null | undefined,
|
||||
resetSecs: number | null | undefined,
|
||||
updatedAt: number | null | undefined
|
||||
updatedAt: number | null | undefined,
|
||||
usedPercent: number | null | undefined
|
||||
): 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.isCritical) return 'text-destructive font-medium animate-pulse'
|
||||
if (status.isUrgent) return 'text-amber-500 dark:text-amber-400'
|
||||
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 {
|
||||
const days = Math.floor(seconds / 86400)
|
||||
|
||||
@@ -2428,7 +2428,7 @@ function getQuotaProgressLabel(label: string): string {
|
||||
|
||||
function getQuotaProgressCountdown(item: QuotaProgressItem) {
|
||||
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 {
|
||||
@@ -2438,6 +2438,9 @@ function getQuotaProgressCountdownText(item: QuotaProgressItem): string {
|
||||
}
|
||||
|
||||
function getQuotaProgressTooltip(item: QuotaProgressItem): string {
|
||||
if ((item.label === '5H' || item.label === '周') && item.remainingPercent >= 100) {
|
||||
return ''
|
||||
}
|
||||
const detail = item.detail?.trim() || ''
|
||||
const countdownText = getQuotaProgressCountdownText(item)
|
||||
if (countdownText) {
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
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()
|
||||
|
||||
|
||||
@@ -379,7 +405,10 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
||||
func.count(RequestCandidate.id).label("count"),
|
||||
)
|
||||
.join(RequestCandidate, ProviderEndpoint.id == RequestCandidate.endpoint_id)
|
||||
.join(Provider, ProviderEndpoint.provider_id == Provider.id)
|
||||
.filter(
|
||||
ProviderEndpoint.is_active.is_(True),
|
||||
Provider.is_active.is_(True),
|
||||
RequestCandidate.created_at >= since,
|
||||
RequestCandidate.status.in_(final_statuses),
|
||||
)
|
||||
@@ -395,40 +424,15 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
||||
status_counts[fmt] = {"success": 0, "failed": 0, "skipped": 0}
|
||||
status_counts[fmt][status] = count
|
||||
|
||||
# 3. 获取最近一段时间的 RequestCandidate(限制数量)
|
||||
# 使用上面定义的 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. 为所有活跃格式生成监控数据(包括没有请求记录的)
|
||||
# 3. 为所有活跃格式生成监控数据(包括没有请求记录的)
|
||||
monitors: list[ApiFormatHealthMonitor] = []
|
||||
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
|
||||
# 中间状态(available, pending, used, started)不计入统计
|
||||
|
||||
@@ -485,7 +485,7 @@ class BaseMessageHandler:
|
||||
api_format: str | None = None,
|
||||
request_headers: dict[str, Any] | None = None,
|
||||
request_body: dict[str, Any] | None = None,
|
||||
) -> None:
|
||||
) -> bool:
|
||||
"""在请求开始时创建 pending 状态的 Usage 记录
|
||||
|
||||
让前端可以立即看到"处理中"的请求,提升用户体验。
|
||||
@@ -498,6 +498,8 @@ class BaseMessageHandler:
|
||||
api_format: API 格式
|
||||
request_headers: 原始请求头
|
||||
request_body: 原始请求体
|
||||
Returns:
|
||||
bool: True 表示已成功创建;False 表示创建失败(调用方可按需回退处理)。
|
||||
"""
|
||||
try:
|
||||
UsageService.create_pending_usage(
|
||||
@@ -512,9 +514,11 @@ class BaseMessageHandler:
|
||||
request_headers=request_headers,
|
||||
request_body=request_body,
|
||||
)
|
||||
return True
|
||||
except Exception as 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:
|
||||
"""更新 Usage 状态为 streaming(流式传输开始时调用)
|
||||
|
||||
@@ -487,7 +487,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
model = getattr(converted_request, "model", original_request_body.get("model", "unknown"))
|
||||
|
||||
# 提前创建 pending 记录,让前端可以立即看到"处理中"
|
||||
self._create_pending_usage(
|
||||
pending_usage_created = self._create_pending_usage(
|
||||
model=model,
|
||||
is_stream=True,
|
||||
request_type="chat",
|
||||
@@ -579,6 +579,8 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
request_body_state=request_state,
|
||||
request_headers=original_headers,
|
||||
request_body=original_request_body,
|
||||
# 预创建失败时,回退到 TaskService 侧创建,避免丢失 pending 状态。
|
||||
create_pending_usage=not pending_usage_created,
|
||||
)
|
||||
stream_generator = exec_result.response
|
||||
provider_name = exec_result.provider_name or "unknown"
|
||||
|
||||
@@ -113,7 +113,7 @@ class ChatSyncExecutor:
|
||||
api_format = handler.allowed_api_formats[0]
|
||||
|
||||
# 提前创建 pending 记录,让前端可以立即看到"处理中"
|
||||
handler._create_pending_usage(
|
||||
pending_usage_created = handler._create_pending_usage(
|
||||
model=model,
|
||||
is_stream=False,
|
||||
request_type="chat",
|
||||
@@ -176,6 +176,8 @@ class ChatSyncExecutor:
|
||||
request_body_state=request_state,
|
||||
request_headers=original_headers,
|
||||
request_body=original_request_body,
|
||||
# 预创建失败时,回退到 TaskService 侧创建,避免丢失 pending 状态。
|
||||
create_pending_usage=not pending_usage_created,
|
||||
)
|
||||
actual_provider_name = exec_result.provider_name or "unknown"
|
||||
ctx.provider_id = exec_result.provider_id
|
||||
|
||||
@@ -3,6 +3,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import time
|
||||
from collections.abc import AsyncGenerator
|
||||
from typing import TYPE_CHECKING, Any
|
||||
@@ -29,12 +30,23 @@ if TYPE_CHECKING:
|
||||
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:
|
||||
"""监控和统计相关方法的 Mixin"""
|
||||
|
||||
# CancelledError 归因时,断连检查参数(秒)
|
||||
CANCEL_DISCONNECT_CHECK_TIMEOUT_SECONDS = 0.5
|
||||
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(
|
||||
self,
|
||||
@@ -115,8 +127,29 @@ class CliMonitorMixin:
|
||||
|
||||
last_chunk_time = time_module.time()
|
||||
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:
|
||||
idle_watch_task = asyncio.create_task(watch_stream_idle_timeout())
|
||||
if http_request is not None:
|
||||
# 使用后台任务检测断连,完全不阻塞流式传输
|
||||
disconnected = False
|
||||
@@ -149,6 +182,7 @@ class CliMonitorMixin:
|
||||
ctx.status_code = 499
|
||||
ctx.error_message = "client_disconnected"
|
||||
break
|
||||
stream_started = True
|
||||
last_chunk_time = time_module.time()
|
||||
chunk_count += 1
|
||||
yield chunk
|
||||
@@ -158,14 +192,33 @@ class CliMonitorMixin:
|
||||
await check_task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
if idle_watch_task is not None:
|
||||
idle_watch_task.cancel()
|
||||
try:
|
||||
await idle_watch_task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
else:
|
||||
# 无 http_request,仅被动监控
|
||||
async for chunk in stream_generator:
|
||||
last_chunk_time = time_module.time()
|
||||
chunk_count += 1
|
||||
yield chunk
|
||||
try:
|
||||
async for chunk in stream_generator:
|
||||
stream_started = True
|
||||
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:
|
||||
# 防御性清理:正常路径中 idle_watch_task 已由内部 finally 取消,
|
||||
# 但若 CancelledError 在异常路径传播,确保不留孤儿 task。
|
||||
if idle_watch_task is not None and not idle_watch_task.done():
|
||||
idle_watch_task.cancel()
|
||||
# 注意:CancelledError 不等于"用户手动取消",它既可能是客户端断连触发,
|
||||
# 也可能是服务端(重载/关停/内部取消)导致的协程取消。
|
||||
# 这里尽量做一次"断连归因":仅当能确认客户端已断开时才记为 499 cancelled。
|
||||
@@ -173,6 +226,28 @@ class CliMonitorMixin:
|
||||
if not ctx.has_completion:
|
||||
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
|
||||
disconnect_check_uncertain = False
|
||||
if http_request is not None:
|
||||
@@ -229,10 +304,14 @@ class CliMonitorMixin:
|
||||
)
|
||||
raise
|
||||
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.error_message = str(e)
|
||||
raise
|
||||
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.error_message = str(e)
|
||||
raise
|
||||
@@ -294,59 +373,149 @@ class CliMonitorMixin:
|
||||
ctx, ctx.provider_request_body or original_request_body
|
||||
)
|
||||
|
||||
response_body = ctx.build_response_body(response_time_ms)
|
||||
client_response_body = ctx.build_client_response_body(response_time_ms)
|
||||
with ctx.managed_recorded_bodies(response_time_ms) as recorded_bodies:
|
||||
# 根据状态码决定记录成功还是失败
|
||||
# 499 = 客户端取消(不算系统失败);其他 4xx/5xx 视为失败
|
||||
if ctx.status_code and ctx.status_code >= 400:
|
||||
client_response_headers = ctx.client_response_headers or {
|
||||
"content-type": "application/json"
|
||||
}
|
||||
|
||||
# 根据状态码决定记录成功还是失败
|
||||
# 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():
|
||||
# 客户端取消:记录为 cancelled(不算系统失败)
|
||||
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_cancelled(
|
||||
provider=ctx.provider_name or "unknown",
|
||||
model=ctx.model,
|
||||
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,
|
||||
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,
|
||||
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,
|
||||
has_format_conversion=ctx.has_format_conversion,
|
||||
target_model=ctx.mapped_model,
|
||||
request_metadata=request_metadata,
|
||||
)
|
||||
logger.debug("{} 流式响应被客户端取消", self.FORMAT_ID)
|
||||
logger.info(
|
||||
f"[CANCEL] {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}"
|
||||
)
|
||||
if ctx.is_client_disconnected():
|
||||
# 客户端取消:记录为 cancelled(不算系统失败)
|
||||
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_cancelled(
|
||||
provider=ctx.provider_name or "unknown",
|
||||
model=ctx.model,
|
||||
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,
|
||||
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,
|
||||
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(
|
||||
"[CANCEL] {} | {} | {} | {}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:
|
||||
# 服务端/上游异常:记录为失败
|
||||
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:
|
||||
# 服务端/上游异常:记录为失败
|
||||
# 在记录统计前,允许子类从 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(
|
||||
{"perf": ctx.perf_metrics} if ctx.perf_metrics else None,
|
||||
selected_key_id=ctx.key_id,
|
||||
@@ -354,129 +523,60 @@ class CliMonitorMixin:
|
||||
pool_summary=ctx.pool_summary,
|
||||
fallback_from_request=True,
|
||||
)
|
||||
await bg_telemetry.record_failure(
|
||||
provider=ctx.provider_name or "unknown",
|
||||
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,
|
||||
error_message=ctx.error_message or f"HTTP {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=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,
|
||||
provider_request_headers=ctx.provider_request_headers,
|
||||
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=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,
|
||||
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("{} 流式响应中断", 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(
|
||||
"[{}] 流式转换完成: {}->{}, total_events={}",
|
||||
self.request_id[:8],
|
||||
ctx.provider_api_format,
|
||||
ctx.client_api_format,
|
||||
ctx.stream_conversion_event_count,
|
||||
"[{}] 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,
|
||||
)
|
||||
|
||||
# 流未正常完成(如上游截断/连接中断)且无 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 会在流开始时过早地标记成功(只记录了连接建立的时间)
|
||||
@@ -556,6 +656,9 @@ class CliMonitorMixin:
|
||||
|
||||
except Exception as e:
|
||||
logger.exception("记录流式统计信息时出错")
|
||||
finally:
|
||||
# 遥测写入完成后主动释放大对象列表,降低高并发长流的内存滞留。
|
||||
ctx.release_recorded_chunks()
|
||||
|
||||
async def _record_stream_failure(
|
||||
self,
|
||||
@@ -591,26 +694,30 @@ class CliMonitorMixin:
|
||||
pool_summary=ctx.pool_summary,
|
||||
fallback_from_request=True,
|
||||
)
|
||||
await self.telemetry.record_failure(
|
||||
provider=ctx.provider_name or "unknown",
|
||||
model=ctx.model,
|
||||
response_time_ms=response_time_ms,
|
||||
status_code=status_code,
|
||||
error_message=extract_client_error_message(error),
|
||||
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,
|
||||
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,
|
||||
)
|
||||
try:
|
||||
await self.telemetry.record_failure(
|
||||
provider=ctx.provider_name or "unknown",
|
||||
model=ctx.model,
|
||||
response_time_ms=response_time_ms,
|
||||
status_code=status_code,
|
||||
error_message=extract_client_error_message(error),
|
||||
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,
|
||||
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,
|
||||
)
|
||||
finally:
|
||||
# 失败路径同样可能持有 chunk 审计数据,及时释放。
|
||||
ctx.release_recorded_chunks()
|
||||
|
||||
@@ -81,7 +81,7 @@ class CliHandlerProtocol(Protocol):
|
||||
api_format: str | None = ...,
|
||||
request_headers: dict[str, Any] | None = ...,
|
||||
request_body: dict[str, Any] | None = ...,
|
||||
) -> None: ...
|
||||
) -> bool: ...
|
||||
|
||||
def _build_request_metadata(
|
||||
self,
|
||||
|
||||
@@ -92,7 +92,7 @@ class CliStreamMixin:
|
||||
client_api_format = self.primary_api_format
|
||||
|
||||
# 提前创建 pending 记录,让前端可以立即看到"处理中"
|
||||
self._create_pending_usage(
|
||||
pending_usage_created = self._create_pending_usage(
|
||||
model=model,
|
||||
is_stream=True,
|
||||
request_type="chat",
|
||||
@@ -168,6 +168,8 @@ class CliStreamMixin:
|
||||
request_body_state=request_state,
|
||||
request_headers=original_headers,
|
||||
request_body=original_request_body,
|
||||
# 预创建失败时,回退到 TaskService 侧创建,避免丢失 pending 状态。
|
||||
create_pending_usage=not pending_usage_created,
|
||||
)
|
||||
stream_generator = exec_result.response
|
||||
provider_name = exec_result.provider_name or "unknown"
|
||||
@@ -260,8 +262,7 @@ class CliStreamMixin:
|
||||
) -> AsyncGenerator[bytes]:
|
||||
"""执行流式请求并返回流生成器"""
|
||||
# 重置上下文状态(重试时清除之前的数据,避免累积)
|
||||
ctx.parsed_chunks = []
|
||||
ctx.provider_parsed_chunks = []
|
||||
ctx.release_recorded_chunks()
|
||||
ctx.chunk_count = 0
|
||||
ctx.data_count = 0
|
||||
ctx.has_completion = False
|
||||
|
||||
@@ -75,7 +75,7 @@ class CliSyncMixin:
|
||||
sync_start_time = time.time()
|
||||
|
||||
# 提前创建 pending 记录,让前端可以立即看到"处理中"
|
||||
self._create_pending_usage(
|
||||
pending_usage_created = self._create_pending_usage(
|
||||
model=model,
|
||||
is_stream=False,
|
||||
request_type="chat",
|
||||
@@ -390,6 +390,8 @@ class CliSyncMixin:
|
||||
request_body_state=request_state,
|
||||
request_headers=original_headers,
|
||||
request_body=original_request_body,
|
||||
# 预创建失败时,回退到 TaskService 侧创建,避免丢失 pending 状态。
|
||||
create_pending_usage=not pending_usage_created,
|
||||
)
|
||||
result = exec_result.response
|
||||
actual_provider_name = exec_result.provider_name or "unknown"
|
||||
|
||||
@@ -12,8 +12,9 @@ from __future__ import annotations
|
||||
|
||||
import json
|
||||
import time
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass, field
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from typing import TYPE_CHECKING, Any, Iterator
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.core.api_format.conversion.stream_state import StreamState
|
||||
@@ -49,6 +50,21 @@ def is_format_converted(
|
||||
_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
|
||||
class StreamContext:
|
||||
"""
|
||||
@@ -160,8 +176,7 @@ class StreamContext:
|
||||
在故障转移重试时调用,清除之前的数据避免累积。
|
||||
保留 model 和 api_format,重置其他所有状态。
|
||||
"""
|
||||
self.parsed_chunks = []
|
||||
self.provider_parsed_chunks = []
|
||||
self.release_recorded_chunks()
|
||||
self.chunk_count = 0
|
||||
self.data_count = 0
|
||||
self.has_completion = False
|
||||
@@ -194,6 +209,32 @@ class StreamContext:
|
||||
self.needs_conversion = False
|
||||
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
|
||||
def collected_text(self) -> str:
|
||||
"""已收集的文本内容(按需拼接,避免在流式过程中频繁做字符串拷贝)"""
|
||||
|
||||
@@ -99,6 +99,7 @@ class StreamTelemetryRecorder:
|
||||
try:
|
||||
writer = await self._get_telemetry_writer(bg_db, ctx, response_time_ms)
|
||||
if writer is None:
|
||||
ctx.release_recorded_chunks()
|
||||
return
|
||||
# 兜底估算:流未正常完成且 token 均为 0 时,从请求体粗略估算。
|
||||
# 覆盖成功但缺少 completion,以及已传出部分数据后被中断的场景。
|
||||
@@ -114,53 +115,47 @@ class StreamTelemetryRecorder:
|
||||
if isinstance(writer, QueueTelemetryWriter)
|
||||
else should_log_body
|
||||
)
|
||||
response_body = (
|
||||
ctx.build_response_body(response_time_ms) if include_bodies else None
|
||||
)
|
||||
client_response_body = (
|
||||
ctx.build_client_response_body(response_time_ms) if include_bodies else None
|
||||
)
|
||||
|
||||
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(
|
||||
with ctx.managed_recorded_bodies(
|
||||
response_time_ms, include_bodies=include_bodies
|
||||
) as recorded_bodies:
|
||||
try:
|
||||
await self._dispatch_record(
|
||||
bg_db,
|
||||
status=self._get_status_from_ctx(ctx),
|
||||
response_time_ms=response_time_ms,
|
||||
status_code=ctx.status_code,
|
||||
writer,
|
||||
ctx,
|
||||
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)
|
||||
@@ -175,6 +170,9 @@ class StreamTelemetryRecorder:
|
||||
response_time_ms=response_time_ms,
|
||||
error_message=f"记录统计信息失败: {str(e)[:200]}",
|
||||
)
|
||||
finally:
|
||||
# 遥测写入后主动释放大列表,避免长流式请求对象滞留在 worker 堆中。
|
||||
ctx.release_recorded_chunks()
|
||||
|
||||
async def _record_success(
|
||||
self,
|
||||
|
||||
@@ -47,6 +47,32 @@ router = APIRouter(prefix="/api/public", tags=["System Catalog"])
|
||||
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")
|
||||
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)
|
||||
|
||||
# 2. 获取最近一段时间的 RequestCandidate(限制数量)
|
||||
# 只查询最终状态的记录:success, failed, skipped
|
||||
# 2. 统计窗口内每个 API 格式的真实状态分布
|
||||
final_statuses = ["success", "failed", "skipped"]
|
||||
limit_rows = max(500, self.per_format_limit * 10)
|
||||
rows = (
|
||||
status_counts_query = (
|
||||
db.query(
|
||||
RequestCandidate,
|
||||
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(
|
||||
ProviderEndpoint.is_active.is_(True),
|
||||
Provider.is_active.is_(True),
|
||||
RequestCandidate.created_at >= since,
|
||||
RequestCandidate.status.in_(final_statuses),
|
||||
)
|
||||
.order_by(RequestCandidate.created_at.desc())
|
||||
.limit(limit_rows)
|
||||
.group_by(ProviderEndpoint.api_format, RequestCandidate.status)
|
||||
.all()
|
||||
)
|
||||
|
||||
grouped_candidates: dict[str, list[RequestCandidate]] = {}
|
||||
|
||||
for candidate, api_format_enum in rows:
|
||||
status_counts: dict[str, dict[str, int]] = {}
|
||||
for api_format_enum, status, count in status_counts_query:
|
||||
api_format = (
|
||||
api_format_enum.value if hasattr(api_format_enum, "value") else str(api_format_enum)
|
||||
)
|
||||
if api_format not in grouped_candidates:
|
||||
grouped_candidates[api_format] = []
|
||||
|
||||
if len(grouped_candidates[api_format]) < self.per_format_limit:
|
||||
grouped_candidates[api_format].append(candidate)
|
||||
if api_format not in status_counts:
|
||||
status_counts[api_format] = {"success": 0, "failed": 0, "skipped": 0}
|
||||
status_counts[api_format][status] = count
|
||||
|
||||
# 3. 为所有活跃格式生成监控数据
|
||||
monitors: list[PublicApiFormatHealthMonitor] = []
|
||||
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,
|
||||
)
|
||||
|
||||
# 统计
|
||||
success_count = sum(1 for c in candidates if c.status == "success")
|
||||
failed_count = sum(1 for c in candidates if c.status == "failed")
|
||||
skipped_count = sum(1 for c in candidates if c.status == "skipped")
|
||||
total_attempts = len(candidates)
|
||||
# 统计使用窗口内真实总数,events 仅保留最近样本用于展示。
|
||||
format_stats = status_counts.get(api_format, {"success": 0, "failed": 0, "skipped": 0})
|
||||
success_count = format_stats.get("success", 0)
|
||||
failed_count = format_stats.get("failed", 0)
|
||||
skipped_count = format_stats.get("skipped", 0)
|
||||
total_attempts = success_count + failed_count + skipped_count
|
||||
|
||||
# 计算成功率 = success / (success + failed)
|
||||
actual_completed = success_count + failed_count
|
||||
|
||||
@@ -49,7 +49,8 @@ class PluginManager:
|
||||
# notification 默认不加载,避免未配置插件(如 email)在启动时初始化失败并占用内存。
|
||||
DEFAULT_ENABLED_PLUGIN_MODULES: dict[str, tuple[str, ...]] = {
|
||||
"auth": ("api_key",),
|
||||
"rate_limit": ("sliding_window",),
|
||||
# 默认切到 token_bucket;全局 Redis 就绪时会自动切到分布式后端。
|
||||
"rate_limit": ("token_bucket",),
|
||||
"cache": ("memory",),
|
||||
"monitor": ("prometheus",),
|
||||
"token": ("claude",),
|
||||
|
||||
@@ -29,6 +29,7 @@ class TokenBucket:
|
||||
self.refill_rate = refill_rate
|
||||
self.tokens = capacity
|
||||
self.last_refill = time.time()
|
||||
self.last_access_time = self.last_refill
|
||||
|
||||
def _refill(self) -> None:
|
||||
"""补充令牌"""
|
||||
@@ -36,6 +37,7 @@ class TokenBucket:
|
||||
time_passed = now - self.last_refill
|
||||
tokens_to_add = time_passed * self.refill_rate
|
||||
|
||||
self.last_access_time = now
|
||||
if tokens_to_add > 0:
|
||||
self.tokens = min(self.capacity, self.tokens + tokens_to_add)
|
||||
self.last_refill = now
|
||||
@@ -64,7 +66,7 @@ class TokenBucket:
|
||||
|
||||
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)
|
||||
|
||||
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:
|
||||
super().__init__("token_bucket")
|
||||
self.buckets: dict[str, TokenBucket] = {}
|
||||
@@ -90,11 +96,46 @@ class TokenBucketStrategy(RateLimitStrategy):
|
||||
# 默认配置
|
||||
self.default_capacity = 100 # 默认桶容量
|
||||
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 后端
|
||||
self._redis_backend: RedisTokenBucketBackend | None = None
|
||||
self._redis_checked = False
|
||||
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:
|
||||
"""
|
||||
@@ -107,41 +148,88 @@ class TokenBucketStrategy(RateLimitStrategy):
|
||||
Returns:
|
||||
令牌桶实例
|
||||
"""
|
||||
if key not in self.buckets:
|
||||
# 如果提供了rate_limit参数(来自数据库),优先使用
|
||||
if rate_limit is not None:
|
||||
# rate_limit 是每分钟请求数,转换为令牌桶参数
|
||||
capacity = rate_limit # 桶容量等于每分钟限制
|
||||
refill_rate = rate_limit / 60.0 # 每秒补充的令牌数
|
||||
# 否则根据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
|
||||
capacity, refill_rate = self._resolve_bucket_config(key, rate_limit)
|
||||
bucket = self.buckets.get(key)
|
||||
if bucket is None:
|
||||
bucket = TokenBucket(capacity, refill_rate)
|
||||
self.buckets[key] = bucket
|
||||
return bucket
|
||||
|
||||
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:
|
||||
return self._backend_mode in {"auto", "redis"}
|
||||
|
||||
async def _ensure_backend(self) -> None:
|
||||
if self._redis_checked:
|
||||
if self._redis_backend is not None:
|
||||
return
|
||||
self._redis_checked = True
|
||||
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
|
||||
redis_client = get_redis_client_sync()
|
||||
if redis_client:
|
||||
self._redis_backend = RedisTokenBucketBackend(redis_client)
|
||||
self._redis_checked = True
|
||||
self._next_redis_probe_time = 0.0
|
||||
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 客户端不可用,回退到内存桶")
|
||||
|
||||
async def check_limit(self, key: str, **kwargs: Any) -> RateLimitResult:
|
||||
@@ -155,11 +243,14 @@ class TokenBucketStrategy(RateLimitStrategy):
|
||||
Returns:
|
||||
速率限制检查结果
|
||||
"""
|
||||
await self._ensure_backend()
|
||||
|
||||
rate_limit = kwargs.get("rate_limit")
|
||||
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:
|
||||
return await self._redis_backend.peek(
|
||||
key=key,
|
||||
@@ -169,6 +260,7 @@ class TokenBucketStrategy(RateLimitStrategy):
|
||||
)
|
||||
|
||||
async with self._lock:
|
||||
await self._maybe_cleanup()
|
||||
bucket = self._get_bucket(key, rate_limit)
|
||||
remaining = bucket.get_remaining()
|
||||
reset_at = bucket.get_reset_time()
|
||||
@@ -203,13 +295,17 @@ class TokenBucketStrategy(RateLimitStrategy):
|
||||
Returns:
|
||||
是否成功消费
|
||||
"""
|
||||
rate_limit = kwargs.get("rate_limit")
|
||||
if self._is_unlimited_rate_limit(rate_limit):
|
||||
return True
|
||||
|
||||
await self._ensure_backend()
|
||||
|
||||
if self._redis_backend:
|
||||
success, remaining = await self._redis_backend.consume(
|
||||
key=key,
|
||||
capacity=self._resolve_capacity(key, kwargs.get("rate_limit")),
|
||||
refill_rate=self._resolve_refill_rate(key, kwargs.get("rate_limit")),
|
||||
capacity=self._resolve_capacity(key, rate_limit),
|
||||
refill_rate=self._resolve_refill_rate(key, rate_limit),
|
||||
amount=amount,
|
||||
)
|
||||
if success:
|
||||
@@ -219,7 +315,8 @@ class TokenBucketStrategy(RateLimitStrategy):
|
||||
return success
|
||||
|
||||
async with self._lock:
|
||||
bucket = self._get_bucket(key)
|
||||
await self._maybe_cleanup()
|
||||
bucket = self._get_bucket(key, rate_limit)
|
||||
success = bucket.consume(amount)
|
||||
|
||||
if success:
|
||||
@@ -270,6 +367,7 @@ class TokenBucketStrategy(RateLimitStrategy):
|
||||
)
|
||||
|
||||
async with self._lock:
|
||||
await self._maybe_cleanup()
|
||||
bucket = self._get_bucket(key)
|
||||
return {
|
||||
"strategy": "token_bucket",
|
||||
@@ -293,24 +391,17 @@ class TokenBucketStrategy(RateLimitStrategy):
|
||||
super().configure(config)
|
||||
self.default_capacity = config.get("default_capacity", self.default_capacity)
|
||||
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:
|
||||
if rate_limit is not None:
|
||||
return rate_limit
|
||||
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
|
||||
capacity, _ = self._resolve_bucket_config(key, rate_limit)
|
||||
return capacity
|
||||
|
||||
def _resolve_refill_rate(self, key: str, rate_limit: int | None = None) -> float:
|
||||
if rate_limit is not None:
|
||||
return rate_limit / 60.0
|
||||
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
|
||||
_, refill_rate = self._resolve_bucket_config(key, rate_limit)
|
||||
return refill_rate
|
||||
|
||||
|
||||
class RedisTokenBucketBackend:
|
||||
@@ -365,6 +456,9 @@ class RedisTokenBucketBackend:
|
||||
refill_rate: float,
|
||||
amount: int,
|
||||
) -> RateLimitResult:
|
||||
if capacity <= 0 or refill_rate <= 0:
|
||||
return RateLimitResult(allowed=True, remaining=0)
|
||||
|
||||
bucket_key = self._redis_key(key)
|
||||
data = await self.redis.hmget(bucket_key, "tokens", "timestamp")
|
||||
tokens = data[0]
|
||||
@@ -372,7 +466,7 @@ class RedisTokenBucketBackend:
|
||||
|
||||
if tokens is None or last_refill is None:
|
||||
remaining = capacity
|
||||
reset_at = datetime.now(timezone.utc) + timedelta(seconds=capacity / refill_rate)
|
||||
reset_at = datetime.now(timezone.utc)
|
||||
else:
|
||||
tokens_value = float(tokens)
|
||||
last_refill_value = float(last_refill)
|
||||
@@ -407,6 +501,9 @@ class RedisTokenBucketBackend:
|
||||
refill_rate: float,
|
||||
amount: int,
|
||||
) -> tuple[bool, int]:
|
||||
if capacity <= 0 or refill_rate <= 0:
|
||||
return True, 0
|
||||
|
||||
result = await self._consume_script(
|
||||
keys=[self._redis_key(key)],
|
||||
args=[time.time(), capacity, refill_rate, amount],
|
||||
|
||||
@@ -69,7 +69,13 @@ class EndpointHealthService:
|
||||
|
||||
# 查询所有活跃的端点(一次性获取所有需要的数据)
|
||||
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
|
||||
@@ -155,16 +161,13 @@ class EndpointHealthService:
|
||||
format_stats[api_format]["health_scores"].append(health_score)
|
||||
|
||||
# 批量生成所有格式的时间线数据
|
||||
all_key_ids = []
|
||||
format_key_mapping: dict[str, list[str]] = {}
|
||||
format_endpoint_mapping: dict[str, list[str]] = {}
|
||||
for api_format, stats in format_stats.items():
|
||||
key_ids = stats["key_ids"]
|
||||
format_key_mapping[api_format] = key_ids
|
||||
all_key_ids.extend(key_ids)
|
||||
format_endpoint_mapping[api_format] = stats["endpoint_ids"]
|
||||
|
||||
# 一次性查询所有时间线数据
|
||||
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
|
||||
def _generate_timeline_batch(
|
||||
db: Session,
|
||||
format_key_mapping: dict[str, list[str]],
|
||||
format_endpoint_mapping: dict[str, list[str]],
|
||||
now: datetime,
|
||||
lookback_hours: int,
|
||||
segments: int = 100,
|
||||
@@ -249,7 +252,7 @@ class EndpointHealthService:
|
||||
|
||||
Args:
|
||||
db: 数据库会话
|
||||
format_key_mapping: API格式 -> key_ids 的映射
|
||||
format_endpoint_mapping: API格式 -> endpoint_ids 的映射
|
||||
now: 当前时间
|
||||
lookback_hours: 回溯小时数
|
||||
segments: 时间段数量
|
||||
@@ -257,19 +260,19 @@ class EndpointHealthService:
|
||||
Returns:
|
||||
API格式 -> 时间线数据的映射
|
||||
"""
|
||||
# 收集所有 key_ids
|
||||
all_key_ids = []
|
||||
for key_ids in format_key_mapping.values():
|
||||
all_key_ids.extend(key_ids)
|
||||
# 收集所有 endpoint_ids
|
||||
all_endpoint_ids = []
|
||||
for endpoint_ids in format_endpoint_mapping.values():
|
||||
all_endpoint_ids.extend(endpoint_ids)
|
||||
|
||||
if not all_key_ids:
|
||||
if not all_endpoint_ids:
|
||||
return {
|
||||
api_format: {
|
||||
"timeline": ["unknown"] * 100,
|
||||
"time_range_start": 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) 保证,这里做防御性检查)
|
||||
@@ -293,7 +296,7 @@ class EndpointHealthService:
|
||||
|
||||
candidate_stats = (
|
||||
db.query(
|
||||
RequestCandidate.key_id,
|
||||
RequestCandidate.endpoint_id,
|
||||
segment_expr,
|
||||
func.count(RequestCandidate.id).label("total_count"),
|
||||
func.sum(case((RequestCandidate.status == "success", 1), else_=0)).label(
|
||||
@@ -306,20 +309,20 @@ class EndpointHealthService:
|
||||
func.max(RequestCandidate.created_at).label("max_time"),
|
||||
)
|
||||
.filter(
|
||||
RequestCandidate.key_id.in_(all_key_ids),
|
||||
RequestCandidate.endpoint_id.in_(all_endpoint_ids),
|
||||
RequestCandidate.created_at >= start_time,
|
||||
RequestCandidate.created_at <= now,
|
||||
RequestCandidate.status.in_(final_statuses),
|
||||
)
|
||||
.group_by(RequestCandidate.key_id, segment_expr)
|
||||
.group_by(RequestCandidate.endpoint_id, segment_expr)
|
||||
.all()
|
||||
)
|
||||
|
||||
# 构建 key_id -> api_format 的反向映射
|
||||
key_to_format: dict[str, str] = {}
|
||||
for api_format, key_ids in format_key_mapping.items():
|
||||
for key_id in key_ids:
|
||||
key_to_format[key_id] = api_format
|
||||
# 构建 endpoint_id -> api_format 的反向映射
|
||||
endpoint_to_format: dict[str, str] = {}
|
||||
for api_format, endpoint_ids in format_endpoint_mapping.items():
|
||||
for endpoint_id in endpoint_ids:
|
||||
endpoint_to_format[endpoint_id] = api_format
|
||||
|
||||
# 按 api_format 和 segment 聚合数据
|
||||
format_segment_data: dict[str, dict[int, dict]] = defaultdict(
|
||||
@@ -335,9 +338,9 @@ class EndpointHealthService:
|
||||
)
|
||||
|
||||
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
|
||||
api_format = key_to_format.get(key_id)
|
||||
api_format = endpoint_to_format.get(endpoint_id)
|
||||
|
||||
if api_format and 0 <= segment_idx < segments:
|
||||
seg_data = format_segment_data[api_format][segment_idx]
|
||||
@@ -355,7 +358,7 @@ class EndpointHealthService:
|
||||
# 生成各格式的时间线
|
||||
result: dict[str, dict[str, Any]] = {}
|
||||
|
||||
for api_format in format_key_mapping.keys():
|
||||
for api_format in format_endpoint_mapping.keys():
|
||||
timeline = []
|
||||
earliest_time = None
|
||||
latest_time = None
|
||||
@@ -427,46 +430,10 @@ class EndpointHealthService:
|
||||
"time_range_end": None,
|
||||
}
|
||||
|
||||
# 基于 endpoint_ids 反推 provider_ids 与 api_format,再选出支持该格式的 keys
|
||||
endpoint_rows = (
|
||||
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}
|
||||
# 直接按 endpoint_id 聚合,避免共享 Key 时把不同 API 格式串桶。
|
||||
format_endpoint_mapping = {"_single": endpoint_ids}
|
||||
result = EndpointHealthService._generate_timeline_batch(
|
||||
db, format_key_mapping, now, lookback_hours, segments
|
||||
db, format_endpoint_mapping, now, lookback_hours, segments
|
||||
)
|
||||
|
||||
return result.get(
|
||||
|
||||
@@ -140,6 +140,13 @@ def _extract_codex_weekly_reset_seconds(metadata: dict[str, Any]) -> float | Non
|
||||
if not isinstance(codex, dict):
|
||||
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()
|
||||
|
||||
# 优先绝对时间戳,避免 reset_seconds 快照随时间漂移。
|
||||
|
||||
@@ -69,6 +69,14 @@ def _format_quota_value(value: float) -> str:
|
||||
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:
|
||||
seconds = _to_float(seconds_raw)
|
||||
if seconds is None:
|
||||
@@ -223,7 +231,11 @@ class CodexQuotaReader(PoolQuotaReader):
|
||||
primary_used = _to_float(self._data.get("primary_used_percent"))
|
||||
if primary_used is not None:
|
||||
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:
|
||||
part = f"{part} ({reset_text})"
|
||||
parts.append(part)
|
||||
@@ -231,7 +243,11 @@ class CodexQuotaReader(PoolQuotaReader):
|
||||
secondary_used = _to_float(self._data.get("secondary_used_percent"))
|
||||
if secondary_used is not None:
|
||||
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:
|
||||
part = f"{part} ({reset_text})"
|
||||
parts.append(part)
|
||||
|
||||
@@ -4,7 +4,7 @@ Sub2API 余额查询操作
|
||||
|
||||
import asyncio
|
||||
import time
|
||||
from typing import Any
|
||||
from typing import Any, cast
|
||||
|
||||
import httpx
|
||||
|
||||
@@ -40,10 +40,13 @@ class Sub2ApiBalanceAction(BalanceAction):
|
||||
sub_endpoint = self.config.get("subscription_endpoint", "/api/v1/subscriptions/summary")
|
||||
|
||||
try:
|
||||
me_resp, sub_resp = await asyncio.gather(
|
||||
client.get(me_endpoint),
|
||||
client.get(sub_endpoint),
|
||||
return_exceptions=True,
|
||||
me_resp, sub_resp = cast(
|
||||
tuple[httpx.Response | BaseException, httpx.Response | BaseException],
|
||||
await asyncio.gather(
|
||||
client.get(me_endpoint),
|
||||
client.get(sub_endpoint),
|
||||
return_exceptions=True,
|
||||
),
|
||||
)
|
||||
response_time_ms = int((time.time() - start_time) * 1000)
|
||||
|
||||
|
||||
@@ -4,7 +4,7 @@ YesCode 余额查询操作
|
||||
|
||||
import asyncio
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Any
|
||||
from typing import Any, cast
|
||||
|
||||
import httpx
|
||||
|
||||
@@ -37,8 +37,9 @@ async def fetch_yescode_combined_data(
|
||||
balance_task = client.get(f"{base_url}/api/v1/user/balance")
|
||||
profile_task = client.get(f"{base_url}/api/v1/auth/profile")
|
||||
|
||||
balance_resp, profile_resp = await asyncio.gather(
|
||||
balance_task, profile_task, return_exceptions=True
|
||||
balance_resp, profile_resp = cast(
|
||||
tuple[httpx.Response | BaseException, httpx.Response | BaseException],
|
||||
await asyncio.gather(balance_task, profile_task, return_exceptions=True),
|
||||
)
|
||||
|
||||
# 解析 balance 接口
|
||||
|
||||
@@ -975,21 +975,23 @@ class MaintenanceScheduler:
|
||||
try:
|
||||
now = datetime.now(timezone.utc)
|
||||
|
||||
# 1. 压缩详细日志 (body 字段 -> 压缩字段)
|
||||
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_cleaned = self._cleanup_compressed_fields(compressed_cutoff, batch_size)
|
||||
|
||||
# 3. 清理请求头
|
||||
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)
|
||||
|
||||
# 先删最老的整行,再按窗口处理剩余记录,避免同一行在一轮里被重复改写。
|
||||
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
|
||||
keys_db = create_session()
|
||||
@@ -1005,7 +1007,7 @@ class MaintenanceScheduler:
|
||||
|
||||
logger.info(
|
||||
f"清理完成: 压缩 {body_compressed} 条, "
|
||||
f"清理压缩字段 {compressed_cleaned} 条, "
|
||||
f"清理body {body_cleaned} 条, "
|
||||
f"清理header {header_cleaned} 条, "
|
||||
f"删除记录 {records_deleted} 条, "
|
||||
f"清理过期Keys {keys_cleaned} 条"
|
||||
@@ -1016,10 +1018,16 @@ class MaintenanceScheduler:
|
||||
loop = asyncio.get_running_loop()
|
||||
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 字段到压缩字段
|
||||
|
||||
逐条处理,确保每条记录都正确更新(同步方法,在线程池中调用)
|
||||
仅处理指定时间窗口内仍保留原始 body 的记录,避免对更老记录重复写放大。
|
||||
"""
|
||||
from sqlalchemy import null, update
|
||||
|
||||
@@ -1027,84 +1035,45 @@ class MaintenanceScheduler:
|
||||
no_progress_count = 0
|
||||
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:
|
||||
batch_db = create_session()
|
||||
try:
|
||||
record_ids = [
|
||||
row.id
|
||||
for row in (
|
||||
batch_db.query(Usage.id)
|
||||
.filter(Usage.created_at < cutoff_time)
|
||||
.filter(
|
||||
(Usage.request_body.isnot(None))
|
||||
| (Usage.response_body.isnot(None))
|
||||
| (Usage.provider_request_body.isnot(None))
|
||||
| (Usage.client_response_body.isnot(None))
|
||||
)
|
||||
.order_by(Usage.created_at.asc(), Usage.id.asc())
|
||||
.limit(memory_safe_batch_size)
|
||||
.all()
|
||||
query = batch_db.query(
|
||||
Usage.id,
|
||||
Usage.request_body,
|
||||
Usage.response_body,
|
||||
Usage.provider_request_body,
|
||||
Usage.client_response_body,
|
||||
).filter(Usage.created_at < cutoff_time)
|
||||
if newer_than is not None:
|
||||
query = query.filter(Usage.created_at >= newer_than)
|
||||
records = (
|
||||
query.filter(
|
||||
(Usage.request_body.isnot(None))
|
||||
| (Usage.response_body.isnot(None))
|
||||
| (Usage.provider_request_body.isnot(None))
|
||||
| (Usage.client_response_body.isnot(None))
|
||||
)
|
||||
]
|
||||
except Exception as e:
|
||||
logger.exception("加载待压缩 body 记录 ID 失败: {}", e)
|
||||
try:
|
||||
batch_db.rollback()
|
||||
except Exception:
|
||||
pass
|
||||
break
|
||||
finally:
|
||||
batch_db.close()
|
||||
.order_by(Usage.created_at.asc(), Usage.id.asc())
|
||||
.limit(memory_safe_batch_size)
|
||||
.all()
|
||||
)
|
||||
|
||||
if not record_ids:
|
||||
break
|
||||
if not records:
|
||||
break
|
||||
|
||||
batch_success = 0
|
||||
batch_progress = False
|
||||
|
||||
for record_id in record_ids:
|
||||
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(
|
||||
batch_success = 0
|
||||
batch_progress = False
|
||||
for record in records:
|
||||
result = batch_db.execute(
|
||||
update(Usage)
|
||||
.where(Usage.id == record.id)
|
||||
.values(
|
||||
@@ -1133,20 +1102,20 @@ class MaintenanceScheduler:
|
||||
)
|
||||
.execution_options(synchronize_session=False)
|
||||
)
|
||||
record_db.commit()
|
||||
|
||||
if result.rowcount > 0:
|
||||
batch_success += 1
|
||||
batch_progress = True
|
||||
except Exception as e:
|
||||
logger.warning("压缩记录 {} 失败: {}", record_id, e)
|
||||
try:
|
||||
record_db.rollback()
|
||||
except Exception:
|
||||
pass
|
||||
continue
|
||||
finally:
|
||||
record_db.close()
|
||||
|
||||
batch_db.commit()
|
||||
except Exception as e:
|
||||
logger.warning("压缩 body 批次失败: {}", e)
|
||||
try:
|
||||
batch_db.rollback()
|
||||
except Exception:
|
||||
pass
|
||||
break
|
||||
finally:
|
||||
batch_db.close()
|
||||
|
||||
if not batch_progress:
|
||||
no_progress_count += 1
|
||||
@@ -1166,27 +1135,47 @@ class MaintenanceScheduler:
|
||||
|
||||
return total_compressed
|
||||
|
||||
def _cleanup_compressed_fields(self, cutoff_time: datetime, batch_size: int) -> int:
|
||||
"""清理压缩字段(删除压缩的body)
|
||||
def _cleanup_stale_body_fields(
|
||||
self,
|
||||
cutoff_time: datetime,
|
||||
batch_size: int,
|
||||
*,
|
||||
newer_than: datetime | None = None,
|
||||
) -> int:
|
||||
"""清理已超过压缩保留期的 body 字段
|
||||
|
||||
每批使用短生命周期 session(同步方法,在线程池中调用)
|
||||
直接清空 raw/compressed body,避免更老记录先压缩再马上被清掉。
|
||||
"""
|
||||
from sqlalchemy import null, update
|
||||
|
||||
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:
|
||||
batch_db = create_session()
|
||||
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 = (
|
||||
batch_db.query(Usage.id)
|
||||
.filter(Usage.created_at < cutoff_time)
|
||||
.filter(
|
||||
(Usage.request_body_compressed.isnot(None))
|
||||
query.filter(
|
||||
(Usage.request_body.isnot(None))
|
||||
| (Usage.response_body.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.provider_request_body_compressed.isnot(None))
|
||||
| (Usage.client_response_body_compressed.isnot(None))
|
||||
)
|
||||
.order_by(Usage.created_at.asc(), Usage.id.asc())
|
||||
.limit(batch_size)
|
||||
.all()
|
||||
)
|
||||
@@ -1200,6 +1189,10 @@ class MaintenanceScheduler:
|
||||
update(Usage)
|
||||
.where(Usage.id.in_(record_ids))
|
||||
.values(
|
||||
request_body=null(),
|
||||
response_body=null(),
|
||||
provider_request_body=null(),
|
||||
client_response_body=null(),
|
||||
request_body_compressed=null(),
|
||||
response_body_compressed=null(),
|
||||
provider_request_body_compressed=null(),
|
||||
@@ -1211,14 +1204,14 @@ class MaintenanceScheduler:
|
||||
batch_db.commit()
|
||||
|
||||
if rows_updated == 0:
|
||||
logger.warning("清理压缩字段: rowcount=0,可能存在问题")
|
||||
logger.warning("清理 body 字段: rowcount=0,可能存在问题")
|
||||
break
|
||||
|
||||
total_cleaned += rows_updated
|
||||
logger.debug(f"已清理 {rows_updated} 条记录的压缩字段,累计 {total_cleaned} 条")
|
||||
logger.debug(f"已清理 {rows_updated} 条记录的 body 字段,累计 {total_cleaned} 条")
|
||||
|
||||
except Exception as e:
|
||||
logger.exception(f"清理压缩字段失败: {e}")
|
||||
logger.exception(f"清理 body 字段失败: {e}")
|
||||
try:
|
||||
batch_db.rollback()
|
||||
except Exception:
|
||||
@@ -1229,7 +1222,13 @@ class MaintenanceScheduler:
|
||||
|
||||
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 字段
|
||||
|
||||
每批使用短生命周期 session(同步方法,在线程池中调用)
|
||||
@@ -1238,17 +1237,28 @@ class MaintenanceScheduler:
|
||||
|
||||
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:
|
||||
batch_db = create_session()
|
||||
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 = (
|
||||
batch_db.query(Usage.id)
|
||||
.filter(Usage.created_at < cutoff_time)
|
||||
.filter(
|
||||
query.filter(
|
||||
(Usage.request_headers.isnot(None))
|
||||
| (Usage.response_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)
|
||||
.all()
|
||||
)
|
||||
@@ -1265,6 +1275,7 @@ class MaintenanceScheduler:
|
||||
request_headers=null(),
|
||||
response_headers=null(),
|
||||
provider_request_headers=null(),
|
||||
client_response_headers=null(),
|
||||
)
|
||||
)
|
||||
|
||||
@@ -1300,6 +1311,7 @@ class MaintenanceScheduler:
|
||||
records_to_delete = (
|
||||
batch_db.query(Usage.id)
|
||||
.filter(Usage.created_at < cutoff_time)
|
||||
.order_by(Usage.created_at.asc(), Usage.id.asc())
|
||||
.limit(batch_size)
|
||||
.all()
|
||||
)
|
||||
|
||||
@@ -69,6 +69,7 @@ class SyncTaskExecutionService:
|
||||
request_body_state: RequestBodyState | None,
|
||||
request_headers: dict[str, Any] | None,
|
||||
request_body: dict[str, Any] | None,
|
||||
create_pending_usage: bool = True,
|
||||
) -> ExecutionResult:
|
||||
"""
|
||||
Unified candidate traversal loop for SYNC.
|
||||
@@ -143,26 +144,32 @@ class SyncTaskExecutionService:
|
||||
affinity_key = str(user_api_key.id)
|
||||
user_id = str(user_api_key.user_id)
|
||||
api_format_norm = normalize_endpoint_signature(api_format)
|
||||
user: User | None = None
|
||||
username_snapshot = None
|
||||
api_key_name_snapshot = getattr(user_api_key, "name", None)
|
||||
|
||||
# Keep pending usage creation behavior consistent with previous behavior
|
||||
try:
|
||||
user = self.db.query(User).filter(User.id == user_api_key.user_id).first()
|
||||
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:
|
||||
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(
|
||||
api_format=api_format_norm,
|
||||
|
||||
@@ -126,6 +126,7 @@ class TaskService:
|
||||
supported_auth_types: set[str] | None = None,
|
||||
allow_format_conversion: bool = False,
|
||||
max_candidates: int | None = None,
|
||||
create_pending_usage: bool = True,
|
||||
) -> ExecutionResult:
|
||||
"""兼容入口:默认绑定到 TaskService 内部执行路由。"""
|
||||
return await self._execute_facade_ops.execute(
|
||||
@@ -146,6 +147,7 @@ class TaskService:
|
||||
supported_auth_types=supported_auth_types,
|
||||
allow_format_conversion=allow_format_conversion,
|
||||
max_candidates=max_candidates,
|
||||
create_pending_usage=create_pending_usage,
|
||||
)
|
||||
|
||||
async def _execute_internal(
|
||||
@@ -168,6 +170,7 @@ class TaskService:
|
||||
supported_auth_types: set[str] | None = None,
|
||||
allow_format_conversion: bool = False,
|
||||
max_candidates: int | None = None,
|
||||
create_pending_usage: bool = True,
|
||||
) -> ExecutionResult:
|
||||
if task_mode == TaskMode.ASYNC:
|
||||
if extract_external_task_id is None:
|
||||
@@ -250,6 +253,7 @@ class TaskService:
|
||||
request_body_state=request_body_state,
|
||||
request_headers=request_headers,
|
||||
request_body=request_body,
|
||||
create_pending_usage=create_pending_usage,
|
||||
)
|
||||
|
||||
async def execute_sync_candidates(
|
||||
|
||||
@@ -39,6 +39,17 @@ async def _yield_once_then_cancel(ctx: StreamContext) -> AsyncGenerator[bytes, N
|
||||
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
|
||||
async def test_create_monitored_stream_marks_client_disconnected_when_confirmed() -> None:
|
||||
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.output_tokens == expected_output_tokens
|
||||
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
|
||||
|
||||
@@ -23,7 +23,7 @@ class _DummySyncHandler(CliSyncMixin):
|
||||
) -> str:
|
||||
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)
|
||||
raise _StopExecution()
|
||||
|
||||
@@ -41,7 +41,7 @@ class _DummyStreamHandler(CliStreamMixin):
|
||||
) -> str:
|
||||
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)
|
||||
raise _StopExecution()
|
||||
|
||||
|
||||
@@ -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.stream_context import StreamContext
|
||||
|
||||
@@ -50,7 +56,36 @@ def test_reset_for_retry_clears_state() -> 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")
|
||||
start_time = 100.0
|
||||
@@ -63,7 +98,7 @@ def test_record_first_byte_time(monkeypatch) -> None:
|
||||
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")
|
||||
start_time = 100.0
|
||||
@@ -82,7 +117,7 @@ def test_record_first_byte_time_idempotent(monkeypatch) -> None:
|
||||
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")
|
||||
start_time = 100.0
|
||||
|
||||
@@ -1,11 +1,15 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import time
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
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.stream_context import StreamContext
|
||||
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.output_tokens == max(1, len("partial output") // 4)
|
||||
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 == []
|
||||
|
||||
146
tests/api/test_health_monitor_api_formats.py
Normal file
146
tests/api/test_health_monitor_api_formats.py
Normal 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)
|
||||
50
tests/plugins/test_manager.py
Normal file
50
tests/plugins/test_manager.py
Normal 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
|
||||
111
tests/plugins/test_token_bucket.py
Normal file
111
tests/plugins/test_token_bucket.py
Normal 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
|
||||
@@ -36,6 +36,22 @@ def test_codex_reader_preserves_summary_formats() -> None:
|
||||
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:
|
||||
reader = get_quota_reader(
|
||||
"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)
|
||||
|
||||
|
||||
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:
|
||||
monkeypatch.setattr("src.services.provider.pool.dimensions._helpers.time.time", lambda: 1000.0)
|
||||
key_obj = SimpleNamespace(
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import inspect
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from types import SimpleNamespace
|
||||
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
|
||||
|
||||
|
||||
def test_cleanup_body_fields_loads_ids_then_processes_records_individually(
|
||||
def test_cleanup_body_fields_batches_records_with_single_commit(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
scheduler = MaintenanceScheduler()
|
||||
|
||||
class _IdBatchSession:
|
||||
def __init__(self, ids: list[str]) -> None:
|
||||
self.ids = ids
|
||||
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
|
||||
class _BatchSession:
|
||||
def __init__(self, records: list[SimpleNamespace]) -> None:
|
||||
self.records = records
|
||||
self.closed = False
|
||||
self.committed = False
|
||||
self.executed = 0
|
||||
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]
|
||||
self.query_args = args
|
||||
return self.query_obj
|
||||
|
||||
def execute(self, _statement): # type: ignore[no-untyped-def]
|
||||
self.executed += 1
|
||||
return SimpleNamespace(rowcount=1)
|
||||
|
||||
def commit(self) -> None:
|
||||
@@ -173,27 +159,26 @@ def test_cleanup_body_fields_loads_ids_then_processes_records_individually(
|
||||
def close(self) -> None:
|
||||
self.closed = True
|
||||
|
||||
batch_one = _IdBatchSession(["usage-1", "usage-2"])
|
||||
record_one = _RecordSession(
|
||||
SimpleNamespace(
|
||||
id="usage-1",
|
||||
request_body={"hello": "world"},
|
||||
response_body=None,
|
||||
provider_request_body=None,
|
||||
client_response_body=None,
|
||||
)
|
||||
batch_one = _BatchSession(
|
||||
[
|
||||
SimpleNamespace(
|
||||
id="usage-1",
|
||||
request_body={"hello": "world"},
|
||||
response_body=None,
|
||||
provider_request_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(
|
||||
SimpleNamespace(
|
||||
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])
|
||||
batch_two = _BatchSession([])
|
||||
sessions = iter([batch_one, batch_two])
|
||||
|
||||
monkeypatch.setattr(
|
||||
maintenance_scheduler_module,
|
||||
@@ -212,12 +197,173 @@ def test_cleanup_body_fields_loads_ids_then_processes_records_individually(
|
||||
)
|
||||
|
||||
assert compressed == 2
|
||||
assert batch_one.query_args == (maintenance_scheduler_module.Usage.id,)
|
||||
assert len(record_one.query_args) == 5
|
||||
assert len(record_two.query_args) == 5
|
||||
assert len(batch_one.query_args) == 5
|
||||
assert batch_one.executed == 2
|
||||
assert batch_one.committed is True
|
||||
assert batch_one.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
|
||||
assert record_two.closed is True
|
||||
|
||||
|
||||
def test_cleanup_header_fields_clears_client_response_headers(
|
||||
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),
|
||||
]
|
||||
|
||||
173
tests/unit/test_endpoint_health_timeline.py
Normal file
173
tests/unit/test_endpoint_health_timeline.py
Normal 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)
|
||||
Reference in New Issue
Block a user