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

Close #233

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

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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