feat(agent): 增加最终模型请求预算观测 (#6290)

This commit is contained in:
InfinityPacer
2026-08-13 13:52:30 +08:00
committed by GitHub
parent ce3508730b
commit f979274c07
6 changed files with 1526 additions and 68 deletions
+229 -9
View File
@@ -9,19 +9,25 @@ from langchain.agents.middleware.types import (
ResponseT,
)
from langchain_core.messages import AIMessage
from langchain_core.messages.utils import count_tokens_approximately
from app.log import logger
class UsageMiddleware(AgentMiddleware):
"""记录模型调用 usage 信息并回传给外部会话"""
"""观察最终模型请求预算,并记录模型返回的真实 usage"""
def __init__(
self,
*,
on_usage: Callable[[dict[str, Any]], None] | None = None,
on_request_budget: Callable[[dict[str, Any]], None] | None = None,
next_request_sequence: Callable[[], int] | None = None,
) -> None:
self.on_usage = on_usage
self.on_request_budget = on_request_budget
self.next_request_sequence = next_request_sequence
self._request_sequence = 0
@staticmethod
def _coerce_int(value: Any) -> int | None:
@@ -32,6 +38,37 @@ class UsageMiddleware(AgentMiddleware):
except (TypeError, ValueError):
return None
@staticmethod
def _coerce_positive_int(value: Any) -> int | None:
"""仅接受模型 profile 和请求设置声明的非 bool 正整数。"""
if isinstance(value, bool) or not isinstance(value, int) or value <= 0:
return None
return value
@classmethod
def _lookup_positive_int(cls, container: Any, *keys: str) -> int | None:
"""按字段优先级读取 token 上限,拒绝隐式数值转换。"""
if not container:
return None
getter = getattr(container, "get", None)
if callable(getter):
for key in keys:
value = getter(key)
if value is not None:
normalized = cls._coerce_positive_int(value)
if normalized is not None:
return normalized
for key in keys:
value = getattr(container, key, None)
if value is not None:
normalized = cls._coerce_positive_int(value)
if normalized is not None:
return normalized
return None
@classmethod
def _lookup_int(cls, container: Any, *keys: str) -> int | None:
if not container:
@@ -65,18 +102,158 @@ class UsageMiddleware(AgentMiddleware):
@classmethod
def _extract_model_name(cls, model: Any) -> str | None:
return (
getattr(model, "model", None)
or getattr(model, "model_name", None)
or getattr(model, "model_id", None)
)
for field in ("model", "model_name", "model_id"):
try:
value = getattr(model, field, None)
except Exception:
continue
if value:
return value
return None
@classmethod
def _extract_context_window_tokens(cls, model: Any) -> int | None:
profile = getattr(model, "profile", None)
try:
profile = getattr(model, "profile", None)
except Exception:
return None
if not profile:
return None
return cls._lookup_int(profile, "max_input_tokens", "input_token_limit")
try:
return cls._lookup_positive_int(
profile, "max_input_tokens", "input_token_limit"
)
except Exception:
return None
@classmethod
def _extract_model_max_output_tokens(cls, model: Any) -> int | None:
"""读取模型输出能力上限;该值不代表单次请求已经预留的输出空间。"""
try:
profile = getattr(model, "profile", None)
except Exception:
return None
if not profile:
return None
try:
return cls._lookup_positive_int(
profile, "max_output_tokens", "output_token_limit"
)
except Exception:
return None
def _next_request_sequence(self) -> int | None:
"""优先使用会话级序号,使图重建后的请求仍保持单调顺序。"""
if callable(self.next_request_sequence):
try:
return self.next_request_sequence()
except Exception as error:
logger.debug(
"分配会话级模型请求序号失败: error_type=%s",
type(error).__name__,
)
# 无法证明顺序的请求仍可累计 usage,但不能参与最近请求快照竞争。
return None
self._request_sequence += 1
return self._request_sequence
@classmethod
def _extract_configured_output_limit_tokens(
cls, request: ModelRequest
) -> int | None:
"""读取最终请求显式配置的单次输出上限。"""
model_settings = request.model_settings or {}
value = cls._lookup_positive_int(
model_settings,
"max_completion_tokens",
"max_tokens",
"max_output_tokens",
)
return value
@staticmethod
def _count_multimodal_blocks(messages: list[Any]) -> tuple[int, int]:
"""统计图片和未知多模态块,不保留块内容或资源地址。"""
image_count = 0
unknown_count = 0
for message in messages:
content = getattr(message, "content", None)
if not isinstance(content, list):
continue
for block in content:
if isinstance(block, str):
continue
if not isinstance(block, dict):
unknown_count += 1
continue
block_type = block.get("type")
if block_type in {"image", "image_url"}:
image_count += 1
elif block_type != "text":
unknown_count += 1
return image_count, unknown_count
@classmethod
def estimate_request(cls, request: ModelRequest) -> dict[str, Any]:
"""估算最终模型输入组成,仅返回可安全暴露的聚合数字。"""
messages = list(request.messages or [])
system_messages = [request.system_message] if request.system_message else []
tools = list(request.tools or [])
message_tokens = count_tokens_approximately(
messages,
use_usage_metadata_scaling=False,
)
system_tokens = count_tokens_approximately(
system_messages,
use_usage_metadata_scaling=False,
)
tool_tokens = count_tokens_approximately(
[],
tools=tools,
use_usage_metadata_scaling=False,
)
estimated_input_tokens = message_tokens + system_tokens + tool_tokens
context_window_tokens = cls._extract_context_window_tokens(request.model)
model_max_output_tokens = cls._extract_model_max_output_tokens(request.model)
configured_output_limit_tokens = cls._extract_configured_output_limit_tokens(
request
)
image_count, unknown_multimodal_count = cls._count_multimodal_blocks(
[*system_messages, *messages]
)
estimated_input_ratio = (
estimated_input_tokens / context_window_tokens
if context_window_tokens
else None
)
return {
"has_estimate": True,
"model": cls._extract_model_name(request.model),
"message_count": len(messages),
"tool_count": len(tools),
"image_count": image_count,
"unknown_multimodal_count": unknown_multimodal_count,
"message_tokens": message_tokens,
"system_tokens": system_tokens,
"tool_tokens": tool_tokens,
# 该成本已经包含在 message_tokens 中,只单独暴露组成,不能再次汇总。
"multimodal_tokens": image_count * 85,
"estimated_input_tokens": estimated_input_tokens,
"context_window_tokens": context_window_tokens,
"estimated_remaining_input_tokens": (
context_window_tokens - estimated_input_tokens
if context_window_tokens
else None
),
"estimated_input_ratio": estimated_input_ratio,
"estimated_over_input_limit": (
estimated_input_tokens > context_window_tokens
if context_window_tokens
else None
),
"model_max_output_tokens": model_max_output_tokens,
"configured_output_limit_tokens": configured_output_limit_tokens,
}
@classmethod
def _extract_usage(cls, ai_message: AIMessage) -> dict[str, Any]:
@@ -290,6 +467,7 @@ class UsageMiddleware(AgentMiddleware):
cache_miss_tokens,
)
)
input_usage_available = input_tokens is not None
resolved_input = input_tokens or 0
resolved_output = output_tokens or 0
resolved_total = (
@@ -315,6 +493,7 @@ class UsageMiddleware(AgentMiddleware):
return {
"has_usage": has_usage,
"input_usage_available": input_usage_available,
"cache_usage_available": has_cache_usage,
"input_tokens": resolved_input,
"output_tokens": resolved_output,
@@ -332,6 +511,39 @@ class UsageMiddleware(AgentMiddleware):
[ModelRequest[ContextT]], Awaitable[ModelResponse[ResponseT]]
],
) -> ModelResponse[ResponseT]:
request_sequence = self._next_request_sequence()
request_budget = None
try:
request_budget = {
"request_sequence": request_sequence,
**self.estimate_request(request),
}
except Exception as error:
logger.debug(
"估算最终模型请求预算失败: error_type=%s",
type(error).__name__,
)
request_budget = {
"request_sequence": request_sequence,
"has_estimate": False,
"model": self._extract_model_name(request.model),
"context_window_tokens": self._extract_context_window_tokens(
request.model
),
}
if callable(self.on_request_budget):
request_budget_recorded = False
try:
self.on_request_budget(request_budget)
request_budget_recorded = True
except Exception as error:
logger.debug(
"记录最终模型请求预算失败: error_type=%s",
type(error).__name__,
)
else:
request_budget_recorded = False
response = await handler(request)
if not callable(self.on_usage):
@@ -351,6 +563,7 @@ class UsageMiddleware(AgentMiddleware):
if ai_message
else {
"has_usage": False,
"input_usage_available": False,
"cache_usage_available": False,
"input_tokens": 0,
"output_tokens": 0,
@@ -363,11 +576,18 @@ class UsageMiddleware(AgentMiddleware):
)
context_window_tokens = self._extract_context_window_tokens(request.model)
context_usage_ratio = None
if context_window_tokens and usage["has_usage"]:
if context_window_tokens and usage["input_usage_available"]:
context_usage_ratio = usage["input_tokens"] / context_window_tokens
self.on_usage(
{
"request_sequence": request_sequence,
"request_budget_recorded": request_budget_recorded,
"estimated_input_tokens": (
request_budget.get("estimated_input_tokens")
if request_budget
else None
),
"model": self._extract_model_name(request.model),
"context_window_tokens": context_window_tokens,
"context_usage_ratio": context_usage_ratio,