diff --git a/app/agent/__init__.py b/app/agent/__init__.py index 58993e980..dfbbbbad3 100644 --- a/app/agent/__init__.py +++ b/app/agent/__init__.py @@ -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 diff --git a/app/agent/middleware/usage.py b/app/agent/middleware/usage.py index 3d3fdeede..de8d21dd4 100644 --- a/app/agent/middleware/usage.py +++ b/app/agent/middleware/usage.py @@ -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, diff --git a/app/chain/message.py b/app/chain/message.py index cfe6c97d9..0e09b4d89 100644 --- a/app/chain/message.py +++ b/app/chain/message.py @@ -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( diff --git a/tests/test_agent_graph_cache.py b/tests/test_agent_graph_cache.py index a4772b144..d4b003169 100644 --- a/tests/test_agent_graph_cache.py +++ b/tests/test_agent_graph_cache.py @@ -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"] diff --git a/tests/test_agent_request_budget.py b/tests/test_agent_request_budget.py new file mode 100644 index 000000000..1a9124543 --- /dev/null +++ b/tests/test_agent_request_budget.py @@ -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"] diff --git a/tests/test_agent_session_status.py b/tests/test_agent_session_status.py index 338967ee6..f4883500c 100644 --- a/tests/test_agent_session_status.py +++ b/tests/test_agent_session_status.py @@ -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)