mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +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-jitter 会自动设置为 MAX_REQUESTS/20 (5%)
|
||||||
# MAX_REQUESTS=4000
|
# MAX_REQUESTS=4000
|
||||||
|
|
||||||
|
# glibc malloc arena 上限(默认 2)
|
||||||
|
# 降低 malloc 内存碎片,减少 gunicorn worker RSS
|
||||||
|
# MALLOC_ARENA_MAX=2
|
||||||
|
|
||||||
# HTTP 连接池上限(默认总预算约 200,按 worker 平分)
|
# HTTP 连接池上限(默认总预算约 200,按 worker 平分)
|
||||||
# 如果容器内存偏高,可继续下调;例如 2 worker 时设为 80-100
|
# 如果容器内存偏高,可继续下调;例如 2 worker 时设为 80-100
|
||||||
# HTTP_MAX_CONNECTIONS=100
|
# HTTP_MAX_CONNECTIONS=100
|
||||||
@@ -74,6 +78,11 @@ ADMIN_PASSWORD=admin123456
|
|||||||
# - 需要更多调试上下文:4
|
# - 需要更多调试上下文:4
|
||||||
# RESPONSE_CHUNKS_MAX_SIZE_MB=2
|
# RESPONSE_CHUNKS_MAX_SIZE_MB=2
|
||||||
|
|
||||||
|
# 流式空闲超时(单位秒,默认 30)
|
||||||
|
# 当流已经开始但连续一段时间没有任何新 chunk 时,提前中断并返回 504,
|
||||||
|
# 避免一直等到 worker 超时(如 300s)
|
||||||
|
# STREAM_IDLE_TIMEOUT_SECONDS=30
|
||||||
|
|
||||||
# API Key 前缀(默认 sk)
|
# API Key 前缀(默认 sk)
|
||||||
# API_KEY_PREFIX=sk
|
# API_KEY_PREFIX=sk
|
||||||
|
|
||||||
|
|||||||
@@ -25,7 +25,12 @@ RUN --mount=type=cache,target=/var/cache/apt,sharing=locked \
|
|||||||
nginx \
|
nginx \
|
||||||
supervisor \
|
supervisor \
|
||||||
libpq5 \
|
libpq5 \
|
||||||
curl
|
curl \
|
||||||
|
libjemalloc2
|
||||||
|
RUN set -eux; \
|
||||||
|
jemalloc_path="$(find /usr/lib -type f -name 'libjemalloc.so.2' | head -n1)"; \
|
||||||
|
[ -n "$jemalloc_path" ]; \
|
||||||
|
ln -sf "$jemalloc_path" /usr/local/lib/libjemalloc.so.2
|
||||||
# 从 base 镜像复制 Python 包
|
# 从 base 镜像复制 Python 包
|
||||||
COPY --from=builder /usr/local/lib/python3.13/site-packages /usr/local/lib/python3.13/site-packages
|
COPY --from=builder /usr/local/lib/python3.13/site-packages /usr/local/lib/python3.13/site-packages
|
||||||
# 只复制需要的 Python 可执行文件
|
# 只复制需要的 Python 可执行文件
|
||||||
@@ -247,7 +252,7 @@ RUN printf '%s\n' \
|
|||||||
'stdout_logfile_maxbytes=0' \
|
'stdout_logfile_maxbytes=0' \
|
||||||
'stderr_logfile=/dev/stderr' \
|
'stderr_logfile=/dev/stderr' \
|
||||||
'stderr_logfile_maxbytes=0' \
|
'stderr_logfile_maxbytes=0' \
|
||||||
'environment=PYTHONUNBUFFERED=1,PYTHONIOENCODING=utf-8,LANG=C.UTF-8,LC_ALL=C.UTF-8,DOCKER_CONTAINER=true' \
|
'environment=PYTHONUNBUFFERED=1,PYTHONIOENCODING=utf-8,LANG=C.UTF-8,LC_ALL=C.UTF-8,DOCKER_CONTAINER=true,LD_PRELOAD=/usr/local/lib/libjemalloc.so.2,MALLOC_CONF="background_thread:true,dirty_decay_ms:5000,muzzy_decay_ms:5000"' \
|
||||||
'' \
|
'' \
|
||||||
'[program:tunnel-hub]' \
|
'[program:tunnel-hub]' \
|
||||||
'command=/usr/local/bin/aether-hub --bind 0.0.0.0:8085' \
|
'command=/usr/local/bin/aether-hub --bind 0.0.0.0:8085' \
|
||||||
@@ -268,6 +273,8 @@ ENV PYTHONUNBUFFERED=1 \
|
|||||||
PYTHONIOENCODING=utf-8 \
|
PYTHONIOENCODING=utf-8 \
|
||||||
LANG=C.UTF-8 \
|
LANG=C.UTF-8 \
|
||||||
LC_ALL=C.UTF-8 \
|
LC_ALL=C.UTF-8 \
|
||||||
|
LD_PRELOAD=/usr/local/lib/libjemalloc.so.2 \
|
||||||
|
MALLOC_CONF=background_thread:true,dirty_decay_ms:5000,muzzy_decay_ms:5000 \
|
||||||
PORT=8084 \
|
PORT=8084 \
|
||||||
GUNICORN_WORKERS=2 \
|
GUNICORN_WORKERS=2 \
|
||||||
MAX_REQUESTS=4000
|
MAX_REQUESTS=4000
|
||||||
|
|||||||
@@ -34,7 +34,12 @@ RUN --mount=type=cache,target=/var/cache/apt,sharing=locked \
|
|||||||
nginx \
|
nginx \
|
||||||
supervisor \
|
supervisor \
|
||||||
libpq5 \
|
libpq5 \
|
||||||
curl
|
curl \
|
||||||
|
libjemalloc2
|
||||||
|
RUN set -eux; \
|
||||||
|
jemalloc_path="$(find /usr/lib -type f -name 'libjemalloc.so.2' | head -n1)"; \
|
||||||
|
[ -n "$jemalloc_path" ]; \
|
||||||
|
ln -sf "$jemalloc_path" /usr/local/lib/libjemalloc.so.2
|
||||||
|
|
||||||
# 从 base 镜像复制 Python 包
|
# 从 base 镜像复制 Python 包
|
||||||
COPY --from=builder /usr/local/lib/python3.13/site-packages /usr/local/lib/python3.13/site-packages
|
COPY --from=builder /usr/local/lib/python3.13/site-packages /usr/local/lib/python3.13/site-packages
|
||||||
@@ -270,7 +275,7 @@ RUN printf '%s\n' \
|
|||||||
'stdout_logfile_maxbytes=0' \
|
'stdout_logfile_maxbytes=0' \
|
||||||
'stderr_logfile=/dev/stderr' \
|
'stderr_logfile=/dev/stderr' \
|
||||||
'stderr_logfile_maxbytes=0' \
|
'stderr_logfile_maxbytes=0' \
|
||||||
'environment=PYTHONUNBUFFERED=1,PYTHONIOENCODING=utf-8,LANG=C.UTF-8,LC_ALL=C.UTF-8,DOCKER_CONTAINER=true' \
|
'environment=PYTHONUNBUFFERED=1,PYTHONIOENCODING=utf-8,LANG=C.UTF-8,LC_ALL=C.UTF-8,DOCKER_CONTAINER=true,LD_PRELOAD=/usr/local/lib/libjemalloc.so.2,MALLOC_CONF="background_thread:true,dirty_decay_ms:5000,muzzy_decay_ms:5000"' \
|
||||||
'' \
|
'' \
|
||||||
'[program:tunnel-hub]' \
|
'[program:tunnel-hub]' \
|
||||||
'command=/usr/local/bin/aether-hub --bind 0.0.0.0:8085' \
|
'command=/usr/local/bin/aether-hub --bind 0.0.0.0:8085' \
|
||||||
@@ -294,6 +299,8 @@ ENV PYTHONUNBUFFERED=1 \
|
|||||||
PYTHONIOENCODING=utf-8 \
|
PYTHONIOENCODING=utf-8 \
|
||||||
LANG=C.UTF-8 \
|
LANG=C.UTF-8 \
|
||||||
LC_ALL=C.UTF-8 \
|
LC_ALL=C.UTF-8 \
|
||||||
|
LD_PRELOAD=/usr/local/lib/libjemalloc.so.2 \
|
||||||
|
MALLOC_CONF=background_thread:true,dirty_decay_ms:5000,muzzy_decay_ms:5000 \
|
||||||
PORT=8084 \
|
PORT=8084 \
|
||||||
GUNICORN_WORKERS=2 \
|
GUNICORN_WORKERS=2 \
|
||||||
MAX_REQUESTS=4000
|
MAX_REQUESTS=4000
|
||||||
|
|||||||
@@ -84,15 +84,24 @@ export interface CodexResetStatus {
|
|||||||
* @param resetSecs 相对剩余秒数(用于 fallback)
|
* @param resetSecs 相对剩余秒数(用于 fallback)
|
||||||
* @param updatedAt 元数据更新时间(Unix 秒)
|
* @param updatedAt 元数据更新时间(Unix 秒)
|
||||||
* @param _tick 响应式触发器(传入 tick.value 以触发响应式更新)
|
* @param _tick 响应式触发器(传入 tick.value 以触发响应式更新)
|
||||||
|
* @param remainingPercent 当前窗口剩余额度百分比(0-100,100 表示满额不启动倒计时)
|
||||||
*/
|
*/
|
||||||
export function getCodexResetCountdown(
|
export function getCodexResetCountdown(
|
||||||
resetAt: number | null | undefined,
|
resetAt: number | null | undefined,
|
||||||
resetSecs: number | null | undefined,
|
resetSecs: number | null | undefined,
|
||||||
updatedAt: number | null | undefined,
|
updatedAt: number | null | undefined,
|
||||||
_tick: number
|
_tick: number,
|
||||||
|
remainingPercent?: number | null
|
||||||
): CodexResetStatus | null {
|
): CodexResetStatus | null {
|
||||||
void _tick
|
void _tick
|
||||||
|
|
||||||
|
if (remainingPercent != null) {
|
||||||
|
const normalizedRemaining = Number(remainingPercent)
|
||||||
|
if (Number.isFinite(normalizedRemaining) && normalizedRemaining >= 100) {
|
||||||
|
return null
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
const nowSec = Math.floor(Date.now() / 1000)
|
const nowSec = Math.floor(Date.now() / 1000)
|
||||||
let remaining: number
|
let remaining: number
|
||||||
|
|
||||||
|
|||||||
@@ -573,18 +573,20 @@
|
|||||||
/>
|
/>
|
||||||
</div>
|
</div>
|
||||||
<div
|
<div
|
||||||
v-if="key.upstream_metadata.codex.primary_reset_at || key.upstream_metadata.codex.primary_reset_seconds"
|
v-if="(key.upstream_metadata.codex.primary_reset_at || key.upstream_metadata.codex.primary_reset_seconds) && shouldStartCodexResetCountdown(key.upstream_metadata.codex.primary_used_percent)"
|
||||||
class="text-[9px] mt-0.5 tabular-nums"
|
class="text-[9px] mt-0.5 tabular-nums"
|
||||||
:class="getResetCountdownClass(
|
:class="getResetCountdownClass(
|
||||||
key.upstream_metadata.codex.primary_reset_at,
|
key.upstream_metadata.codex.primary_reset_at,
|
||||||
key.upstream_metadata.codex.primary_reset_seconds,
|
key.upstream_metadata.codex.primary_reset_seconds,
|
||||||
key.upstream_metadata.codex.updated_at
|
key.upstream_metadata.codex.updated_at,
|
||||||
|
key.upstream_metadata.codex.primary_used_percent
|
||||||
)"
|
)"
|
||||||
>
|
>
|
||||||
{{ getResetCountdownText(
|
{{ getResetCountdownText(
|
||||||
key.upstream_metadata.codex.primary_reset_at,
|
key.upstream_metadata.codex.primary_reset_at,
|
||||||
key.upstream_metadata.codex.primary_reset_seconds,
|
key.upstream_metadata.codex.primary_reset_seconds,
|
||||||
key.upstream_metadata.codex.updated_at
|
key.upstream_metadata.codex.updated_at,
|
||||||
|
key.upstream_metadata.codex.primary_used_percent
|
||||||
) }}
|
) }}
|
||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
@@ -604,18 +606,21 @@
|
|||||||
/>
|
/>
|
||||||
</div>
|
</div>
|
||||||
<div
|
<div
|
||||||
|
v-if="shouldStartCodexResetCountdown(key.upstream_metadata.codex.secondary_used_percent)"
|
||||||
class="text-[9px] mt-0.5 tabular-nums"
|
class="text-[9px] mt-0.5 tabular-nums"
|
||||||
:class="getResetCountdownClass(
|
:class="getResetCountdownClass(
|
||||||
key.upstream_metadata.codex.secondary_reset_at,
|
key.upstream_metadata.codex.secondary_reset_at,
|
||||||
key.upstream_metadata.codex.secondary_reset_seconds,
|
key.upstream_metadata.codex.secondary_reset_seconds,
|
||||||
key.upstream_metadata.codex.updated_at
|
key.upstream_metadata.codex.updated_at,
|
||||||
|
key.upstream_metadata.codex.secondary_used_percent
|
||||||
)"
|
)"
|
||||||
>
|
>
|
||||||
<template v-if="key.upstream_metadata.codex.secondary_reset_at || key.upstream_metadata.codex.secondary_reset_seconds">
|
<template v-if="key.upstream_metadata.codex.secondary_reset_at || key.upstream_metadata.codex.secondary_reset_seconds">
|
||||||
{{ getResetCountdownText(
|
{{ getResetCountdownText(
|
||||||
key.upstream_metadata.codex.secondary_reset_at,
|
key.upstream_metadata.codex.secondary_reset_at,
|
||||||
key.upstream_metadata.codex.secondary_reset_seconds,
|
key.upstream_metadata.codex.secondary_reset_seconds,
|
||||||
key.upstream_metadata.codex.updated_at
|
key.upstream_metadata.codex.updated_at,
|
||||||
|
key.upstream_metadata.codex.secondary_used_percent
|
||||||
) }}
|
) }}
|
||||||
</template>
|
</template>
|
||||||
<template v-else>
|
<template v-else>
|
||||||
@@ -2549,9 +2554,16 @@ function getAntigravityQuotaSummary(metadata: UpstreamMetadata | null | undefine
|
|||||||
function getResetCountdownText(
|
function getResetCountdownText(
|
||||||
resetAt: number | null | undefined,
|
resetAt: number | null | undefined,
|
||||||
resetSecs: number | null | undefined,
|
resetSecs: number | null | undefined,
|
||||||
updatedAt: number | null | undefined
|
updatedAt: number | null | undefined,
|
||||||
|
usedPercent: number | null | undefined
|
||||||
): string {
|
): string {
|
||||||
const status = getCodexResetCountdown(resetAt, resetSecs, updatedAt, countdownTick.value)
|
const status = getCodexResetCountdown(
|
||||||
|
resetAt,
|
||||||
|
resetSecs,
|
||||||
|
updatedAt,
|
||||||
|
countdownTick.value,
|
||||||
|
toCodexRemainingPercent(usedPercent)
|
||||||
|
)
|
||||||
if (!status) return ''
|
if (!status) return ''
|
||||||
return status.isExpired ? status.text : `${status.text} 后重置`
|
return status.isExpired ? status.text : `${status.text} 后重置`
|
||||||
}
|
}
|
||||||
@@ -2559,15 +2571,35 @@ function getResetCountdownText(
|
|||||||
function getResetCountdownClass(
|
function getResetCountdownClass(
|
||||||
resetAt: number | null | undefined,
|
resetAt: number | null | undefined,
|
||||||
resetSecs: number | null | undefined,
|
resetSecs: number | null | undefined,
|
||||||
updatedAt: number | null | undefined
|
updatedAt: number | null | undefined,
|
||||||
|
usedPercent: number | null | undefined
|
||||||
): string {
|
): string {
|
||||||
const status = getCodexResetCountdown(resetAt, resetSecs, updatedAt, countdownTick.value)
|
const status = getCodexResetCountdown(
|
||||||
|
resetAt,
|
||||||
|
resetSecs,
|
||||||
|
updatedAt,
|
||||||
|
countdownTick.value,
|
||||||
|
toCodexRemainingPercent(usedPercent)
|
||||||
|
)
|
||||||
if (!status || status.isExpired) return 'text-muted-foreground/70'
|
if (!status || status.isExpired) return 'text-muted-foreground/70'
|
||||||
if (status.isCritical) return 'text-destructive font-medium animate-pulse'
|
if (status.isCritical) return 'text-destructive font-medium animate-pulse'
|
||||||
if (status.isUrgent) return 'text-amber-500 dark:text-amber-400'
|
if (status.isUrgent) return 'text-amber-500 dark:text-amber-400'
|
||||||
return 'text-muted-foreground/70'
|
return 'text-muted-foreground/70'
|
||||||
}
|
}
|
||||||
|
|
||||||
|
function toCodexRemainingPercent(usedPercent: number | null | undefined): number | null {
|
||||||
|
const normalizedUsed = Number(usedPercent)
|
||||||
|
if (!Number.isFinite(normalizedUsed)) return null
|
||||||
|
const clampedUsed = Math.min(Math.max(normalizedUsed, 0), 100)
|
||||||
|
return Math.max(100 - clampedUsed, 0)
|
||||||
|
}
|
||||||
|
|
||||||
|
function shouldStartCodexResetCountdown(usedPercent: number | null | undefined): boolean {
|
||||||
|
const remainingPercent = toCodexRemainingPercent(usedPercent)
|
||||||
|
if (remainingPercent == null) return true
|
||||||
|
return remainingPercent < 100
|
||||||
|
}
|
||||||
|
|
||||||
// 格式化重置时间
|
// 格式化重置时间
|
||||||
function formatResetTime(seconds: number): string {
|
function formatResetTime(seconds: number): string {
|
||||||
const days = Math.floor(seconds / 86400)
|
const days = Math.floor(seconds / 86400)
|
||||||
|
|||||||
@@ -2428,7 +2428,7 @@ function getQuotaProgressLabel(label: string): string {
|
|||||||
|
|
||||||
function getQuotaProgressCountdown(item: QuotaProgressItem) {
|
function getQuotaProgressCountdown(item: QuotaProgressItem) {
|
||||||
if ((item.label !== '5H' && item.label !== '周') || item.resetAtSeconds == null) return null
|
if ((item.label !== '5H' && item.label !== '周') || item.resetAtSeconds == null) return null
|
||||||
return getCodexResetCountdown(item.resetAtSeconds, null, null, countdownTick.value)
|
return getCodexResetCountdown(item.resetAtSeconds, null, null, countdownTick.value, item.remainingPercent)
|
||||||
}
|
}
|
||||||
|
|
||||||
function getQuotaProgressCountdownText(item: QuotaProgressItem): string {
|
function getQuotaProgressCountdownText(item: QuotaProgressItem): string {
|
||||||
@@ -2438,6 +2438,9 @@ function getQuotaProgressCountdownText(item: QuotaProgressItem): string {
|
|||||||
}
|
}
|
||||||
|
|
||||||
function getQuotaProgressTooltip(item: QuotaProgressItem): string {
|
function getQuotaProgressTooltip(item: QuotaProgressItem): string {
|
||||||
|
if ((item.label === '5H' || item.label === '周') && item.remainingPercent >= 100) {
|
||||||
|
return ''
|
||||||
|
}
|
||||||
const detail = item.detail?.trim() || ''
|
const detail = item.detail?.trim() || ''
|
||||||
const countdownText = getQuotaProgressCountdownText(item)
|
const countdownText = getQuotaProgressCountdownText(item)
|
||||||
if (countdownText) {
|
if (countdownText) {
|
||||||
|
|||||||
@@ -93,6 +93,32 @@ def _format_str(api_format_enum: Any) -> str:
|
|||||||
return api_format_enum.value if hasattr(api_format_enum, "value") else str(api_format_enum)
|
return api_format_enum.value if hasattr(api_format_enum, "value") else str(api_format_enum)
|
||||||
|
|
||||||
|
|
||||||
|
def _fetch_recent_attempts_for_api_format(
|
||||||
|
db: Session,
|
||||||
|
*,
|
||||||
|
api_format: str,
|
||||||
|
since: datetime,
|
||||||
|
per_format_limit: int,
|
||||||
|
) -> list[RequestCandidate]:
|
||||||
|
"""获取单个 API 格式最近的最终态请求,用于事件展示。"""
|
||||||
|
final_statuses = ["success", "failed", "skipped"]
|
||||||
|
return (
|
||||||
|
db.query(RequestCandidate)
|
||||||
|
.join(ProviderEndpoint, RequestCandidate.endpoint_id == ProviderEndpoint.id)
|
||||||
|
.join(Provider, ProviderEndpoint.provider_id == Provider.id)
|
||||||
|
.filter(
|
||||||
|
ProviderEndpoint.is_active.is_(True),
|
||||||
|
Provider.is_active.is_(True),
|
||||||
|
ProviderEndpoint.api_format == api_format,
|
||||||
|
RequestCandidate.created_at >= since,
|
||||||
|
RequestCandidate.status.in_(final_statuses),
|
||||||
|
)
|
||||||
|
.order_by(RequestCandidate.created_at.desc())
|
||||||
|
.limit(per_format_limit)
|
||||||
|
.all()
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
pipeline = get_pipeline()
|
pipeline = get_pipeline()
|
||||||
|
|
||||||
|
|
||||||
@@ -379,7 +405,10 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
|||||||
func.count(RequestCandidate.id).label("count"),
|
func.count(RequestCandidate.id).label("count"),
|
||||||
)
|
)
|
||||||
.join(RequestCandidate, ProviderEndpoint.id == RequestCandidate.endpoint_id)
|
.join(RequestCandidate, ProviderEndpoint.id == RequestCandidate.endpoint_id)
|
||||||
|
.join(Provider, ProviderEndpoint.provider_id == Provider.id)
|
||||||
.filter(
|
.filter(
|
||||||
|
ProviderEndpoint.is_active.is_(True),
|
||||||
|
Provider.is_active.is_(True),
|
||||||
RequestCandidate.created_at >= since,
|
RequestCandidate.created_at >= since,
|
||||||
RequestCandidate.status.in_(final_statuses),
|
RequestCandidate.status.in_(final_statuses),
|
||||||
)
|
)
|
||||||
@@ -395,40 +424,15 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
|||||||
status_counts[fmt] = {"success": 0, "failed": 0, "skipped": 0}
|
status_counts[fmt] = {"success": 0, "failed": 0, "skipped": 0}
|
||||||
status_counts[fmt][status] = count
|
status_counts[fmt][status] = count
|
||||||
|
|
||||||
# 3. 获取最近一段时间的 RequestCandidate(限制数量)
|
# 3. 为所有活跃格式生成监控数据(包括没有请求记录的)
|
||||||
# 使用上面定义的 final_statuses,排除中间状态
|
|
||||||
limit_rows = max(500, self.per_format_limit * 10)
|
|
||||||
rows = (
|
|
||||||
db.query(
|
|
||||||
RequestCandidate,
|
|
||||||
ProviderEndpoint.api_format,
|
|
||||||
ProviderEndpoint.provider_id,
|
|
||||||
)
|
|
||||||
.join(ProviderEndpoint, RequestCandidate.endpoint_id == ProviderEndpoint.id)
|
|
||||||
.filter(
|
|
||||||
RequestCandidate.created_at >= since,
|
|
||||||
RequestCandidate.status.in_(final_statuses),
|
|
||||||
)
|
|
||||||
.order_by(RequestCandidate.created_at.desc())
|
|
||||||
.limit(limit_rows)
|
|
||||||
.all()
|
|
||||||
)
|
|
||||||
|
|
||||||
grouped_attempts: dict[str, list[RequestCandidate]] = {}
|
|
||||||
|
|
||||||
for attempt, api_format_enum, provider_id in rows:
|
|
||||||
fmt = _format_str(api_format_enum)
|
|
||||||
if fmt not in grouped_attempts:
|
|
||||||
grouped_attempts[fmt] = []
|
|
||||||
|
|
||||||
# 只保留每个 API 格式最近 per_format_limit 条记录
|
|
||||||
if len(grouped_attempts[fmt]) < self.per_format_limit:
|
|
||||||
grouped_attempts[fmt].append(attempt)
|
|
||||||
|
|
||||||
# 4. 为所有活跃格式生成监控数据(包括没有请求记录的)
|
|
||||||
monitors: list[ApiFormatHealthMonitor] = []
|
monitors: list[ApiFormatHealthMonitor] = []
|
||||||
for api_format in all_formats:
|
for api_format in all_formats:
|
||||||
attempts = grouped_attempts.get(api_format, [])
|
attempts = _fetch_recent_attempts_for_api_format(
|
||||||
|
db=db,
|
||||||
|
api_format=api_format,
|
||||||
|
since=since,
|
||||||
|
per_format_limit=self.per_format_limit,
|
||||||
|
)
|
||||||
# 获取窗口内的真实统计数据
|
# 获取窗口内的真实统计数据
|
||||||
# 只统计最终状态:success, failed, skipped
|
# 只统计最终状态:success, failed, skipped
|
||||||
# 中间状态(available, pending, used, started)不计入统计
|
# 中间状态(available, pending, used, started)不计入统计
|
||||||
|
|||||||
@@ -485,7 +485,7 @@ class BaseMessageHandler:
|
|||||||
api_format: str | None = None,
|
api_format: str | None = None,
|
||||||
request_headers: dict[str, Any] | None = None,
|
request_headers: dict[str, Any] | None = None,
|
||||||
request_body: dict[str, Any] | None = None,
|
request_body: dict[str, Any] | None = None,
|
||||||
) -> None:
|
) -> bool:
|
||||||
"""在请求开始时创建 pending 状态的 Usage 记录
|
"""在请求开始时创建 pending 状态的 Usage 记录
|
||||||
|
|
||||||
让前端可以立即看到"处理中"的请求,提升用户体验。
|
让前端可以立即看到"处理中"的请求,提升用户体验。
|
||||||
@@ -498,6 +498,8 @@ class BaseMessageHandler:
|
|||||||
api_format: API 格式
|
api_format: API 格式
|
||||||
request_headers: 原始请求头
|
request_headers: 原始请求头
|
||||||
request_body: 原始请求体
|
request_body: 原始请求体
|
||||||
|
Returns:
|
||||||
|
bool: True 表示已成功创建;False 表示创建失败(调用方可按需回退处理)。
|
||||||
"""
|
"""
|
||||||
try:
|
try:
|
||||||
UsageService.create_pending_usage(
|
UsageService.create_pending_usage(
|
||||||
@@ -512,9 +514,11 @@ class BaseMessageHandler:
|
|||||||
request_headers=request_headers,
|
request_headers=request_headers,
|
||||||
request_body=request_body,
|
request_body=request_body,
|
||||||
)
|
)
|
||||||
|
return True
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
# 创建失败不影响主流程
|
# 创建失败不影响主流程
|
||||||
logger.warning(f"[{self.request_id}] Failed to create pending usage: {exc}")
|
logger.warning(f"[{self.request_id}] Failed to create pending usage: {exc}")
|
||||||
|
return False
|
||||||
|
|
||||||
def _update_usage_to_streaming(self, request_id: str | None = None) -> None:
|
def _update_usage_to_streaming(self, request_id: str | None = None) -> None:
|
||||||
"""更新 Usage 状态为 streaming(流式传输开始时调用)
|
"""更新 Usage 状态为 streaming(流式传输开始时调用)
|
||||||
|
|||||||
@@ -487,7 +487,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
model = getattr(converted_request, "model", original_request_body.get("model", "unknown"))
|
model = getattr(converted_request, "model", original_request_body.get("model", "unknown"))
|
||||||
|
|
||||||
# 提前创建 pending 记录,让前端可以立即看到"处理中"
|
# 提前创建 pending 记录,让前端可以立即看到"处理中"
|
||||||
self._create_pending_usage(
|
pending_usage_created = self._create_pending_usage(
|
||||||
model=model,
|
model=model,
|
||||||
is_stream=True,
|
is_stream=True,
|
||||||
request_type="chat",
|
request_type="chat",
|
||||||
@@ -579,6 +579,8 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
request_body_state=request_state,
|
request_body_state=request_state,
|
||||||
request_headers=original_headers,
|
request_headers=original_headers,
|
||||||
request_body=original_request_body,
|
request_body=original_request_body,
|
||||||
|
# 预创建失败时,回退到 TaskService 侧创建,避免丢失 pending 状态。
|
||||||
|
create_pending_usage=not pending_usage_created,
|
||||||
)
|
)
|
||||||
stream_generator = exec_result.response
|
stream_generator = exec_result.response
|
||||||
provider_name = exec_result.provider_name or "unknown"
|
provider_name = exec_result.provider_name or "unknown"
|
||||||
|
|||||||
@@ -113,7 +113,7 @@ class ChatSyncExecutor:
|
|||||||
api_format = handler.allowed_api_formats[0]
|
api_format = handler.allowed_api_formats[0]
|
||||||
|
|
||||||
# 提前创建 pending 记录,让前端可以立即看到"处理中"
|
# 提前创建 pending 记录,让前端可以立即看到"处理中"
|
||||||
handler._create_pending_usage(
|
pending_usage_created = handler._create_pending_usage(
|
||||||
model=model,
|
model=model,
|
||||||
is_stream=False,
|
is_stream=False,
|
||||||
request_type="chat",
|
request_type="chat",
|
||||||
@@ -176,6 +176,8 @@ class ChatSyncExecutor:
|
|||||||
request_body_state=request_state,
|
request_body_state=request_state,
|
||||||
request_headers=original_headers,
|
request_headers=original_headers,
|
||||||
request_body=original_request_body,
|
request_body=original_request_body,
|
||||||
|
# 预创建失败时,回退到 TaskService 侧创建,避免丢失 pending 状态。
|
||||||
|
create_pending_usage=not pending_usage_created,
|
||||||
)
|
)
|
||||||
actual_provider_name = exec_result.provider_name or "unknown"
|
actual_provider_name = exec_result.provider_name or "unknown"
|
||||||
ctx.provider_id = exec_result.provider_id
|
ctx.provider_id = exec_result.provider_id
|
||||||
|
|||||||
@@ -3,6 +3,7 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
|
import os
|
||||||
import time
|
import time
|
||||||
from collections.abc import AsyncGenerator
|
from collections.abc import AsyncGenerator
|
||||||
from typing import TYPE_CHECKING, Any
|
from typing import TYPE_CHECKING, Any
|
||||||
@@ -29,12 +30,23 @@ if TYPE_CHECKING:
|
|||||||
from src.api.handlers.base.cli_protocol import CliHandlerProtocol
|
from src.api.handlers.base.cli_protocol import CliHandlerProtocol
|
||||||
|
|
||||||
|
|
||||||
|
def _read_stream_idle_timeout_seconds() -> float:
|
||||||
|
raw_value = os.getenv("STREAM_IDLE_TIMEOUT_SECONDS", "30")
|
||||||
|
try:
|
||||||
|
parsed = float(raw_value)
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
return 30.0
|
||||||
|
return parsed if parsed > 0 else 30.0
|
||||||
|
|
||||||
|
|
||||||
class CliMonitorMixin:
|
class CliMonitorMixin:
|
||||||
"""监控和统计相关方法的 Mixin"""
|
"""监控和统计相关方法的 Mixin"""
|
||||||
|
|
||||||
# CancelledError 归因时,断连检查参数(秒)
|
# CancelledError 归因时,断连检查参数(秒)
|
||||||
CANCEL_DISCONNECT_CHECK_TIMEOUT_SECONDS = 0.5
|
CANCEL_DISCONNECT_CHECK_TIMEOUT_SECONDS = 0.5
|
||||||
CANCEL_DISCONNECT_RETRY_DELAYS_SECONDS = (0.1, 0.2)
|
CANCEL_DISCONNECT_RETRY_DELAYS_SECONDS = (0.1, 0.2)
|
||||||
|
# 流式传输过程中若长时间没有任何 chunk,提前判定为 idle timeout,避免一直等到 worker 超时
|
||||||
|
STREAM_IDLE_TIMEOUT_SECONDS = _read_stream_idle_timeout_seconds()
|
||||||
|
|
||||||
async def _probe_client_disconnect(
|
async def _probe_client_disconnect(
|
||||||
self,
|
self,
|
||||||
@@ -115,8 +127,29 @@ class CliMonitorMixin:
|
|||||||
|
|
||||||
last_chunk_time = time_module.time()
|
last_chunk_time = time_module.time()
|
||||||
chunk_count = 0
|
chunk_count = 0
|
||||||
|
stream_started = False
|
||||||
|
idle_timeout_triggered = False
|
||||||
|
idle_timeout = self.STREAM_IDLE_TIMEOUT_SECONDS
|
||||||
|
parent_task = asyncio.current_task()
|
||||||
|
idle_watch_task: asyncio.Task[None] | None = None
|
||||||
|
|
||||||
|
async def watch_stream_idle_timeout() -> None:
|
||||||
|
nonlocal idle_timeout_triggered
|
||||||
|
if parent_task is None:
|
||||||
|
return
|
||||||
|
poll_interval = min(1.0, max(0.1, idle_timeout / 5))
|
||||||
|
while not ctx.has_completion:
|
||||||
|
await asyncio.sleep(poll_interval)
|
||||||
|
if not stream_started:
|
||||||
|
continue
|
||||||
|
if (time_module.time() - last_chunk_time) <= idle_timeout:
|
||||||
|
continue
|
||||||
|
idle_timeout_triggered = True
|
||||||
|
parent_task.cancel()
|
||||||
|
return
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
idle_watch_task = asyncio.create_task(watch_stream_idle_timeout())
|
||||||
if http_request is not None:
|
if http_request is not None:
|
||||||
# 使用后台任务检测断连,完全不阻塞流式传输
|
# 使用后台任务检测断连,完全不阻塞流式传输
|
||||||
disconnected = False
|
disconnected = False
|
||||||
@@ -149,6 +182,7 @@ class CliMonitorMixin:
|
|||||||
ctx.status_code = 499
|
ctx.status_code = 499
|
||||||
ctx.error_message = "client_disconnected"
|
ctx.error_message = "client_disconnected"
|
||||||
break
|
break
|
||||||
|
stream_started = True
|
||||||
last_chunk_time = time_module.time()
|
last_chunk_time = time_module.time()
|
||||||
chunk_count += 1
|
chunk_count += 1
|
||||||
yield chunk
|
yield chunk
|
||||||
@@ -158,14 +192,33 @@ class CliMonitorMixin:
|
|||||||
await check_task
|
await check_task
|
||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
pass
|
pass
|
||||||
|
if idle_watch_task is not None:
|
||||||
|
idle_watch_task.cancel()
|
||||||
|
try:
|
||||||
|
await idle_watch_task
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
pass
|
||||||
else:
|
else:
|
||||||
# 无 http_request,仅被动监控
|
# 无 http_request,仅被动监控
|
||||||
async for chunk in stream_generator:
|
try:
|
||||||
last_chunk_time = time_module.time()
|
async for chunk in stream_generator:
|
||||||
chunk_count += 1
|
stream_started = True
|
||||||
yield chunk
|
last_chunk_time = time_module.time()
|
||||||
|
chunk_count += 1
|
||||||
|
yield chunk
|
||||||
|
finally:
|
||||||
|
if idle_watch_task is not None:
|
||||||
|
idle_watch_task.cancel()
|
||||||
|
try:
|
||||||
|
await idle_watch_task
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
pass
|
||||||
|
|
||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
|
# 防御性清理:正常路径中 idle_watch_task 已由内部 finally 取消,
|
||||||
|
# 但若 CancelledError 在异常路径传播,确保不留孤儿 task。
|
||||||
|
if idle_watch_task is not None and not idle_watch_task.done():
|
||||||
|
idle_watch_task.cancel()
|
||||||
# 注意:CancelledError 不等于"用户手动取消",它既可能是客户端断连触发,
|
# 注意:CancelledError 不等于"用户手动取消",它既可能是客户端断连触发,
|
||||||
# 也可能是服务端(重载/关停/内部取消)导致的协程取消。
|
# 也可能是服务端(重载/关停/内部取消)导致的协程取消。
|
||||||
# 这里尽量做一次"断连归因":仅当能确认客户端已断开时才记为 499 cancelled。
|
# 这里尽量做一次"断连归因":仅当能确认客户端已断开时才记为 499 cancelled。
|
||||||
@@ -173,6 +226,28 @@ class CliMonitorMixin:
|
|||||||
if not ctx.has_completion:
|
if not ctx.has_completion:
|
||||||
ctx.ensure_estimated_output_tokens()
|
ctx.ensure_estimated_output_tokens()
|
||||||
|
|
||||||
|
if not ctx.has_completion and idle_timeout_triggered:
|
||||||
|
ctx.status_code = 504
|
||||||
|
ctx.error_message = "stream_idle_timeout"
|
||||||
|
cancel_origin = "stream_idle_timeout"
|
||||||
|
logger.warning(
|
||||||
|
f"ID:{ctx.request_id} | Stream idle timeout: "
|
||||||
|
f"idle_timeout={idle_timeout:g}s, "
|
||||||
|
f"chunks={chunk_count}, "
|
||||||
|
f"has_completion={ctx.has_completion}, "
|
||||||
|
f"time_since_last_chunk={time_since_last_chunk:.2f}s, "
|
||||||
|
f"output_tokens={ctx.output_tokens}"
|
||||||
|
)
|
||||||
|
ctx.upstream_response = (
|
||||||
|
f"cancel_origin={cancel_origin}, "
|
||||||
|
f"chunks={chunk_count}, "
|
||||||
|
f"has_completion={ctx.has_completion}, "
|
||||||
|
f"time_since_last_chunk={time_since_last_chunk:.2f}s, "
|
||||||
|
f"output_tokens={ctx.output_tokens}, "
|
||||||
|
f"idle_timeout={idle_timeout:g}s"
|
||||||
|
)
|
||||||
|
raise
|
||||||
|
|
||||||
is_client_disconnected = False
|
is_client_disconnected = False
|
||||||
disconnect_check_uncertain = False
|
disconnect_check_uncertain = False
|
||||||
if http_request is not None:
|
if http_request is not None:
|
||||||
@@ -229,10 +304,14 @@ class CliMonitorMixin:
|
|||||||
)
|
)
|
||||||
raise
|
raise
|
||||||
except httpx.TimeoutException as e:
|
except httpx.TimeoutException as e:
|
||||||
|
if idle_watch_task is not None and not idle_watch_task.done():
|
||||||
|
idle_watch_task.cancel()
|
||||||
ctx.status_code = 504
|
ctx.status_code = 504
|
||||||
ctx.error_message = str(e)
|
ctx.error_message = str(e)
|
||||||
raise
|
raise
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
|
if idle_watch_task is not None and not idle_watch_task.done():
|
||||||
|
idle_watch_task.cancel()
|
||||||
ctx.status_code = 500
|
ctx.status_code = 500
|
||||||
ctx.error_message = str(e)
|
ctx.error_message = str(e)
|
||||||
raise
|
raise
|
||||||
@@ -294,59 +373,149 @@ class CliMonitorMixin:
|
|||||||
ctx, ctx.provider_request_body or original_request_body
|
ctx, ctx.provider_request_body or original_request_body
|
||||||
)
|
)
|
||||||
|
|
||||||
response_body = ctx.build_response_body(response_time_ms)
|
with ctx.managed_recorded_bodies(response_time_ms) as recorded_bodies:
|
||||||
client_response_body = ctx.build_client_response_body(response_time_ms)
|
# 根据状态码决定记录成功还是失败
|
||||||
|
# 499 = 客户端取消(不算系统失败);其他 4xx/5xx 视为失败
|
||||||
|
if ctx.status_code and ctx.status_code >= 400:
|
||||||
|
client_response_headers = ctx.client_response_headers or {
|
||||||
|
"content-type": "application/json"
|
||||||
|
}
|
||||||
|
|
||||||
# 根据状态码决定记录成功还是失败
|
if ctx.is_client_disconnected():
|
||||||
# 499 = 客户端取消(不算系统失败);其他 4xx/5xx 视为失败
|
# 客户端取消:记录为 cancelled(不算系统失败)
|
||||||
if ctx.status_code and ctx.status_code >= 400:
|
request_metadata = self._merge_scheduling_metadata(
|
||||||
client_response_headers = ctx.client_response_headers or {
|
{"perf": ctx.perf_metrics} if ctx.perf_metrics else None,
|
||||||
"content-type": "application/json"
|
selected_key_id=ctx.key_id,
|
||||||
}
|
candidate_keys=ctx.candidate_keys,
|
||||||
|
pool_summary=ctx.pool_summary,
|
||||||
if ctx.is_client_disconnected():
|
fallback_from_request=True,
|
||||||
# 客户端取消:记录为 cancelled(不算系统失败)
|
)
|
||||||
request_metadata = self._merge_scheduling_metadata(
|
await bg_telemetry.record_cancelled(
|
||||||
{"perf": ctx.perf_metrics} if ctx.perf_metrics else None,
|
provider=ctx.provider_name or "unknown",
|
||||||
selected_key_id=ctx.key_id,
|
model=ctx.model,
|
||||||
candidate_keys=ctx.candidate_keys,
|
response_time_ms=response_time_ms,
|
||||||
pool_summary=ctx.pool_summary,
|
first_byte_time_ms=ctx.first_byte_time_ms,
|
||||||
fallback_from_request=True,
|
status_code=ctx.status_code,
|
||||||
)
|
request_headers=original_headers,
|
||||||
await bg_telemetry.record_cancelled(
|
request_body=original_request_body,
|
||||||
provider=ctx.provider_name or "unknown",
|
is_stream=True,
|
||||||
model=ctx.model,
|
api_format=ctx.api_format,
|
||||||
response_time_ms=response_time_ms,
|
api_family=self.api_family,
|
||||||
first_byte_time_ms=ctx.first_byte_time_ms,
|
endpoint_kind=self.endpoint_kind,
|
||||||
status_code=ctx.status_code,
|
provider_request_headers=ctx.provider_request_headers,
|
||||||
request_headers=original_headers,
|
provider_request_body=ctx.provider_request_body,
|
||||||
request_body=original_request_body,
|
input_tokens=ctx.input_tokens,
|
||||||
is_stream=True,
|
output_tokens=ctx.output_tokens,
|
||||||
api_format=ctx.api_format,
|
cache_creation_tokens=ctx.cache_creation_tokens,
|
||||||
api_family=self.api_family,
|
cache_read_tokens=ctx.cached_tokens,
|
||||||
endpoint_kind=self.endpoint_kind,
|
response_body=recorded_bodies.response_body,
|
||||||
provider_request_headers=ctx.provider_request_headers,
|
client_response_body=recorded_bodies.client_response_body,
|
||||||
provider_request_body=ctx.provider_request_body,
|
response_headers=ctx.response_headers,
|
||||||
input_tokens=ctx.input_tokens,
|
client_response_headers=client_response_headers,
|
||||||
output_tokens=ctx.output_tokens,
|
endpoint_api_format=ctx.provider_api_format or None,
|
||||||
cache_creation_tokens=ctx.cache_creation_tokens,
|
has_format_conversion=ctx.has_format_conversion,
|
||||||
cache_read_tokens=ctx.cached_tokens,
|
target_model=ctx.mapped_model,
|
||||||
response_body=response_body,
|
request_metadata=request_metadata,
|
||||||
client_response_body=client_response_body,
|
)
|
||||||
response_headers=ctx.response_headers,
|
logger.debug("{} 流式响应被客户端取消", self.FORMAT_ID)
|
||||||
client_response_headers=client_response_headers,
|
logger.info(
|
||||||
endpoint_api_format=ctx.provider_api_format or None,
|
"[CANCEL] {} | {} | {} | {}ms | {} | in:{} out:{} cache:{}",
|
||||||
has_format_conversion=ctx.has_format_conversion,
|
self.request_id[:8],
|
||||||
target_model=ctx.mapped_model,
|
ctx.model,
|
||||||
request_metadata=request_metadata,
|
ctx.provider_name,
|
||||||
)
|
response_time_ms,
|
||||||
logger.debug("{} 流式响应被客户端取消", self.FORMAT_ID)
|
ctx.status_code,
|
||||||
logger.info(
|
ctx.input_tokens,
|
||||||
f"[CANCEL] {self.request_id[:8]} | {ctx.model} | {ctx.provider_name} | {response_time_ms}ms | "
|
ctx.output_tokens,
|
||||||
f"{ctx.status_code} | in:{ctx.input_tokens} out:{ctx.output_tokens} cache:{ctx.cached_tokens}"
|
ctx.cached_tokens,
|
||||||
)
|
)
|
||||||
|
else:
|
||||||
|
# 服务端/上游异常:记录为失败
|
||||||
|
request_metadata = self._merge_scheduling_metadata(
|
||||||
|
{"perf": ctx.perf_metrics} if ctx.perf_metrics else None,
|
||||||
|
selected_key_id=ctx.key_id,
|
||||||
|
candidate_keys=ctx.candidate_keys,
|
||||||
|
pool_summary=ctx.pool_summary,
|
||||||
|
fallback_from_request=True,
|
||||||
|
)
|
||||||
|
await bg_telemetry.record_failure(
|
||||||
|
provider=ctx.provider_name or "unknown",
|
||||||
|
model=ctx.model,
|
||||||
|
response_time_ms=response_time_ms,
|
||||||
|
status_code=ctx.status_code,
|
||||||
|
error_message=ctx.error_message or f"HTTP {ctx.status_code}",
|
||||||
|
request_headers=original_headers,
|
||||||
|
request_body=original_request_body,
|
||||||
|
is_stream=True,
|
||||||
|
api_format=ctx.api_format,
|
||||||
|
api_family=self.api_family,
|
||||||
|
endpoint_kind=self.endpoint_kind,
|
||||||
|
provider_request_headers=ctx.provider_request_headers,
|
||||||
|
provider_request_body=ctx.provider_request_body,
|
||||||
|
# 预估 token 信息(来自 message_start 事件)
|
||||||
|
input_tokens=ctx.input_tokens,
|
||||||
|
output_tokens=ctx.output_tokens,
|
||||||
|
cache_creation_tokens=ctx.cache_creation_tokens,
|
||||||
|
cache_read_tokens=ctx.cached_tokens,
|
||||||
|
response_body=recorded_bodies.response_body,
|
||||||
|
client_response_body=recorded_bodies.client_response_body,
|
||||||
|
response_headers=ctx.response_headers,
|
||||||
|
client_response_headers=client_response_headers,
|
||||||
|
# 格式转换追踪
|
||||||
|
endpoint_api_format=ctx.provider_api_format or None,
|
||||||
|
has_format_conversion=ctx.has_format_conversion,
|
||||||
|
# 模型映射信息
|
||||||
|
target_model=ctx.mapped_model,
|
||||||
|
request_metadata=request_metadata,
|
||||||
|
)
|
||||||
|
logger.debug("{} 流式响应中断", self.FORMAT_ID)
|
||||||
|
logger.info(
|
||||||
|
"[FAIL] {} | {} | {} | {}ms | {} | in:{} out:{} cache:{}",
|
||||||
|
self.request_id[:8],
|
||||||
|
ctx.model,
|
||||||
|
ctx.provider_name,
|
||||||
|
response_time_ms,
|
||||||
|
ctx.status_code,
|
||||||
|
ctx.input_tokens,
|
||||||
|
ctx.output_tokens,
|
||||||
|
ctx.cached_tokens,
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
# 服务端/上游异常:记录为失败
|
# 在记录统计前,允许子类从 parsed_chunks 中提取额外的元数据
|
||||||
|
self._finalize_stream_metadata(ctx)
|
||||||
|
|
||||||
|
# 流式格式转换汇总日志
|
||||||
|
if ctx.stream_conversion_event_count > 0:
|
||||||
|
logger.debug(
|
||||||
|
"[{}] 流式转换完成: {}->{}, total_events={}",
|
||||||
|
self.request_id[:8],
|
||||||
|
ctx.provider_api_format,
|
||||||
|
ctx.client_api_format,
|
||||||
|
ctx.stream_conversion_event_count,
|
||||||
|
)
|
||||||
|
|
||||||
|
# 流未正常完成(如上游截断/连接中断)且无 token 数据时,
|
||||||
|
# 从已收集的文本和请求体估算 tokens,避免 usage 记录为 0
|
||||||
|
# 流式成功时,返回给客户端的是提供商响应头 + SSE 必需头
|
||||||
|
client_response_headers = filter_proxy_response_headers(
|
||||||
|
ctx.response_headers
|
||||||
|
)
|
||||||
|
client_response_headers.update(
|
||||||
|
{
|
||||||
|
"Cache-Control": "no-cache, no-transform",
|
||||||
|
"X-Accel-Buffering": "no",
|
||||||
|
"content-type": "text/event-stream",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
logger.debug(
|
||||||
|
"[{}] 开始记录 Usage: provider={}, model={}, in={}, out={}",
|
||||||
|
ctx.request_id,
|
||||||
|
ctx.provider_name,
|
||||||
|
ctx.model,
|
||||||
|
ctx.input_tokens,
|
||||||
|
ctx.output_tokens,
|
||||||
|
)
|
||||||
request_metadata = self._merge_scheduling_metadata(
|
request_metadata = self._merge_scheduling_metadata(
|
||||||
{"perf": ctx.perf_metrics} if ctx.perf_metrics else None,
|
{"perf": ctx.perf_metrics} if ctx.perf_metrics else None,
|
||||||
selected_key_id=ctx.key_id,
|
selected_key_id=ctx.key_id,
|
||||||
@@ -354,129 +523,60 @@ class CliMonitorMixin:
|
|||||||
pool_summary=ctx.pool_summary,
|
pool_summary=ctx.pool_summary,
|
||||||
fallback_from_request=True,
|
fallback_from_request=True,
|
||||||
)
|
)
|
||||||
await bg_telemetry.record_failure(
|
total_cost = await bg_telemetry.record_success(
|
||||||
provider=ctx.provider_name or "unknown",
|
provider=ctx.provider_name,
|
||||||
model=ctx.model,
|
model=ctx.model,
|
||||||
|
input_tokens=ctx.input_tokens,
|
||||||
|
output_tokens=ctx.output_tokens,
|
||||||
response_time_ms=response_time_ms,
|
response_time_ms=response_time_ms,
|
||||||
|
first_byte_time_ms=ctx.first_byte_time_ms, # 传递首字时间
|
||||||
status_code=ctx.status_code,
|
status_code=ctx.status_code,
|
||||||
error_message=ctx.error_message or f"HTTP {ctx.status_code}",
|
|
||||||
request_headers=original_headers,
|
request_headers=original_headers,
|
||||||
request_body=original_request_body,
|
request_body=original_request_body,
|
||||||
|
response_headers=ctx.response_headers,
|
||||||
|
client_response_headers=client_response_headers,
|
||||||
|
response_body=recorded_bodies.response_body,
|
||||||
|
client_response_body=recorded_bodies.client_response_body,
|
||||||
|
provider_request_body=ctx.provider_request_body,
|
||||||
|
cache_creation_tokens=ctx.cache_creation_tokens,
|
||||||
|
cache_read_tokens=ctx.cached_tokens,
|
||||||
is_stream=True,
|
is_stream=True,
|
||||||
|
provider_request_headers=ctx.provider_request_headers,
|
||||||
api_format=ctx.api_format,
|
api_format=ctx.api_format,
|
||||||
api_family=self.api_family,
|
api_family=self.api_family,
|
||||||
endpoint_kind=self.endpoint_kind,
|
endpoint_kind=self.endpoint_kind,
|
||||||
provider_request_headers=ctx.provider_request_headers,
|
|
||||||
provider_request_body=ctx.provider_request_body,
|
|
||||||
# 预估 token 信息(来自 message_start 事件)
|
|
||||||
input_tokens=ctx.input_tokens,
|
|
||||||
output_tokens=ctx.output_tokens,
|
|
||||||
cache_creation_tokens=ctx.cache_creation_tokens,
|
|
||||||
cache_read_tokens=ctx.cached_tokens,
|
|
||||||
response_body=response_body,
|
|
||||||
client_response_body=client_response_body,
|
|
||||||
response_headers=ctx.response_headers,
|
|
||||||
client_response_headers=client_response_headers,
|
|
||||||
# 格式转换追踪
|
# 格式转换追踪
|
||||||
endpoint_api_format=ctx.provider_api_format or None,
|
endpoint_api_format=ctx.provider_api_format or None,
|
||||||
has_format_conversion=ctx.has_format_conversion,
|
has_format_conversion=ctx.has_format_conversion,
|
||||||
|
# Provider 侧追踪信息(用于记录真实成本)
|
||||||
|
provider_id=ctx.provider_id,
|
||||||
|
provider_endpoint_id=ctx.endpoint_id,
|
||||||
|
provider_api_key_id=ctx.key_id,
|
||||||
# 模型映射信息
|
# 模型映射信息
|
||||||
target_model=ctx.mapped_model,
|
target_model=ctx.mapped_model,
|
||||||
|
# Provider 响应元数据(如 Gemini 的 modelVersion)
|
||||||
|
response_metadata=(
|
||||||
|
ctx.response_metadata if ctx.response_metadata else None
|
||||||
|
),
|
||||||
request_metadata=request_metadata,
|
request_metadata=request_metadata,
|
||||||
)
|
)
|
||||||
logger.debug("{} 流式响应中断", self.FORMAT_ID)
|
|
||||||
logger.info(
|
|
||||||
f"[FAIL] {self.request_id[:8]} | {ctx.model} | {ctx.provider_name} | {response_time_ms}ms | "
|
|
||||||
f"{ctx.status_code} | in:{ctx.input_tokens} out:{ctx.output_tokens} cache:{ctx.cached_tokens}"
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
# 在记录统计前,允许子类从 parsed_chunks 中提取额外的元数据
|
|
||||||
self._finalize_stream_metadata(ctx)
|
|
||||||
|
|
||||||
# 流式格式转换汇总日志
|
|
||||||
if ctx.stream_conversion_event_count > 0:
|
|
||||||
logger.debug(
|
logger.debug(
|
||||||
"[{}] 流式转换完成: {}->{}, total_events={}",
|
"[{}] Usage 记录完成: cost=${:.6f}", ctx.request_id, total_cost
|
||||||
self.request_id[:8],
|
)
|
||||||
ctx.provider_api_format,
|
# 简洁的请求完成摘要(两行格式)
|
||||||
ctx.client_api_format,
|
ttfb_part = (
|
||||||
ctx.stream_conversion_event_count,
|
f" | TTFB: {ctx.first_byte_time_ms}ms" if ctx.first_byte_time_ms else ""
|
||||||
|
)
|
||||||
|
logger.info(
|
||||||
|
"[OK] {} | {} | {}{}\n Total: {}ms | in:{} out:{}",
|
||||||
|
self.request_id[:8],
|
||||||
|
ctx.model,
|
||||||
|
ctx.provider_name,
|
||||||
|
ttfb_part,
|
||||||
|
response_time_ms,
|
||||||
|
ctx.input_tokens or 0,
|
||||||
|
ctx.output_tokens or 0,
|
||||||
)
|
)
|
||||||
|
|
||||||
# 流未正常完成(如上游截断/连接中断)且无 token 数据时,
|
|
||||||
# 从已收集的文本和请求体估算 tokens,避免 usage 记录为 0
|
|
||||||
# 流式成功时,返回给客户端的是提供商响应头 + SSE 必需头
|
|
||||||
client_response_headers = filter_proxy_response_headers(ctx.response_headers)
|
|
||||||
client_response_headers.update(
|
|
||||||
{
|
|
||||||
"Cache-Control": "no-cache, no-transform",
|
|
||||||
"X-Accel-Buffering": "no",
|
|
||||||
"content-type": "text/event-stream",
|
|
||||||
}
|
|
||||||
)
|
|
||||||
|
|
||||||
logger.debug(
|
|
||||||
f"[{ctx.request_id}] 开始记录 Usage: "
|
|
||||||
f"provider={ctx.provider_name}, model={ctx.model}, "
|
|
||||||
f"in={ctx.input_tokens}, out={ctx.output_tokens}"
|
|
||||||
)
|
|
||||||
request_metadata = self._merge_scheduling_metadata(
|
|
||||||
{"perf": ctx.perf_metrics} if ctx.perf_metrics else None,
|
|
||||||
selected_key_id=ctx.key_id,
|
|
||||||
candidate_keys=ctx.candidate_keys,
|
|
||||||
pool_summary=ctx.pool_summary,
|
|
||||||
fallback_from_request=True,
|
|
||||||
)
|
|
||||||
total_cost = await bg_telemetry.record_success(
|
|
||||||
provider=ctx.provider_name,
|
|
||||||
model=ctx.model,
|
|
||||||
input_tokens=ctx.input_tokens,
|
|
||||||
output_tokens=ctx.output_tokens,
|
|
||||||
response_time_ms=response_time_ms,
|
|
||||||
first_byte_time_ms=ctx.first_byte_time_ms, # 传递首字时间
|
|
||||||
status_code=ctx.status_code,
|
|
||||||
request_headers=original_headers,
|
|
||||||
request_body=original_request_body,
|
|
||||||
response_headers=ctx.response_headers,
|
|
||||||
client_response_headers=client_response_headers,
|
|
||||||
response_body=response_body,
|
|
||||||
client_response_body=client_response_body,
|
|
||||||
provider_request_body=ctx.provider_request_body,
|
|
||||||
cache_creation_tokens=ctx.cache_creation_tokens,
|
|
||||||
cache_read_tokens=ctx.cached_tokens,
|
|
||||||
is_stream=True,
|
|
||||||
provider_request_headers=ctx.provider_request_headers,
|
|
||||||
api_format=ctx.api_format,
|
|
||||||
api_family=self.api_family,
|
|
||||||
endpoint_kind=self.endpoint_kind,
|
|
||||||
# 格式转换追踪
|
|
||||||
endpoint_api_format=ctx.provider_api_format or None,
|
|
||||||
has_format_conversion=ctx.has_format_conversion,
|
|
||||||
# Provider 侧追踪信息(用于记录真实成本)
|
|
||||||
provider_id=ctx.provider_id,
|
|
||||||
provider_endpoint_id=ctx.endpoint_id,
|
|
||||||
provider_api_key_id=ctx.key_id,
|
|
||||||
# 模型映射信息
|
|
||||||
target_model=ctx.mapped_model,
|
|
||||||
# Provider 响应元数据(如 Gemini 的 modelVersion)
|
|
||||||
response_metadata=ctx.response_metadata if ctx.response_metadata else None,
|
|
||||||
request_metadata=request_metadata,
|
|
||||||
)
|
|
||||||
logger.debug("[{}] Usage 记录完成: cost=${:.6f}", ctx.request_id, total_cost)
|
|
||||||
# 简洁的请求完成摘要(两行格式)
|
|
||||||
ttfb_part = (
|
|
||||||
f" | TTFB: {ctx.first_byte_time_ms}ms" if ctx.first_byte_time_ms else ""
|
|
||||||
)
|
|
||||||
logger.info(
|
|
||||||
"[OK] {} | {} | {}{}\n Total: {}ms | in:{} out:{}",
|
|
||||||
self.request_id[:8],
|
|
||||||
ctx.model,
|
|
||||||
ctx.provider_name,
|
|
||||||
ttfb_part,
|
|
||||||
response_time_ms,
|
|
||||||
ctx.input_tokens or 0,
|
|
||||||
ctx.output_tokens or 0,
|
|
||||||
)
|
|
||||||
|
|
||||||
# 更新候选记录的最终状态和延迟时间
|
# 更新候选记录的最终状态和延迟时间
|
||||||
# 注意:RequestExecutor 会在流开始时过早地标记成功(只记录了连接建立的时间)
|
# 注意:RequestExecutor 会在流开始时过早地标记成功(只记录了连接建立的时间)
|
||||||
@@ -556,6 +656,9 @@ class CliMonitorMixin:
|
|||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.exception("记录流式统计信息时出错")
|
logger.exception("记录流式统计信息时出错")
|
||||||
|
finally:
|
||||||
|
# 遥测写入完成后主动释放大对象列表,降低高并发长流的内存滞留。
|
||||||
|
ctx.release_recorded_chunks()
|
||||||
|
|
||||||
async def _record_stream_failure(
|
async def _record_stream_failure(
|
||||||
self,
|
self,
|
||||||
@@ -591,26 +694,30 @@ class CliMonitorMixin:
|
|||||||
pool_summary=ctx.pool_summary,
|
pool_summary=ctx.pool_summary,
|
||||||
fallback_from_request=True,
|
fallback_from_request=True,
|
||||||
)
|
)
|
||||||
await self.telemetry.record_failure(
|
try:
|
||||||
provider=ctx.provider_name or "unknown",
|
await self.telemetry.record_failure(
|
||||||
model=ctx.model,
|
provider=ctx.provider_name or "unknown",
|
||||||
response_time_ms=response_time_ms,
|
model=ctx.model,
|
||||||
status_code=status_code,
|
response_time_ms=response_time_ms,
|
||||||
error_message=extract_client_error_message(error),
|
status_code=status_code,
|
||||||
request_headers=original_headers,
|
error_message=extract_client_error_message(error),
|
||||||
request_body=original_request_body,
|
request_headers=original_headers,
|
||||||
is_stream=True,
|
request_body=original_request_body,
|
||||||
api_format=ctx.api_format,
|
is_stream=True,
|
||||||
api_family=self.api_family,
|
api_format=ctx.api_format,
|
||||||
endpoint_kind=self.endpoint_kind,
|
api_family=self.api_family,
|
||||||
provider_request_headers=ctx.provider_request_headers,
|
endpoint_kind=self.endpoint_kind,
|
||||||
provider_request_body=ctx.provider_request_body,
|
provider_request_headers=ctx.provider_request_headers,
|
||||||
response_headers=ctx.response_headers,
|
provider_request_body=ctx.provider_request_body,
|
||||||
client_response_headers=client_response_headers,
|
response_headers=ctx.response_headers,
|
||||||
# 格式转换追踪
|
client_response_headers=client_response_headers,
|
||||||
endpoint_api_format=ctx.provider_api_format or None,
|
# 格式转换追踪
|
||||||
has_format_conversion=ctx.has_format_conversion,
|
endpoint_api_format=ctx.provider_api_format or None,
|
||||||
# 模型映射信息
|
has_format_conversion=ctx.has_format_conversion,
|
||||||
target_model=ctx.mapped_model,
|
# 模型映射信息
|
||||||
request_metadata=request_metadata,
|
target_model=ctx.mapped_model,
|
||||||
)
|
request_metadata=request_metadata,
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
# 失败路径同样可能持有 chunk 审计数据,及时释放。
|
||||||
|
ctx.release_recorded_chunks()
|
||||||
|
|||||||
@@ -81,7 +81,7 @@ class CliHandlerProtocol(Protocol):
|
|||||||
api_format: str | None = ...,
|
api_format: str | None = ...,
|
||||||
request_headers: dict[str, Any] | None = ...,
|
request_headers: dict[str, Any] | None = ...,
|
||||||
request_body: dict[str, Any] | None = ...,
|
request_body: dict[str, Any] | None = ...,
|
||||||
) -> None: ...
|
) -> bool: ...
|
||||||
|
|
||||||
def _build_request_metadata(
|
def _build_request_metadata(
|
||||||
self,
|
self,
|
||||||
|
|||||||
@@ -92,7 +92,7 @@ class CliStreamMixin:
|
|||||||
client_api_format = self.primary_api_format
|
client_api_format = self.primary_api_format
|
||||||
|
|
||||||
# 提前创建 pending 记录,让前端可以立即看到"处理中"
|
# 提前创建 pending 记录,让前端可以立即看到"处理中"
|
||||||
self._create_pending_usage(
|
pending_usage_created = self._create_pending_usage(
|
||||||
model=model,
|
model=model,
|
||||||
is_stream=True,
|
is_stream=True,
|
||||||
request_type="chat",
|
request_type="chat",
|
||||||
@@ -168,6 +168,8 @@ class CliStreamMixin:
|
|||||||
request_body_state=request_state,
|
request_body_state=request_state,
|
||||||
request_headers=original_headers,
|
request_headers=original_headers,
|
||||||
request_body=original_request_body,
|
request_body=original_request_body,
|
||||||
|
# 预创建失败时,回退到 TaskService 侧创建,避免丢失 pending 状态。
|
||||||
|
create_pending_usage=not pending_usage_created,
|
||||||
)
|
)
|
||||||
stream_generator = exec_result.response
|
stream_generator = exec_result.response
|
||||||
provider_name = exec_result.provider_name or "unknown"
|
provider_name = exec_result.provider_name or "unknown"
|
||||||
@@ -260,8 +262,7 @@ class CliStreamMixin:
|
|||||||
) -> AsyncGenerator[bytes]:
|
) -> AsyncGenerator[bytes]:
|
||||||
"""执行流式请求并返回流生成器"""
|
"""执行流式请求并返回流生成器"""
|
||||||
# 重置上下文状态(重试时清除之前的数据,避免累积)
|
# 重置上下文状态(重试时清除之前的数据,避免累积)
|
||||||
ctx.parsed_chunks = []
|
ctx.release_recorded_chunks()
|
||||||
ctx.provider_parsed_chunks = []
|
|
||||||
ctx.chunk_count = 0
|
ctx.chunk_count = 0
|
||||||
ctx.data_count = 0
|
ctx.data_count = 0
|
||||||
ctx.has_completion = False
|
ctx.has_completion = False
|
||||||
|
|||||||
@@ -75,7 +75,7 @@ class CliSyncMixin:
|
|||||||
sync_start_time = time.time()
|
sync_start_time = time.time()
|
||||||
|
|
||||||
# 提前创建 pending 记录,让前端可以立即看到"处理中"
|
# 提前创建 pending 记录,让前端可以立即看到"处理中"
|
||||||
self._create_pending_usage(
|
pending_usage_created = self._create_pending_usage(
|
||||||
model=model,
|
model=model,
|
||||||
is_stream=False,
|
is_stream=False,
|
||||||
request_type="chat",
|
request_type="chat",
|
||||||
@@ -390,6 +390,8 @@ class CliSyncMixin:
|
|||||||
request_body_state=request_state,
|
request_body_state=request_state,
|
||||||
request_headers=original_headers,
|
request_headers=original_headers,
|
||||||
request_body=original_request_body,
|
request_body=original_request_body,
|
||||||
|
# 预创建失败时,回退到 TaskService 侧创建,避免丢失 pending 状态。
|
||||||
|
create_pending_usage=not pending_usage_created,
|
||||||
)
|
)
|
||||||
result = exec_result.response
|
result = exec_result.response
|
||||||
actual_provider_name = exec_result.provider_name or "unknown"
|
actual_provider_name = exec_result.provider_name or "unknown"
|
||||||
|
|||||||
@@ -12,8 +12,9 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import json
|
import json
|
||||||
import time
|
import time
|
||||||
|
from contextlib import contextmanager
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from typing import TYPE_CHECKING, Any
|
from typing import TYPE_CHECKING, Any, Iterator
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from src.core.api_format.conversion.stream_state import StreamState
|
from src.core.api_format.conversion.stream_state import StreamState
|
||||||
@@ -49,6 +50,21 @@ def is_format_converted(
|
|||||||
_MAX_COLLECTED_TEXT_CHARS = 16 * 1024
|
_MAX_COLLECTED_TEXT_CHARS = 16 * 1024
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class RecordedStreamBodies:
|
||||||
|
"""统一封装 telemetry/usage 使用的流式响应体引用。"""
|
||||||
|
|
||||||
|
response_body: dict[str, Any] | None
|
||||||
|
client_response_body: dict[str, Any] | None
|
||||||
|
|
||||||
|
def ensure_populated(self, ctx: StreamContext, response_time_ms: int) -> None:
|
||||||
|
"""在 fallback 到需要 body 的路径时按需补建响应体。"""
|
||||||
|
if self.response_body is None:
|
||||||
|
self.response_body = ctx.build_response_body(response_time_ms)
|
||||||
|
if self.client_response_body is None:
|
||||||
|
self.client_response_body = ctx.build_client_response_body(response_time_ms)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class StreamContext:
|
class StreamContext:
|
||||||
"""
|
"""
|
||||||
@@ -160,8 +176,7 @@ class StreamContext:
|
|||||||
在故障转移重试时调用,清除之前的数据避免累积。
|
在故障转移重试时调用,清除之前的数据避免累积。
|
||||||
保留 model 和 api_format,重置其他所有状态。
|
保留 model 和 api_format,重置其他所有状态。
|
||||||
"""
|
"""
|
||||||
self.parsed_chunks = []
|
self.release_recorded_chunks()
|
||||||
self.provider_parsed_chunks = []
|
|
||||||
self.chunk_count = 0
|
self.chunk_count = 0
|
||||||
self.data_count = 0
|
self.data_count = 0
|
||||||
self.has_completion = False
|
self.has_completion = False
|
||||||
@@ -194,6 +209,32 @@ class StreamContext:
|
|||||||
self.needs_conversion = False
|
self.needs_conversion = False
|
||||||
self.selected_base_url = None
|
self.selected_base_url = None
|
||||||
|
|
||||||
|
def release_recorded_chunks(self) -> None:
|
||||||
|
"""释放 telemetry/usage 已消费完的 chunk 列表,避免后台任务继续持有大对象。"""
|
||||||
|
self.parsed_chunks = []
|
||||||
|
self.provider_parsed_chunks = []
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def managed_recorded_bodies(
|
||||||
|
self,
|
||||||
|
response_time_ms: int,
|
||||||
|
*,
|
||||||
|
include_bodies: bool = True,
|
||||||
|
) -> Iterator[RecordedStreamBodies]:
|
||||||
|
"""统一管理响应体构建与 chunk 释放,避免 telemetry 路径重复写 finally。"""
|
||||||
|
recorded_bodies = RecordedStreamBodies(
|
||||||
|
response_body=self.build_response_body(response_time_ms) if include_bodies else None,
|
||||||
|
client_response_body=(
|
||||||
|
self.build_client_response_body(response_time_ms) if include_bodies else None
|
||||||
|
),
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
yield recorded_bodies
|
||||||
|
finally:
|
||||||
|
self.release_recorded_chunks()
|
||||||
|
recorded_bodies.response_body = None
|
||||||
|
recorded_bodies.client_response_body = None
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def collected_text(self) -> str:
|
def collected_text(self) -> str:
|
||||||
"""已收集的文本内容(按需拼接,避免在流式过程中频繁做字符串拷贝)"""
|
"""已收集的文本内容(按需拼接,避免在流式过程中频繁做字符串拷贝)"""
|
||||||
|
|||||||
@@ -99,6 +99,7 @@ class StreamTelemetryRecorder:
|
|||||||
try:
|
try:
|
||||||
writer = await self._get_telemetry_writer(bg_db, ctx, response_time_ms)
|
writer = await self._get_telemetry_writer(bg_db, ctx, response_time_ms)
|
||||||
if writer is None:
|
if writer is None:
|
||||||
|
ctx.release_recorded_chunks()
|
||||||
return
|
return
|
||||||
# 兜底估算:流未正常完成且 token 均为 0 时,从请求体粗略估算。
|
# 兜底估算:流未正常完成且 token 均为 0 时,从请求体粗略估算。
|
||||||
# 覆盖成功但缺少 completion,以及已传出部分数据后被中断的场景。
|
# 覆盖成功但缺少 completion,以及已传出部分数据后被中断的场景。
|
||||||
@@ -114,53 +115,47 @@ class StreamTelemetryRecorder:
|
|||||||
if isinstance(writer, QueueTelemetryWriter)
|
if isinstance(writer, QueueTelemetryWriter)
|
||||||
else should_log_body
|
else should_log_body
|
||||||
)
|
)
|
||||||
response_body = (
|
with ctx.managed_recorded_bodies(
|
||||||
ctx.build_response_body(response_time_ms) if include_bodies else None
|
response_time_ms, include_bodies=include_bodies
|
||||||
)
|
) as recorded_bodies:
|
||||||
client_response_body = (
|
try:
|
||||||
ctx.build_client_response_body(response_time_ms) if include_bodies else None
|
await self._dispatch_record(
|
||||||
)
|
|
||||||
|
|
||||||
try:
|
|
||||||
await self._dispatch_record(
|
|
||||||
bg_db,
|
|
||||||
writer,
|
|
||||||
ctx,
|
|
||||||
original_headers,
|
|
||||||
original_request_body,
|
|
||||||
response_body,
|
|
||||||
response_time_ms,
|
|
||||||
client_response_body=client_response_body,
|
|
||||||
)
|
|
||||||
except Exception as writer_error:
|
|
||||||
if not isinstance(writer, QueueTelemetryWriter):
|
|
||||||
raise
|
|
||||||
logger.warning(
|
|
||||||
f"[{self.request_id}] Queue writer failed, falling back to DB: {writer_error}"
|
|
||||||
)
|
|
||||||
db_writer = self._build_db_writer(bg_db)
|
|
||||||
if db_writer is None:
|
|
||||||
await self._update_usage_status_directly(
|
|
||||||
bg_db,
|
bg_db,
|
||||||
status=self._get_status_from_ctx(ctx),
|
writer,
|
||||||
response_time_ms=response_time_ms,
|
ctx,
|
||||||
status_code=ctx.status_code,
|
original_headers,
|
||||||
|
original_request_body,
|
||||||
|
recorded_bodies.response_body,
|
||||||
|
response_time_ms,
|
||||||
|
client_response_body=recorded_bodies.client_response_body,
|
||||||
|
)
|
||||||
|
except Exception as writer_error:
|
||||||
|
if not isinstance(writer, QueueTelemetryWriter):
|
||||||
|
raise
|
||||||
|
logger.warning(
|
||||||
|
f"[{self.request_id}] Queue writer failed, falling back to DB: {writer_error}"
|
||||||
|
)
|
||||||
|
db_writer = self._build_db_writer(bg_db)
|
||||||
|
if db_writer is None:
|
||||||
|
await self._update_usage_status_directly(
|
||||||
|
bg_db,
|
||||||
|
status=self._get_status_from_ctx(ctx),
|
||||||
|
response_time_ms=response_time_ms,
|
||||||
|
status_code=ctx.status_code,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
if should_log_body:
|
||||||
|
recorded_bodies.ensure_populated(ctx, response_time_ms)
|
||||||
|
await self._dispatch_record(
|
||||||
|
bg_db,
|
||||||
|
db_writer,
|
||||||
|
ctx,
|
||||||
|
original_headers,
|
||||||
|
original_request_body,
|
||||||
|
recorded_bodies.response_body,
|
||||||
|
response_time_ms,
|
||||||
|
client_response_body=recorded_bodies.client_response_body,
|
||||||
)
|
)
|
||||||
return
|
|
||||||
if response_body is None and should_log_body:
|
|
||||||
response_body = ctx.build_response_body(response_time_ms)
|
|
||||||
if client_response_body is None and should_log_body:
|
|
||||||
client_response_body = ctx.build_client_response_body(response_time_ms)
|
|
||||||
await self._dispatch_record(
|
|
||||||
bg_db,
|
|
||||||
db_writer,
|
|
||||||
ctx,
|
|
||||||
original_headers,
|
|
||||||
original_request_body,
|
|
||||||
response_body,
|
|
||||||
response_time_ms,
|
|
||||||
client_response_body=client_response_body,
|
|
||||||
)
|
|
||||||
|
|
||||||
# 更新候选记录状态
|
# 更新候选记录状态
|
||||||
await self._update_candidate_status(bg_db, ctx, response_time_ms, start_time)
|
await self._update_candidate_status(bg_db, ctx, response_time_ms, start_time)
|
||||||
@@ -175,6 +170,9 @@ class StreamTelemetryRecorder:
|
|||||||
response_time_ms=response_time_ms,
|
response_time_ms=response_time_ms,
|
||||||
error_message=f"记录统计信息失败: {str(e)[:200]}",
|
error_message=f"记录统计信息失败: {str(e)[:200]}",
|
||||||
)
|
)
|
||||||
|
finally:
|
||||||
|
# 遥测写入后主动释放大列表,避免长流式请求对象滞留在 worker 堆中。
|
||||||
|
ctx.release_recorded_chunks()
|
||||||
|
|
||||||
async def _record_success(
|
async def _record_success(
|
||||||
self,
|
self,
|
||||||
|
|||||||
@@ -47,6 +47,32 @@ router = APIRouter(prefix="/api/public", tags=["System Catalog"])
|
|||||||
pipeline = get_pipeline()
|
pipeline = get_pipeline()
|
||||||
|
|
||||||
|
|
||||||
|
def _fetch_recent_public_health_attempts_for_api_format(
|
||||||
|
db: Session,
|
||||||
|
*,
|
||||||
|
api_format: str,
|
||||||
|
since: datetime,
|
||||||
|
per_format_limit: int,
|
||||||
|
) -> list[RequestCandidate]:
|
||||||
|
"""获取单个 API 格式最近的最终态请求,用于公开监控事件列表。"""
|
||||||
|
final_statuses = ["success", "failed", "skipped"]
|
||||||
|
return (
|
||||||
|
db.query(RequestCandidate)
|
||||||
|
.join(ProviderEndpoint, RequestCandidate.endpoint_id == ProviderEndpoint.id)
|
||||||
|
.join(Provider, ProviderEndpoint.provider_id == Provider.id)
|
||||||
|
.filter(
|
||||||
|
ProviderEndpoint.is_active.is_(True),
|
||||||
|
Provider.is_active.is_(True),
|
||||||
|
ProviderEndpoint.api_format == api_format,
|
||||||
|
RequestCandidate.created_at >= since,
|
||||||
|
RequestCandidate.status.in_(final_statuses),
|
||||||
|
)
|
||||||
|
.order_by(RequestCandidate.created_at.desc())
|
||||||
|
.limit(per_format_limit)
|
||||||
|
.all()
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@router.get("/site-info")
|
@router.get("/site-info")
|
||||||
def get_site_info(db: Session = Depends(get_db)) -> dict[str, str]:
|
def get_site_info(db: Session = Depends(get_db)) -> dict[str, str]:
|
||||||
"""获取站点基本信息(公开接口,无需认证)"""
|
"""获取站点基本信息(公开接口,无需认证)"""
|
||||||
@@ -640,47 +666,51 @@ class PublicApiFormatHealthMonitorAdapter(PublicApiAdapter):
|
|||||||
)
|
)
|
||||||
endpoint_map[api_format].append(endpoint_id)
|
endpoint_map[api_format].append(endpoint_id)
|
||||||
|
|
||||||
# 2. 获取最近一段时间的 RequestCandidate(限制数量)
|
# 2. 统计窗口内每个 API 格式的真实状态分布
|
||||||
# 只查询最终状态的记录:success, failed, skipped
|
|
||||||
final_statuses = ["success", "failed", "skipped"]
|
final_statuses = ["success", "failed", "skipped"]
|
||||||
limit_rows = max(500, self.per_format_limit * 10)
|
status_counts_query = (
|
||||||
rows = (
|
|
||||||
db.query(
|
db.query(
|
||||||
RequestCandidate,
|
|
||||||
ProviderEndpoint.api_format,
|
ProviderEndpoint.api_format,
|
||||||
|
RequestCandidate.status,
|
||||||
|
func.count(RequestCandidate.id).label("count"),
|
||||||
)
|
)
|
||||||
.join(ProviderEndpoint, RequestCandidate.endpoint_id == ProviderEndpoint.id)
|
.join(RequestCandidate, ProviderEndpoint.id == RequestCandidate.endpoint_id)
|
||||||
|
.join(Provider, ProviderEndpoint.provider_id == Provider.id)
|
||||||
.filter(
|
.filter(
|
||||||
|
ProviderEndpoint.is_active.is_(True),
|
||||||
|
Provider.is_active.is_(True),
|
||||||
RequestCandidate.created_at >= since,
|
RequestCandidate.created_at >= since,
|
||||||
RequestCandidate.status.in_(final_statuses),
|
RequestCandidate.status.in_(final_statuses),
|
||||||
)
|
)
|
||||||
.order_by(RequestCandidate.created_at.desc())
|
.group_by(ProviderEndpoint.api_format, RequestCandidate.status)
|
||||||
.limit(limit_rows)
|
|
||||||
.all()
|
.all()
|
||||||
)
|
)
|
||||||
|
|
||||||
grouped_candidates: dict[str, list[RequestCandidate]] = {}
|
status_counts: dict[str, dict[str, int]] = {}
|
||||||
|
for api_format_enum, status, count in status_counts_query:
|
||||||
for candidate, api_format_enum in rows:
|
|
||||||
api_format = (
|
api_format = (
|
||||||
api_format_enum.value if hasattr(api_format_enum, "value") else str(api_format_enum)
|
api_format_enum.value if hasattr(api_format_enum, "value") else str(api_format_enum)
|
||||||
)
|
)
|
||||||
if api_format not in grouped_candidates:
|
if api_format not in status_counts:
|
||||||
grouped_candidates[api_format] = []
|
status_counts[api_format] = {"success": 0, "failed": 0, "skipped": 0}
|
||||||
|
status_counts[api_format][status] = count
|
||||||
if len(grouped_candidates[api_format]) < self.per_format_limit:
|
|
||||||
grouped_candidates[api_format].append(candidate)
|
|
||||||
|
|
||||||
# 3. 为所有活跃格式生成监控数据
|
# 3. 为所有活跃格式生成监控数据
|
||||||
monitors: list[PublicApiFormatHealthMonitor] = []
|
monitors: list[PublicApiFormatHealthMonitor] = []
|
||||||
for api_format in all_formats:
|
for api_format in all_formats:
|
||||||
candidates = grouped_candidates.get(api_format, [])
|
candidates = _fetch_recent_public_health_attempts_for_api_format(
|
||||||
|
db=db,
|
||||||
|
api_format=api_format,
|
||||||
|
since=since,
|
||||||
|
per_format_limit=self.per_format_limit,
|
||||||
|
)
|
||||||
|
|
||||||
# 统计
|
# 统计使用窗口内真实总数,events 仅保留最近样本用于展示。
|
||||||
success_count = sum(1 for c in candidates if c.status == "success")
|
format_stats = status_counts.get(api_format, {"success": 0, "failed": 0, "skipped": 0})
|
||||||
failed_count = sum(1 for c in candidates if c.status == "failed")
|
success_count = format_stats.get("success", 0)
|
||||||
skipped_count = sum(1 for c in candidates if c.status == "skipped")
|
failed_count = format_stats.get("failed", 0)
|
||||||
total_attempts = len(candidates)
|
skipped_count = format_stats.get("skipped", 0)
|
||||||
|
total_attempts = success_count + failed_count + skipped_count
|
||||||
|
|
||||||
# 计算成功率 = success / (success + failed)
|
# 计算成功率 = success / (success + failed)
|
||||||
actual_completed = success_count + failed_count
|
actual_completed = success_count + failed_count
|
||||||
|
|||||||
@@ -49,7 +49,8 @@ class PluginManager:
|
|||||||
# notification 默认不加载,避免未配置插件(如 email)在启动时初始化失败并占用内存。
|
# notification 默认不加载,避免未配置插件(如 email)在启动时初始化失败并占用内存。
|
||||||
DEFAULT_ENABLED_PLUGIN_MODULES: dict[str, tuple[str, ...]] = {
|
DEFAULT_ENABLED_PLUGIN_MODULES: dict[str, tuple[str, ...]] = {
|
||||||
"auth": ("api_key",),
|
"auth": ("api_key",),
|
||||||
"rate_limit": ("sliding_window",),
|
# 默认切到 token_bucket;全局 Redis 就绪时会自动切到分布式后端。
|
||||||
|
"rate_limit": ("token_bucket",),
|
||||||
"cache": ("memory",),
|
"cache": ("memory",),
|
||||||
"monitor": ("prometheus",),
|
"monitor": ("prometheus",),
|
||||||
"token": ("claude",),
|
"token": ("claude",),
|
||||||
|
|||||||
@@ -29,6 +29,7 @@ class TokenBucket:
|
|||||||
self.refill_rate = refill_rate
|
self.refill_rate = refill_rate
|
||||||
self.tokens = capacity
|
self.tokens = capacity
|
||||||
self.last_refill = time.time()
|
self.last_refill = time.time()
|
||||||
|
self.last_access_time = self.last_refill
|
||||||
|
|
||||||
def _refill(self) -> None:
|
def _refill(self) -> None:
|
||||||
"""补充令牌"""
|
"""补充令牌"""
|
||||||
@@ -36,6 +37,7 @@ class TokenBucket:
|
|||||||
time_passed = now - self.last_refill
|
time_passed = now - self.last_refill
|
||||||
tokens_to_add = time_passed * self.refill_rate
|
tokens_to_add = time_passed * self.refill_rate
|
||||||
|
|
||||||
|
self.last_access_time = now
|
||||||
if tokens_to_add > 0:
|
if tokens_to_add > 0:
|
||||||
self.tokens = min(self.capacity, self.tokens + tokens_to_add)
|
self.tokens = min(self.capacity, self.tokens + tokens_to_add)
|
||||||
self.last_refill = now
|
self.last_refill = now
|
||||||
@@ -64,7 +66,7 @@ class TokenBucket:
|
|||||||
|
|
||||||
def get_reset_time(self) -> datetime:
|
def get_reset_time(self) -> datetime:
|
||||||
"""获取下次完全恢复的时间"""
|
"""获取下次完全恢复的时间"""
|
||||||
if self.tokens >= self.capacity:
|
if self.tokens >= self.capacity or self.refill_rate <= 0:
|
||||||
return datetime.now(timezone.utc)
|
return datetime.now(timezone.utc)
|
||||||
|
|
||||||
tokens_needed = self.capacity - self.tokens
|
tokens_needed = self.capacity - self.tokens
|
||||||
@@ -82,6 +84,10 @@ class TokenBucketStrategy(RateLimitStrategy):
|
|||||||
- 适合处理不均匀的流量模式
|
- 适合处理不均匀的流量模式
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
DEFAULT_MAX_BUCKETS = 10000
|
||||||
|
DEFAULT_BUCKET_EXPIRY = 3600
|
||||||
|
DEFAULT_REDIS_RETRY_INTERVAL = 30.0
|
||||||
|
|
||||||
def __init__(self) -> None:
|
def __init__(self) -> None:
|
||||||
super().__init__("token_bucket")
|
super().__init__("token_bucket")
|
||||||
self.buckets: dict[str, TokenBucket] = {}
|
self.buckets: dict[str, TokenBucket] = {}
|
||||||
@@ -90,11 +96,46 @@ class TokenBucketStrategy(RateLimitStrategy):
|
|||||||
# 默认配置
|
# 默认配置
|
||||||
self.default_capacity = 100 # 默认桶容量
|
self.default_capacity = 100 # 默认桶容量
|
||||||
self.default_refill_rate = 10 # 默认每秒补充10个令牌
|
self.default_refill_rate = 10 # 默认每秒补充10个令牌
|
||||||
|
self.max_buckets = self.DEFAULT_MAX_BUCKETS
|
||||||
|
self.bucket_expiry = self.DEFAULT_BUCKET_EXPIRY
|
||||||
|
self._last_cleanup_time: float = time.time()
|
||||||
|
self._cleanup_interval = 300 # 每 5 分钟检查一次清理
|
||||||
|
|
||||||
# 可选的 Redis 后端
|
# 可选的 Redis 后端
|
||||||
self._redis_backend: RedisTokenBucketBackend | None = None
|
self._redis_backend: RedisTokenBucketBackend | None = None
|
||||||
self._redis_checked = False
|
self._redis_checked = False
|
||||||
self._backend_mode = os.getenv("RATE_LIMIT_BACKEND", "auto").lower()
|
self._backend_mode = os.getenv("RATE_LIMIT_BACKEND", "auto").lower()
|
||||||
|
self._redis_retry_interval = self.DEFAULT_REDIS_RETRY_INTERVAL
|
||||||
|
self._next_redis_probe_time = 0.0
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _is_unlimited_rate_limit(rate_limit: Any) -> bool:
|
||||||
|
"""显式传入 0/负数时,按“不限流”处理。"""
|
||||||
|
if rate_limit is None:
|
||||||
|
return False
|
||||||
|
try:
|
||||||
|
return int(rate_limit) <= 0
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
return False
|
||||||
|
|
||||||
|
def _resolve_bucket_config(self, key: str, rate_limit: int | None = None) -> tuple[int, float]:
|
||||||
|
"""解析指定 key 当前应使用的桶容量和补充速率。"""
|
||||||
|
if rate_limit is not None:
|
||||||
|
normalized_rate_limit = int(rate_limit)
|
||||||
|
if normalized_rate_limit <= 0:
|
||||||
|
return 0, 0.0
|
||||||
|
return normalized_rate_limit, normalized_rate_limit / 60.0
|
||||||
|
if key.startswith("api_key:"):
|
||||||
|
return (
|
||||||
|
self.config.get("api_key_capacity", self.default_capacity),
|
||||||
|
self.config.get("api_key_refill_rate", self.default_refill_rate),
|
||||||
|
)
|
||||||
|
if key.startswith("user:"):
|
||||||
|
return (
|
||||||
|
self.config.get("user_capacity", self.default_capacity * 2),
|
||||||
|
self.config.get("user_refill_rate", self.default_refill_rate * 2),
|
||||||
|
)
|
||||||
|
return self.default_capacity, self.default_refill_rate
|
||||||
|
|
||||||
def _get_bucket(self, key: str, rate_limit: int | None = None) -> TokenBucket:
|
def _get_bucket(self, key: str, rate_limit: int | None = None) -> TokenBucket:
|
||||||
"""
|
"""
|
||||||
@@ -107,41 +148,88 @@ class TokenBucketStrategy(RateLimitStrategy):
|
|||||||
Returns:
|
Returns:
|
||||||
令牌桶实例
|
令牌桶实例
|
||||||
"""
|
"""
|
||||||
if key not in self.buckets:
|
capacity, refill_rate = self._resolve_bucket_config(key, rate_limit)
|
||||||
# 如果提供了rate_limit参数(来自数据库),优先使用
|
bucket = self.buckets.get(key)
|
||||||
if rate_limit is not None:
|
if bucket is None:
|
||||||
# rate_limit 是每分钟请求数,转换为令牌桶参数
|
bucket = TokenBucket(capacity, refill_rate)
|
||||||
capacity = rate_limit # 桶容量等于每分钟限制
|
self.buckets[key] = bucket
|
||||||
refill_rate = rate_limit / 60.0 # 每秒补充的令牌数
|
return bucket
|
||||||
# 否则根据key的不同前缀使用不同的配置
|
|
||||||
elif key.startswith("api_key:"):
|
|
||||||
capacity = self.config.get("api_key_capacity", self.default_capacity)
|
|
||||||
refill_rate = self.config.get("api_key_refill_rate", self.default_refill_rate)
|
|
||||||
elif key.startswith("user:"):
|
|
||||||
capacity = self.config.get("user_capacity", self.default_capacity * 2)
|
|
||||||
refill_rate = self.config.get("user_refill_rate", self.default_refill_rate * 2)
|
|
||||||
else:
|
|
||||||
capacity = self.default_capacity
|
|
||||||
refill_rate = self.default_refill_rate
|
|
||||||
|
|
||||||
self.buckets[key] = TokenBucket(capacity, refill_rate)
|
if bucket.capacity != capacity or bucket.refill_rate != refill_rate:
|
||||||
|
bucket._refill()
|
||||||
|
bucket.capacity = capacity
|
||||||
|
bucket.refill_rate = refill_rate
|
||||||
|
bucket.tokens = min(bucket.tokens, capacity)
|
||||||
|
|
||||||
return self.buckets[key]
|
return bucket
|
||||||
|
|
||||||
|
def _cleanup_expired_buckets(self) -> int:
|
||||||
|
"""清理长时间未访问的桶,避免 key 集合无限增长。"""
|
||||||
|
current_time = time.time()
|
||||||
|
expired_keys = [
|
||||||
|
key
|
||||||
|
for key, bucket in self.buckets.items()
|
||||||
|
if current_time - bucket.last_access_time > self.bucket_expiry
|
||||||
|
]
|
||||||
|
|
||||||
|
for key in expired_keys:
|
||||||
|
del self.buckets[key]
|
||||||
|
|
||||||
|
if expired_keys:
|
||||||
|
logger.info("清理了 {} 个过期的令牌桶", len(expired_keys))
|
||||||
|
|
||||||
|
return len(expired_keys)
|
||||||
|
|
||||||
|
def _evict_lru_buckets(self, count: int) -> int:
|
||||||
|
"""达到容量上限时淘汰最久未使用的桶。"""
|
||||||
|
if not self.buckets or count <= 0:
|
||||||
|
return 0
|
||||||
|
|
||||||
|
sorted_keys = sorted(self.buckets, key=lambda key: self.buckets[key].last_access_time)
|
||||||
|
evicted = 0
|
||||||
|
for key in sorted_keys[:count]:
|
||||||
|
del self.buckets[key]
|
||||||
|
evicted += 1
|
||||||
|
|
||||||
|
if evicted:
|
||||||
|
logger.warning("LRU 淘汰了 {} 个令牌桶(达到容量上限)", evicted)
|
||||||
|
|
||||||
|
return evicted
|
||||||
|
|
||||||
|
async def _maybe_cleanup(self) -> None:
|
||||||
|
"""定期清理或淘汰桶,控制进程内桶数量。"""
|
||||||
|
current_time = time.time()
|
||||||
|
if current_time - self._last_cleanup_time > self._cleanup_interval:
|
||||||
|
self._cleanup_expired_buckets()
|
||||||
|
self._last_cleanup_time = current_time
|
||||||
|
|
||||||
|
if len(self.buckets) >= self.max_buckets:
|
||||||
|
evict_count = max(1, self.max_buckets // 10)
|
||||||
|
self._evict_lru_buckets(evict_count)
|
||||||
|
|
||||||
def _want_redis_backend(self) -> bool:
|
def _want_redis_backend(self) -> bool:
|
||||||
return self._backend_mode in {"auto", "redis"}
|
return self._backend_mode in {"auto", "redis"}
|
||||||
|
|
||||||
async def _ensure_backend(self) -> None:
|
async def _ensure_backend(self) -> None:
|
||||||
if self._redis_checked:
|
if self._redis_backend is not None:
|
||||||
return
|
return
|
||||||
self._redis_checked = True
|
|
||||||
if not self._want_redis_backend():
|
if not self._want_redis_backend():
|
||||||
|
self._redis_checked = True
|
||||||
|
return
|
||||||
|
current_time = time.time()
|
||||||
|
if self._redis_checked and current_time < self._next_redis_probe_time:
|
||||||
return
|
return
|
||||||
redis_client = get_redis_client_sync()
|
redis_client = get_redis_client_sync()
|
||||||
if redis_client:
|
if redis_client:
|
||||||
self._redis_backend = RedisTokenBucketBackend(redis_client)
|
self._redis_backend = RedisTokenBucketBackend(redis_client)
|
||||||
|
self._redis_checked = True
|
||||||
|
self._next_redis_probe_time = 0.0
|
||||||
logger.info("速率限制改用 Redis 令牌桶后端")
|
logger.info("速率限制改用 Redis 令牌桶后端")
|
||||||
elif self._backend_mode == "redis":
|
return
|
||||||
|
|
||||||
|
self._redis_checked = True
|
||||||
|
self._next_redis_probe_time = current_time + self._redis_retry_interval
|
||||||
|
if self._backend_mode == "redis":
|
||||||
logger.warning("RATE_LIMIT_BACKEND=redis 但 Redis 客户端不可用,回退到内存桶")
|
logger.warning("RATE_LIMIT_BACKEND=redis 但 Redis 客户端不可用,回退到内存桶")
|
||||||
|
|
||||||
async def check_limit(self, key: str, **kwargs: Any) -> RateLimitResult:
|
async def check_limit(self, key: str, **kwargs: Any) -> RateLimitResult:
|
||||||
@@ -155,11 +243,14 @@ class TokenBucketStrategy(RateLimitStrategy):
|
|||||||
Returns:
|
Returns:
|
||||||
速率限制检查结果
|
速率限制检查结果
|
||||||
"""
|
"""
|
||||||
await self._ensure_backend()
|
|
||||||
|
|
||||||
rate_limit = kwargs.get("rate_limit")
|
rate_limit = kwargs.get("rate_limit")
|
||||||
amount = kwargs.get("amount", 1)
|
amount = kwargs.get("amount", 1)
|
||||||
|
|
||||||
|
if self._is_unlimited_rate_limit(rate_limit):
|
||||||
|
return RateLimitResult(allowed=True, remaining=0)
|
||||||
|
|
||||||
|
await self._ensure_backend()
|
||||||
|
|
||||||
if self._redis_backend:
|
if self._redis_backend:
|
||||||
return await self._redis_backend.peek(
|
return await self._redis_backend.peek(
|
||||||
key=key,
|
key=key,
|
||||||
@@ -169,6 +260,7 @@ class TokenBucketStrategy(RateLimitStrategy):
|
|||||||
)
|
)
|
||||||
|
|
||||||
async with self._lock:
|
async with self._lock:
|
||||||
|
await self._maybe_cleanup()
|
||||||
bucket = self._get_bucket(key, rate_limit)
|
bucket = self._get_bucket(key, rate_limit)
|
||||||
remaining = bucket.get_remaining()
|
remaining = bucket.get_remaining()
|
||||||
reset_at = bucket.get_reset_time()
|
reset_at = bucket.get_reset_time()
|
||||||
@@ -203,13 +295,17 @@ class TokenBucketStrategy(RateLimitStrategy):
|
|||||||
Returns:
|
Returns:
|
||||||
是否成功消费
|
是否成功消费
|
||||||
"""
|
"""
|
||||||
|
rate_limit = kwargs.get("rate_limit")
|
||||||
|
if self._is_unlimited_rate_limit(rate_limit):
|
||||||
|
return True
|
||||||
|
|
||||||
await self._ensure_backend()
|
await self._ensure_backend()
|
||||||
|
|
||||||
if self._redis_backend:
|
if self._redis_backend:
|
||||||
success, remaining = await self._redis_backend.consume(
|
success, remaining = await self._redis_backend.consume(
|
||||||
key=key,
|
key=key,
|
||||||
capacity=self._resolve_capacity(key, kwargs.get("rate_limit")),
|
capacity=self._resolve_capacity(key, rate_limit),
|
||||||
refill_rate=self._resolve_refill_rate(key, kwargs.get("rate_limit")),
|
refill_rate=self._resolve_refill_rate(key, rate_limit),
|
||||||
amount=amount,
|
amount=amount,
|
||||||
)
|
)
|
||||||
if success:
|
if success:
|
||||||
@@ -219,7 +315,8 @@ class TokenBucketStrategy(RateLimitStrategy):
|
|||||||
return success
|
return success
|
||||||
|
|
||||||
async with self._lock:
|
async with self._lock:
|
||||||
bucket = self._get_bucket(key)
|
await self._maybe_cleanup()
|
||||||
|
bucket = self._get_bucket(key, rate_limit)
|
||||||
success = bucket.consume(amount)
|
success = bucket.consume(amount)
|
||||||
|
|
||||||
if success:
|
if success:
|
||||||
@@ -270,6 +367,7 @@ class TokenBucketStrategy(RateLimitStrategy):
|
|||||||
)
|
)
|
||||||
|
|
||||||
async with self._lock:
|
async with self._lock:
|
||||||
|
await self._maybe_cleanup()
|
||||||
bucket = self._get_bucket(key)
|
bucket = self._get_bucket(key)
|
||||||
return {
|
return {
|
||||||
"strategy": "token_bucket",
|
"strategy": "token_bucket",
|
||||||
@@ -293,24 +391,17 @@ class TokenBucketStrategy(RateLimitStrategy):
|
|||||||
super().configure(config)
|
super().configure(config)
|
||||||
self.default_capacity = config.get("default_capacity", self.default_capacity)
|
self.default_capacity = config.get("default_capacity", self.default_capacity)
|
||||||
self.default_refill_rate = config.get("default_refill_rate", self.default_refill_rate)
|
self.default_refill_rate = config.get("default_refill_rate", self.default_refill_rate)
|
||||||
|
self.max_buckets = int(config.get("max_buckets", self.max_buckets))
|
||||||
|
self.bucket_expiry = int(config.get("bucket_expiry", self.bucket_expiry))
|
||||||
|
self._cleanup_interval = int(config.get("cleanup_interval", self._cleanup_interval))
|
||||||
|
|
||||||
def _resolve_capacity(self, key: str, rate_limit: int | None = None) -> int:
|
def _resolve_capacity(self, key: str, rate_limit: int | None = None) -> int:
|
||||||
if rate_limit is not None:
|
capacity, _ = self._resolve_bucket_config(key, rate_limit)
|
||||||
return rate_limit
|
return capacity
|
||||||
if key.startswith("api_key:"):
|
|
||||||
return self.config.get("api_key_capacity", self.default_capacity)
|
|
||||||
if key.startswith("user:"):
|
|
||||||
return self.config.get("user_capacity", self.default_capacity * 2)
|
|
||||||
return self.default_capacity
|
|
||||||
|
|
||||||
def _resolve_refill_rate(self, key: str, rate_limit: int | None = None) -> float:
|
def _resolve_refill_rate(self, key: str, rate_limit: int | None = None) -> float:
|
||||||
if rate_limit is not None:
|
_, refill_rate = self._resolve_bucket_config(key, rate_limit)
|
||||||
return rate_limit / 60.0
|
return refill_rate
|
||||||
if key.startswith("api_key:"):
|
|
||||||
return self.config.get("api_key_refill_rate", self.default_refill_rate)
|
|
||||||
if key.startswith("user:"):
|
|
||||||
return self.config.get("user_refill_rate", self.default_refill_rate * 2)
|
|
||||||
return self.default_refill_rate
|
|
||||||
|
|
||||||
|
|
||||||
class RedisTokenBucketBackend:
|
class RedisTokenBucketBackend:
|
||||||
@@ -365,6 +456,9 @@ class RedisTokenBucketBackend:
|
|||||||
refill_rate: float,
|
refill_rate: float,
|
||||||
amount: int,
|
amount: int,
|
||||||
) -> RateLimitResult:
|
) -> RateLimitResult:
|
||||||
|
if capacity <= 0 or refill_rate <= 0:
|
||||||
|
return RateLimitResult(allowed=True, remaining=0)
|
||||||
|
|
||||||
bucket_key = self._redis_key(key)
|
bucket_key = self._redis_key(key)
|
||||||
data = await self.redis.hmget(bucket_key, "tokens", "timestamp")
|
data = await self.redis.hmget(bucket_key, "tokens", "timestamp")
|
||||||
tokens = data[0]
|
tokens = data[0]
|
||||||
@@ -372,7 +466,7 @@ class RedisTokenBucketBackend:
|
|||||||
|
|
||||||
if tokens is None or last_refill is None:
|
if tokens is None or last_refill is None:
|
||||||
remaining = capacity
|
remaining = capacity
|
||||||
reset_at = datetime.now(timezone.utc) + timedelta(seconds=capacity / refill_rate)
|
reset_at = datetime.now(timezone.utc)
|
||||||
else:
|
else:
|
||||||
tokens_value = float(tokens)
|
tokens_value = float(tokens)
|
||||||
last_refill_value = float(last_refill)
|
last_refill_value = float(last_refill)
|
||||||
@@ -407,6 +501,9 @@ class RedisTokenBucketBackend:
|
|||||||
refill_rate: float,
|
refill_rate: float,
|
||||||
amount: int,
|
amount: int,
|
||||||
) -> tuple[bool, int]:
|
) -> tuple[bool, int]:
|
||||||
|
if capacity <= 0 or refill_rate <= 0:
|
||||||
|
return True, 0
|
||||||
|
|
||||||
result = await self._consume_script(
|
result = await self._consume_script(
|
||||||
keys=[self._redis_key(key)],
|
keys=[self._redis_key(key)],
|
||||||
args=[time.time(), capacity, refill_rate, amount],
|
args=[time.time(), capacity, refill_rate, amount],
|
||||||
|
|||||||
@@ -69,7 +69,13 @@ class EndpointHealthService:
|
|||||||
|
|
||||||
# 查询所有活跃的端点(一次性获取所有需要的数据)
|
# 查询所有活跃的端点(一次性获取所有需要的数据)
|
||||||
endpoints = (
|
endpoints = (
|
||||||
db.query(ProviderEndpoint).join(Provider).filter(Provider.is_active.is_(True)).all()
|
db.query(ProviderEndpoint)
|
||||||
|
.join(Provider)
|
||||||
|
.filter(
|
||||||
|
ProviderEndpoint.is_active.is_(True),
|
||||||
|
Provider.is_active.is_(True),
|
||||||
|
)
|
||||||
|
.all()
|
||||||
)
|
)
|
||||||
|
|
||||||
# 收集所有 provider_ids
|
# 收集所有 provider_ids
|
||||||
@@ -155,16 +161,13 @@ class EndpointHealthService:
|
|||||||
format_stats[api_format]["health_scores"].append(health_score)
|
format_stats[api_format]["health_scores"].append(health_score)
|
||||||
|
|
||||||
# 批量生成所有格式的时间线数据
|
# 批量生成所有格式的时间线数据
|
||||||
all_key_ids = []
|
format_endpoint_mapping: dict[str, list[str]] = {}
|
||||||
format_key_mapping: dict[str, list[str]] = {}
|
|
||||||
for api_format, stats in format_stats.items():
|
for api_format, stats in format_stats.items():
|
||||||
key_ids = stats["key_ids"]
|
format_endpoint_mapping[api_format] = stats["endpoint_ids"]
|
||||||
format_key_mapping[api_format] = key_ids
|
|
||||||
all_key_ids.extend(key_ids)
|
|
||||||
|
|
||||||
# 一次性查询所有时间线数据
|
# 一次性查询所有时间线数据
|
||||||
timeline_data_map = EndpointHealthService._generate_timeline_batch(
|
timeline_data_map = EndpointHealthService._generate_timeline_batch(
|
||||||
db, format_key_mapping, now, lookback_hours
|
db, format_endpoint_mapping, now, lookback_hours
|
||||||
)
|
)
|
||||||
|
|
||||||
# 生成结果
|
# 生成结果
|
||||||
@@ -234,7 +237,7 @@ class EndpointHealthService:
|
|||||||
@staticmethod
|
@staticmethod
|
||||||
def _generate_timeline_batch(
|
def _generate_timeline_batch(
|
||||||
db: Session,
|
db: Session,
|
||||||
format_key_mapping: dict[str, list[str]],
|
format_endpoint_mapping: dict[str, list[str]],
|
||||||
now: datetime,
|
now: datetime,
|
||||||
lookback_hours: int,
|
lookback_hours: int,
|
||||||
segments: int = 100,
|
segments: int = 100,
|
||||||
@@ -249,7 +252,7 @@ class EndpointHealthService:
|
|||||||
|
|
||||||
Args:
|
Args:
|
||||||
db: 数据库会话
|
db: 数据库会话
|
||||||
format_key_mapping: API格式 -> key_ids 的映射
|
format_endpoint_mapping: API格式 -> endpoint_ids 的映射
|
||||||
now: 当前时间
|
now: 当前时间
|
||||||
lookback_hours: 回溯小时数
|
lookback_hours: 回溯小时数
|
||||||
segments: 时间段数量
|
segments: 时间段数量
|
||||||
@@ -257,19 +260,19 @@ class EndpointHealthService:
|
|||||||
Returns:
|
Returns:
|
||||||
API格式 -> 时间线数据的映射
|
API格式 -> 时间线数据的映射
|
||||||
"""
|
"""
|
||||||
# 收集所有 key_ids
|
# 收集所有 endpoint_ids
|
||||||
all_key_ids = []
|
all_endpoint_ids = []
|
||||||
for key_ids in format_key_mapping.values():
|
for endpoint_ids in format_endpoint_mapping.values():
|
||||||
all_key_ids.extend(key_ids)
|
all_endpoint_ids.extend(endpoint_ids)
|
||||||
|
|
||||||
if not all_key_ids:
|
if not all_endpoint_ids:
|
||||||
return {
|
return {
|
||||||
api_format: {
|
api_format: {
|
||||||
"timeline": ["unknown"] * 100,
|
"timeline": ["unknown"] * 100,
|
||||||
"time_range_start": None,
|
"time_range_start": None,
|
||||||
"time_range_end": None,
|
"time_range_end": None,
|
||||||
}
|
}
|
||||||
for api_format in format_key_mapping.keys()
|
for api_format in format_endpoint_mapping.keys()
|
||||||
}
|
}
|
||||||
|
|
||||||
# 参数校验(API 层已通过 Query(ge=1) 保证,这里做防御性检查)
|
# 参数校验(API 层已通过 Query(ge=1) 保证,这里做防御性检查)
|
||||||
@@ -293,7 +296,7 @@ class EndpointHealthService:
|
|||||||
|
|
||||||
candidate_stats = (
|
candidate_stats = (
|
||||||
db.query(
|
db.query(
|
||||||
RequestCandidate.key_id,
|
RequestCandidate.endpoint_id,
|
||||||
segment_expr,
|
segment_expr,
|
||||||
func.count(RequestCandidate.id).label("total_count"),
|
func.count(RequestCandidate.id).label("total_count"),
|
||||||
func.sum(case((RequestCandidate.status == "success", 1), else_=0)).label(
|
func.sum(case((RequestCandidate.status == "success", 1), else_=0)).label(
|
||||||
@@ -306,20 +309,20 @@ class EndpointHealthService:
|
|||||||
func.max(RequestCandidate.created_at).label("max_time"),
|
func.max(RequestCandidate.created_at).label("max_time"),
|
||||||
)
|
)
|
||||||
.filter(
|
.filter(
|
||||||
RequestCandidate.key_id.in_(all_key_ids),
|
RequestCandidate.endpoint_id.in_(all_endpoint_ids),
|
||||||
RequestCandidate.created_at >= start_time,
|
RequestCandidate.created_at >= start_time,
|
||||||
RequestCandidate.created_at <= now,
|
RequestCandidate.created_at <= now,
|
||||||
RequestCandidate.status.in_(final_statuses),
|
RequestCandidate.status.in_(final_statuses),
|
||||||
)
|
)
|
||||||
.group_by(RequestCandidate.key_id, segment_expr)
|
.group_by(RequestCandidate.endpoint_id, segment_expr)
|
||||||
.all()
|
.all()
|
||||||
)
|
)
|
||||||
|
|
||||||
# 构建 key_id -> api_format 的反向映射
|
# 构建 endpoint_id -> api_format 的反向映射
|
||||||
key_to_format: dict[str, str] = {}
|
endpoint_to_format: dict[str, str] = {}
|
||||||
for api_format, key_ids in format_key_mapping.items():
|
for api_format, endpoint_ids in format_endpoint_mapping.items():
|
||||||
for key_id in key_ids:
|
for endpoint_id in endpoint_ids:
|
||||||
key_to_format[key_id] = api_format
|
endpoint_to_format[endpoint_id] = api_format
|
||||||
|
|
||||||
# 按 api_format 和 segment 聚合数据
|
# 按 api_format 和 segment 聚合数据
|
||||||
format_segment_data: dict[str, dict[int, dict]] = defaultdict(
|
format_segment_data: dict[str, dict[int, dict]] = defaultdict(
|
||||||
@@ -335,9 +338,9 @@ class EndpointHealthService:
|
|||||||
)
|
)
|
||||||
|
|
||||||
for row in candidate_stats:
|
for row in candidate_stats:
|
||||||
key_id = row.key_id
|
endpoint_id = row.endpoint_id
|
||||||
segment_idx = int(row.segment_idx) if row.segment_idx is not None else 0
|
segment_idx = int(row.segment_idx) if row.segment_idx is not None else 0
|
||||||
api_format = key_to_format.get(key_id)
|
api_format = endpoint_to_format.get(endpoint_id)
|
||||||
|
|
||||||
if api_format and 0 <= segment_idx < segments:
|
if api_format and 0 <= segment_idx < segments:
|
||||||
seg_data = format_segment_data[api_format][segment_idx]
|
seg_data = format_segment_data[api_format][segment_idx]
|
||||||
@@ -355,7 +358,7 @@ class EndpointHealthService:
|
|||||||
# 生成各格式的时间线
|
# 生成各格式的时间线
|
||||||
result: dict[str, dict[str, Any]] = {}
|
result: dict[str, dict[str, Any]] = {}
|
||||||
|
|
||||||
for api_format in format_key_mapping.keys():
|
for api_format in format_endpoint_mapping.keys():
|
||||||
timeline = []
|
timeline = []
|
||||||
earliest_time = None
|
earliest_time = None
|
||||||
latest_time = None
|
latest_time = None
|
||||||
@@ -427,46 +430,10 @@ class EndpointHealthService:
|
|||||||
"time_range_end": None,
|
"time_range_end": None,
|
||||||
}
|
}
|
||||||
|
|
||||||
# 基于 endpoint_ids 反推 provider_ids 与 api_format,再选出支持该格式的 keys
|
# 直接按 endpoint_id 聚合,避免共享 Key 时把不同 API 格式串桶。
|
||||||
endpoint_rows = (
|
format_endpoint_mapping = {"_single": endpoint_ids}
|
||||||
db.query(ProviderEndpoint.provider_id, ProviderEndpoint.api_format)
|
|
||||||
.filter(ProviderEndpoint.id.in_(endpoint_ids))
|
|
||||||
.all()
|
|
||||||
)
|
|
||||||
|
|
||||||
if not endpoint_rows:
|
|
||||||
return {
|
|
||||||
"timeline": ["unknown"] * 100,
|
|
||||||
"time_range_start": None,
|
|
||||||
"time_range_end": None,
|
|
||||||
}
|
|
||||||
|
|
||||||
provider_ids = {str(pid) for pid, _fmt in endpoint_rows}
|
|
||||||
# 同一调用中 endpoint_ids 来自同一 api_format(上层已按格式分组)
|
|
||||||
api_format = (
|
|
||||||
endpoint_rows[0][1].value
|
|
||||||
if hasattr(endpoint_rows[0][1], "value")
|
|
||||||
else str(endpoint_rows[0][1])
|
|
||||||
)
|
|
||||||
|
|
||||||
keys = (
|
|
||||||
db.query(ProviderAPIKey.id, ProviderAPIKey.api_formats)
|
|
||||||
.filter(ProviderAPIKey.provider_id.in_(provider_ids))
|
|
||||||
.all()
|
|
||||||
)
|
|
||||||
key_ids = [str(key_id) for key_id, formats in keys if api_format in (formats or [])]
|
|
||||||
|
|
||||||
if not key_ids:
|
|
||||||
return {
|
|
||||||
"timeline": ["unknown"] * 100,
|
|
||||||
"time_range_start": None,
|
|
||||||
"time_range_end": None,
|
|
||||||
}
|
|
||||||
|
|
||||||
# 使用批量查询
|
|
||||||
format_key_mapping = {"_single": key_ids}
|
|
||||||
result = EndpointHealthService._generate_timeline_batch(
|
result = EndpointHealthService._generate_timeline_batch(
|
||||||
db, format_key_mapping, now, lookback_hours, segments
|
db, format_endpoint_mapping, now, lookback_hours, segments
|
||||||
)
|
)
|
||||||
|
|
||||||
return result.get(
|
return result.get(
|
||||||
|
|||||||
@@ -140,6 +140,13 @@ def _extract_codex_weekly_reset_seconds(metadata: dict[str, Any]) -> float | Non
|
|||||||
if not isinstance(codex, dict):
|
if not isinstance(codex, dict):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
weekly_used_percent = safe_float(codex.get("primary_used_percent"))
|
||||||
|
if weekly_used_percent is not None:
|
||||||
|
clamped_used = max(0.0, min(weekly_used_percent, 100.0))
|
||||||
|
if clamped_used <= 1e-6:
|
||||||
|
# 周额度仍为满额时,不启用周窗口重置倒计时。
|
||||||
|
return None
|
||||||
|
|
||||||
now = time.time()
|
now = time.time()
|
||||||
|
|
||||||
# 优先绝对时间戳,避免 reset_seconds 快照随时间漂移。
|
# 优先绝对时间戳,避免 reset_seconds 快照随时间漂移。
|
||||||
|
|||||||
@@ -69,6 +69,14 @@ def _format_quota_value(value: float) -> str:
|
|||||||
return f"{value:.1f}"
|
return f"{value:.1f}"
|
||||||
|
|
||||||
|
|
||||||
|
def _has_quota_consumption(used_percent_raw: Any) -> bool:
|
||||||
|
used = _to_float(used_percent_raw)
|
||||||
|
if used is None:
|
||||||
|
return False
|
||||||
|
clamped_used = max(0.0, min(used, 100.0))
|
||||||
|
return clamped_used > 1e-6
|
||||||
|
|
||||||
|
|
||||||
def _format_reset_after(seconds_raw: Any) -> str | None:
|
def _format_reset_after(seconds_raw: Any) -> str | None:
|
||||||
seconds = _to_float(seconds_raw)
|
seconds = _to_float(seconds_raw)
|
||||||
if seconds is None:
|
if seconds is None:
|
||||||
@@ -223,7 +231,11 @@ class CodexQuotaReader(PoolQuotaReader):
|
|||||||
primary_used = _to_float(self._data.get("primary_used_percent"))
|
primary_used = _to_float(self._data.get("primary_used_percent"))
|
||||||
if primary_used is not None:
|
if primary_used is not None:
|
||||||
part = f"周剩余 {_format_percent(100.0 - primary_used)}"
|
part = f"周剩余 {_format_percent(100.0 - primary_used)}"
|
||||||
reset_text = _format_reset_after(self._data.get("primary_reset_seconds"))
|
reset_text = (
|
||||||
|
_format_reset_after(self._data.get("primary_reset_seconds"))
|
||||||
|
if _has_quota_consumption(primary_used)
|
||||||
|
else None
|
||||||
|
)
|
||||||
if reset_text:
|
if reset_text:
|
||||||
part = f"{part} ({reset_text})"
|
part = f"{part} ({reset_text})"
|
||||||
parts.append(part)
|
parts.append(part)
|
||||||
@@ -231,7 +243,11 @@ class CodexQuotaReader(PoolQuotaReader):
|
|||||||
secondary_used = _to_float(self._data.get("secondary_used_percent"))
|
secondary_used = _to_float(self._data.get("secondary_used_percent"))
|
||||||
if secondary_used is not None:
|
if secondary_used is not None:
|
||||||
part = f"5H剩余 {_format_percent(100.0 - secondary_used)}"
|
part = f"5H剩余 {_format_percent(100.0 - secondary_used)}"
|
||||||
reset_text = _format_reset_after(self._data.get("secondary_reset_seconds"))
|
reset_text = (
|
||||||
|
_format_reset_after(self._data.get("secondary_reset_seconds"))
|
||||||
|
if _has_quota_consumption(secondary_used)
|
||||||
|
else None
|
||||||
|
)
|
||||||
if reset_text:
|
if reset_text:
|
||||||
part = f"{part} ({reset_text})"
|
part = f"{part} ({reset_text})"
|
||||||
parts.append(part)
|
parts.append(part)
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ Sub2API 余额查询操作
|
|||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import time
|
import time
|
||||||
from typing import Any
|
from typing import Any, cast
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
|
|
||||||
@@ -40,10 +40,13 @@ class Sub2ApiBalanceAction(BalanceAction):
|
|||||||
sub_endpoint = self.config.get("subscription_endpoint", "/api/v1/subscriptions/summary")
|
sub_endpoint = self.config.get("subscription_endpoint", "/api/v1/subscriptions/summary")
|
||||||
|
|
||||||
try:
|
try:
|
||||||
me_resp, sub_resp = await asyncio.gather(
|
me_resp, sub_resp = cast(
|
||||||
client.get(me_endpoint),
|
tuple[httpx.Response | BaseException, httpx.Response | BaseException],
|
||||||
client.get(sub_endpoint),
|
await asyncio.gather(
|
||||||
return_exceptions=True,
|
client.get(me_endpoint),
|
||||||
|
client.get(sub_endpoint),
|
||||||
|
return_exceptions=True,
|
||||||
|
),
|
||||||
)
|
)
|
||||||
response_time_ms = int((time.time() - start_time) * 1000)
|
response_time_ms = int((time.time() - start_time) * 1000)
|
||||||
|
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ YesCode 余额查询操作
|
|||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
from datetime import datetime, timedelta
|
from datetime import datetime, timedelta
|
||||||
from typing import Any
|
from typing import Any, cast
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
|
|
||||||
@@ -37,8 +37,9 @@ async def fetch_yescode_combined_data(
|
|||||||
balance_task = client.get(f"{base_url}/api/v1/user/balance")
|
balance_task = client.get(f"{base_url}/api/v1/user/balance")
|
||||||
profile_task = client.get(f"{base_url}/api/v1/auth/profile")
|
profile_task = client.get(f"{base_url}/api/v1/auth/profile")
|
||||||
|
|
||||||
balance_resp, profile_resp = await asyncio.gather(
|
balance_resp, profile_resp = cast(
|
||||||
balance_task, profile_task, return_exceptions=True
|
tuple[httpx.Response | BaseException, httpx.Response | BaseException],
|
||||||
|
await asyncio.gather(balance_task, profile_task, return_exceptions=True),
|
||||||
)
|
)
|
||||||
|
|
||||||
# 解析 balance 接口
|
# 解析 balance 接口
|
||||||
|
|||||||
@@ -975,21 +975,23 @@ class MaintenanceScheduler:
|
|||||||
try:
|
try:
|
||||||
now = datetime.now(timezone.utc)
|
now = datetime.now(timezone.utc)
|
||||||
|
|
||||||
# 1. 压缩详细日志 (body 字段 -> 压缩字段)
|
|
||||||
detail_cutoff = now - timedelta(days=detail_retention)
|
detail_cutoff = now - timedelta(days=detail_retention)
|
||||||
body_compressed = self._cleanup_body_fields(detail_cutoff, batch_size)
|
|
||||||
|
|
||||||
# 2. 清理压缩字段
|
|
||||||
compressed_cutoff = now - timedelta(days=compressed_retention)
|
compressed_cutoff = now - timedelta(days=compressed_retention)
|
||||||
compressed_cleaned = self._cleanup_compressed_fields(compressed_cutoff, batch_size)
|
|
||||||
|
|
||||||
# 3. 清理请求头
|
|
||||||
header_cutoff = now - timedelta(days=header_retention)
|
header_cutoff = now - timedelta(days=header_retention)
|
||||||
header_cleaned = self._cleanup_header_fields(header_cutoff, batch_size)
|
|
||||||
|
|
||||||
# 4. 删除过期记录
|
|
||||||
log_cutoff = now - timedelta(days=log_retention)
|
log_cutoff = now - timedelta(days=log_retention)
|
||||||
|
|
||||||
|
# 先删最老的整行,再按窗口处理剩余记录,避免同一行在一轮里被重复改写。
|
||||||
records_deleted = self._delete_old_records(log_cutoff, batch_size)
|
records_deleted = self._delete_old_records(log_cutoff, batch_size)
|
||||||
|
header_cleaned = self._cleanup_header_fields(
|
||||||
|
header_cutoff, batch_size, newer_than=log_cutoff
|
||||||
|
)
|
||||||
|
body_cleaned = self._cleanup_stale_body_fields(
|
||||||
|
compressed_cutoff, batch_size, newer_than=log_cutoff
|
||||||
|
)
|
||||||
|
# 仅压缩 7-30 天窗口内的 body;更老记录直接清空 body,不再先压缩再清理。
|
||||||
|
body_compressed = self._cleanup_body_fields(
|
||||||
|
detail_cutoff, batch_size, newer_than=compressed_cutoff
|
||||||
|
)
|
||||||
|
|
||||||
# 5. 清理过期的API Keys
|
# 5. 清理过期的API Keys
|
||||||
keys_db = create_session()
|
keys_db = create_session()
|
||||||
@@ -1005,7 +1007,7 @@ class MaintenanceScheduler:
|
|||||||
|
|
||||||
logger.info(
|
logger.info(
|
||||||
f"清理完成: 压缩 {body_compressed} 条, "
|
f"清理完成: 压缩 {body_compressed} 条, "
|
||||||
f"清理压缩字段 {compressed_cleaned} 条, "
|
f"清理body {body_cleaned} 条, "
|
||||||
f"清理header {header_cleaned} 条, "
|
f"清理header {header_cleaned} 条, "
|
||||||
f"删除记录 {records_deleted} 条, "
|
f"删除记录 {records_deleted} 条, "
|
||||||
f"清理过期Keys {keys_cleaned} 条"
|
f"清理过期Keys {keys_cleaned} 条"
|
||||||
@@ -1016,10 +1018,16 @@ class MaintenanceScheduler:
|
|||||||
loop = asyncio.get_running_loop()
|
loop = asyncio.get_running_loop()
|
||||||
await loop.run_in_executor(None, _do_cleanup)
|
await loop.run_in_executor(None, _do_cleanup)
|
||||||
|
|
||||||
def _cleanup_body_fields(self, cutoff_time: datetime, batch_size: int) -> int:
|
def _cleanup_body_fields(
|
||||||
|
self,
|
||||||
|
cutoff_time: datetime,
|
||||||
|
batch_size: int,
|
||||||
|
*,
|
||||||
|
newer_than: datetime | None = None,
|
||||||
|
) -> int:
|
||||||
"""压缩 request_body 和 response_body 字段到压缩字段
|
"""压缩 request_body 和 response_body 字段到压缩字段
|
||||||
|
|
||||||
逐条处理,确保每条记录都正确更新(同步方法,在线程池中调用)
|
仅处理指定时间窗口内仍保留原始 body 的记录,避免对更老记录重复写放大。
|
||||||
"""
|
"""
|
||||||
from sqlalchemy import null, update
|
from sqlalchemy import null, update
|
||||||
|
|
||||||
@@ -1027,84 +1035,45 @@ class MaintenanceScheduler:
|
|||||||
no_progress_count = 0
|
no_progress_count = 0
|
||||||
memory_safe_batch_size = max(1, min(batch_size, 25))
|
memory_safe_batch_size = max(1, min(batch_size, 25))
|
||||||
|
|
||||||
|
if newer_than is not None and newer_than >= cutoff_time:
|
||||||
|
logger.warning(
|
||||||
|
"压缩 body 字段跳过: 无效时间窗口 newer_than={} cutoff_time={}",
|
||||||
|
newer_than,
|
||||||
|
cutoff_time,
|
||||||
|
)
|
||||||
|
return 0
|
||||||
|
|
||||||
while True:
|
while True:
|
||||||
batch_db = create_session()
|
batch_db = create_session()
|
||||||
try:
|
try:
|
||||||
record_ids = [
|
query = batch_db.query(
|
||||||
row.id
|
Usage.id,
|
||||||
for row in (
|
Usage.request_body,
|
||||||
batch_db.query(Usage.id)
|
Usage.response_body,
|
||||||
.filter(Usage.created_at < cutoff_time)
|
Usage.provider_request_body,
|
||||||
.filter(
|
Usage.client_response_body,
|
||||||
(Usage.request_body.isnot(None))
|
).filter(Usage.created_at < cutoff_time)
|
||||||
| (Usage.response_body.isnot(None))
|
if newer_than is not None:
|
||||||
| (Usage.provider_request_body.isnot(None))
|
query = query.filter(Usage.created_at >= newer_than)
|
||||||
| (Usage.client_response_body.isnot(None))
|
records = (
|
||||||
)
|
query.filter(
|
||||||
.order_by(Usage.created_at.asc(), Usage.id.asc())
|
(Usage.request_body.isnot(None))
|
||||||
.limit(memory_safe_batch_size)
|
| (Usage.response_body.isnot(None))
|
||||||
.all()
|
| (Usage.provider_request_body.isnot(None))
|
||||||
|
| (Usage.client_response_body.isnot(None))
|
||||||
)
|
)
|
||||||
]
|
.order_by(Usage.created_at.asc(), Usage.id.asc())
|
||||||
except Exception as e:
|
.limit(memory_safe_batch_size)
|
||||||
logger.exception("加载待压缩 body 记录 ID 失败: {}", e)
|
.all()
|
||||||
try:
|
)
|
||||||
batch_db.rollback()
|
|
||||||
except Exception:
|
|
||||||
pass
|
|
||||||
break
|
|
||||||
finally:
|
|
||||||
batch_db.close()
|
|
||||||
|
|
||||||
if not record_ids:
|
if not records:
|
||||||
break
|
break
|
||||||
|
|
||||||
batch_success = 0
|
batch_success = 0
|
||||||
batch_progress = False
|
batch_progress = False
|
||||||
|
for record in records:
|
||||||
for record_id in record_ids:
|
result = batch_db.execute(
|
||||||
record_db = create_session()
|
|
||||||
try:
|
|
||||||
record = (
|
|
||||||
record_db.query(
|
|
||||||
Usage.id,
|
|
||||||
Usage.request_body,
|
|
||||||
Usage.response_body,
|
|
||||||
Usage.provider_request_body,
|
|
||||||
Usage.client_response_body,
|
|
||||||
)
|
|
||||||
.filter(Usage.id == record_id)
|
|
||||||
.first()
|
|
||||||
)
|
|
||||||
|
|
||||||
if record is None:
|
|
||||||
continue
|
|
||||||
|
|
||||||
has_body_payload = (
|
|
||||||
record.request_body is not None
|
|
||||||
or record.response_body is not None
|
|
||||||
or record.provider_request_body is not None
|
|
||||||
or record.client_response_body is not None
|
|
||||||
)
|
|
||||||
|
|
||||||
if not has_body_payload:
|
|
||||||
result = record_db.execute(
|
|
||||||
update(Usage)
|
|
||||||
.where(Usage.id == record.id)
|
|
||||||
.values(
|
|
||||||
request_body=null(),
|
|
||||||
response_body=null(),
|
|
||||||
provider_request_body=null(),
|
|
||||||
client_response_body=null(),
|
|
||||||
)
|
|
||||||
.execution_options(synchronize_session=False)
|
|
||||||
)
|
|
||||||
record_db.commit()
|
|
||||||
if result.rowcount > 0:
|
|
||||||
batch_progress = True
|
|
||||||
continue
|
|
||||||
|
|
||||||
result = record_db.execute(
|
|
||||||
update(Usage)
|
update(Usage)
|
||||||
.where(Usage.id == record.id)
|
.where(Usage.id == record.id)
|
||||||
.values(
|
.values(
|
||||||
@@ -1133,20 +1102,20 @@ class MaintenanceScheduler:
|
|||||||
)
|
)
|
||||||
.execution_options(synchronize_session=False)
|
.execution_options(synchronize_session=False)
|
||||||
)
|
)
|
||||||
record_db.commit()
|
|
||||||
|
|
||||||
if result.rowcount > 0:
|
if result.rowcount > 0:
|
||||||
batch_success += 1
|
batch_success += 1
|
||||||
batch_progress = True
|
batch_progress = True
|
||||||
except Exception as e:
|
|
||||||
logger.warning("压缩记录 {} 失败: {}", record_id, e)
|
batch_db.commit()
|
||||||
try:
|
except Exception as e:
|
||||||
record_db.rollback()
|
logger.warning("压缩 body 批次失败: {}", e)
|
||||||
except Exception:
|
try:
|
||||||
pass
|
batch_db.rollback()
|
||||||
continue
|
except Exception:
|
||||||
finally:
|
pass
|
||||||
record_db.close()
|
break
|
||||||
|
finally:
|
||||||
|
batch_db.close()
|
||||||
|
|
||||||
if not batch_progress:
|
if not batch_progress:
|
||||||
no_progress_count += 1
|
no_progress_count += 1
|
||||||
@@ -1166,27 +1135,47 @@ class MaintenanceScheduler:
|
|||||||
|
|
||||||
return total_compressed
|
return total_compressed
|
||||||
|
|
||||||
def _cleanup_compressed_fields(self, cutoff_time: datetime, batch_size: int) -> int:
|
def _cleanup_stale_body_fields(
|
||||||
"""清理压缩字段(删除压缩的body)
|
self,
|
||||||
|
cutoff_time: datetime,
|
||||||
|
batch_size: int,
|
||||||
|
*,
|
||||||
|
newer_than: datetime | None = None,
|
||||||
|
) -> int:
|
||||||
|
"""清理已超过压缩保留期的 body 字段
|
||||||
|
|
||||||
每批使用短生命周期 session(同步方法,在线程池中调用)
|
直接清空 raw/compressed body,避免更老记录先压缩再马上被清掉。
|
||||||
"""
|
"""
|
||||||
from sqlalchemy import null, update
|
from sqlalchemy import null, update
|
||||||
|
|
||||||
total_cleaned = 0
|
total_cleaned = 0
|
||||||
|
|
||||||
|
if newer_than is not None and newer_than >= cutoff_time:
|
||||||
|
logger.warning(
|
||||||
|
"清理 body 字段跳过: 无效时间窗口 newer_than={} cutoff_time={}",
|
||||||
|
newer_than,
|
||||||
|
cutoff_time,
|
||||||
|
)
|
||||||
|
return 0
|
||||||
|
|
||||||
while True:
|
while True:
|
||||||
batch_db = create_session()
|
batch_db = create_session()
|
||||||
try:
|
try:
|
||||||
|
query = batch_db.query(Usage.id).filter(Usage.created_at < cutoff_time)
|
||||||
|
if newer_than is not None:
|
||||||
|
query = query.filter(Usage.created_at >= newer_than)
|
||||||
records_to_clean = (
|
records_to_clean = (
|
||||||
batch_db.query(Usage.id)
|
query.filter(
|
||||||
.filter(Usage.created_at < cutoff_time)
|
(Usage.request_body.isnot(None))
|
||||||
.filter(
|
| (Usage.response_body.isnot(None))
|
||||||
(Usage.request_body_compressed.isnot(None))
|
| (Usage.provider_request_body.isnot(None))
|
||||||
|
| (Usage.client_response_body.isnot(None))
|
||||||
|
| (Usage.request_body_compressed.isnot(None))
|
||||||
| (Usage.response_body_compressed.isnot(None))
|
| (Usage.response_body_compressed.isnot(None))
|
||||||
| (Usage.provider_request_body_compressed.isnot(None))
|
| (Usage.provider_request_body_compressed.isnot(None))
|
||||||
| (Usage.client_response_body_compressed.isnot(None))
|
| (Usage.client_response_body_compressed.isnot(None))
|
||||||
)
|
)
|
||||||
|
.order_by(Usage.created_at.asc(), Usage.id.asc())
|
||||||
.limit(batch_size)
|
.limit(batch_size)
|
||||||
.all()
|
.all()
|
||||||
)
|
)
|
||||||
@@ -1200,6 +1189,10 @@ class MaintenanceScheduler:
|
|||||||
update(Usage)
|
update(Usage)
|
||||||
.where(Usage.id.in_(record_ids))
|
.where(Usage.id.in_(record_ids))
|
||||||
.values(
|
.values(
|
||||||
|
request_body=null(),
|
||||||
|
response_body=null(),
|
||||||
|
provider_request_body=null(),
|
||||||
|
client_response_body=null(),
|
||||||
request_body_compressed=null(),
|
request_body_compressed=null(),
|
||||||
response_body_compressed=null(),
|
response_body_compressed=null(),
|
||||||
provider_request_body_compressed=null(),
|
provider_request_body_compressed=null(),
|
||||||
@@ -1211,14 +1204,14 @@ class MaintenanceScheduler:
|
|||||||
batch_db.commit()
|
batch_db.commit()
|
||||||
|
|
||||||
if rows_updated == 0:
|
if rows_updated == 0:
|
||||||
logger.warning("清理压缩字段: rowcount=0,可能存在问题")
|
logger.warning("清理 body 字段: rowcount=0,可能存在问题")
|
||||||
break
|
break
|
||||||
|
|
||||||
total_cleaned += rows_updated
|
total_cleaned += rows_updated
|
||||||
logger.debug(f"已清理 {rows_updated} 条记录的压缩字段,累计 {total_cleaned} 条")
|
logger.debug(f"已清理 {rows_updated} 条记录的 body 字段,累计 {total_cleaned} 条")
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.exception(f"清理压缩字段失败: {e}")
|
logger.exception(f"清理 body 字段失败: {e}")
|
||||||
try:
|
try:
|
||||||
batch_db.rollback()
|
batch_db.rollback()
|
||||||
except Exception:
|
except Exception:
|
||||||
@@ -1229,7 +1222,13 @@ class MaintenanceScheduler:
|
|||||||
|
|
||||||
return total_cleaned
|
return total_cleaned
|
||||||
|
|
||||||
def _cleanup_header_fields(self, cutoff_time: datetime, batch_size: int) -> int:
|
def _cleanup_header_fields(
|
||||||
|
self,
|
||||||
|
cutoff_time: datetime,
|
||||||
|
batch_size: int,
|
||||||
|
*,
|
||||||
|
newer_than: datetime | None = None,
|
||||||
|
) -> int:
|
||||||
"""清理 request_headers, response_headers 和 provider_request_headers 字段
|
"""清理 request_headers, response_headers 和 provider_request_headers 字段
|
||||||
|
|
||||||
每批使用短生命周期 session(同步方法,在线程池中调用)
|
每批使用短生命周期 session(同步方法,在线程池中调用)
|
||||||
@@ -1238,17 +1237,28 @@ class MaintenanceScheduler:
|
|||||||
|
|
||||||
total_cleaned = 0
|
total_cleaned = 0
|
||||||
|
|
||||||
|
if newer_than is not None and newer_than >= cutoff_time:
|
||||||
|
logger.warning(
|
||||||
|
"清理 header 字段跳过: 无效时间窗口 newer_than={} cutoff_time={}",
|
||||||
|
newer_than,
|
||||||
|
cutoff_time,
|
||||||
|
)
|
||||||
|
return 0
|
||||||
|
|
||||||
while True:
|
while True:
|
||||||
batch_db = create_session()
|
batch_db = create_session()
|
||||||
try:
|
try:
|
||||||
|
query = batch_db.query(Usage.id).filter(Usage.created_at < cutoff_time)
|
||||||
|
if newer_than is not None:
|
||||||
|
query = query.filter(Usage.created_at >= newer_than)
|
||||||
records_to_clean = (
|
records_to_clean = (
|
||||||
batch_db.query(Usage.id)
|
query.filter(
|
||||||
.filter(Usage.created_at < cutoff_time)
|
|
||||||
.filter(
|
|
||||||
(Usage.request_headers.isnot(None))
|
(Usage.request_headers.isnot(None))
|
||||||
| (Usage.response_headers.isnot(None))
|
| (Usage.response_headers.isnot(None))
|
||||||
| (Usage.provider_request_headers.isnot(None))
|
| (Usage.provider_request_headers.isnot(None))
|
||||||
|
| (Usage.client_response_headers.isnot(None))
|
||||||
)
|
)
|
||||||
|
.order_by(Usage.created_at.asc(), Usage.id.asc())
|
||||||
.limit(batch_size)
|
.limit(batch_size)
|
||||||
.all()
|
.all()
|
||||||
)
|
)
|
||||||
@@ -1265,6 +1275,7 @@ class MaintenanceScheduler:
|
|||||||
request_headers=null(),
|
request_headers=null(),
|
||||||
response_headers=null(),
|
response_headers=null(),
|
||||||
provider_request_headers=null(),
|
provider_request_headers=null(),
|
||||||
|
client_response_headers=null(),
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -1300,6 +1311,7 @@ class MaintenanceScheduler:
|
|||||||
records_to_delete = (
|
records_to_delete = (
|
||||||
batch_db.query(Usage.id)
|
batch_db.query(Usage.id)
|
||||||
.filter(Usage.created_at < cutoff_time)
|
.filter(Usage.created_at < cutoff_time)
|
||||||
|
.order_by(Usage.created_at.asc(), Usage.id.asc())
|
||||||
.limit(batch_size)
|
.limit(batch_size)
|
||||||
.all()
|
.all()
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -69,6 +69,7 @@ class SyncTaskExecutionService:
|
|||||||
request_body_state: RequestBodyState | None,
|
request_body_state: RequestBodyState | None,
|
||||||
request_headers: dict[str, Any] | None,
|
request_headers: dict[str, Any] | None,
|
||||||
request_body: dict[str, Any] | None,
|
request_body: dict[str, Any] | None,
|
||||||
|
create_pending_usage: bool = True,
|
||||||
) -> ExecutionResult:
|
) -> ExecutionResult:
|
||||||
"""
|
"""
|
||||||
Unified candidate traversal loop for SYNC.
|
Unified candidate traversal loop for SYNC.
|
||||||
@@ -143,26 +144,32 @@ class SyncTaskExecutionService:
|
|||||||
affinity_key = str(user_api_key.id)
|
affinity_key = str(user_api_key.id)
|
||||||
user_id = str(user_api_key.user_id)
|
user_id = str(user_api_key.user_id)
|
||||||
api_format_norm = normalize_endpoint_signature(api_format)
|
api_format_norm = normalize_endpoint_signature(api_format)
|
||||||
|
user: User | None = None
|
||||||
username_snapshot = None
|
username_snapshot = None
|
||||||
api_key_name_snapshot = getattr(user_api_key, "name", None)
|
api_key_name_snapshot = getattr(user_api_key, "name", None)
|
||||||
|
|
||||||
# Keep pending usage creation behavior consistent with previous behavior
|
|
||||||
try:
|
try:
|
||||||
user = self.db.query(User).filter(User.id == user_api_key.user_id).first()
|
user = self.db.query(User).filter(User.id == user_api_key.user_id).first()
|
||||||
username_snapshot = getattr(user, "username", None) if user else None
|
username_snapshot = getattr(user, "username", None) if user else None
|
||||||
UsageService.create_pending_usage(
|
|
||||||
db=self.db,
|
|
||||||
request_id=request_id,
|
|
||||||
user=user,
|
|
||||||
api_key=user_api_key,
|
|
||||||
model=model_name,
|
|
||||||
is_stream=is_stream,
|
|
||||||
api_format=api_format_norm,
|
|
||||||
request_headers=request_headers,
|
|
||||||
request_body=request_body,
|
|
||||||
)
|
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
logger.warning("创建 pending 使用记录失败: {}", str(exc))
|
# username 仅用于审计快照,不应阻塞主请求链路。
|
||||||
|
logger.warning("查询用户快照失败: {}", str(exc))
|
||||||
|
|
||||||
|
# 默认由 TaskService 创建 pending 使用记录;已预创建的调用方可关闭。
|
||||||
|
if create_pending_usage:
|
||||||
|
try:
|
||||||
|
UsageService.create_pending_usage(
|
||||||
|
db=self.db,
|
||||||
|
request_id=request_id,
|
||||||
|
user=user,
|
||||||
|
api_key=user_api_key,
|
||||||
|
model=model_name,
|
||||||
|
is_stream=is_stream,
|
||||||
|
api_format=api_format_norm,
|
||||||
|
request_headers=request_headers,
|
||||||
|
request_body=request_body,
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning("创建 pending 使用记录失败: {}", str(exc))
|
||||||
|
|
||||||
all_candidates, global_model_id = await candidate_resolver.fetch_candidates(
|
all_candidates, global_model_id = await candidate_resolver.fetch_candidates(
|
||||||
api_format=api_format_norm,
|
api_format=api_format_norm,
|
||||||
|
|||||||
@@ -126,6 +126,7 @@ class TaskService:
|
|||||||
supported_auth_types: set[str] | None = None,
|
supported_auth_types: set[str] | None = None,
|
||||||
allow_format_conversion: bool = False,
|
allow_format_conversion: bool = False,
|
||||||
max_candidates: int | None = None,
|
max_candidates: int | None = None,
|
||||||
|
create_pending_usage: bool = True,
|
||||||
) -> ExecutionResult:
|
) -> ExecutionResult:
|
||||||
"""兼容入口:默认绑定到 TaskService 内部执行路由。"""
|
"""兼容入口:默认绑定到 TaskService 内部执行路由。"""
|
||||||
return await self._execute_facade_ops.execute(
|
return await self._execute_facade_ops.execute(
|
||||||
@@ -146,6 +147,7 @@ class TaskService:
|
|||||||
supported_auth_types=supported_auth_types,
|
supported_auth_types=supported_auth_types,
|
||||||
allow_format_conversion=allow_format_conversion,
|
allow_format_conversion=allow_format_conversion,
|
||||||
max_candidates=max_candidates,
|
max_candidates=max_candidates,
|
||||||
|
create_pending_usage=create_pending_usage,
|
||||||
)
|
)
|
||||||
|
|
||||||
async def _execute_internal(
|
async def _execute_internal(
|
||||||
@@ -168,6 +170,7 @@ class TaskService:
|
|||||||
supported_auth_types: set[str] | None = None,
|
supported_auth_types: set[str] | None = None,
|
||||||
allow_format_conversion: bool = False,
|
allow_format_conversion: bool = False,
|
||||||
max_candidates: int | None = None,
|
max_candidates: int | None = None,
|
||||||
|
create_pending_usage: bool = True,
|
||||||
) -> ExecutionResult:
|
) -> ExecutionResult:
|
||||||
if task_mode == TaskMode.ASYNC:
|
if task_mode == TaskMode.ASYNC:
|
||||||
if extract_external_task_id is None:
|
if extract_external_task_id is None:
|
||||||
@@ -250,6 +253,7 @@ class TaskService:
|
|||||||
request_body_state=request_body_state,
|
request_body_state=request_body_state,
|
||||||
request_headers=request_headers,
|
request_headers=request_headers,
|
||||||
request_body=request_body,
|
request_body=request_body,
|
||||||
|
create_pending_usage=create_pending_usage,
|
||||||
)
|
)
|
||||||
|
|
||||||
async def execute_sync_candidates(
|
async def execute_sync_candidates(
|
||||||
|
|||||||
@@ -39,6 +39,17 @@ async def _yield_once_then_cancel(ctx: StreamContext) -> AsyncGenerator[bytes, N
|
|||||||
raise asyncio.CancelledError()
|
raise asyncio.CancelledError()
|
||||||
|
|
||||||
|
|
||||||
|
async def _yield_once_then_hang(ctx: StreamContext) -> AsyncGenerator[bytes, None]:
|
||||||
|
ctx.append_text("partial output")
|
||||||
|
yield b"data: chunk\n\n"
|
||||||
|
await asyncio.sleep(3600)
|
||||||
|
|
||||||
|
|
||||||
|
async def _yield_after_delay_then_complete() -> AsyncGenerator[bytes, None]:
|
||||||
|
await asyncio.sleep(0.25)
|
||||||
|
yield b"data: first\n\n"
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_create_monitored_stream_marks_client_disconnected_when_confirmed() -> None:
|
async def test_create_monitored_stream_marks_client_disconnected_when_confirmed() -> None:
|
||||||
monitor = _DummyMonitor()
|
monitor = _DummyMonitor()
|
||||||
@@ -113,3 +124,38 @@ async def test_create_monitored_stream_estimates_output_tokens_before_unknown_ca
|
|||||||
assert ctx.error_message == "cancelled_unknown"
|
assert ctx.error_message == "cancelled_unknown"
|
||||||
assert ctx.output_tokens == expected_output_tokens
|
assert ctx.output_tokens == expected_output_tokens
|
||||||
assert f"output_tokens={expected_output_tokens}" in (ctx.upstream_response or "")
|
assert f"output_tokens={expected_output_tokens}" in (ctx.upstream_response or "")
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_create_monitored_stream_marks_idle_timeout_before_worker_timeout() -> None:
|
||||||
|
monitor = _DummyMonitor()
|
||||||
|
monitor.CANCEL_DISCONNECT_RETRY_DELAYS_SECONDS = ()
|
||||||
|
monitor.STREAM_IDLE_TIMEOUT_SECONDS = 1.0
|
||||||
|
ctx = StreamContext(model="test-model", api_format="openai:cli", request_id="req-idle-timeout")
|
||||||
|
|
||||||
|
monitored = monitor._create_monitored_stream(ctx, _yield_once_then_hang(ctx), None)
|
||||||
|
|
||||||
|
with pytest.raises(asyncio.CancelledError):
|
||||||
|
async for _ in monitored:
|
||||||
|
pass
|
||||||
|
|
||||||
|
expected_output_tokens = max(1, len("partial output") // 4)
|
||||||
|
assert ctx.status_code == 504
|
||||||
|
assert ctx.error_message == "stream_idle_timeout"
|
||||||
|
assert ctx.output_tokens == expected_output_tokens
|
||||||
|
assert "cancel_origin=stream_idle_timeout" in (ctx.upstream_response or "")
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_create_monitored_stream_does_not_idle_timeout_before_first_chunk() -> None:
|
||||||
|
monitor = _DummyMonitor()
|
||||||
|
monitor.CANCEL_DISCONNECT_RETRY_DELAYS_SECONDS = ()
|
||||||
|
monitor.STREAM_IDLE_TIMEOUT_SECONDS = 0.05
|
||||||
|
ctx = StreamContext(model="test-model", api_format="openai:cli", request_id="req-first-chunk")
|
||||||
|
|
||||||
|
monitored = monitor._create_monitored_stream(ctx, _yield_after_delay_then_complete(), None)
|
||||||
|
chunks = [chunk async for chunk in monitored]
|
||||||
|
|
||||||
|
assert chunks == [b"data: first\n\n"]
|
||||||
|
assert ctx.status_code == 200
|
||||||
|
assert ctx.error_message is None
|
||||||
|
|||||||
@@ -23,7 +23,7 @@ class _DummySyncHandler(CliSyncMixin):
|
|||||||
) -> str:
|
) -> str:
|
||||||
return str(request_body.get("model") or "unknown")
|
return str(request_body.get("model") or "unknown")
|
||||||
|
|
||||||
def _create_pending_usage(self, **kwargs: object) -> None:
|
def _create_pending_usage(self, **kwargs: object) -> bool:
|
||||||
self.pending_calls.append(kwargs)
|
self.pending_calls.append(kwargs)
|
||||||
raise _StopExecution()
|
raise _StopExecution()
|
||||||
|
|
||||||
@@ -41,7 +41,7 @@ class _DummyStreamHandler(CliStreamMixin):
|
|||||||
) -> str:
|
) -> str:
|
||||||
return str(request_body.get("model") or "unknown")
|
return str(request_body.get("model") or "unknown")
|
||||||
|
|
||||||
def _create_pending_usage(self, **kwargs: object) -> None:
|
def _create_pending_usage(self, **kwargs: object) -> bool:
|
||||||
self.pending_calls.append(kwargs)
|
self.pending_calls.append(kwargs)
|
||||||
raise _StopExecution()
|
raise _StopExecution()
|
||||||
|
|
||||||
|
|||||||
@@ -1,3 +1,9 @@
|
|||||||
|
import os
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
os.environ.setdefault("JWT_SECRET_KEY", "test-secret-key")
|
||||||
|
|
||||||
from src.api.handlers.base import stream_context
|
from src.api.handlers.base import stream_context
|
||||||
from src.api.handlers.base.stream_context import StreamContext
|
from src.api.handlers.base.stream_context import StreamContext
|
||||||
|
|
||||||
@@ -50,7 +56,36 @@ def test_reset_for_retry_clears_state() -> None:
|
|||||||
assert ctx.error_message is None
|
assert ctx.error_message is None
|
||||||
|
|
||||||
|
|
||||||
def test_record_first_byte_time(monkeypatch) -> None:
|
def test_release_recorded_chunks_clears_both_chunk_lists() -> None:
|
||||||
|
ctx = StreamContext(model="test-model", api_format="openai:chat")
|
||||||
|
ctx.parsed_chunks.append({"type": "client"})
|
||||||
|
ctx.provider_parsed_chunks.append({"type": "provider"})
|
||||||
|
|
||||||
|
ctx.release_recorded_chunks()
|
||||||
|
|
||||||
|
assert ctx.parsed_chunks == []
|
||||||
|
assert ctx.provider_parsed_chunks == []
|
||||||
|
|
||||||
|
|
||||||
|
def test_managed_recorded_bodies_builds_then_releases_chunks() -> None:
|
||||||
|
ctx = StreamContext(model="test-model", api_format="openai:chat")
|
||||||
|
ctx.parsed_chunks.append({"type": "client"})
|
||||||
|
ctx.provider_parsed_chunks.append({"type": "provider"})
|
||||||
|
ctx.data_count = 1
|
||||||
|
|
||||||
|
with ctx.managed_recorded_bodies(123) as recorded_bodies:
|
||||||
|
assert recorded_bodies.response_body is not None
|
||||||
|
assert recorded_bodies.response_body["chunks"] == [{"type": "provider"}]
|
||||||
|
assert recorded_bodies.client_response_body is not None
|
||||||
|
assert recorded_bodies.client_response_body["chunks"] == [{"type": "client"}]
|
||||||
|
|
||||||
|
assert ctx.parsed_chunks == []
|
||||||
|
assert ctx.provider_parsed_chunks == []
|
||||||
|
assert recorded_bodies.response_body is None
|
||||||
|
assert recorded_bodies.client_response_body is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_record_first_byte_time(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||||
"""测试记录首字时间"""
|
"""测试记录首字时间"""
|
||||||
ctx = StreamContext(model="claude-3", api_format="claude_messages")
|
ctx = StreamContext(model="claude-3", api_format="claude_messages")
|
||||||
start_time = 100.0
|
start_time = 100.0
|
||||||
@@ -63,7 +98,7 @@ def test_record_first_byte_time(monkeypatch) -> None:
|
|||||||
assert ctx.first_byte_time_ms == 12
|
assert ctx.first_byte_time_ms == 12
|
||||||
|
|
||||||
|
|
||||||
def test_record_first_byte_time_idempotent(monkeypatch) -> None:
|
def test_record_first_byte_time_idempotent(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||||
"""测试首字时间只记录一次"""
|
"""测试首字时间只记录一次"""
|
||||||
ctx = StreamContext(model="claude-3", api_format="claude_messages")
|
ctx = StreamContext(model="claude-3", api_format="claude_messages")
|
||||||
start_time = 100.0
|
start_time = 100.0
|
||||||
@@ -82,7 +117,7 @@ def test_record_first_byte_time_idempotent(monkeypatch) -> None:
|
|||||||
assert first_value == second_value
|
assert first_value == second_value
|
||||||
|
|
||||||
|
|
||||||
def test_reset_for_retry_clears_first_byte_time(monkeypatch) -> None:
|
def test_reset_for_retry_clears_first_byte_time(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||||
"""测试重试时清除首字时间"""
|
"""测试重试时清除首字时间"""
|
||||||
ctx = StreamContext(model="claude-3", api_format="claude_messages")
|
ctx = StreamContext(model="claude-3", api_format="claude_messages")
|
||||||
start_time = 100.0
|
start_time = 100.0
|
||||||
|
|||||||
@@ -1,11 +1,15 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import os
|
||||||
import time
|
import time
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
|
from typing import Any
|
||||||
from unittest.mock import AsyncMock
|
from unittest.mock import AsyncMock
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
|
os.environ.setdefault("JWT_SECRET_KEY", "test-secret-key")
|
||||||
|
|
||||||
from src.api.handlers.base import stream_telemetry as stream_telemetry_module
|
from src.api.handlers.base import stream_telemetry as stream_telemetry_module
|
||||||
from src.api.handlers.base.stream_context import StreamContext
|
from src.api.handlers.base.stream_context import StreamContext
|
||||||
from src.api.handlers.base.stream_telemetry import StreamTelemetryRecorder
|
from src.api.handlers.base.stream_telemetry import StreamTelemetryRecorder
|
||||||
@@ -64,3 +68,69 @@ async def test_record_stream_stats_estimates_tokens_for_failed_partial_stream(
|
|||||||
assert ctx.input_tokens > 0
|
assert ctx.input_tokens > 0
|
||||||
assert ctx.output_tokens == max(1, len("partial output") // 4)
|
assert ctx.output_tokens == max(1, len("partial output") // 4)
|
||||||
recorder._dispatch_record.assert_awaited_once() # type: ignore[attr-defined]
|
recorder._dispatch_record.assert_awaited_once() # type: ignore[attr-defined]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_record_stream_stats_releases_parsed_chunks_after_dispatch(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
recorder = StreamTelemetryRecorder(
|
||||||
|
request_id="req-release",
|
||||||
|
user_id="1",
|
||||||
|
api_key_id="2",
|
||||||
|
client_ip="127.0.0.1",
|
||||||
|
format_id="openai:chat",
|
||||||
|
)
|
||||||
|
recorder._get_telemetry_writer = AsyncMock( # type: ignore[method-assign]
|
||||||
|
return_value=SimpleNamespace(include_bodies=True)
|
||||||
|
)
|
||||||
|
recorder._update_candidate_status = AsyncMock() # type: ignore[method-assign]
|
||||||
|
|
||||||
|
dispatch_payloads: list[dict[str, Any]] = []
|
||||||
|
|
||||||
|
async def _capture_dispatch(*args: Any, **kwargs: Any) -> None:
|
||||||
|
dispatch_payloads.append(
|
||||||
|
{
|
||||||
|
"response_body": args[5],
|
||||||
|
"client_response_body": kwargs.get("client_response_body"),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
recorder._dispatch_record = _capture_dispatch # type: ignore[method-assign]
|
||||||
|
|
||||||
|
ctx = StreamContext(
|
||||||
|
model="test-model",
|
||||||
|
api_format="openai:chat",
|
||||||
|
request_id="req-release",
|
||||||
|
user_id=1,
|
||||||
|
api_key_id=2,
|
||||||
|
)
|
||||||
|
ctx.provider_name = "test-provider"
|
||||||
|
ctx.parsed_chunks.extend([{"type": "chunk-1"}, {"type": "chunk-2"}])
|
||||||
|
ctx.data_count = 2
|
||||||
|
ctx.chunk_count = 2
|
||||||
|
|
||||||
|
monkeypatch.setattr(stream_telemetry_module, "get_db", lambda: iter([_DummyDb()]))
|
||||||
|
monkeypatch.setattr(
|
||||||
|
stream_telemetry_module.SystemConfigService,
|
||||||
|
"should_log_body",
|
||||||
|
lambda _db: True,
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(stream_telemetry_module.config, "stream_stats_delay", 0)
|
||||||
|
|
||||||
|
await recorder.record_stream_stats(
|
||||||
|
ctx,
|
||||||
|
original_headers={},
|
||||||
|
original_request_body={"input": [{"content": "hello world"}]},
|
||||||
|
start_time=time.time(),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert len(dispatch_payloads) == 1
|
||||||
|
response_body = dispatch_payloads[0]["response_body"]
|
||||||
|
assert response_body["chunks"] == [{"type": "chunk-1"}, {"type": "chunk-2"}]
|
||||||
|
assert response_body["metadata"]["stream"] is True
|
||||||
|
assert response_body["metadata"]["total_chunks"] == 2
|
||||||
|
assert response_body["metadata"]["data_count"] == 2
|
||||||
|
assert dispatch_payloads[0]["client_response_body"] is None
|
||||||
|
assert ctx.parsed_chunks == []
|
||||||
|
assert ctx.provider_parsed_chunks == []
|
||||||
|
|||||||
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"
|
assert credits_reader.display_summary() == "积分 12.35"
|
||||||
|
|
||||||
|
|
||||||
|
def test_codex_reader_hides_reset_countdown_when_remaining_is_full() -> None:
|
||||||
|
reader = get_quota_reader(
|
||||||
|
"codex",
|
||||||
|
{
|
||||||
|
"codex": {
|
||||||
|
"primary_used_percent": 0.0,
|
||||||
|
"primary_reset_seconds": 266400,
|
||||||
|
"secondary_used_percent": 0.0,
|
||||||
|
"secondary_reset_seconds": 3600,
|
||||||
|
}
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
assert reader.display_summary() == "周剩余 100.0% | 5H剩余 100.0%"
|
||||||
|
|
||||||
|
|
||||||
def test_antigravity_reader_keeps_used_percent_fallbacks() -> None:
|
def test_antigravity_reader_keeps_used_percent_fallbacks() -> None:
|
||||||
reader = get_quota_reader(
|
reader = get_quota_reader(
|
||||||
"antigravity",
|
"antigravity",
|
||||||
@@ -88,6 +104,20 @@ def test_extract_reset_seconds_uses_codex_weekly_reset_for_codex_provider() -> N
|
|||||||
assert extract_reset_seconds(key_obj) == pytest.approx(1800.0)
|
assert extract_reset_seconds(key_obj) == pytest.approx(1800.0)
|
||||||
|
|
||||||
|
|
||||||
|
def test_extract_reset_seconds_codex_weekly_full_quota_returns_none() -> None:
|
||||||
|
key_obj = SimpleNamespace(
|
||||||
|
provider_type="codex",
|
||||||
|
upstream_metadata={
|
||||||
|
"codex": {
|
||||||
|
"primary_used_percent": 0.0,
|
||||||
|
"primary_reset_seconds": 1800.0,
|
||||||
|
}
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
assert extract_reset_seconds(key_obj) is None
|
||||||
|
|
||||||
|
|
||||||
def test_extract_reset_seconds_prefers_codex_reset_at(monkeypatch: pytest.MonkeyPatch) -> None:
|
def test_extract_reset_seconds_prefers_codex_reset_at(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||||
monkeypatch.setattr("src.services.provider.pool.dimensions._helpers.time.time", lambda: 1000.0)
|
monkeypatch.setattr("src.services.provider.pool.dimensions._helpers.time.time", lambda: 1000.0)
|
||||||
key_obj = SimpleNamespace(
|
key_obj = SimpleNamespace(
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import inspect
|
import inspect
|
||||||
|
from datetime import datetime, timedelta, timezone
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
from unittest.mock import MagicMock
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
@@ -124,44 +125,29 @@ async def test_candidate_cleanup_uses_dedicated_retention_and_batch_settings(
|
|||||||
assert batch_two.closed is True
|
assert batch_two.closed is True
|
||||||
|
|
||||||
|
|
||||||
def test_cleanup_body_fields_loads_ids_then_processes_records_individually(
|
def test_cleanup_body_fields_batches_records_with_single_commit(
|
||||||
monkeypatch: pytest.MonkeyPatch,
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
) -> None:
|
) -> None:
|
||||||
scheduler = MaintenanceScheduler()
|
scheduler = MaintenanceScheduler()
|
||||||
|
|
||||||
class _IdBatchSession:
|
class _BatchSession:
|
||||||
def __init__(self, ids: list[str]) -> None:
|
def __init__(self, records: list[SimpleNamespace]) -> None:
|
||||||
self.ids = ids
|
self.records = records
|
||||||
self.closed = False
|
|
||||||
self.query_obj = MagicMock()
|
|
||||||
filtered = self.query_obj.filter.return_value
|
|
||||||
filtered.filter.return_value.order_by.return_value.limit.return_value.all.return_value = [
|
|
||||||
SimpleNamespace(id=value) for value in ids
|
|
||||||
]
|
|
||||||
|
|
||||||
def query(self, *args): # type: ignore[no-untyped-def]
|
|
||||||
self.query_args = args
|
|
||||||
return self.query_obj
|
|
||||||
|
|
||||||
def rollback(self) -> None:
|
|
||||||
raise AssertionError("rollback should not be called")
|
|
||||||
|
|
||||||
def close(self) -> None:
|
|
||||||
self.closed = True
|
|
||||||
|
|
||||||
class _RecordSession:
|
|
||||||
def __init__(self, record: SimpleNamespace) -> None:
|
|
||||||
self.record = record
|
|
||||||
self.closed = False
|
self.closed = False
|
||||||
self.committed = False
|
self.committed = False
|
||||||
|
self.executed = 0
|
||||||
self.query_obj = MagicMock()
|
self.query_obj = MagicMock()
|
||||||
self.query_obj.filter.return_value.first.return_value = record
|
filtered = self.query_obj.filter.return_value
|
||||||
|
filtered.filter.return_value.order_by.return_value.limit.return_value.all.return_value = (
|
||||||
|
records
|
||||||
|
)
|
||||||
|
|
||||||
def query(self, *args): # type: ignore[no-untyped-def]
|
def query(self, *args): # type: ignore[no-untyped-def]
|
||||||
self.query_args = args
|
self.query_args = args
|
||||||
return self.query_obj
|
return self.query_obj
|
||||||
|
|
||||||
def execute(self, _statement): # type: ignore[no-untyped-def]
|
def execute(self, _statement): # type: ignore[no-untyped-def]
|
||||||
|
self.executed += 1
|
||||||
return SimpleNamespace(rowcount=1)
|
return SimpleNamespace(rowcount=1)
|
||||||
|
|
||||||
def commit(self) -> None:
|
def commit(self) -> None:
|
||||||
@@ -173,27 +159,26 @@ def test_cleanup_body_fields_loads_ids_then_processes_records_individually(
|
|||||||
def close(self) -> None:
|
def close(self) -> None:
|
||||||
self.closed = True
|
self.closed = True
|
||||||
|
|
||||||
batch_one = _IdBatchSession(["usage-1", "usage-2"])
|
batch_one = _BatchSession(
|
||||||
record_one = _RecordSession(
|
[
|
||||||
SimpleNamespace(
|
SimpleNamespace(
|
||||||
id="usage-1",
|
id="usage-1",
|
||||||
request_body={"hello": "world"},
|
request_body={"hello": "world"},
|
||||||
response_body=None,
|
response_body=None,
|
||||||
provider_request_body=None,
|
provider_request_body=None,
|
||||||
client_response_body=None,
|
client_response_body=None,
|
||||||
)
|
),
|
||||||
|
SimpleNamespace(
|
||||||
|
id="usage-2",
|
||||||
|
request_body=None,
|
||||||
|
response_body={"ok": True},
|
||||||
|
provider_request_body=None,
|
||||||
|
client_response_body=None,
|
||||||
|
),
|
||||||
|
]
|
||||||
)
|
)
|
||||||
record_two = _RecordSession(
|
batch_two = _BatchSession([])
|
||||||
SimpleNamespace(
|
sessions = iter([batch_one, batch_two])
|
||||||
id="usage-2",
|
|
||||||
request_body=None,
|
|
||||||
response_body={"ok": True},
|
|
||||||
provider_request_body=None,
|
|
||||||
client_response_body=None,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
batch_two = _IdBatchSession([])
|
|
||||||
sessions = iter([batch_one, record_one, record_two, batch_two])
|
|
||||||
|
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr(
|
||||||
maintenance_scheduler_module,
|
maintenance_scheduler_module,
|
||||||
@@ -212,12 +197,173 @@ def test_cleanup_body_fields_loads_ids_then_processes_records_individually(
|
|||||||
)
|
)
|
||||||
|
|
||||||
assert compressed == 2
|
assert compressed == 2
|
||||||
assert batch_one.query_args == (maintenance_scheduler_module.Usage.id,)
|
assert len(batch_one.query_args) == 5
|
||||||
assert len(record_one.query_args) == 5
|
assert batch_one.executed == 2
|
||||||
assert len(record_two.query_args) == 5
|
assert batch_one.committed is True
|
||||||
assert batch_one.closed is True
|
assert batch_one.closed is True
|
||||||
assert batch_two.closed is True
|
assert batch_two.closed is True
|
||||||
assert record_one.committed is True
|
|
||||||
assert record_two.committed is True
|
|
||||||
assert record_one.closed is True
|
def test_cleanup_header_fields_clears_client_response_headers(
|
||||||
assert record_two.closed is True
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
scheduler = MaintenanceScheduler()
|
||||||
|
|
||||||
|
class _BatchSession:
|
||||||
|
def __init__(self, ids: list[str]) -> None:
|
||||||
|
self.ids = ids
|
||||||
|
self.closed = False
|
||||||
|
self.committed = False
|
||||||
|
self.query_obj = MagicMock()
|
||||||
|
self.filtered_by_time = MagicMock()
|
||||||
|
self.filtered_by_headers = MagicMock()
|
||||||
|
self.query_obj.filter.return_value = self.filtered_by_time
|
||||||
|
self.filtered_by_time.filter.return_value = self.filtered_by_headers
|
||||||
|
self.filtered_by_headers.order_by.return_value.limit.return_value.all.return_value = [
|
||||||
|
SimpleNamespace(id=value) for value in ids
|
||||||
|
]
|
||||||
|
self.executed_statements: list[str] = []
|
||||||
|
|
||||||
|
def query(self, *args): # type: ignore[no-untyped-def]
|
||||||
|
self.query_args = args
|
||||||
|
return self.query_obj
|
||||||
|
|
||||||
|
def execute(self, statement): # type: ignore[no-untyped-def]
|
||||||
|
self.executed_statements.append(str(statement))
|
||||||
|
return SimpleNamespace(rowcount=len(self.ids))
|
||||||
|
|
||||||
|
def commit(self) -> None:
|
||||||
|
self.committed = True
|
||||||
|
|
||||||
|
def rollback(self) -> None:
|
||||||
|
raise AssertionError("rollback should not be called")
|
||||||
|
|
||||||
|
def close(self) -> None:
|
||||||
|
self.closed = True
|
||||||
|
|
||||||
|
batch_one = _BatchSession(["usage-1"])
|
||||||
|
batch_two = _BatchSession([])
|
||||||
|
sessions = iter([batch_one, batch_two])
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
maintenance_scheduler_module,
|
||||||
|
"create_session",
|
||||||
|
lambda: next(sessions),
|
||||||
|
)
|
||||||
|
|
||||||
|
cleaned = scheduler._cleanup_header_fields(
|
||||||
|
cutoff_time=SimpleNamespace(), # type: ignore[arg-type]
|
||||||
|
batch_size=1000,
|
||||||
|
)
|
||||||
|
|
||||||
|
header_filter = str(batch_one.filtered_by_time.filter.call_args.args[0])
|
||||||
|
|
||||||
|
assert cleaned == 1
|
||||||
|
assert batch_one.query_args == (maintenance_scheduler_module.Usage.id,)
|
||||||
|
assert "client_response_headers" in header_filter
|
||||||
|
assert "client_response_headers" in batch_one.executed_statements[0]
|
||||||
|
assert batch_one.committed is True
|
||||||
|
assert batch_one.closed is True
|
||||||
|
assert batch_two.closed is True
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_perform_cleanup_deletes_first_and_uses_non_overlapping_windows(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
scheduler = MaintenanceScheduler()
|
||||||
|
fixed_now = datetime(2026, 3, 18, 3, 0, 0, tzinfo=timezone.utc)
|
||||||
|
|
||||||
|
class _FakeDateTime(datetime):
|
||||||
|
@classmethod
|
||||||
|
def now(cls, tz=None): # type: ignore[override]
|
||||||
|
if tz is None:
|
||||||
|
return fixed_now.replace(tzinfo=None)
|
||||||
|
return fixed_now.astimezone(tz)
|
||||||
|
|
||||||
|
class _FakeLoop:
|
||||||
|
async def run_in_executor(self, _executor, func): # type: ignore[no-untyped-def]
|
||||||
|
return func()
|
||||||
|
|
||||||
|
class _ConfigSession:
|
||||||
|
def close(self) -> None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
calls: list[tuple[str, datetime, int, datetime | None]] = []
|
||||||
|
|
||||||
|
def _record(name: str, count: int):
|
||||||
|
def _inner(
|
||||||
|
cutoff_time: datetime,
|
||||||
|
batch_size: int,
|
||||||
|
*,
|
||||||
|
newer_than: datetime | None = None,
|
||||||
|
) -> int:
|
||||||
|
calls.append((name, cutoff_time, batch_size, newer_than))
|
||||||
|
return count
|
||||||
|
|
||||||
|
return _inner
|
||||||
|
|
||||||
|
config_values = {
|
||||||
|
"enable_auto_cleanup": True,
|
||||||
|
"detail_log_retention_days": 7,
|
||||||
|
"compressed_log_retention_days": 30,
|
||||||
|
"header_retention_days": 90,
|
||||||
|
"log_retention_days": 365,
|
||||||
|
"cleanup_batch_size": 123,
|
||||||
|
"auto_delete_expired_keys": False,
|
||||||
|
}
|
||||||
|
|
||||||
|
monkeypatch.setattr(maintenance_scheduler_module, "datetime", _FakeDateTime)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
maintenance_scheduler_module.asyncio, "get_running_loop", lambda: _FakeLoop()
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
maintenance_scheduler_module,
|
||||||
|
"create_session",
|
||||||
|
lambda: _ConfigSession(),
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
maintenance_scheduler_module.SystemConfigService,
|
||||||
|
"get_config",
|
||||||
|
lambda _db, key, default=None: config_values.get(key, default),
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
scheduler,
|
||||||
|
"_delete_old_records",
|
||||||
|
lambda cutoff_time, batch_size: calls.append(("delete", cutoff_time, batch_size, None))
|
||||||
|
or 5,
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
scheduler,
|
||||||
|
"_cleanup_header_fields",
|
||||||
|
_record("header", 4),
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
scheduler,
|
||||||
|
"_cleanup_stale_body_fields",
|
||||||
|
_record("body", 3),
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
scheduler,
|
||||||
|
"_cleanup_body_fields",
|
||||||
|
_record("compress", 2),
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
maintenance_scheduler_module.ApiKeyService,
|
||||||
|
"cleanup_expired_keys",
|
||||||
|
lambda _db, auto_delete=False: 0,
|
||||||
|
)
|
||||||
|
|
||||||
|
await scheduler._perform_cleanup()
|
||||||
|
|
||||||
|
detail_cutoff = fixed_now - timedelta(days=7)
|
||||||
|
compressed_cutoff = fixed_now - timedelta(days=30)
|
||||||
|
header_cutoff = fixed_now - timedelta(days=90)
|
||||||
|
log_cutoff = fixed_now - timedelta(days=365)
|
||||||
|
|
||||||
|
assert calls == [
|
||||||
|
("delete", log_cutoff, 123, None),
|
||||||
|
("header", header_cutoff, 123, log_cutoff),
|
||||||
|
("body", compressed_cutoff, 123, log_cutoff),
|
||||||
|
("compress", detail_cutoff, 123, compressed_cutoff),
|
||||||
|
]
|
||||||
|
|||||||
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