feat: 性能监控基础设施、解密缓存及计费简化

- 新增 PerfRecorder 性能记录工具,支持采样率与慢请求日志
- 在请求管道中埋点:auth、body_read、json_parse、context_build、authorize、handle
- 流处理器增加 parse/conversion 耗时追踪与 perf_metrics 落库
- 解密服务添加 LRU 缓存,降低高频解密 CPU 开销
- 格式转换分层开关设计:全局 OFF 时回退到端点配置,而非一刀切拒绝
- 移除 shadow billing 模块,统一使用新计费引擎
- 新增 Codex 网关请求适配器(store=false、role 映射、include 补齐)
- endpoint 创建接口支持 body_rules 参数
This commit is contained in:
fawney19
2026-02-05 14:22:11 +08:00
parent e72e5370c4
commit ed2ff5c1d7
25 changed files with 836 additions and 963 deletions

View File

@@ -320,6 +320,7 @@ class AdminCreateProviderEndpointAdapter(AdminApiAdapter):
base_url=self.endpoint_data.base_url,
custom_path=self.endpoint_data.custom_path,
header_rules=self.endpoint_data.header_rules,
body_rules=self.endpoint_data.body_rules,
max_retries=self.endpoint_data.max_retries,
is_active=True,
config=self.endpoint_data.config,

View File

@@ -11,6 +11,7 @@ from sqlalchemy.orm import Session
from src.core.logger import logger
from src.models.database import ApiKey, ManagementToken, User
from src.utils.perf import PerfRecorder
from src.utils.request_utils import get_client_ip
@@ -55,9 +56,32 @@ class ApiRequestContext:
if not self.raw_body:
raise HTTPException(status_code=400, detail="请求体不能为空")
perf_metrics = getattr(self.request.state, "perf_metrics", None)
perf_sampled = isinstance(perf_metrics, dict) and bool(perf_metrics)
parse_start = PerfRecorder.start(force=perf_sampled)
def _record_parse_duration(duration: float | None) -> None:
if duration is None:
return
if not isinstance(perf_metrics, dict):
return
perf_metrics.setdefault("pipeline", {})["json_parse_ms"] = int(duration * 1000)
try:
self.json_body = json.loads(self.raw_body.decode("utf-8"))
parse_duration = PerfRecorder.stop(
parse_start,
"pipeline_json_parse",
labels={"mode": self.mode},
)
_record_parse_duration(parse_duration)
except json.JSONDecodeError as exc:
parse_duration = PerfRecorder.stop(
parse_start,
"pipeline_json_parse",
labels={"mode": self.mode},
)
_record_parse_duration(parse_duration)
logger.warning(f"解析JSON失败: {exc}")
raise HTTPException(status_code=400, detail="请求体必须是合法的JSON") from exc
@@ -112,6 +136,10 @@ class ApiRequestContext:
path_params=path_params or {},
)
perf_metrics = getattr(request.state, "perf_metrics", None)
if isinstance(perf_metrics, dict) and perf_metrics:
context.extra["perf"] = perf_metrics
# 便于插件/日志引用
request.state.request_id = request_id
if user:

View File

@@ -15,6 +15,7 @@ from src.models.database import ApiKey, AuditEventType, User
from src.services.auth.service import AuthService
from src.services.system.audit import AuditService
from src.services.usage.service import UsageService
from src.utils.perf import PerfRecorder
if TYPE_CHECKING:
from src.models.database import ManagementToken
@@ -55,6 +56,28 @@ class ApiRequestPipeline:
api_format_hint: str | None = None,
path_params: dict[str, Any] | None = None,
) -> Any:
perf_labels = {
"mode": getattr(mode, "value", str(mode)),
"adapter": adapter.name,
}
perf_sampled = PerfRecorder.should_store_sample()
if perf_sampled:
setattr(http_request.state, "perf_sampled", True)
setattr(
http_request.state,
"perf_metrics",
{"pipeline": {}, "sample_rate": getattr(config, "perf_store_sample_rate", 1.0)},
)
def _record_perf_metric(key: str, duration: float | None) -> None:
if duration is None:
return
perf_metrics = getattr(http_request.state, "perf_metrics", None)
if not isinstance(perf_metrics, dict):
return
bucket = perf_metrics.setdefault("pipeline", {})
bucket[key] = int(duration * 1000)
# 高频轮询端点抑制 debug 日志
is_quiet = http_request.url.path in QUIET_POLLING_PATHS
if not is_quiet:
@@ -66,26 +89,31 @@ class ApiRequestPipeline:
adapter.mode,
http_request.url.path,
)
if mode == ApiMode.ADMIN:
user, management_token = await self._authenticate_admin(http_request, db)
api_key = None
elif mode == ApiMode.USER:
user, management_token = await self._authenticate_user(http_request, db)
api_key = None
elif mode == ApiMode.PUBLIC:
user = None
api_key = None
management_token = None
elif mode == ApiMode.MANAGEMENT:
user, management_token = await self._authenticate_management(http_request, db)
api_key = None
else:
if not is_quiet:
logger.debug("[Pipeline] 调用 _authenticate_client")
user, api_key = self._authenticate_client(http_request, db, adapter, quiet=is_quiet)
management_token = None
if not is_quiet:
logger.debug("[Pipeline] 认证完成 | user={}", user.username if user else None)
auth_start = PerfRecorder.start(force=perf_sampled)
try:
if mode == ApiMode.ADMIN:
user, management_token = await self._authenticate_admin(http_request, db)
api_key = None
elif mode == ApiMode.USER:
user, management_token = await self._authenticate_user(http_request, db)
api_key = None
elif mode == ApiMode.PUBLIC:
user = None
api_key = None
management_token = None
elif mode == ApiMode.MANAGEMENT:
user, management_token = await self._authenticate_management(http_request, db)
api_key = None
else:
if not is_quiet:
logger.debug("[Pipeline] 调用 _authenticate_client")
user, api_key = self._authenticate_client(http_request, db, adapter, quiet=is_quiet)
management_token = None
if not is_quiet:
logger.debug("[Pipeline] 认证完成 | user={}", user.username if user else None)
finally:
auth_duration = PerfRecorder.stop(auth_start, "pipeline_auth", labels=perf_labels)
_record_perf_metric("auth_ms", auth_duration)
raw_body = None
if http_request.method in {"POST", "PUT", "PATCH"}:
@@ -93,9 +121,25 @@ class ApiRequestPipeline:
import asyncio
# 添加超时防止卡死
raw_body = await asyncio.wait_for(
http_request.body(), timeout=config.request_body_timeout
)
body_start = PerfRecorder.start(force=perf_sampled)
body_size = 0
try:
raw_body = await asyncio.wait_for(
http_request.body(), timeout=config.request_body_timeout
)
body_size = len(raw_body) if raw_body is not None else 0
finally:
body_duration = PerfRecorder.stop(
body_start,
"pipeline_body_read",
labels=perf_labels,
log_hint=f"size={body_size}",
)
_record_perf_metric("body_read_ms", body_duration)
if perf_sampled:
perf_metrics = getattr(http_request.state, "perf_metrics", None)
if isinstance(perf_metrics, dict):
perf_metrics.setdefault("pipeline", {})["body_bytes"] = int(body_size)
if not is_quiet:
logger.debug(
"[Pipeline] Raw body读取完成 | size={} bytes",
@@ -112,6 +156,7 @@ class ApiRequestPipeline:
if not is_quiet:
logger.debug("[Pipeline] 非写请求跳过读取Body | method={}", http_request.method)
context_start = PerfRecorder.start(force=perf_sampled)
context = ApiRequestContext.build(
request=http_request,
db=db,
@@ -122,6 +167,10 @@ class ApiRequestPipeline:
api_format_hint=api_format_hint,
path_params=path_params,
)
context_duration = PerfRecorder.stop(
context_start, "pipeline_context_build", labels=perf_labels
)
_record_perf_metric("context_build_ms", context_duration)
# 存储 management_token 到 context用于权限检查
if management_token:
context.management_token = management_token
@@ -145,16 +194,28 @@ class ApiRequestPipeline:
context.user,
)
# authorize 可能是异步的,需要检查并 await
authorize_result = adapter.authorize(context)
if hasattr(authorize_result, "__await__"):
await authorize_result
authorize_start = PerfRecorder.start(force=perf_sampled)
try:
authorize_result = adapter.authorize(context)
if hasattr(authorize_result, "__await__"):
await authorize_result
finally:
authorize_duration = PerfRecorder.stop(
authorize_start, "pipeline_authorize", labels=perf_labels
)
_record_perf_metric("authorize_ms", authorize_duration)
try:
handle_start = PerfRecorder.start(force=perf_sampled)
response = await adapter.handle(context)
handle_duration = PerfRecorder.stop(handle_start, "pipeline_handle", labels=perf_labels)
_record_perf_metric("handle_ms", handle_duration)
status_code = getattr(response, "status_code", None)
self._record_audit_event(context, adapter, success=True, status_code=status_code)
return response
except HTTPException as exc:
handle_duration = PerfRecorder.stop(handle_start, "pipeline_handle", labels=perf_labels)
_record_perf_metric("handle_ms", handle_duration)
err_detail = exc.detail if isinstance(exc.detail, str) else str(exc.detail)
self._record_audit_event(
context,
@@ -165,6 +226,8 @@ class ApiRequestPipeline:
)
raise
except Exception as exc:
handle_duration = PerfRecorder.stop(handle_start, "pipeline_handle", labels=perf_labels)
_record_perf_metric("handle_ms", handle_duration)
self._record_audit_event(
context,
adapter,

View File

@@ -125,7 +125,16 @@ class MessageTelemetry:
target_model: str | None = None,
# Provider 响应元数据(如 Gemini 的 modelVersion
response_metadata: dict[str, Any] | None = None,
# 请求元数据(用于性能与调试记录)
request_metadata: dict[str, Any] | None = None,
) -> float:
metadata = response_metadata
if request_metadata:
merged = dict(request_metadata)
if response_metadata:
merged.setdefault("response", response_metadata)
metadata = merged
usage = await UsageService.record_usage(
db=self.db,
user=self.user,
@@ -157,8 +166,8 @@ class MessageTelemetry:
provider_api_key_id=provider_api_key_id,
# 模型映射信息
target_model=target_model,
# Provider 响应元数据
metadata=response_metadata,
# Provider 响应元数据/请求元数据
metadata=metadata,
)
total_cost = float(getattr(usage, "total_cost_usd", 0.0) or 0.0)
@@ -207,6 +216,8 @@ class MessageTelemetry:
has_format_conversion: bool = False,
# 模型映射信息
target_model: str | None = None,
# 请求元数据(用于性能与调试记录)
request_metadata: dict[str, Any] | None = None,
) -> None:
"""
记录失败请求
@@ -257,6 +268,8 @@ class MessageTelemetry:
request_id=self.request_id,
# 模型映射信息
target_model=target_model,
# 请求元数据
metadata=request_metadata,
)
async def record_cancelled(
@@ -283,6 +296,8 @@ class MessageTelemetry:
endpoint_api_format: str | None = None,
has_format_conversion: bool = False,
target_model: str | None = None,
# 请求元数据(用于性能与调试记录)
request_metadata: dict[str, Any] | None = None,
) -> None:
"""
记录客户端取消的请求
@@ -318,6 +333,7 @@ class MessageTelemetry:
response_body=response_body or {},
request_id=self.request_id,
target_model=target_model,
metadata=request_metadata,
)
@@ -375,6 +391,7 @@ class BaseMessageHandler:
start_time: float,
allowed_api_formats: list[str] | None = None,
adapter_detector: AdapterDetectorType | None = None,
perf_metrics: dict[str, Any] | None = None,
) -> None:
self.db = db
self.user = user
@@ -387,6 +404,7 @@ class BaseMessageHandler:
self.allowed_api_formats = allowed_api_formats or ["claude:chat"]
self.primary_api_format = normalize_endpoint_signature(self.allowed_api_formats[0])
self.adapter_detector = adapter_detector
self.perf_metrics = perf_metrics
redis_client = get_redis_client_sync()
self.redis = redis_client
@@ -395,6 +413,11 @@ class BaseMessageHandler:
def elapsed_ms(self) -> int:
return int((time.time() - self.start_time) * 1000)
def _build_request_metadata(self, http_request: Request | None = None) -> dict[str, Any] | None:
if not isinstance(self.perf_metrics, dict) or not self.perf_metrics:
return None
return {"perf": self.perf_metrics}
def _resolve_capability_requirements(
self,
model_name: str,

View File

@@ -189,6 +189,7 @@ class ChatAdapterBase(ApiAdapter):
client_ip=client_ip,
user_agent=user_agent,
start_time=start_time,
perf_metrics=context.extra.get("perf"),
)
# 处理请求
@@ -269,6 +270,7 @@ class ChatAdapterBase(ApiAdapter):
client_ip: str,
user_agent: str,
start_time: float,
perf_metrics: dict[str, Any] | None = None,
) -> Any:
"""创建 Handler 实例 - 子类可覆盖"""
return self.HANDLER_CLASS(
@@ -281,6 +283,7 @@ class ChatAdapterBase(ApiAdapter):
start_time=start_time,
allowed_api_formats=self.allowed_api_formats,
adapter_detector=self.detect_capability_requirements,
perf_metrics=perf_metrics,
)
def _merge_path_params(
@@ -679,6 +682,7 @@ class ChatAdapterBase(ApiAdapter):
if header_rules:
# 获取认证头名称,防止被规则覆盖
from src.core.api_format import get_auth_config_for_endpoint
auth_header, _ = get_auth_config_for_endpoint(cls.FORMAT_ID)
protected_keys = {auth_header.lower(), "content-type"}

View File

@@ -245,6 +245,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
adapter_detector: None | (
Callable[[dict[str, str], dict[str, Any] | None], dict[str, bool]]
) = None,
perf_metrics: dict[str, Any] | None = None,
):
allowed = allowed_api_formats or [self.FORMAT_ID]
super().__init__(
@@ -257,6 +258,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
start_time=start_time,
allowed_api_formats=allowed,
adapter_detector=adapter_detector,
perf_metrics=perf_metrics,
)
self._parser: ResponseParser | None = None
self._request_builder = PassthroughRequestBuilder()
@@ -554,6 +556,10 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
)
# 仅在 FULL 级别才需要保留 parsed_chunks避免长流式响应导致的内存占用
ctx.record_parsed_chunks = SystemConfigService.should_log_body(self.db)
request_metadata = self._build_request_metadata()
if request_metadata and isinstance(request_metadata.get("perf"), dict):
ctx.perf_sampled = True
ctx.perf_metrics.update(request_metadata["perf"])
# 创建更新状态的回调闭包(可以访问 ctx
def update_streaming_status() -> None:
@@ -1318,6 +1324,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
client_response_headers = filter_proxy_response_headers(response_headers)
client_response_headers["content-type"] = "application/json"
request_metadata = self._build_request_metadata()
total_cost = await self.telemetry.record_success(
provider=provider_name,
model=model,
@@ -1343,6 +1350,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
provider_api_key_id=key_id,
# 模型映射信息
target_model=mapped_model_result,
request_metadata=request_metadata,
)
logger.debug(f"{self.FORMAT_ID} 非流式响应完成")
@@ -1365,6 +1373,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
# 记录实际发送给 Provider 的请求体,便于排查问题根因
response_time_ms = self.elapsed_ms()
actual_request_body = provider_request_body or original_request_body
request_metadata = self._build_request_metadata()
await self.telemetry.record_failure(
provider=provider_name or "unknown",
model=model,
@@ -1374,6 +1383,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
request_body=actual_request_body,
error_message=str(e),
is_stream=False,
request_metadata=request_metadata,
)
client_format = (client_api_format_for_error or "").upper()
provider_format = (provider_api_format_for_error or client_format).upper()
@@ -1388,6 +1398,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
except UpstreamClientException as e:
response_time_ms = self.elapsed_ms()
actual_request_body = provider_request_body or original_request_body
request_metadata = self._build_request_metadata()
await self.telemetry.record_failure(
provider=provider_name or "unknown",
model=model,
@@ -1405,6 +1416,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
endpoint_api_format=provider_api_format_for_error or None,
has_format_conversion=needs_conversion_for_error,
target_model=mapped_model_result,
request_metadata=request_metadata,
)
client_format = (client_api_format_for_error or "").upper()
provider_format = (provider_api_format_for_error or client_format).upper()
@@ -1436,6 +1448,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
elif isinstance(e, httpx.HTTPStatusError) and hasattr(e, "response"):
error_response_headers = dict(e.response.headers)
request_metadata = self._build_request_metadata()
await self.telemetry.record_failure(
provider=provider_name or "unknown",
model=model,
@@ -1455,6 +1468,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
has_format_conversion=needs_conversion_for_error,
# 模型映射信息
target_model=mapped_model_result,
request_metadata=request_metadata,
)
raise

View File

@@ -192,6 +192,7 @@ class CliAdapterBase(ApiAdapter):
start_time=start_time,
allowed_api_formats=self.allowed_api_formats,
adapter_detector=self.detect_capability_requirements,
perf_metrics=context.extra.get("perf"),
)
# 处理请求
@@ -661,6 +662,7 @@ class CliAdapterBase(ApiAdapter):
if header_rules:
# 获取认证头名称,防止被规则覆盖
from src.core.api_format import get_auth_config_for_endpoint
auth_header, _ = get_auth_config_for_endpoint(cls.FORMAT_ID)
protected_keys = {auth_header.lower(), "content-type"}

View File

@@ -213,6 +213,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
adapter_detector: None | (
Callable[[dict[str, str], dict[str, Any] | None], dict[str, bool]]
) = None,
perf_metrics: dict[str, Any] | None = None,
):
allowed = allowed_api_formats or [self.FORMAT_ID]
super().__init__(
@@ -225,6 +226,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
start_time=start_time,
allowed_api_formats=allowed,
adapter_detector=adapter_detector,
perf_metrics=perf_metrics,
)
self._parser: ResponseParser | None = None
self._request_builder = PassthroughRequestBuilder()
@@ -561,6 +563,10 @@ class CliMessageHandlerBase(BaseMessageHandler):
)
# 仅在 FULL 级别才需要保留 parsed_chunks避免长流式响应导致的内存占用
ctx.record_parsed_chunks = SystemConfigService.should_log_body(self.db)
request_metadata = self._build_request_metadata(http_request)
if request_metadata and isinstance(request_metadata.get("perf"), dict):
ctx.perf_sampled = True
ctx.perf_metrics.update(request_metadata["perf"])
# 定义请求函数
async def stream_request_func(
@@ -1972,6 +1978,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
if ctx.is_client_disconnected():
# 客户端取消:记录为 cancelled不算系统失败
request_metadata = {"perf": ctx.perf_metrics} if ctx.perf_metrics else None
await bg_telemetry.record_cancelled(
provider=ctx.provider_name or "unknown",
model=ctx.model,
@@ -1993,6 +2000,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
endpoint_api_format=ctx.provider_api_format or None,
has_format_conversion=ctx.needs_conversion,
target_model=ctx.mapped_model,
request_metadata=request_metadata,
)
logger.debug(f"{self.FORMAT_ID} 流式响应被客户端取消")
logger.info(
@@ -2001,6 +2009,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
)
else:
# 服务端/上游异常:记录为失败
request_metadata = {"perf": ctx.perf_metrics} if ctx.perf_metrics else None
await bg_telemetry.record_failure(
provider=ctx.provider_name or "unknown",
model=ctx.model,
@@ -2025,6 +2034,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
has_format_conversion=ctx.needs_conversion,
# 模型映射信息
target_model=ctx.mapped_model,
request_metadata=request_metadata,
)
logger.debug(f"{self.FORMAT_ID} 流式响应中断")
logger.info(
@@ -2050,6 +2060,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
f"provider={ctx.provider_name}, model={ctx.model}, "
f"in={ctx.input_tokens}, out={ctx.output_tokens}"
)
request_metadata = {"perf": ctx.perf_metrics} if ctx.perf_metrics else None
total_cost = await bg_telemetry.record_success(
provider=ctx.provider_name,
model=ctx.model,
@@ -2079,6 +2090,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
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(f"[{ctx.request_id}] Usage 记录完成: cost=${total_cost:.6f}")
# 简洁的请求完成摘要(两行格式)
@@ -2210,6 +2222,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
# 失败时返回给客户端的是 JSON 错误响应
client_response_headers = {"content-type": "application/json"}
request_metadata = {"perf": ctx.perf_metrics} if ctx.perf_metrics else None
await self.telemetry.record_failure(
provider=ctx.provider_name or "unknown",
model=ctx.model,
@@ -2228,6 +2241,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
has_format_conversion=ctx.needs_conversion,
# 模型映射信息
target_model=ctx.mapped_model,
request_metadata=request_metadata,
)
# _update_usage_to_streaming 方法已移至基类 BaseMessageHandler
@@ -2551,6 +2565,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
client_response_headers = filter_proxy_response_headers(response_headers)
client_response_headers["content-type"] = "application/json"
request_metadata = self._build_request_metadata()
total_cost = await self.telemetry.record_success(
provider=provider_name,
model=model,
@@ -2579,6 +2594,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
target_model=mapped_model_result,
# Provider 响应元数据(如 Gemini 的 modelVersion
response_metadata=response_metadata_result if response_metadata_result else None,
request_metadata=request_metadata,
)
logger.info(f"{self.FORMAT_ID} 非流式响应处理完成")
@@ -2595,6 +2611,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
# 记录实际发送给 Provider 的请求体,便于排查问题根因
response_time_ms = int((time.time() - sync_start_time) * 1000)
actual_request_body = provider_request_body or original_request_body
request_metadata = self._build_request_metadata()
await self.telemetry.record_failure(
provider=provider_name or "unknown",
model=model,
@@ -2605,6 +2622,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
error_message=str(e),
is_stream=False,
api_format=api_format,
request_metadata=request_metadata,
)
raise
@@ -2629,6 +2647,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
elif isinstance(e, httpx.HTTPStatusError) and hasattr(e, "response"):
error_response_headers = dict(e.response.headers)
request_metadata = self._build_request_metadata()
await self.telemetry.record_failure(
provider=provider_name or "unknown",
model=model,
@@ -2648,6 +2667,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
has_format_conversion=needs_conversion,
# 模型映射信息
target_model=mapped_model_result,
request_metadata=request_metadata,
)
raise

View File

@@ -94,6 +94,10 @@ class StreamContext:
# 是否记录 parsed_chunks可用于降低高并发/长流式响应的内存占用)
record_parsed_chunks: bool = True
# 性能采集(可选)
perf_sampled: bool = False
perf_metrics: dict[str, Any] = field(default_factory=dict)
# 流式格式转换状态(跨 chunk 追踪)
stream_conversion_state: StreamState | None = None

View File

@@ -14,6 +14,7 @@ from __future__ import annotations
import asyncio
import codecs
import json
import time
from collections.abc import AsyncGenerator, Callable
from dataclasses import dataclass
from typing import Any
@@ -43,6 +44,7 @@ from src.core.exceptions import (
)
from src.core.logger import logger
from src.models.database import Provider, ProviderEndpoint
from src.utils.perf import PerfRecorder
from src.utils.sse_parser import SSEEventParser
from src.utils.timeout import read_first_chunk_with_ttfb_timeout
@@ -413,16 +415,20 @@ class StreamProcessor:
buffer = b""
# 使用增量解码器处理跨 chunk 的 UTF-8 字符
decoder = codecs.getincrementaldecoder("utf-8")(errors="replace")
metrics_enabled = PerfRecorder.enabled()
perf_capture = metrics_enabled or ctx.perf_sampled
parse_time = 0.0
convert_time = 0.0
_api_format_str = str(ctx.api_format or "")
client_format = (ctx.client_api_format or _api_format_str).strip().lower()
provider_format = (ctx.provider_api_format or _api_format_str).strip().lower()
client_family = (
client_format.split(":", 1)[0] if ":" in client_format else client_format
)
) or "unknown"
provider_family = (
provider_format.split(":", 1)[0] if ":" in provider_format else provider_format
)
) or "unknown"
# 使用 handler 层预计算的 needs_conversion由 candidate 决定)
needs_conversion = ctx.needs_conversion
@@ -446,6 +452,15 @@ class StreamProcessor:
self.on_streaming_start()
streaming_started = True
def _process_line_with_perf(line: str, *, skip_record: bool = False) -> None:
nonlocal parse_time
if perf_capture:
t0 = time.perf_counter()
self._process_line(ctx, sse_parser, line, skip_record=skip_record)
parse_time += time.perf_counter() - t0
return
self._process_line(ctx, sse_parser, line, skip_record=skip_record)
def _build_stream_error_payload(message: str) -> dict:
if client_family == "openai":
return {
@@ -477,103 +492,112 @@ class StreamProcessor:
message_id=ctx.response_id or ctx.request_id or "",
)
skip_next_blank_line = False
empty_yield_count = 0 # 空转计数(防护异常情况)
openai_done_sent = (
False # 统一为 OpenAI 客户端补齐 [DONE](避免不同 Provider 行为差异)
)
# 转换状态变量(在 needs_conversion 块内统一初始化,确保作用域正确)
skip_next_blank_line = False
empty_yield_count = 0 # 空转计数(防护异常情况)
openai_done_sent = (
False # 统一为 OpenAI 客户端补齐 [DONE](避免不同 Provider 行为差异)
)
def _emit_converted_line(normalized_line: str) -> list[bytes]:
nonlocal skip_next_blank_line, openai_done_sent
def _emit_converted_line(normalized_line: str) -> list[bytes]:
nonlocal skip_next_blank_line, openai_done_sent, convert_time
# 空行:事件分隔符(避免重复输出)
if normalized_line == "":
if skip_next_blank_line:
skip_next_blank_line = False
return []
return [b"\n"]
# 丢弃 Provider 的 event 行,避免泄漏/污染目标格式
if normalized_line.startswith("event:"):
# 空行:事件分隔符(避免重复输出)
if normalized_line == "":
if skip_next_blank_line:
skip_next_blank_line = False
return []
return [b"\n"]
# OpenAI done 信号(仅用于 OpenAI 客户端)
if (
normalized_line.startswith("data:")
and normalized_line[5:].strip() == "[DONE]"
):
skip_next_blank_line = True
if client_family == "openai":
openai_done_sent = True
return [b"data: [DONE]\n\n"]
return []
# 默认只处理 SSE 的 data 行;但 Gemini 上游可能返回 JSON-array/chunks无 data 前缀)
is_data_line = normalized_line.startswith("data:")
if not is_data_line:
if provider_family != "gemini":
return []
data_content = normalized_line.strip()
else:
data_content = normalized_line[5:].strip()
# Gemini 可能包含 JSON 数组包装符,直接忽略
if data_content in ("", "[", "]", ","):
return []
# JSON-array/chunks 可能带前后逗号(对象分隔符),做一次保守清理
data_content = data_content.lstrip(",").rstrip(",").strip()
if data_content in ("", "[", "]", ","):
return []
try:
data_obj = json.loads(data_content)
except json.JSONDecodeError:
# 跨格式转换时JSON 解析失败应跳过而不是透传(避免泄漏 Provider 格式)
logger.warning(
f"[{self.request_id}] JSON 解析失败,跳过该行: {data_content[:100]}"
)
return []
if not isinstance(data_obj, dict):
return []
try:
converted_events = registry.convert_stream_chunk(
data_obj,
provider_format,
client_format,
state=ctx.stream_conversion_state,
)
except Exception as conv_err:
# 首字节后无法 failover输出目标格式错误事件并终止流
# 使用 502 表示上游返回了非预期格式Bad Gateway
ctx.status_code = 502
ctx.error_message = "format_conversion_failed"
# 日志记录完整错误(内部排查),客户端只返回脱敏消息
logger.warning(f"[{self.request_id}] 流式格式转换失败: {conv_err}")
payload = _build_stream_error_payload("响应格式转换失败,请稍后重试")
error_bytes = (
f"data: {json.dumps(payload, ensure_ascii=False)}\n\n".encode()
)
done_bytes = b"data: [DONE]\n\n" if client_family == "openai" else b""
if done_bytes:
openai_done_sent = True
return [error_bytes, done_bytes]
# 丢弃 Provider 的 event 行,避免泄漏/污染目标格式
if normalized_line.startswith("event:"):
return []
# OpenAI done 信号(仅用于 OpenAI 客户端)
if (
normalized_line.startswith("data:")
and normalized_line[5:].strip() == "[DONE]"
):
skip_next_blank_line = True
out: list[bytes] = []
if client_family == "openai":
openai_done_sent = True
return [b"data: [DONE]\n\n"]
return []
for evt in converted_events:
# 记录转换后的数据到 parsed_chunks这是客户端实际收到的格式
if isinstance(evt, dict):
ctx.data_count += 1
if ctx.record_parsed_chunks:
ctx.parsed_chunks.append(evt)
# 默认只处理 SSE 的 data 行;但 Gemini 上游可能返回 JSON-array/chunks无 data 前缀)
is_data_line = normalized_line.startswith("data:")
if not is_data_line:
if provider_family != "gemini":
return []
data_content = normalized_line.strip()
else:
data_content = normalized_line[5:].strip()
# 统一使用 SSE 格式输出Gemini streamGenerateContent 也使用 SSE
# 参考: https://ai.google.dev/api/generate-content
out.append(f"data: {json.dumps(evt, ensure_ascii=False)}\n\n".encode())
return out
# Gemini 可能包含 JSON 数组包装符,直接忽略
if data_content in ("", "[", "]", ","):
return []
# JSON-array/chunks 可能带前后逗号(对象分隔符),做一次保守清理
data_content = data_content.lstrip(",").rstrip(",").strip()
if data_content in ("", "[", "]", ","):
return []
convert_start = time.perf_counter() if perf_capture else None
try:
data_obj = json.loads(data_content)
except json.JSONDecodeError:
if perf_capture and convert_start is not None:
convert_time += time.perf_counter() - convert_start
# 跨格式转换时JSON 解析失败应跳过而不是透传(避免泄漏 Provider 格式)
logger.warning(
f"[{self.request_id}] JSON 解析失败,跳过该行: {data_content[:100]}"
)
return []
if not isinstance(data_obj, dict):
return []
try:
converted_events = registry.convert_stream_chunk(
data_obj,
provider_format,
client_format,
state=ctx.stream_conversion_state,
)
except Exception as conv_err:
# 首字节后无法 failover输出目标格式错误事件并终止流
# 使用 502 表示上游返回了非预期格式Bad Gateway
ctx.status_code = 502
ctx.error_message = "format_conversion_failed"
# 日志记录完整错误(内部排查),客户端只返回脱敏消息
logger.warning(f"[{self.request_id}] 流式格式转换失败: {conv_err}")
payload = _build_stream_error_payload("响应格式转换失败,请稍后重试")
error_bytes = (
f"data: {json.dumps(payload, ensure_ascii=False)}\n\n".encode()
)
done_bytes = b"data: [DONE]\n\n" if client_family == "openai" else b""
if done_bytes:
openai_done_sent = True
if perf_capture and convert_start is not None:
convert_time += time.perf_counter() - convert_start
return [error_bytes, done_bytes]
if perf_capture and convert_start is not None:
convert_time += time.perf_counter() - convert_start
skip_next_blank_line = True
out: list[bytes] = []
for evt in converted_events:
# 记录转换后的数据到 parsed_chunks这是客户端实际收到的格式
if isinstance(evt, dict):
ctx.data_count += 1
if ctx.record_parsed_chunks:
ctx.parsed_chunks.append(evt)
# 统一使用 SSE 格式输出Gemini streamGenerateContent 也使用 SSE
# 参考: https://ai.google.dev/api/generate-content
out.append(f"data: {json.dumps(evt, ensure_ascii=False)}\n\n".encode())
return out
# 统一处理 prefetched + iterator
if prefetched_chunks:
@@ -591,7 +615,7 @@ class StreamProcessor:
if line:
# 需要格式转换时,跳过记录原始数据(由 _emit_converted_line 记录转换后的数据)
self._process_line(ctx, sse_parser, line, skip_record=True)
_process_line_with_perf(line, skip_record=True)
normalized_line = line.rstrip("\r\n") if line else ""
out_chunks = _emit_converted_line(normalized_line)
if not out_chunks:
@@ -625,7 +649,7 @@ class StreamProcessor:
if line:
# 需要格式转换时,跳过记录原始数据(由 _emit_converted_line 记录转换后的数据)
self._process_line(ctx, sse_parser, line, skip_record=True)
_process_line_with_perf(line, skip_record=True)
normalized_line = line.rstrip("\r\n") if line else ""
out_chunks = _emit_converted_line(normalized_line)
if not out_chunks:
@@ -655,7 +679,7 @@ class StreamProcessor:
line = ""
if line:
# 需要格式转换时,跳过记录原始数据
self._process_line(ctx, sse_parser, line, skip_record=True)
_process_line_with_perf(line, skip_record=True)
normalized_line = line.rstrip("\r\n")
out_chunks = _emit_converted_line(normalized_line)
for out in out_chunks:
@@ -684,7 +708,7 @@ class StreamProcessor:
try:
# 使用增量解码器,可以正确处理跨 chunk 的多字节字符
line = decoder.decode(line_bytes + b"\n", False)
self._process_line(ctx, sse_parser, line)
_process_line_with_perf(line)
except Exception as e:
# 解码失败,记录警告但继续处理
logger.warning(
@@ -708,7 +732,7 @@ class StreamProcessor:
try:
# 使用增量解码器,可以正确处理跨 chunk 的多字节字符
line = decoder.decode(line_bytes + b"\n", False)
self._process_line(ctx, sse_parser, line)
_process_line_with_perf(line)
except Exception as e:
# 解码失败,记录警告但继续处理
logger.warning(
@@ -722,7 +746,7 @@ class StreamProcessor:
try:
# 使用 final=True 处理最后的不完整字符
line = decoder.decode(buffer, True)
self._process_line(ctx, sse_parser, line)
_process_line_with_perf(line)
except Exception as e:
logger.warning(
f"[{self.request_id}] 处理剩余缓冲区失败: {e}, bytes={buffer[:50]!r}"
@@ -735,6 +759,33 @@ class StreamProcessor:
except GeneratorExit:
raise
finally:
if metrics_enabled:
labels = {
"format": client_family or "unknown",
"provider": str(ctx.provider_name or "unknown"),
"conversion": "true" if ctx.needs_conversion else "false",
}
if parse_time > 0:
PerfRecorder.record_timing("stream_parse", parse_time, labels=labels)
if convert_time > 0:
PerfRecorder.record_timing("stream_conversion", convert_time, labels=labels)
if ctx.chunk_count:
PerfRecorder.record_counter(
"stream_chunks_total", ctx.chunk_count, labels=labels
)
if ctx.data_count:
PerfRecorder.record_counter(
"stream_data_events_total", ctx.data_count, labels=labels
)
if ctx.perf_sampled:
if parse_time > 0:
ctx.perf_metrics["stream_parse_ms"] = int(parse_time * 1000)
if convert_time > 0:
ctx.perf_metrics["stream_conversion_ms"] = int(convert_time * 1000)
if ctx.chunk_count:
ctx.perf_metrics["stream_chunks"] = int(ctx.chunk_count)
if ctx.data_count:
ctx.perf_metrics["stream_data_events"] = int(ctx.data_count)
await self._cleanup(response_ctx, http_client)
def _process_line(

View File

@@ -184,6 +184,9 @@ class StreamTelemetryRecorder:
"content-type": "text/event-stream",
}
)
metadata: dict[str, Any] = {"stream": True, "content_length": ctx.data_count}
if ctx.perf_metrics:
metadata["perf"] = ctx.perf_metrics
await writer.record_success(
provider=ctx.provider_name or "unknown",
@@ -208,7 +211,7 @@ class StreamTelemetryRecorder:
provider_api_key_id=ctx.key_id,
target_model=ctx.mapped_model,
request_type="chat",
metadata={"stream": True, "content_length": ctx.data_count},
metadata=metadata,
endpoint_api_format=ctx.provider_api_format,
has_format_conversion=ctx.needs_conversion,
)
@@ -230,6 +233,9 @@ class StreamTelemetryRecorder:
client_response_headers = ctx.client_response_headers or {
"content-type": "application/json"
}
metadata: dict[str, Any] = {"stream": True, "content_length": ctx.data_count}
if ctx.perf_metrics:
metadata["perf"] = ctx.perf_metrics
await writer.record_failure(
provider=ctx.provider_name or "unknown",
@@ -251,7 +257,7 @@ class StreamTelemetryRecorder:
client_response_headers=client_response_headers,
target_model=ctx.mapped_model,
request_type="chat",
metadata={"stream": True, "content_length": ctx.data_count},
metadata=metadata,
endpoint_api_format=ctx.provider_api_format,
has_format_conversion=ctx.needs_conversion,
)
@@ -274,6 +280,9 @@ class StreamTelemetryRecorder:
client_response_headers = ctx.client_response_headers or {
"content-type": "application/json"
}
metadata: dict[str, Any] = {"stream": True, "content_length": ctx.data_count}
if ctx.perf_metrics:
metadata["perf"] = ctx.perf_metrics
await writer.record_cancelled(
provider=ctx.provider_name or "unknown",
@@ -295,7 +304,7 @@ class StreamTelemetryRecorder:
client_response_headers=client_response_headers,
target_model=ctx.mapped_model,
request_type="chat",
metadata={"stream": True, "content_length": ctx.data_count},
metadata=metadata,
endpoint_api_format=ctx.provider_api_format,
has_format_conversion=ctx.needs_conversion,
)