mirror of
https://github.com/jxxghp/MoviePilot.git
synced 2026-08-15 19:14:01 +08:00
feat(agent): 增加最终模型请求预算观测 (#6290)
This commit is contained in:
@@ -139,10 +139,30 @@ async def _async_finish_processing_status(
|
||||
class _SessionUsageSnapshot:
|
||||
model: Optional[str] = None
|
||||
context_window_tokens: Optional[int] = None
|
||||
last_input_tokens: int = 0
|
||||
last_output_tokens: int = 0
|
||||
last_total_tokens: int = 0
|
||||
last_input_usage_available: bool = False
|
||||
last_input_tokens: Optional[int] = None
|
||||
last_output_tokens: Optional[int] = None
|
||||
last_total_tokens: Optional[int] = None
|
||||
last_context_usage_ratio: Optional[float] = None
|
||||
last_request_sequence: int = 0
|
||||
last_request_estimate_available: bool = False
|
||||
last_estimated_input_tokens: Optional[int] = None
|
||||
last_estimated_message_tokens: Optional[int] = None
|
||||
last_estimated_system_tokens: Optional[int] = None
|
||||
last_estimated_tool_tokens: Optional[int] = None
|
||||
last_estimated_multimodal_tokens: Optional[int] = None
|
||||
last_estimated_input_ratio: Optional[float] = None
|
||||
last_estimated_remaining_input_tokens: Optional[int] = None
|
||||
last_estimated_over_input_limit: Optional[bool] = None
|
||||
last_message_count: int = 0
|
||||
last_tool_count: int = 0
|
||||
last_image_count: int = 0
|
||||
last_unknown_multimodal_count: int = 0
|
||||
model_max_output_tokens: Optional[int] = None
|
||||
configured_output_limit_tokens: Optional[int] = None
|
||||
last_actual_input_tokens: Optional[int] = None
|
||||
last_estimate_error_tokens: Optional[int] = None
|
||||
last_estimate_error_ratio: Optional[float] = None
|
||||
last_cache_usage_available: bool = False
|
||||
last_cache_read_input_tokens: int = 0
|
||||
last_cache_write_input_tokens: int = 0
|
||||
@@ -163,10 +183,30 @@ class _SessionUsageSnapshot:
|
||||
"session_id": session_id,
|
||||
"model": self.model,
|
||||
"context_window_tokens": self.context_window_tokens,
|
||||
"last_input_usage_available": self.last_input_usage_available,
|
||||
"last_input_tokens": self.last_input_tokens,
|
||||
"last_output_tokens": self.last_output_tokens,
|
||||
"last_total_tokens": self.last_total_tokens,
|
||||
"last_context_usage_ratio": self.last_context_usage_ratio,
|
||||
"last_request_sequence": self.last_request_sequence,
|
||||
"last_request_estimate_available": self.last_request_estimate_available,
|
||||
"last_estimated_input_tokens": self.last_estimated_input_tokens,
|
||||
"last_estimated_message_tokens": self.last_estimated_message_tokens,
|
||||
"last_estimated_system_tokens": self.last_estimated_system_tokens,
|
||||
"last_estimated_tool_tokens": self.last_estimated_tool_tokens,
|
||||
"last_estimated_multimodal_tokens": self.last_estimated_multimodal_tokens,
|
||||
"last_estimated_input_ratio": self.last_estimated_input_ratio,
|
||||
"last_estimated_remaining_input_tokens": self.last_estimated_remaining_input_tokens,
|
||||
"last_estimated_over_input_limit": self.last_estimated_over_input_limit,
|
||||
"last_message_count": self.last_message_count,
|
||||
"last_tool_count": self.last_tool_count,
|
||||
"last_image_count": self.last_image_count,
|
||||
"last_unknown_multimodal_count": self.last_unknown_multimodal_count,
|
||||
"model_max_output_tokens": self.model_max_output_tokens,
|
||||
"configured_output_limit_tokens": self.configured_output_limit_tokens,
|
||||
"last_actual_input_tokens": self.last_actual_input_tokens,
|
||||
"last_estimate_error_tokens": self.last_estimate_error_tokens,
|
||||
"last_estimate_error_ratio": self.last_estimate_error_ratio,
|
||||
"last_cache_usage_available": self.last_cache_usage_available,
|
||||
"last_cache_read_input_tokens": self.last_cache_read_input_tokens,
|
||||
"last_cache_write_input_tokens": self.last_cache_write_input_tokens,
|
||||
@@ -355,6 +395,7 @@ class MoviePilotAgent:
|
||||
self._pending_secret_confirmation: Optional[_PendingSecretConfirmation] = None
|
||||
self._streamed_output = ""
|
||||
self._session_usage = _SessionUsageSnapshot()
|
||||
self._request_sequence = 0
|
||||
self._llm_runtime_config: Optional[Dict[str, Any]] = None
|
||||
self._llm_provider_selection: Dict[str, Any] = {}
|
||||
self._agent_started_at: Optional[datetime] = None
|
||||
@@ -550,6 +591,13 @@ class MoviePilotAgent:
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _coerce_positive_int(value: Any) -> Optional[int]:
|
||||
"""仅接受模型 profile 声明的非 bool 正整数。"""
|
||||
if isinstance(value, bool) or not isinstance(value, int) or value <= 0:
|
||||
return None
|
||||
return value
|
||||
|
||||
@staticmethod
|
||||
def _get_recursion_limit() -> int:
|
||||
"""读取 LangGraph 递归上限,防止模型持续循环调用工具。"""
|
||||
@@ -573,49 +621,88 @@ class MoviePilotAgent:
|
||||
if not profile:
|
||||
return None
|
||||
if isinstance(profile, dict):
|
||||
return cls._coerce_int(
|
||||
profile.get("max_input_tokens") or profile.get("input_token_limit")
|
||||
candidates = (
|
||||
profile.get("max_input_tokens"),
|
||||
profile.get("input_token_limit"),
|
||||
)
|
||||
return cls._coerce_int(
|
||||
getattr(profile, "max_input_tokens", None)
|
||||
or getattr(profile, "input_token_limit", None)
|
||||
)
|
||||
else:
|
||||
candidates = (
|
||||
getattr(profile, "max_input_tokens", None),
|
||||
getattr(profile, "input_token_limit", None),
|
||||
)
|
||||
for candidate in candidates:
|
||||
normalized = cls._coerce_positive_int(candidate)
|
||||
if normalized is not None:
|
||||
return normalized
|
||||
return None
|
||||
|
||||
def _sync_model_profile(self, model: Any) -> None:
|
||||
model_name = self._get_model_name(model)
|
||||
context_window_tokens = self._get_context_window_tokens(model)
|
||||
if model_name:
|
||||
self._session_usage.model = model_name
|
||||
if context_window_tokens:
|
||||
self._session_usage.context_window_tokens = context_window_tokens
|
||||
self._session_usage.context_window_tokens = context_window_tokens
|
||||
|
||||
def _next_request_sequence(self) -> int:
|
||||
"""为当前会话中的模型请求分配跨 Agent 图单调递增的序号。"""
|
||||
self._request_sequence += 1
|
||||
return self._request_sequence
|
||||
|
||||
def _record_usage(self, usage: dict[str, Any]) -> None:
|
||||
if not usage:
|
||||
return
|
||||
|
||||
model_name = usage.get("model")
|
||||
context_window_tokens = self._coerce_int(usage.get("context_window_tokens"))
|
||||
if model_name:
|
||||
self._session_usage.model = model_name
|
||||
if context_window_tokens:
|
||||
self._session_usage.context_window_tokens = context_window_tokens
|
||||
|
||||
self._session_usage.model_call_count += 1
|
||||
self._session_usage.last_updated_at = datetime.now()
|
||||
|
||||
has_request_sequence = "request_sequence" in usage
|
||||
request_sequence = self._coerce_int(usage.get("request_sequence"))
|
||||
if (
|
||||
usage.get("request_budget_recorded") is False
|
||||
and request_sequence is not None
|
||||
and request_sequence >= self._session_usage.last_request_sequence
|
||||
):
|
||||
self._record_request_budget(
|
||||
{
|
||||
"request_sequence": request_sequence,
|
||||
"has_estimate": False,
|
||||
}
|
||||
)
|
||||
is_current_request = (
|
||||
not has_request_sequence
|
||||
or request_sequence == self._session_usage.last_request_sequence
|
||||
)
|
||||
|
||||
if is_current_request:
|
||||
model_name = usage.get("model")
|
||||
context_window_tokens = self._coerce_positive_int(
|
||||
usage.get("context_window_tokens")
|
||||
)
|
||||
if model_name:
|
||||
self._session_usage.model = model_name
|
||||
self._session_usage.context_window_tokens = context_window_tokens
|
||||
|
||||
if not usage.get("has_usage"):
|
||||
if is_current_request:
|
||||
self._session_usage.last_input_usage_available = False
|
||||
self._session_usage.last_input_tokens = None
|
||||
self._session_usage.last_output_tokens = None
|
||||
self._session_usage.last_total_tokens = None
|
||||
self._session_usage.last_context_usage_ratio = None
|
||||
self._session_usage.last_cache_usage_available = False
|
||||
self._session_usage.last_cache_read_input_tokens = 0
|
||||
self._session_usage.last_cache_write_input_tokens = 0
|
||||
self._session_usage.last_uncached_input_tokens = 0
|
||||
self._session_usage.last_cache_hit_ratio = None
|
||||
return
|
||||
|
||||
input_usage_available = usage.get("input_usage_available") is True
|
||||
input_tokens = self._coerce_int(usage.get("input_tokens")) or 0
|
||||
output_tokens = self._coerce_int(usage.get("output_tokens")) or 0
|
||||
total_tokens = self._coerce_int(usage.get("total_tokens"))
|
||||
if total_tokens is None:
|
||||
total_tokens = input_tokens + output_tokens
|
||||
|
||||
self._session_usage.last_input_tokens = input_tokens
|
||||
self._session_usage.last_output_tokens = output_tokens
|
||||
self._session_usage.last_total_tokens = total_tokens
|
||||
self._session_usage.last_context_usage_ratio = usage.get("context_usage_ratio")
|
||||
cache_usage_available = bool(usage.get("cache_usage_available"))
|
||||
cache_read_input_tokens = self._coerce_int(
|
||||
usage.get("cache_read_input_tokens")
|
||||
@@ -631,11 +718,6 @@ class MoviePilotAgent:
|
||||
input_tokens - cache_read_input_tokens - cache_write_input_tokens,
|
||||
0,
|
||||
)
|
||||
self._session_usage.last_cache_usage_available = cache_usage_available
|
||||
self._session_usage.last_cache_read_input_tokens = cache_read_input_tokens
|
||||
self._session_usage.last_cache_write_input_tokens = cache_write_input_tokens
|
||||
self._session_usage.last_uncached_input_tokens = uncached_input_tokens
|
||||
self._session_usage.last_cache_hit_ratio = usage.get("cache_hit_ratio")
|
||||
self._session_usage.total_input_tokens += input_tokens
|
||||
self._session_usage.total_output_tokens += output_tokens
|
||||
self._session_usage.total_tokens += total_tokens
|
||||
@@ -644,10 +726,119 @@ class MoviePilotAgent:
|
||||
self._session_usage.total_uncached_input_tokens += uncached_input_tokens
|
||||
self._session_usage.cache_usage_available |= cache_usage_available
|
||||
|
||||
if not is_current_request:
|
||||
return
|
||||
|
||||
self._session_usage.last_input_usage_available = input_usage_available
|
||||
self._session_usage.last_input_tokens = (
|
||||
input_tokens if input_usage_available else None
|
||||
)
|
||||
self._session_usage.last_output_tokens = output_tokens
|
||||
self._session_usage.last_total_tokens = total_tokens
|
||||
self._session_usage.last_context_usage_ratio = usage.get("context_usage_ratio")
|
||||
self._session_usage.last_cache_usage_available = cache_usage_available
|
||||
self._session_usage.last_cache_read_input_tokens = cache_read_input_tokens
|
||||
self._session_usage.last_cache_write_input_tokens = cache_write_input_tokens
|
||||
self._session_usage.last_uncached_input_tokens = uncached_input_tokens
|
||||
self._session_usage.last_cache_hit_ratio = usage.get("cache_hit_ratio")
|
||||
|
||||
estimated_input_tokens = self._coerce_int(
|
||||
usage.get("estimated_input_tokens")
|
||||
)
|
||||
if (
|
||||
usage.get("request_budget_recorded") is True
|
||||
and input_usage_available
|
||||
and estimated_input_tokens is not None
|
||||
and estimated_input_tokens
|
||||
== self._session_usage.last_estimated_input_tokens
|
||||
):
|
||||
estimate_error_tokens = input_tokens - estimated_input_tokens
|
||||
self._session_usage.last_actual_input_tokens = input_tokens
|
||||
self._session_usage.last_estimate_error_tokens = estimate_error_tokens
|
||||
self._session_usage.last_estimate_error_ratio = (
|
||||
estimate_error_tokens / input_tokens if input_tokens else None
|
||||
)
|
||||
|
||||
def _record_request_budget(self, budget: dict[str, Any]) -> None:
|
||||
"""保存最终请求的脱敏估算,并清除不属于本轮的旧校准结果。"""
|
||||
if not budget:
|
||||
return
|
||||
request_sequence = self._coerce_int(budget.get("request_sequence"))
|
||||
if "request_sequence" in budget and request_sequence is None:
|
||||
return
|
||||
request_sequence = request_sequence or 0
|
||||
if request_sequence < self._session_usage.last_request_sequence:
|
||||
return
|
||||
self._session_usage.last_request_sequence = request_sequence
|
||||
estimate_available = bool(budget.get("has_estimate"))
|
||||
self._session_usage.last_request_estimate_available = estimate_available
|
||||
self._session_usage.last_estimated_input_tokens = self._coerce_int(
|
||||
budget.get("estimated_input_tokens")
|
||||
)
|
||||
self._session_usage.last_estimated_message_tokens = self._coerce_int(
|
||||
budget.get("message_tokens")
|
||||
)
|
||||
self._session_usage.last_estimated_system_tokens = self._coerce_int(
|
||||
budget.get("system_tokens")
|
||||
)
|
||||
self._session_usage.last_estimated_tool_tokens = self._coerce_int(
|
||||
budget.get("tool_tokens")
|
||||
)
|
||||
self._session_usage.last_estimated_multimodal_tokens = self._coerce_int(
|
||||
budget.get("multimodal_tokens")
|
||||
)
|
||||
self._session_usage.last_estimated_input_ratio = budget.get(
|
||||
"estimated_input_ratio"
|
||||
)
|
||||
self._session_usage.last_estimated_remaining_input_tokens = self._coerce_int(
|
||||
budget.get("estimated_remaining_input_tokens")
|
||||
)
|
||||
self._session_usage.last_estimated_over_input_limit = budget.get(
|
||||
"estimated_over_input_limit"
|
||||
)
|
||||
self._session_usage.last_message_count = self._coerce_int(
|
||||
budget.get("message_count")
|
||||
) or 0
|
||||
self._session_usage.last_tool_count = self._coerce_int(
|
||||
budget.get("tool_count")
|
||||
) or 0
|
||||
self._session_usage.last_image_count = self._coerce_int(
|
||||
budget.get("image_count")
|
||||
) or 0
|
||||
self._session_usage.last_unknown_multimodal_count = self._coerce_int(
|
||||
budget.get("unknown_multimodal_count")
|
||||
) or 0
|
||||
self._session_usage.model_max_output_tokens = self._coerce_int(
|
||||
budget.get("model_max_output_tokens")
|
||||
)
|
||||
self._session_usage.configured_output_limit_tokens = self._coerce_int(
|
||||
budget.get("configured_output_limit_tokens")
|
||||
)
|
||||
has_model_snapshot = "model" in budget
|
||||
model_name = budget.get("model")
|
||||
context_window_tokens = self._coerce_positive_int(
|
||||
budget.get("context_window_tokens")
|
||||
)
|
||||
# 模型标识与窗口属于同一次最终请求,必须一起替换,不能拼接两轮状态。
|
||||
if has_model_snapshot:
|
||||
self._session_usage.model = model_name if model_name else None
|
||||
self._session_usage.context_window_tokens = context_window_tokens
|
||||
elif estimate_available:
|
||||
self._session_usage.context_window_tokens = context_window_tokens
|
||||
self._session_usage.last_actual_input_tokens = None
|
||||
self._session_usage.last_estimate_error_tokens = None
|
||||
self._session_usage.last_estimate_error_ratio = None
|
||||
|
||||
def get_session_status(self) -> dict[str, Any]:
|
||||
if not self._session_usage.model:
|
||||
if (
|
||||
not self._session_usage.model
|
||||
and self._session_usage.last_request_sequence == 0
|
||||
):
|
||||
self._session_usage.model = settings.LLM_MODEL
|
||||
if not self._session_usage.context_window_tokens:
|
||||
if (
|
||||
not self._session_usage.context_window_tokens
|
||||
and self._session_usage.last_request_sequence == 0
|
||||
):
|
||||
self._session_usage.context_window_tokens = (
|
||||
settings.LLM_MAX_CONTEXT_TOKENS * 1000
|
||||
if settings.LLM_MAX_CONTEXT_TOKENS
|
||||
@@ -1777,8 +1968,6 @@ class MoviePilotAgent:
|
||||
PatchToolCallsMiddleware(),
|
||||
# 子代理委派
|
||||
*subagent_middlewares,
|
||||
# 用量统计
|
||||
UsageMiddleware(on_usage=self._record_usage),
|
||||
]
|
||||
|
||||
# 工具选择
|
||||
@@ -1797,6 +1986,15 @@ class MoviePilotAgent:
|
||||
)
|
||||
)
|
||||
|
||||
# 预算观察器必须位于最内层,才能看到动态 system 和最终筛选后的工具。
|
||||
middlewares.append(
|
||||
UsageMiddleware(
|
||||
on_usage=self._record_usage,
|
||||
on_request_budget=self._record_request_budget,
|
||||
next_request_sequence=self._next_request_sequence,
|
||||
)
|
||||
)
|
||||
|
||||
agent = create_agent(
|
||||
model=agent_model,
|
||||
tools=[
|
||||
@@ -2267,22 +2465,14 @@ class AgentManager:
|
||||
if agent:
|
||||
status = agent.get_session_status()
|
||||
else:
|
||||
status = {
|
||||
"session_id": session_id,
|
||||
"model": settings.LLM_MODEL,
|
||||
"context_window_tokens": settings.LLM_MAX_CONTEXT_TOKENS * 1000
|
||||
if settings.LLM_MAX_CONTEXT_TOKENS
|
||||
else None,
|
||||
"last_input_tokens": 0,
|
||||
"last_output_tokens": 0,
|
||||
"last_total_tokens": 0,
|
||||
"last_context_usage_ratio": None,
|
||||
"total_input_tokens": 0,
|
||||
"total_output_tokens": 0,
|
||||
"total_tokens": 0,
|
||||
"model_call_count": 0,
|
||||
"last_updated_at": None,
|
||||
}
|
||||
status = _SessionUsageSnapshot(
|
||||
model=settings.LLM_MODEL,
|
||||
context_window_tokens=(
|
||||
settings.LLM_MAX_CONTEXT_TOKENS * 1000
|
||||
if settings.LLM_MAX_CONTEXT_TOKENS
|
||||
else None
|
||||
),
|
||||
).to_dict(session_id)
|
||||
|
||||
queue = self._session_queues.get(session_id)
|
||||
status["pending_messages"] = queue.qsize() if queue else 0
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -1386,7 +1386,11 @@ class MessageChain(ChainBase):
|
||||
last_input_tokens = status.get("last_input_tokens")
|
||||
if context_window_tokens and status.get("model_call_count"):
|
||||
context_ratio = status.get("last_context_usage_ratio")
|
||||
if context_ratio is None and last_input_tokens is not None:
|
||||
if (
|
||||
context_ratio is None
|
||||
and status.get("last_input_usage_available") is True
|
||||
and last_input_tokens is not None
|
||||
):
|
||||
context_ratio = last_input_tokens / context_window_tokens
|
||||
context_usage_text = (
|
||||
f"{cls._format_token_count(last_input_tokens)} / "
|
||||
@@ -1405,17 +1409,45 @@ class MessageChain(ChainBase):
|
||||
f"当前模型: {status.get('model') or '未知'}",
|
||||
f"上下文窗口: {cls._format_token_count(context_window_tokens)} tokens",
|
||||
f"最近一次上下文占用: {context_usage_text}",
|
||||
f"最近一次 tokens: 输入 {cls._format_token_count(status.get('last_input_tokens'))} / 输出 {cls._format_token_count(status.get('last_output_tokens'))} / 总计 {cls._format_token_count(status.get('last_total_tokens'))}",
|
||||
f"当前会话累计 tokens: 输入 {cls._format_token_count(status.get('total_input_tokens'))} / 输出 {cls._format_token_count(status.get('total_output_tokens'))} / 总计 {cls._format_token_count(status.get('total_tokens'))}",
|
||||
f"模型调用次数: {status.get('model_call_count', 0)}",
|
||||
f"排队消息数: {status.get('pending_messages', 0)}",
|
||||
f"最后更新: {status.get('last_updated_at') or '暂无'}",
|
||||
]
|
||||
if status.get("cache_usage_available"):
|
||||
if status.get("last_request_estimate_available"):
|
||||
estimated_tokens = status.get("last_estimated_input_tokens")
|
||||
estimated_ratio = status.get("last_estimated_input_ratio")
|
||||
estimate_text = (
|
||||
f"{cls._format_token_count(estimated_tokens)} / "
|
||||
f"{cls._format_token_count(context_window_tokens)}"
|
||||
)
|
||||
if estimated_ratio is not None:
|
||||
estimate_text += f" ({estimated_ratio * 100:.2f}%)"
|
||||
if status.get("last_estimated_over_input_limit"):
|
||||
estimate_text += ",估算已超输入上限"
|
||||
lines.extend(
|
||||
[
|
||||
f"最终请求估算: {estimate_text}",
|
||||
"估算组成: "
|
||||
f"消息 {cls._format_token_count(status.get('last_estimated_message_tokens'))} / "
|
||||
f"系统 {cls._format_token_count(status.get('last_estimated_system_tokens'))} / "
|
||||
f"工具 {cls._format_token_count(status.get('last_estimated_tool_tokens'))} / "
|
||||
f"其中图片固定成本 {cls._format_token_count(status.get('last_estimated_multimodal_tokens'))}",
|
||||
]
|
||||
)
|
||||
actual_input_tokens = status.get("last_actual_input_tokens")
|
||||
estimate_error_tokens = status.get("last_estimate_error_tokens")
|
||||
if actual_input_tokens is not None and estimate_error_tokens is not None:
|
||||
estimate_error_ratio = status.get("last_estimate_error_ratio")
|
||||
error_text = (
|
||||
f"实际 {cls._format_token_count(actual_input_tokens)} / "
|
||||
f"误差 {estimate_error_tokens:+,}"
|
||||
)
|
||||
if estimate_error_ratio is not None:
|
||||
error_text += f" ({estimate_error_ratio:+.2%})"
|
||||
lines.append(f"估算校准: {error_text}")
|
||||
lines.append(
|
||||
f"最近一次 tokens: 输入 {cls._format_token_count(status.get('last_input_tokens'))} / 输出 {cls._format_token_count(status.get('last_output_tokens'))} / 总计 {cls._format_token_count(status.get('last_total_tokens'))}"
|
||||
)
|
||||
if status.get("last_cache_usage_available"):
|
||||
last_cache_ratio = status.get("last_cache_hit_ratio")
|
||||
total_cache_ratio = status.get("total_cache_hit_ratio")
|
||||
lines.insert(
|
||||
6,
|
||||
lines.append(
|
||||
"最近一次缓存: "
|
||||
f"命中 {cls._format_token_count(status.get('last_cache_read_input_tokens'))} / "
|
||||
f"写入 {cls._format_token_count(status.get('last_cache_write_input_tokens'))} / "
|
||||
@@ -1426,8 +1458,9 @@ class MessageChain(ChainBase):
|
||||
else ""
|
||||
),
|
||||
)
|
||||
lines.insert(
|
||||
8,
|
||||
if status.get("cache_usage_available"):
|
||||
total_cache_ratio = status.get("total_cache_hit_ratio")
|
||||
lines.append(
|
||||
"当前会话累计缓存: "
|
||||
f"命中 {cls._format_token_count(status.get('total_cache_read_input_tokens'))} / "
|
||||
f"写入 {cls._format_token_count(status.get('total_cache_write_input_tokens'))} / "
|
||||
@@ -1438,6 +1471,14 @@ class MessageChain(ChainBase):
|
||||
else ""
|
||||
),
|
||||
)
|
||||
lines.extend(
|
||||
[
|
||||
f"当前会话累计 tokens: 输入 {cls._format_token_count(status.get('total_input_tokens'))} / 输出 {cls._format_token_count(status.get('total_output_tokens'))} / 总计 {cls._format_token_count(status.get('total_tokens'))}",
|
||||
f"模型调用次数: {status.get('model_call_count', 0)}",
|
||||
f"排队消息数: {status.get('pending_messages', 0)}",
|
||||
f"最后更新: {status.get('last_updated_at') or '暂无'}",
|
||||
]
|
||||
)
|
||||
return "\n".join(lines)
|
||||
|
||||
def remote_session_status(
|
||||
|
||||
@@ -491,8 +491,10 @@ async def test_graph_keeps_mcp_first_winner_and_catalogs_all_collisions(
|
||||
activity_tool,
|
||||
subagent_task_tool,
|
||||
]
|
||||
assert captured["middlewares"][-2].name == "selector"
|
||||
else:
|
||||
assert "selection_tools" not in captured
|
||||
assert captured["middlewares"][-1].name == "usage"
|
||||
policy_middleware = next(
|
||||
middleware
|
||||
for middleware in captured["middlewares"]
|
||||
|
||||
990
tests/test_agent_request_budget.py
Normal file
990
tests/test_agent_request_budget.py
Normal file
@@ -0,0 +1,990 @@
|
||||
"""Agent 最终模型请求预算测试。"""
|
||||
|
||||
import asyncio
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
|
||||
from langchain.agents import create_agent
|
||||
from langchain.agents.middleware import AgentMiddleware, SummarizationMiddleware
|
||||
from langchain.agents.middleware.types import ModelRequest, ModelResponse
|
||||
from langchain_core.language_models.fake_chat_models import FakeMessagesListChatModel
|
||||
from langchain_core.messages import AIMessage, HumanMessage, SystemMessage
|
||||
from langchain_core.tools import tool
|
||||
|
||||
from app.agent import MoviePilotAgent
|
||||
from app.agent.middleware.tool_selection import ToolSelectorMiddleware
|
||||
from app.agent.middleware.usage import UsageMiddleware
|
||||
from app.chain.message import MessageChain
|
||||
|
||||
|
||||
class _ToolBindingFakeModel(FakeMessagesListChatModel):
|
||||
"""记录最终绑定工具,同时保留固定响应行为。"""
|
||||
|
||||
bound_tool_names: list[str] = []
|
||||
|
||||
def bind_tools(self, tools, **kwargs):
|
||||
"""记录 LangChain 最终交给模型的工具集合。"""
|
||||
self.bound_tool_names = [
|
||||
item.get("function", {}).get("name") or item.get("name")
|
||||
if isinstance(item, dict)
|
||||
else item.name
|
||||
for item in tools
|
||||
]
|
||||
return self
|
||||
|
||||
|
||||
class _DynamicSystemMiddleware(AgentMiddleware):
|
||||
"""模拟 MoviePilot 运行时动态追加系统上下文。"""
|
||||
|
||||
async def awrap_model_call(self, request, handler):
|
||||
current = request.system_message.content if request.system_message else ""
|
||||
return await handler(
|
||||
request.override(
|
||||
system_message=SystemMessage(
|
||||
content=f"{current}\n{'动态系统上下文 ' * 200}"
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _request(
|
||||
*,
|
||||
messages=None,
|
||||
system_message=None,
|
||||
tools=None,
|
||||
max_input_tokens=4096,
|
||||
max_output_tokens=512,
|
||||
model_settings=None,
|
||||
) -> ModelRequest:
|
||||
"""构造带模型窗口和最终请求组成的测试请求。"""
|
||||
return ModelRequest(
|
||||
model=SimpleNamespace(
|
||||
model="small-model",
|
||||
profile={
|
||||
"max_input_tokens": max_input_tokens,
|
||||
"max_output_tokens": max_output_tokens,
|
||||
},
|
||||
),
|
||||
messages=list(messages or []),
|
||||
system_message=system_message,
|
||||
tools=list(tools or []),
|
||||
state={},
|
||||
runtime=None,
|
||||
model_settings=model_settings,
|
||||
)
|
||||
|
||||
|
||||
def test_final_request_can_exceed_window_before_message_fraction_triggers():
|
||||
"""动态系统提示词和工具定义可能在消息摘要阈值前耗尽输入窗口。"""
|
||||
messages = [HumanMessage(content="用户上下文 " * 1000)]
|
||||
model = SimpleNamespace(
|
||||
_llm_type="test-chat",
|
||||
profile={"max_input_tokens": 4096},
|
||||
)
|
||||
summarizer = SummarizationMiddleware(
|
||||
model=model,
|
||||
trigger=("fraction", 0.85),
|
||||
)
|
||||
message_tokens = summarizer.token_counter(messages)
|
||||
assert message_tokens < 4096 * 0.85
|
||||
assert not summarizer._should_summarize(messages, message_tokens)
|
||||
|
||||
request = _request(
|
||||
messages=messages,
|
||||
system_message=SystemMessage(content="动态系统上下文 " * 1000),
|
||||
tools=[
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "large_schema_tool",
|
||||
"description": "工具业务说明 " * 800,
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"value": {
|
||||
"type": "string",
|
||||
"description": "参数约束 " * 400,
|
||||
}
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
snapshot = UsageMiddleware.estimate_request(request)
|
||||
|
||||
assert snapshot["has_estimate"]
|
||||
assert snapshot["message_tokens"] < 4096 * 0.85
|
||||
assert snapshot["system_tokens"] > 0
|
||||
assert snapshot["tool_tokens"] > 0
|
||||
assert snapshot["estimated_input_tokens"] > 4096
|
||||
assert snapshot["estimated_over_input_limit"] is True
|
||||
assert snapshot["model_max_output_tokens"] == 512
|
||||
assert snapshot["configured_output_limit_tokens"] is None
|
||||
assert "output_headroom_tokens" not in snapshot
|
||||
|
||||
|
||||
def test_request_budget_marks_same_request_over_limit_after_switch_to_small_model():
|
||||
"""相同请求切换到小窗口模型后,应仅改变预算判断而不改变估算输入。"""
|
||||
messages = [HumanMessage(content="x" * 12000)]
|
||||
|
||||
large_snapshot = UsageMiddleware.estimate_request(
|
||||
_request(messages=messages, max_input_tokens=128000)
|
||||
)
|
||||
small_snapshot = UsageMiddleware.estimate_request(
|
||||
_request(messages=messages, max_input_tokens=2048)
|
||||
)
|
||||
|
||||
assert large_snapshot["estimated_input_tokens"] == small_snapshot["estimated_input_tokens"]
|
||||
assert large_snapshot["estimated_over_input_limit"] is False
|
||||
assert small_snapshot["estimated_over_input_limit"] is True
|
||||
|
||||
|
||||
def test_request_budget_reads_only_explicit_per_call_output_limit():
|
||||
"""只有最终请求显式配置的输出上限才可视为单次调用限制。"""
|
||||
request = _request(
|
||||
messages=[HumanMessage(content="hello")],
|
||||
max_output_tokens=8192,
|
||||
model_settings={"max_completion_tokens": 1024},
|
||||
)
|
||||
|
||||
snapshot = UsageMiddleware.estimate_request(request)
|
||||
|
||||
assert snapshot["model_max_output_tokens"] == 8192
|
||||
assert snapshot["configured_output_limit_tokens"] == 1024
|
||||
assert snapshot["estimated_input_tokens"] < snapshot["context_window_tokens"]
|
||||
|
||||
|
||||
def test_request_budget_rejects_non_integer_token_limits():
|
||||
"""近似观察不得把 bool、浮点或字符串误报为有效 token 上限。"""
|
||||
for invalid_value in (True, False, 1.5, 0, -1, "1024"):
|
||||
request = _request(
|
||||
messages=[HumanMessage(content="hello")],
|
||||
max_input_tokens=invalid_value,
|
||||
max_output_tokens=invalid_value,
|
||||
model_settings={"max_completion_tokens": invalid_value},
|
||||
)
|
||||
|
||||
snapshot = UsageMiddleware.estimate_request(request)
|
||||
|
||||
assert snapshot["context_window_tokens"] is None
|
||||
assert snapshot["model_max_output_tokens"] is None
|
||||
assert snapshot["configured_output_limit_tokens"] is None
|
||||
assert snapshot["estimated_input_ratio"] is None
|
||||
assert snapshot["estimated_over_input_limit"] is None
|
||||
|
||||
|
||||
def test_request_budget_uses_next_valid_output_limit_alias():
|
||||
"""高优先字段无效时,应继续读取同一请求中的有效兼容字段。"""
|
||||
request = _request(
|
||||
messages=[HumanMessage(content="hello")],
|
||||
model_settings={"max_completion_tokens": True, "max_tokens": 1024},
|
||||
)
|
||||
|
||||
snapshot = UsageMiddleware.estimate_request(request)
|
||||
|
||||
assert snapshot["configured_output_limit_tokens"] == 1024
|
||||
|
||||
|
||||
def test_request_budget_counts_multimodal_input_without_storing_content():
|
||||
"""图片按固定成本计入估算,快照不得保留请求正文或工具定义。"""
|
||||
secret_marker = "REQUEST_BUDGET_SECRET_MARKER"
|
||||
request = _request(
|
||||
messages=[
|
||||
HumanMessage(
|
||||
content=[
|
||||
{"type": "text", "text": secret_marker},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": "data:image/png;base64,hidden"},
|
||||
},
|
||||
{"type": "file", "file_id": secret_marker},
|
||||
]
|
||||
)
|
||||
],
|
||||
system_message=SystemMessage(content=secret_marker),
|
||||
tools=[
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "secret_tool",
|
||||
"description": secret_marker,
|
||||
"parameters": {"type": "object", "properties": {}},
|
||||
},
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
snapshot = UsageMiddleware.estimate_request(request)
|
||||
|
||||
assert snapshot["image_count"] == 1
|
||||
assert snapshot["unknown_multimodal_count"] == 1
|
||||
assert snapshot["multimodal_tokens"] == 85
|
||||
assert snapshot["estimated_input_tokens"] == (
|
||||
snapshot["message_tokens"]
|
||||
+ snapshot["system_tokens"]
|
||||
+ snapshot["tool_tokens"]
|
||||
)
|
||||
assert snapshot["model"] == "small-model"
|
||||
assert secret_marker not in repr(snapshot)
|
||||
assert all(
|
||||
value is None or isinstance(value, (bool, int, float))
|
||||
for key, value in snapshot.items()
|
||||
if key != "model"
|
||||
)
|
||||
|
||||
|
||||
def test_request_budget_counts_each_image_cost_exactly_once():
|
||||
"""LangChain 的消息估算已包含图片固定成本,汇总预算不得重复相加。"""
|
||||
text_only = UsageMiddleware.estimate_request(
|
||||
_request(
|
||||
messages=[HumanMessage(content=[{"type": "text", "text": "hello"}])]
|
||||
)
|
||||
)
|
||||
with_image = UsageMiddleware.estimate_request(
|
||||
_request(
|
||||
messages=[
|
||||
HumanMessage(
|
||||
content=[
|
||||
{"type": "text", "text": "hello"},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": "data:image/png;base64,hidden"},
|
||||
},
|
||||
]
|
||||
)
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
assert with_image["multimodal_tokens"] == 85
|
||||
assert (
|
||||
with_image["estimated_input_tokens"] - text_only["estimated_input_tokens"]
|
||||
== 85
|
||||
)
|
||||
|
||||
|
||||
def test_request_budget_callback_failure_does_not_block_model_call():
|
||||
"""预算观察失败不得改变模型请求和响应。"""
|
||||
request = _request(messages=[HumanMessage(content="hello")])
|
||||
response = ModelResponse(result=[AIMessage(content="ok")])
|
||||
|
||||
def _raise_callback(_snapshot):
|
||||
raise RuntimeError("observer unavailable")
|
||||
|
||||
middleware = UsageMiddleware(on_request_budget=_raise_callback)
|
||||
handled = []
|
||||
|
||||
async def _handler(received: ModelRequest):
|
||||
handled.append(received)
|
||||
return response
|
||||
|
||||
result = asyncio.run(middleware.awrap_model_call(request, _handler))
|
||||
|
||||
assert result is response
|
||||
assert handled == [request]
|
||||
|
||||
|
||||
def test_request_budget_callback_failure_clears_previous_estimate_state():
|
||||
"""预算回调失败时,本轮 usage 不得与上一轮估算拼接。"""
|
||||
agent = MoviePilotAgent(session_id="request-budget", user_id="user-1")
|
||||
callback_count = 0
|
||||
|
||||
def _record_then_fail(snapshot):
|
||||
nonlocal callback_count
|
||||
callback_count += 1
|
||||
if callback_count == 2:
|
||||
raise RuntimeError("observer unavailable")
|
||||
agent._record_request_budget(snapshot)
|
||||
|
||||
middleware = UsageMiddleware(
|
||||
on_request_budget=_record_then_fail,
|
||||
on_usage=agent._record_usage,
|
||||
)
|
||||
responses = iter(
|
||||
[
|
||||
ModelResponse(
|
||||
result=[
|
||||
AIMessage(
|
||||
content="first",
|
||||
usage_metadata={
|
||||
"input_tokens": 10,
|
||||
"output_tokens": 1,
|
||||
"total_tokens": 11,
|
||||
},
|
||||
)
|
||||
]
|
||||
),
|
||||
ModelResponse(
|
||||
result=[
|
||||
AIMessage(
|
||||
content="second",
|
||||
usage_metadata={
|
||||
"input_tokens": 50,
|
||||
"output_tokens": 2,
|
||||
"total_tokens": 52,
|
||||
},
|
||||
)
|
||||
]
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
async def _handler(_request):
|
||||
return next(responses)
|
||||
|
||||
request = _request(messages=[HumanMessage(content="hello")])
|
||||
asyncio.run(middleware.awrap_model_call(request, _handler))
|
||||
assert agent.get_session_status()["last_request_estimate_available"] is True
|
||||
|
||||
asyncio.run(middleware.awrap_model_call(request, _handler))
|
||||
status = agent.get_session_status()
|
||||
|
||||
assert status["last_request_sequence"] == 2
|
||||
assert status["last_request_estimate_available"] is False
|
||||
assert status["last_estimated_input_tokens"] is None
|
||||
assert status["last_actual_input_tokens"] is None
|
||||
assert status["last_estimate_error_tokens"] is None
|
||||
assert status["model"] == "small-model"
|
||||
assert status["context_window_tokens"] == 4096
|
||||
assert status["last_input_tokens"] == 50
|
||||
|
||||
|
||||
def test_request_budget_estimator_failure_reports_empty_snapshot_and_calls_model():
|
||||
"""估算器异常应清除旧观测状态,并继续原模型调用。"""
|
||||
budgets = []
|
||||
middleware = UsageMiddleware(on_request_budget=budgets.append)
|
||||
request = _request(messages=[HumanMessage(content="hello")])
|
||||
response = ModelResponse(result=[AIMessage(content="ok")])
|
||||
|
||||
async def _handler(_request):
|
||||
return response
|
||||
|
||||
with patch.object(
|
||||
UsageMiddleware,
|
||||
"estimate_request",
|
||||
side_effect=RuntimeError("request content must not reach the snapshot"),
|
||||
):
|
||||
result = asyncio.run(middleware.awrap_model_call(request, _handler))
|
||||
|
||||
assert result is response
|
||||
assert budgets == [
|
||||
{
|
||||
"request_sequence": 1,
|
||||
"has_estimate": False,
|
||||
"model": "small-model",
|
||||
"context_window_tokens": 4096,
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
def test_request_budget_metadata_failure_still_calls_model():
|
||||
"""模型元数据属性异常不得让预算观察器阻断真实模型调用。"""
|
||||
|
||||
class _BrokenMetadataModel:
|
||||
@property
|
||||
def model(self):
|
||||
raise RuntimeError("model metadata unavailable")
|
||||
|
||||
@property
|
||||
def model_name(self):
|
||||
raise RuntimeError("model metadata unavailable")
|
||||
|
||||
@property
|
||||
def model_id(self):
|
||||
raise RuntimeError("model metadata unavailable")
|
||||
|
||||
@property
|
||||
def profile(self):
|
||||
raise RuntimeError("profile metadata unavailable")
|
||||
|
||||
budgets = []
|
||||
middleware = UsageMiddleware(on_request_budget=budgets.append)
|
||||
request = ModelRequest(
|
||||
model=_BrokenMetadataModel(),
|
||||
messages=[HumanMessage(content="hello")],
|
||||
tools=[],
|
||||
state={},
|
||||
runtime=None,
|
||||
)
|
||||
response = ModelResponse(result=[AIMessage(content="ok")])
|
||||
handled = []
|
||||
|
||||
async def _handler(received):
|
||||
handled.append(received)
|
||||
return response
|
||||
|
||||
with patch.object(
|
||||
UsageMiddleware,
|
||||
"estimate_request",
|
||||
side_effect=RuntimeError("estimate unavailable"),
|
||||
):
|
||||
result = asyncio.run(middleware.awrap_model_call(request, _handler))
|
||||
|
||||
assert result is response
|
||||
assert handled == [request]
|
||||
assert budgets == [
|
||||
{
|
||||
"request_sequence": 1,
|
||||
"has_estimate": False,
|
||||
"model": None,
|
||||
"context_window_tokens": None,
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
def test_request_sequence_callback_failure_still_calls_model():
|
||||
"""会话序号分配异常时应放弃最近快照竞争,并继续真实模型调用。"""
|
||||
budgets = []
|
||||
|
||||
def _raise_sequence():
|
||||
raise RuntimeError("sequence unavailable")
|
||||
|
||||
middleware = UsageMiddleware(
|
||||
on_request_budget=budgets.append,
|
||||
next_request_sequence=_raise_sequence,
|
||||
)
|
||||
request = _request(messages=[HumanMessage(content="hello")])
|
||||
response = ModelResponse(result=[AIMessage(content="ok")])
|
||||
handled = []
|
||||
|
||||
async def _handler(received):
|
||||
handled.append(received)
|
||||
return response
|
||||
|
||||
result = asyncio.run(middleware.awrap_model_call(request, _handler))
|
||||
|
||||
assert result is response
|
||||
assert handled == [request]
|
||||
assert budgets[0]["request_sequence"] is None
|
||||
|
||||
|
||||
def test_request_budget_and_actual_usage_share_one_request_sequence():
|
||||
"""真实 usage 只能校准同一次成功模型调用产生的估算。"""
|
||||
budgets = []
|
||||
usages = []
|
||||
middleware = UsageMiddleware(
|
||||
on_request_budget=budgets.append,
|
||||
on_usage=usages.append,
|
||||
)
|
||||
request = _request(messages=[HumanMessage(content="hello")])
|
||||
response = ModelResponse(
|
||||
result=[
|
||||
AIMessage(
|
||||
content="ok",
|
||||
usage_metadata={
|
||||
"input_tokens": 23,
|
||||
"output_tokens": 4,
|
||||
"total_tokens": 27,
|
||||
},
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
async def _handler(_request):
|
||||
return response
|
||||
|
||||
asyncio.run(middleware.awrap_model_call(request, _handler))
|
||||
|
||||
assert budgets[0]["request_sequence"] == 1
|
||||
assert usages[0]["request_sequence"] == 1
|
||||
assert usages[0]["request_budget_recorded"] is True
|
||||
assert usages[0]["input_usage_available"] is True
|
||||
assert usages[0]["estimated_input_tokens"] == budgets[0]["estimated_input_tokens"]
|
||||
|
||||
|
||||
def test_partial_usage_without_input_does_not_calibrate_request_estimate():
|
||||
"""仅有输出 usage 时,不得把缺失的真实输入误报为零。"""
|
||||
agent = MoviePilotAgent(session_id="request-budget", user_id="user-1")
|
||||
middleware = UsageMiddleware(
|
||||
on_request_budget=agent._record_request_budget,
|
||||
on_usage=agent._record_usage,
|
||||
)
|
||||
request = _request(messages=[HumanMessage(content="hello")])
|
||||
response = ModelResponse(
|
||||
result=[
|
||||
AIMessage(
|
||||
content="ok",
|
||||
response_metadata={
|
||||
"token_usage": {
|
||||
"completion_tokens": 7,
|
||||
"total_tokens": 7,
|
||||
}
|
||||
},
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
async def _handler(_request):
|
||||
return response
|
||||
|
||||
asyncio.run(middleware.awrap_model_call(request, _handler))
|
||||
status = agent.get_session_status()
|
||||
|
||||
assert status["last_request_estimate_available"] is True
|
||||
assert status["last_input_usage_available"] is False
|
||||
assert status["last_input_tokens"] is None
|
||||
assert status["last_output_tokens"] == 7
|
||||
assert status["last_actual_input_tokens"] is None
|
||||
assert status["last_estimate_error_tokens"] is None
|
||||
assert status["last_estimate_error_ratio"] is None
|
||||
assert status["last_context_usage_ratio"] is None
|
||||
status_text = MessageChain._format_session_status_text(status)
|
||||
assert "最近一次上下文占用: 未知 / 4,096" in status_text
|
||||
assert "最近一次 tokens: 输入 未知 / 输出 7 / 总计 7" in status_text
|
||||
|
||||
|
||||
def test_missing_usage_clears_previous_last_call_values():
|
||||
"""本轮没有 usage 时,最近一次状态不得保留上一轮实际值。"""
|
||||
agent = MoviePilotAgent(session_id="request-budget", user_id="user-1")
|
||||
agent._record_request_budget(
|
||||
{
|
||||
"request_sequence": 1,
|
||||
"has_estimate": True,
|
||||
"estimated_input_tokens": 10,
|
||||
"context_window_tokens": 1000,
|
||||
}
|
||||
)
|
||||
agent._record_usage(
|
||||
{
|
||||
"request_sequence": 1,
|
||||
"request_budget_recorded": True,
|
||||
"has_usage": True,
|
||||
"input_usage_available": True,
|
||||
"input_tokens": 12,
|
||||
"output_tokens": 3,
|
||||
"total_tokens": 15,
|
||||
"cache_usage_available": True,
|
||||
"cache_read_input_tokens": 4,
|
||||
"cache_write_input_tokens": 0,
|
||||
"uncached_input_tokens": 8,
|
||||
}
|
||||
)
|
||||
agent._record_request_budget(
|
||||
{
|
||||
"request_sequence": 2,
|
||||
"has_estimate": True,
|
||||
"estimated_input_tokens": 20,
|
||||
"context_window_tokens": 1000,
|
||||
}
|
||||
)
|
||||
agent._record_usage(
|
||||
{
|
||||
"request_sequence": 2,
|
||||
"request_budget_recorded": True,
|
||||
"has_usage": False,
|
||||
"input_usage_available": False,
|
||||
}
|
||||
)
|
||||
|
||||
status = agent.get_session_status()
|
||||
|
||||
assert status["last_request_sequence"] == 2
|
||||
assert status["last_request_estimate_available"] is True
|
||||
assert status["last_input_usage_available"] is False
|
||||
assert status["last_input_tokens"] is None
|
||||
assert status["last_output_tokens"] is None
|
||||
assert status["last_total_tokens"] is None
|
||||
assert status["last_context_usage_ratio"] is None
|
||||
assert status["last_cache_usage_available"] is False
|
||||
assert status["total_input_tokens"] == 12
|
||||
assert status["total_output_tokens"] == 3
|
||||
assert status["total_tokens"] == 15
|
||||
status_text = MessageChain._format_session_status_text(status)
|
||||
assert "最近一次 tokens: 输入 未知 / 输出 未知 / 总计 未知" in status_text
|
||||
assert "最近一次缓存:" not in status_text
|
||||
assert "当前会话累计缓存: 命中 4 / 写入 0 / 未命中 8" in status_text
|
||||
|
||||
|
||||
def test_out_of_order_responses_keep_latest_request_snapshot():
|
||||
"""较早请求晚返回时,只累计 usage,不得覆盖较新请求的最近状态。"""
|
||||
agent = MoviePilotAgent(session_id="request-budget", user_id="user-1")
|
||||
middleware = UsageMiddleware(
|
||||
on_request_budget=agent._record_request_budget,
|
||||
on_usage=agent._record_usage,
|
||||
)
|
||||
first_request = _request(
|
||||
messages=[HumanMessage(content="first")],
|
||||
max_input_tokens=1000,
|
||||
)
|
||||
first_request.model.model = "first-model"
|
||||
second_request = _request(
|
||||
messages=[HumanMessage(content="second")],
|
||||
max_input_tokens=2000,
|
||||
)
|
||||
second_request.model.model = "second-model"
|
||||
first_started = asyncio.Event()
|
||||
release_first = asyncio.Event()
|
||||
|
||||
async def _handler(request):
|
||||
if request.model.model == "first-model":
|
||||
first_started.set()
|
||||
await release_first.wait()
|
||||
return ModelResponse(
|
||||
result=[
|
||||
AIMessage(
|
||||
content="first",
|
||||
usage_metadata={
|
||||
"input_tokens": 10,
|
||||
"output_tokens": 1,
|
||||
"total_tokens": 11,
|
||||
},
|
||||
)
|
||||
]
|
||||
)
|
||||
return ModelResponse(
|
||||
result=[
|
||||
AIMessage(
|
||||
content="second",
|
||||
usage_metadata={
|
||||
"input_tokens": 50,
|
||||
"output_tokens": 2,
|
||||
"total_tokens": 52,
|
||||
},
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
async def _run_out_of_order():
|
||||
first_task = asyncio.create_task(
|
||||
middleware.awrap_model_call(first_request, _handler)
|
||||
)
|
||||
await first_started.wait()
|
||||
await middleware.awrap_model_call(second_request, _handler)
|
||||
release_first.set()
|
||||
await first_task
|
||||
|
||||
asyncio.run(_run_out_of_order())
|
||||
status = agent.get_session_status()
|
||||
|
||||
assert status["last_request_sequence"] == 2
|
||||
assert status["model"] == "second-model"
|
||||
assert status["context_window_tokens"] == 2000
|
||||
assert status["last_input_tokens"] == 50
|
||||
assert status["last_output_tokens"] == 2
|
||||
assert status["last_total_tokens"] == 52
|
||||
assert status["last_actual_input_tokens"] == 50
|
||||
assert status["total_input_tokens"] == 60
|
||||
assert status["total_output_tokens"] == 3
|
||||
assert status["total_tokens"] == 63
|
||||
assert status["model_call_count"] == 2
|
||||
|
||||
|
||||
def test_sequence_failure_cannot_overwrite_next_request_snapshot():
|
||||
"""无序号请求晚返回时只能累计 usage,不能覆盖后续有序请求。"""
|
||||
agent = MoviePilotAgent(session_id="request-budget", user_id="user-1")
|
||||
allocation_count = 0
|
||||
|
||||
def _next_sequence():
|
||||
nonlocal allocation_count
|
||||
allocation_count += 1
|
||||
if allocation_count == 1:
|
||||
raise RuntimeError("sequence unavailable")
|
||||
return agent._next_request_sequence()
|
||||
|
||||
middleware = UsageMiddleware(
|
||||
on_request_budget=agent._record_request_budget,
|
||||
on_usage=agent._record_usage,
|
||||
next_request_sequence=_next_sequence,
|
||||
)
|
||||
first_request = _request(
|
||||
messages=[HumanMessage(content="first")],
|
||||
max_input_tokens=1000,
|
||||
)
|
||||
first_request.model.model = "first-model"
|
||||
second_request = _request(
|
||||
messages=[HumanMessage(content="second")],
|
||||
max_input_tokens=2000,
|
||||
)
|
||||
second_request.model.model = "second-model"
|
||||
first_started = asyncio.Event()
|
||||
release_first = asyncio.Event()
|
||||
|
||||
async def _handler(request):
|
||||
if request.model.model == "first-model":
|
||||
first_started.set()
|
||||
await release_first.wait()
|
||||
input_tokens = 10
|
||||
else:
|
||||
input_tokens = 50
|
||||
return ModelResponse(
|
||||
result=[
|
||||
AIMessage(
|
||||
content=request.model.model,
|
||||
usage_metadata={
|
||||
"input_tokens": input_tokens,
|
||||
"output_tokens": 1,
|
||||
"total_tokens": input_tokens + 1,
|
||||
},
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
async def _run_out_of_order():
|
||||
first_task = asyncio.create_task(
|
||||
middleware.awrap_model_call(first_request, _handler)
|
||||
)
|
||||
await first_started.wait()
|
||||
await middleware.awrap_model_call(second_request, _handler)
|
||||
release_first.set()
|
||||
await first_task
|
||||
|
||||
asyncio.run(_run_out_of_order())
|
||||
status = agent.get_session_status()
|
||||
|
||||
assert status["last_request_sequence"] == 1
|
||||
assert status["model"] == "second-model"
|
||||
assert status["context_window_tokens"] == 2000
|
||||
assert status["last_input_tokens"] == 50
|
||||
assert status["last_actual_input_tokens"] == 50
|
||||
assert status["total_input_tokens"] == 60
|
||||
assert status["model_call_count"] == 2
|
||||
|
||||
|
||||
def test_failed_model_switch_keeps_model_and_window_from_same_request():
|
||||
"""新模型请求即使失败,状态也不得混合旧模型名称与新窗口。"""
|
||||
agent = MoviePilotAgent(session_id="request-budget", user_id="user-1")
|
||||
agent._sync_model_profile(
|
||||
SimpleNamespace(model="large-model", profile={"max_input_tokens": 128000})
|
||||
)
|
||||
middleware = UsageMiddleware(
|
||||
on_request_budget=agent._record_request_budget,
|
||||
on_usage=agent._record_usage,
|
||||
)
|
||||
request = _request(
|
||||
messages=[HumanMessage(content="small")],
|
||||
max_input_tokens=2048,
|
||||
)
|
||||
request.model.model = "small-model"
|
||||
|
||||
async def _failing_handler(_request):
|
||||
raise RuntimeError("model unavailable")
|
||||
|
||||
async def _run_failure():
|
||||
try:
|
||||
await middleware.awrap_model_call(request, _failing_handler)
|
||||
except RuntimeError:
|
||||
pass
|
||||
|
||||
asyncio.run(_run_failure())
|
||||
status = agent.get_session_status()
|
||||
|
||||
assert status["last_request_sequence"] == 1
|
||||
assert status["model"] == "small-model"
|
||||
assert status["context_window_tokens"] == 2048
|
||||
|
||||
|
||||
def test_unknown_model_name_is_not_replaced_after_request_snapshot():
|
||||
"""请求已开始后,未知模型名称不得与配置默认值拼成虚假快照。"""
|
||||
agent = MoviePilotAgent(session_id="request-budget", user_id="user-1")
|
||||
agent._record_request_budget(
|
||||
{
|
||||
"has_estimate": True,
|
||||
"request_sequence": 1,
|
||||
"model": None,
|
||||
"estimated_input_tokens": 100,
|
||||
"context_window_tokens": 2048,
|
||||
}
|
||||
)
|
||||
|
||||
status = agent.get_session_status()
|
||||
|
||||
assert status["model"] is None
|
||||
assert status["context_window_tokens"] == 2048
|
||||
|
||||
|
||||
def test_request_sequence_remains_monotonic_after_agent_graph_rebuild():
|
||||
"""重建 Agent 图后,新观察器也必须延续当前会话的请求顺序。"""
|
||||
agent = MoviePilotAgent(session_id="request-budget", user_id="user-1")
|
||||
first_graph = UsageMiddleware(
|
||||
on_request_budget=agent._record_request_budget,
|
||||
on_usage=agent._record_usage,
|
||||
next_request_sequence=agent._next_request_sequence,
|
||||
)
|
||||
rebuilt_graph = UsageMiddleware(
|
||||
on_request_budget=agent._record_request_budget,
|
||||
on_usage=agent._record_usage,
|
||||
next_request_sequence=agent._next_request_sequence,
|
||||
)
|
||||
|
||||
async def _handler(request):
|
||||
return ModelResponse(
|
||||
result=[
|
||||
AIMessage(
|
||||
content=request.model.model,
|
||||
usage_metadata={
|
||||
"input_tokens": request.model.profile["max_input_tokens"] // 10,
|
||||
"output_tokens": 1,
|
||||
"total_tokens": request.model.profile["max_input_tokens"] // 10 + 1,
|
||||
},
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
first = _request(
|
||||
messages=[HumanMessage(content="first")],
|
||||
max_input_tokens=1000,
|
||||
)
|
||||
first.model.model = "first-model"
|
||||
second = _request(
|
||||
messages=[HumanMessage(content="second")],
|
||||
max_input_tokens=2000,
|
||||
)
|
||||
second.model.model = "second-model"
|
||||
|
||||
asyncio.run(first_graph.awrap_model_call(first, _handler))
|
||||
asyncio.run(rebuilt_graph.awrap_model_call(second, _handler))
|
||||
status = agent.get_session_status()
|
||||
|
||||
assert status["last_request_sequence"] == 2
|
||||
assert status["model"] == "second-model"
|
||||
assert status["context_window_tokens"] == 2000
|
||||
assert status["last_input_tokens"] == 200
|
||||
|
||||
|
||||
def test_new_request_estimate_clears_stale_calibration_until_success():
|
||||
"""新请求失败时不得沿用上一轮估算误差。"""
|
||||
agent = MoviePilotAgent(session_id="request-budget", user_id="user-1")
|
||||
first = {
|
||||
"has_estimate": True,
|
||||
"request_sequence": 1,
|
||||
"estimated_input_tokens": 100,
|
||||
"context_window_tokens": 1000,
|
||||
}
|
||||
second = {
|
||||
"has_estimate": True,
|
||||
"request_sequence": 2,
|
||||
"estimated_input_tokens": 200,
|
||||
"context_window_tokens": 1000,
|
||||
}
|
||||
|
||||
agent._record_request_budget(first)
|
||||
agent._record_usage(
|
||||
{
|
||||
"request_sequence": 1,
|
||||
"request_budget_recorded": True,
|
||||
"input_usage_available": True,
|
||||
"estimated_input_tokens": 100,
|
||||
"has_usage": True,
|
||||
"input_tokens": 120,
|
||||
"output_tokens": 1,
|
||||
"total_tokens": 121,
|
||||
}
|
||||
)
|
||||
assert agent.get_session_status()["last_estimate_error_tokens"] == 20
|
||||
|
||||
agent._record_request_budget(second)
|
||||
status = agent.get_session_status()
|
||||
|
||||
assert status["last_request_sequence"] == 2
|
||||
assert status["last_actual_input_tokens"] is None
|
||||
assert status["last_estimate_error_tokens"] is None
|
||||
assert status["last_estimate_error_ratio"] is None
|
||||
|
||||
|
||||
def test_failed_estimate_clears_previous_request_components():
|
||||
"""估算失败后状态不得继续展示上一轮的请求组成。"""
|
||||
agent = MoviePilotAgent(session_id="request-budget", user_id="user-1")
|
||||
agent._record_request_budget(
|
||||
{
|
||||
"has_estimate": True,
|
||||
"request_sequence": 1,
|
||||
"estimated_input_tokens": 100,
|
||||
"message_tokens": 70,
|
||||
"system_tokens": 20,
|
||||
"tool_tokens": 10,
|
||||
"message_count": 3,
|
||||
"tool_count": 2,
|
||||
}
|
||||
)
|
||||
|
||||
agent._record_request_budget(
|
||||
{"has_estimate": False, "request_sequence": 2}
|
||||
)
|
||||
status = agent.get_session_status()
|
||||
|
||||
assert status["last_request_sequence"] == 2
|
||||
assert status["last_request_estimate_available"] is False
|
||||
assert status["last_estimated_input_tokens"] is None
|
||||
assert status["last_estimated_message_tokens"] is None
|
||||
assert status["last_estimated_system_tokens"] is None
|
||||
assert status["last_estimated_tool_tokens"] is None
|
||||
assert status["last_message_count"] == 0
|
||||
assert status["last_tool_count"] == 0
|
||||
|
||||
|
||||
def test_new_request_with_unknown_window_clears_previous_model_window():
|
||||
"""切换到窗口未知的模型后,不得继续展示上一模型的窗口。"""
|
||||
agent = MoviePilotAgent(session_id="request-budget", user_id="user-1")
|
||||
agent._sync_model_profile(
|
||||
SimpleNamespace(model="large-model", profile={"max_input_tokens": 128000})
|
||||
)
|
||||
agent._record_request_budget(
|
||||
{
|
||||
"has_estimate": True,
|
||||
"request_sequence": 1,
|
||||
"estimated_input_tokens": 100,
|
||||
"context_window_tokens": 128000,
|
||||
}
|
||||
)
|
||||
|
||||
agent._sync_model_profile(
|
||||
SimpleNamespace(model="unknown-window-model", profile={"max_input_tokens": 0})
|
||||
)
|
||||
agent._record_request_budget(
|
||||
{
|
||||
"has_estimate": True,
|
||||
"request_sequence": 2,
|
||||
"estimated_input_tokens": 100,
|
||||
"context_window_tokens": None,
|
||||
}
|
||||
)
|
||||
status = agent.get_session_status()
|
||||
|
||||
assert status["model"] == "unknown-window-model"
|
||||
assert status["context_window_tokens"] is None
|
||||
assert status["last_estimated_input_ratio"] is None
|
||||
assert status["last_estimated_over_input_limit"] is None
|
||||
|
||||
|
||||
def test_real_agent_observer_sees_dynamic_system_and_selected_tools():
|
||||
"""末尾观察器必须看到动态 system 和 ToolSelector 最终保留的工具。"""
|
||||
|
||||
@tool
|
||||
def kept_tool(value: str) -> str:
|
||||
"""应保留的测试工具。"""
|
||||
return value
|
||||
|
||||
@tool
|
||||
def removed_tool(value: str) -> str:
|
||||
"""应被筛除的测试工具。"""
|
||||
return value
|
||||
|
||||
selection_model = _ToolBindingFakeModel(
|
||||
responses=[AIMessage(content='{"tools": []}')]
|
||||
)
|
||||
main_model = _ToolBindingFakeModel(responses=[AIMessage(content="done")])
|
||||
budgets = []
|
||||
graph = create_agent(
|
||||
model=main_model,
|
||||
tools=[kept_tool, removed_tool],
|
||||
system_prompt="base",
|
||||
middleware=[
|
||||
_DynamicSystemMiddleware(),
|
||||
ToolSelectorMiddleware(
|
||||
model=selection_model,
|
||||
selection_tools=[kept_tool, removed_tool],
|
||||
max_tools=1,
|
||||
always_include=["kept_tool"],
|
||||
),
|
||||
UsageMiddleware(on_request_budget=budgets.append),
|
||||
],
|
||||
)
|
||||
|
||||
asyncio.run(graph.ainvoke({"messages": [HumanMessage(content="hello")]}))
|
||||
|
||||
assert len(budgets) == 1
|
||||
assert budgets[0]["tool_count"] == 1
|
||||
assert budgets[0]["system_tokens"] > 100
|
||||
assert main_model.bound_tool_names == ["kept_tool"]
|
||||
@@ -67,10 +67,22 @@ class TestAgentSessionStatus(unittest.TestCase):
|
||||
"session_id": "session-1",
|
||||
"model": "gpt-4o-mini",
|
||||
"context_window_tokens": 128000,
|
||||
"last_input_usage_available": True,
|
||||
"last_input_tokens": 1200,
|
||||
"last_output_tokens": 300,
|
||||
"last_total_tokens": 1500,
|
||||
"last_context_usage_ratio": 1200 / 128000,
|
||||
"last_request_estimate_available": True,
|
||||
"last_estimated_input_tokens": 1300,
|
||||
"last_estimated_message_tokens": 900,
|
||||
"last_estimated_system_tokens": 250,
|
||||
"last_estimated_tool_tokens": 150,
|
||||
"last_estimated_multimodal_tokens": 0,
|
||||
"last_estimated_input_ratio": 1300 / 128000,
|
||||
"last_estimated_over_input_limit": False,
|
||||
"last_actual_input_tokens": 1200,
|
||||
"last_estimate_error_tokens": -100,
|
||||
"last_estimate_error_ratio": -100 / 1200,
|
||||
"total_input_tokens": 4500,
|
||||
"total_output_tokens": 1500,
|
||||
"total_tokens": 6000,
|
||||
@@ -98,6 +110,9 @@ class TestAgentSessionStatus(unittest.TestCase):
|
||||
self.assertIn("session-1", notification.text)
|
||||
self.assertIn("gpt-4o-mini", notification.text)
|
||||
self.assertIn("1,200 / 128,000 (0.94%)", notification.text)
|
||||
self.assertIn("最终请求估算: 1,300 / 128,000 (1.02%)", notification.text)
|
||||
self.assertIn("消息 900 / 系统 250 / 工具 150 / 其中图片固定成本 0", notification.text)
|
||||
self.assertIn("实际 1,200 / 误差 -100 (-8.33%)", notification.text)
|
||||
self.assertIn("输入 4,500 / 输出 1,500 / 总计 6,000", notification.text)
|
||||
self.assertIn("运行中", notification.text)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user