mirror of
https://github.com/jxxghp/MoviePilot.git
synced 2026-09-09 09:26:55 +08:00
feat(agent): improve prompt cache hit rate
This commit is contained in:
+69
-1
@@ -129,9 +129,18 @@ class _SessionUsageSnapshot:
|
|||||||
last_output_tokens: int = 0
|
last_output_tokens: int = 0
|
||||||
last_total_tokens: int = 0
|
last_total_tokens: int = 0
|
||||||
last_context_usage_ratio: Optional[float] = None
|
last_context_usage_ratio: Optional[float] = None
|
||||||
|
last_cache_usage_available: bool = False
|
||||||
|
last_cache_read_input_tokens: int = 0
|
||||||
|
last_cache_write_input_tokens: int = 0
|
||||||
|
last_uncached_input_tokens: int = 0
|
||||||
|
last_cache_hit_ratio: Optional[float] = None
|
||||||
total_input_tokens: int = 0
|
total_input_tokens: int = 0
|
||||||
total_output_tokens: int = 0
|
total_output_tokens: int = 0
|
||||||
total_tokens: int = 0
|
total_tokens: int = 0
|
||||||
|
total_cache_read_input_tokens: int = 0
|
||||||
|
total_cache_write_input_tokens: int = 0
|
||||||
|
total_uncached_input_tokens: int = 0
|
||||||
|
cache_usage_available: bool = False
|
||||||
model_call_count: int = 0
|
model_call_count: int = 0
|
||||||
last_updated_at: Optional[datetime] = None
|
last_updated_at: Optional[datetime] = None
|
||||||
|
|
||||||
@@ -144,9 +153,23 @@ class _SessionUsageSnapshot:
|
|||||||
"last_output_tokens": self.last_output_tokens,
|
"last_output_tokens": self.last_output_tokens,
|
||||||
"last_total_tokens": self.last_total_tokens,
|
"last_total_tokens": self.last_total_tokens,
|
||||||
"last_context_usage_ratio": self.last_context_usage_ratio,
|
"last_context_usage_ratio": self.last_context_usage_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,
|
||||||
|
"last_uncached_input_tokens": self.last_uncached_input_tokens,
|
||||||
|
"last_cache_hit_ratio": self.last_cache_hit_ratio,
|
||||||
"total_input_tokens": self.total_input_tokens,
|
"total_input_tokens": self.total_input_tokens,
|
||||||
"total_output_tokens": self.total_output_tokens,
|
"total_output_tokens": self.total_output_tokens,
|
||||||
"total_tokens": self.total_tokens,
|
"total_tokens": self.total_tokens,
|
||||||
|
"total_cache_read_input_tokens": self.total_cache_read_input_tokens,
|
||||||
|
"total_cache_write_input_tokens": self.total_cache_write_input_tokens,
|
||||||
|
"total_uncached_input_tokens": self.total_uncached_input_tokens,
|
||||||
|
"cache_usage_available": self.cache_usage_available,
|
||||||
|
"total_cache_hit_ratio": (
|
||||||
|
self.total_cache_read_input_tokens / self.total_input_tokens
|
||||||
|
if self.cache_usage_available and self.total_input_tokens
|
||||||
|
else None
|
||||||
|
),
|
||||||
"model_call_count": self.model_call_count,
|
"model_call_count": self.model_call_count,
|
||||||
"last_updated_at": self.last_updated_at.strftime("%Y-%m-%d %H:%M:%S")
|
"last_updated_at": self.last_updated_at.strftime("%Y-%m-%d %H:%M:%S")
|
||||||
if self.last_updated_at
|
if self.last_updated_at
|
||||||
@@ -554,9 +577,33 @@ class MoviePilotAgent:
|
|||||||
self._session_usage.last_output_tokens = output_tokens
|
self._session_usage.last_output_tokens = output_tokens
|
||||||
self._session_usage.last_total_tokens = total_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_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")
|
||||||
|
) or 0
|
||||||
|
cache_write_input_tokens = self._coerce_int(
|
||||||
|
usage.get("cache_write_input_tokens")
|
||||||
|
) or 0
|
||||||
|
uncached_input_tokens = self._coerce_int(
|
||||||
|
usage.get("uncached_input_tokens")
|
||||||
|
)
|
||||||
|
if uncached_input_tokens is None:
|
||||||
|
uncached_input_tokens = max(
|
||||||
|
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_input_tokens += input_tokens
|
||||||
self._session_usage.total_output_tokens += output_tokens
|
self._session_usage.total_output_tokens += output_tokens
|
||||||
self._session_usage.total_tokens += total_tokens
|
self._session_usage.total_tokens += total_tokens
|
||||||
|
self._session_usage.total_cache_read_input_tokens += cache_read_input_tokens
|
||||||
|
self._session_usage.total_cache_write_input_tokens += cache_write_input_tokens
|
||||||
|
self._session_usage.total_uncached_input_tokens += uncached_input_tokens
|
||||||
|
self._session_usage.cache_usage_available |= cache_usage_available
|
||||||
|
|
||||||
def get_session_status(self) -> dict[str, Any]:
|
def get_session_status(self) -> dict[str, Any]:
|
||||||
if not self._session_usage.model:
|
if not self._session_usage.model:
|
||||||
@@ -590,6 +637,17 @@ class MoviePilotAgent:
|
|||||||
input_tokens=self._session_usage.total_input_tokens,
|
input_tokens=self._session_usage.total_input_tokens,
|
||||||
output_tokens=self._session_usage.total_output_tokens,
|
output_tokens=self._session_usage.total_output_tokens,
|
||||||
total_tokens=self._session_usage.total_tokens,
|
total_tokens=self._session_usage.total_tokens,
|
||||||
|
cache_read_input_tokens=self._session_usage.total_cache_read_input_tokens,
|
||||||
|
cache_write_input_tokens=self._session_usage.total_cache_write_input_tokens,
|
||||||
|
uncached_input_tokens=self._session_usage.total_uncached_input_tokens,
|
||||||
|
cache_hit_ratio=(
|
||||||
|
self._session_usage.total_cache_read_input_tokens
|
||||||
|
/ self._session_usage.total_input_tokens
|
||||||
|
if self._session_usage.cache_usage_available
|
||||||
|
and self._session_usage.total_input_tokens
|
||||||
|
else None
|
||||||
|
),
|
||||||
|
cache_usage_available=self._session_usage.cache_usage_available,
|
||||||
model_call_count=self._session_usage.model_call_count,
|
model_call_count=self._session_usage.model_call_count,
|
||||||
success=success,
|
success=success,
|
||||||
error=error,
|
error=error,
|
||||||
@@ -819,7 +877,17 @@ class MoviePilotAgent:
|
|||||||
:param streaming: 是否启用流式输出
|
:param streaming: 是否启用流式输出
|
||||||
"""
|
"""
|
||||||
runtime_config = await self._resolve_llm_runtime_config()
|
runtime_config = await self._resolve_llm_runtime_config()
|
||||||
return await LLMHelper.get_llm(streaming=streaming, **runtime_config)
|
return await LLMHelper.get_llm(
|
||||||
|
streaming=streaming,
|
||||||
|
prompt_cache_key=self._build_prompt_cache_key(),
|
||||||
|
**runtime_config,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _build_prompt_cache_key(self) -> str:
|
||||||
|
"""生成不暴露用户标识、且在同一会话内稳定的提示词缓存键。"""
|
||||||
|
cache_identity = f"{self.user_id or ''}\x00{self.session_id}"
|
||||||
|
digest = hashlib.sha256(cache_identity.encode("utf-8")).hexdigest()[:32]
|
||||||
|
return f"moviepilot-agent-{digest}"
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def _has_image_input_content(cls, content: Any) -> bool:
|
def _has_image_input_content(cls, content: Any) -> bool:
|
||||||
|
|||||||
+120
-3
@@ -6,6 +6,7 @@ import json
|
|||||||
import time
|
import time
|
||||||
from functools import wraps
|
from functools import wraps
|
||||||
from typing import TYPE_CHECKING, Any, List, Optional
|
from typing import TYPE_CHECKING, Any, List, Optional
|
||||||
|
from urllib.parse import urlsplit
|
||||||
|
|
||||||
from langchain_core.messages import AIMessage, AIMessageChunk
|
from langchain_core.messages import AIMessage, AIMessageChunk
|
||||||
|
|
||||||
@@ -813,6 +814,91 @@ class LLMHelper:
|
|||||||
headers["User-Agent"] = normalized_user_agent
|
headers["User-Agent"] = normalized_user_agent
|
||||||
return headers or None
|
return headers or None
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _matches_endpoint_host(base_url: str | None, expected_host: str) -> bool:
|
||||||
|
"""严格匹配官方 API 主机,避免向兼容端点发送供应商专属参数。"""
|
||||||
|
try:
|
||||||
|
return (urlsplit(str(base_url or "")).hostname or "").lower() == expected_host
|
||||||
|
except ValueError:
|
||||||
|
return False
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _build_openai_prompt_cache_options(
|
||||||
|
cls,
|
||||||
|
*,
|
||||||
|
provider: str,
|
||||||
|
base_url: str | None,
|
||||||
|
use_responses_api: bool | None,
|
||||||
|
prompt_cache_key: str | None,
|
||||||
|
default_headers: dict[str, str] | None,
|
||||||
|
model_kwargs: dict[str, Any],
|
||||||
|
) -> tuple[dict[str, str] | None, dict[str, Any]]:
|
||||||
|
"""为 OpenAI 与 xAI 官方端点构造稳定提示词缓存路由参数。"""
|
||||||
|
cache_key = str(prompt_cache_key or "").strip()
|
||||||
|
headers = dict(default_headers or {})
|
||||||
|
kwargs = dict(model_kwargs)
|
||||||
|
provider_name = str(provider or "").strip().lower()
|
||||||
|
if not cache_key:
|
||||||
|
return headers or None, kwargs
|
||||||
|
|
||||||
|
is_openai = provider_name in {"chatgpt", "openai"} and cls._matches_endpoint_host(
|
||||||
|
base_url,
|
||||||
|
"api.openai.com",
|
||||||
|
)
|
||||||
|
is_xai = provider_name == "xai" and cls._matches_endpoint_host(
|
||||||
|
base_url,
|
||||||
|
"api.x.ai",
|
||||||
|
)
|
||||||
|
if not is_openai and not is_xai:
|
||||||
|
return headers or None, kwargs
|
||||||
|
|
||||||
|
if is_xai and use_responses_api is not True:
|
||||||
|
headers["x-grok-conv-id"] = cache_key
|
||||||
|
return headers, kwargs
|
||||||
|
|
||||||
|
extra_body = dict(kwargs.get("extra_body") or {})
|
||||||
|
extra_body["prompt_cache_key"] = cache_key
|
||||||
|
kwargs["extra_body"] = extra_body
|
||||||
|
return headers or None, kwargs
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _with_prompt_cache_control(
|
||||||
|
model_cls: type,
|
||||||
|
cache_control: dict[str, str],
|
||||||
|
) -> type:
|
||||||
|
"""创建在最终模型绑定阶段保留缓存控制参数的适配类。"""
|
||||||
|
|
||||||
|
class PromptCachingModel(model_cls):
|
||||||
|
"""在 LangChain 工具绑定后仍保留提示词缓存参数的模型适配器。"""
|
||||||
|
|
||||||
|
def bind(self, **kwargs: Any) -> Any:
|
||||||
|
"""绑定调用参数,并补入当前 Provider 的默认缓存控制。"""
|
||||||
|
kwargs.setdefault("cache_control", dict(cache_control))
|
||||||
|
return super().bind(**kwargs)
|
||||||
|
|
||||||
|
PromptCachingModel.__name__ = f"PromptCaching{model_cls.__name__}"
|
||||||
|
return PromptCachingModel
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _use_anthropic_prompt_cache(
|
||||||
|
cls,
|
||||||
|
*,
|
||||||
|
provider: str,
|
||||||
|
runtime: dict[str, Any],
|
||||||
|
prompt_cache_key: str | None,
|
||||||
|
) -> bool:
|
||||||
|
"""判断当前运行时是否为可安全启用缓存的 Anthropic 官方端点。"""
|
||||||
|
return (
|
||||||
|
bool(str(prompt_cache_key or "").strip())
|
||||||
|
and str(provider or "").strip().lower() == "anthropic"
|
||||||
|
and str(runtime.get("runtime") or "").strip().lower()
|
||||||
|
== "anthropic_compatible"
|
||||||
|
and cls._matches_endpoint_host(
|
||||||
|
runtime.get("base_url"),
|
||||||
|
"api.anthropic.com",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def _should_use_openai_responses_api(
|
def _should_use_openai_responses_api(
|
||||||
cls,
|
cls,
|
||||||
@@ -983,6 +1069,7 @@ class LLMHelper:
|
|||||||
use_proxy: bool | None = None,
|
use_proxy: bool | None = None,
|
||||||
api_protocol: str | None = None,
|
api_protocol: str | None = None,
|
||||||
web_search_mode: str | None = None,
|
web_search_mode: str | None = None,
|
||||||
|
prompt_cache_key: str | None = None,
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
获取LLM实例
|
获取LLM实例
|
||||||
@@ -1006,6 +1093,7 @@ class LLMHelper:
|
|||||||
:param web_search_mode: 联网搜索模式
|
:param web_search_mode: 联网搜索模式
|
||||||
(local/builtin/auto/disabled)。未显式传入时使用配置项
|
(local/builtin/auto/disabled)。未显式传入时使用配置项
|
||||||
``LLM_WEB_SEARCH_MODE``。
|
``LLM_WEB_SEARCH_MODE``。
|
||||||
|
:param prompt_cache_key: 同一 Agent 会话内稳定且脱敏的提示词缓存路由键。
|
||||||
:return: LLM实例
|
:return: LLM实例
|
||||||
"""
|
"""
|
||||||
provider_name = str(provider if provider is not None else settings.LLM_PROVIDER).lower()
|
provider_name = str(provider if provider is not None else settings.LLM_PROVIDER).lower()
|
||||||
@@ -1093,6 +1181,14 @@ class LLMHelper:
|
|||||||
runtime=runtime,
|
runtime=runtime,
|
||||||
api_protocol=effective_api_protocol,
|
api_protocol=effective_api_protocol,
|
||||||
)
|
)
|
||||||
|
default_headers, openai_model_kwargs = cls._build_openai_prompt_cache_options(
|
||||||
|
provider=provider_name,
|
||||||
|
base_url=runtime.get("base_url"),
|
||||||
|
use_responses_api=use_responses_api,
|
||||||
|
prompt_cache_key=prompt_cache_key,
|
||||||
|
default_headers=default_headers,
|
||||||
|
model_kwargs=thinking_kwargs,
|
||||||
|
)
|
||||||
llm_proxy = _resolve_llm_proxy(use_proxy)
|
llm_proxy = _resolve_llm_proxy(use_proxy)
|
||||||
|
|
||||||
if runtime["runtime"] == "google":
|
if runtime["runtime"] == "google":
|
||||||
@@ -1146,6 +1242,16 @@ class LLMHelper:
|
|||||||
|
|
||||||
from app.agent.llm.provider import LLMProviderManager
|
from app.agent.llm.provider import LLMProviderManager
|
||||||
|
|
||||||
|
bedrock_model_cls = ChatBedrockConverse
|
||||||
|
if (
|
||||||
|
str(prompt_cache_key or "").strip()
|
||||||
|
and runtime.get("supports_prompt_cache")
|
||||||
|
):
|
||||||
|
bedrock_model_cls = cls._with_prompt_cache_control(
|
||||||
|
ChatBedrockConverse,
|
||||||
|
{"type": "default"},
|
||||||
|
)
|
||||||
|
|
||||||
aws_region = runtime.get("aws_region") or "us-east-1"
|
aws_region = runtime.get("aws_region") or "us-east-1"
|
||||||
aws_auth = runtime.get("aws_auth") or {}
|
aws_auth = runtime.get("aws_auth") or {}
|
||||||
# Bearer 认证需要跳过 SigV4 签名并注入 Authorization 头,SigV4 认证
|
# Bearer 认证需要跳过 SigV4 签名并注入 Authorization 头,SigV4 认证
|
||||||
@@ -1158,7 +1264,7 @@ class LLMHelper:
|
|||||||
use_proxy=use_proxy,
|
use_proxy=use_proxy,
|
||||||
read_timeout=settings.LLM_TOOL_TIMEOUT,
|
read_timeout=settings.LLM_TOOL_TIMEOUT,
|
||||||
)
|
)
|
||||||
model = ChatBedrockConverse(
|
model = bedrock_model_cls(
|
||||||
model_id=model_name,
|
model_id=model_name,
|
||||||
client=bedrock_client,
|
client=bedrock_client,
|
||||||
temperature=temperature_value,
|
temperature=temperature_value,
|
||||||
@@ -1167,7 +1273,18 @@ class LLMHelper:
|
|||||||
elif runtime["runtime"] in {"anthropic_compatible", "copilot_anthropic"}:
|
elif runtime["runtime"] in {"anthropic_compatible", "copilot_anthropic"}:
|
||||||
from langchain_anthropic import ChatAnthropic
|
from langchain_anthropic import ChatAnthropic
|
||||||
|
|
||||||
model = ChatAnthropic(
|
anthropic_model_cls = ChatAnthropic
|
||||||
|
if cls._use_anthropic_prompt_cache(
|
||||||
|
provider=provider_name,
|
||||||
|
runtime=runtime,
|
||||||
|
prompt_cache_key=prompt_cache_key,
|
||||||
|
):
|
||||||
|
anthropic_model_cls = cls._with_prompt_cache_control(
|
||||||
|
ChatAnthropic,
|
||||||
|
{"type": "ephemeral"},
|
||||||
|
)
|
||||||
|
|
||||||
|
model = anthropic_model_cls(
|
||||||
model=model_name,
|
model=model_name,
|
||||||
api_key=runtime["api_key"],
|
api_key=runtime["api_key"],
|
||||||
base_url=runtime["base_url"],
|
base_url=runtime["base_url"],
|
||||||
@@ -1208,7 +1325,7 @@ class LLMHelper:
|
|||||||
default_headers=default_headers,
|
default_headers=default_headers,
|
||||||
use_responses_api=use_responses_api,
|
use_responses_api=use_responses_api,
|
||||||
output_version=("responses/v1" if use_responses_api else None),
|
output_version=("responses/v1" if use_responses_api else None),
|
||||||
**thinking_kwargs,
|
**openai_model_kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
# 优先使用 provider / models.dev 目录中的上下文上限,减少用户手填成本。
|
# 优先使用 provider / models.dev 目录中的上下文上限,减少用户手填成本。
|
||||||
|
|||||||
+40
-10
@@ -1665,6 +1665,20 @@ class LLMProviderManager(metaclass=Singleton):
|
|||||||
await self.get_models_dev_data(use_proxy=use_proxy)
|
await self.get_models_dev_data(use_proxy=use_proxy)
|
||||||
).get(models_dev_provider_id, {}) or {}
|
).get(models_dev_provider_id, {}) or {}
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _models_dev_model_candidates(
|
||||||
|
provider_id: str,
|
||||||
|
model_id: str,
|
||||||
|
) -> tuple[str, ...]:
|
||||||
|
"""生成模型目录查询候选,兼容 Provider 添加的透明模型前缀。"""
|
||||||
|
candidates = [model_id]
|
||||||
|
if model_id.startswith("models/"):
|
||||||
|
candidates.append(model_id.removeprefix("models/"))
|
||||||
|
if provider_id == "amazon-bedrock" and "." in model_id:
|
||||||
|
# Cross-region Inference Profile 会增加 us./eu./global. 等前缀。
|
||||||
|
candidates.append(model_id.split(".", 1)[1])
|
||||||
|
return tuple(dict.fromkeys(candidates))
|
||||||
|
|
||||||
async def _models_dev_model(
|
async def _models_dev_model(
|
||||||
self,
|
self,
|
||||||
provider_id: str,
|
provider_id: str,
|
||||||
@@ -1684,15 +1698,32 @@ class LLMProviderManager(metaclass=Singleton):
|
|||||||
if not isinstance(models, dict):
|
if not isinstance(models, dict):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
candidates = [model_id]
|
for candidate in self._models_dev_model_candidates(provider_id, model_id):
|
||||||
if model_id.startswith("models/"):
|
|
||||||
candidates.append(model_id.removeprefix("models/"))
|
|
||||||
|
|
||||||
for candidate in candidates:
|
|
||||||
if candidate in models:
|
if candidate in models:
|
||||||
return models[candidate]
|
return models[candidate]
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _metadata_supports_prompt_cache(metadata: Any) -> bool:
|
||||||
|
"""从统一模型元数据中判断是否声明了提示词缓存能力。"""
|
||||||
|
if not isinstance(metadata, dict):
|
||||||
|
return False
|
||||||
|
|
||||||
|
explicit_capability = metadata.get("prompt_cache")
|
||||||
|
if isinstance(explicit_capability, bool):
|
||||||
|
return explicit_capability
|
||||||
|
|
||||||
|
capabilities = metadata.get("capabilities")
|
||||||
|
if isinstance(capabilities, dict):
|
||||||
|
explicit_capability = capabilities.get("prompt_cache")
|
||||||
|
if isinstance(explicit_capability, bool):
|
||||||
|
return explicit_capability
|
||||||
|
|
||||||
|
cost = metadata.get("cost")
|
||||||
|
return isinstance(cost, dict) and any(
|
||||||
|
key in cost for key in ("cache_read", "cache_write")
|
||||||
|
)
|
||||||
|
|
||||||
def _cached_models_dev_model(
|
def _cached_models_dev_model(
|
||||||
self,
|
self,
|
||||||
provider_id: str,
|
provider_id: str,
|
||||||
@@ -1719,11 +1750,7 @@ class LLMProviderManager(metaclass=Singleton):
|
|||||||
if not isinstance(models, dict):
|
if not isinstance(models, dict):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
candidates = [model_id]
|
for candidate in self._models_dev_model_candidates(provider_id, model_id):
|
||||||
if model_id.startswith("models/"):
|
|
||||||
candidates.append(model_id.removeprefix("models/"))
|
|
||||||
|
|
||||||
for candidate in candidates:
|
|
||||||
if candidate in models:
|
if candidate in models:
|
||||||
return models[candidate]
|
return models[candidate]
|
||||||
return None
|
return None
|
||||||
@@ -3112,6 +3139,9 @@ class LLMProviderManager(metaclass=Singleton):
|
|||||||
"model_id": model,
|
"model_id": model,
|
||||||
"model_record": model_record,
|
"model_record": model_record,
|
||||||
"model_metadata": model_metadata,
|
"model_metadata": model_metadata,
|
||||||
|
"supports_prompt_cache": self._metadata_supports_prompt_cache(
|
||||||
|
model_metadata
|
||||||
|
),
|
||||||
"default_headers": None,
|
"default_headers": None,
|
||||||
"use_responses_api": None,
|
"use_responses_api": None,
|
||||||
"auth_mode": "api_key",
|
"auth_mode": "api_key",
|
||||||
|
|||||||
@@ -3,7 +3,7 @@
|
|||||||
|
|
||||||
按日期存储在 CONFIG_PATH/agent/activity/YYYY-MM-DD.md 中,
|
按日期存储在 CONFIG_PATH/agent/activity/YYYY-MM-DD.md 中,
|
||||||
每次 Agent 执行完毕后自动调用 LLM 对本轮对话生成简洁的活动摘要,
|
每次 Agent 执行完毕后自动调用 LLM 对本轮对话生成简洁的活动摘要,
|
||||||
并在每次 Agent 启动时注入轻量索引,完整日志由工具按需查询。
|
系统提示词只注入稳定的检索规则,完整日志由工具按需查询。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
@@ -459,12 +459,8 @@ async def _summarize_with_llm(conversation_text: str) -> Optional[str]:
|
|||||||
|
|
||||||
|
|
||||||
ACTIVITY_LOG_SYSTEM_PROMPT = """<activity_log>
|
ACTIVITY_LOG_SYSTEM_PROMPT = """<activity_log>
|
||||||
<activity_log_index>
|
|
||||||
{activity_log_index}
|
|
||||||
</activity_log_index>
|
|
||||||
|
|
||||||
<activity_log_guidelines>
|
<activity_log_guidelines>
|
||||||
The index only shows recent dates and entry counts, not full log contents.
|
Activity log contents and indexes are not included in the default context.
|
||||||
Use `query_activity_log` only when the user references previous work, asks to continue a prior task, or recent activity is clearly relevant.
|
Use `query_activity_log` only when the user references previous work, asks to continue a prior task, or recent activity is clearly relevant.
|
||||||
Activity logs are read-only and retained for {retention_days} days; use MEMORY.md for durable preferences.
|
Activity logs are read-only and retained for {retention_days} days; use MEMORY.md for durable preferences.
|
||||||
</activity_log_guidelines>
|
</activity_log_guidelines>
|
||||||
@@ -473,10 +469,10 @@ ACTIVITY_LOG_SYSTEM_PROMPT = """<activity_log>
|
|||||||
|
|
||||||
|
|
||||||
class ActivityLogMiddleware(AgentMiddleware[ActivityLogState, ContextT, ResponseT]): # noqa
|
class ActivityLogMiddleware(AgentMiddleware[ActivityLogState, ContextT, ResponseT]): # noqa
|
||||||
"""自动记录 Agent 活动日志并注入轻量索引的中间件。
|
"""自动记录 Agent 活动日志并注入稳定检索规则的中间件。
|
||||||
|
|
||||||
- abefore_agent: 加载近几天的活动日志索引
|
- abefore_agent: 加载近几天的活动日志索引
|
||||||
- awrap_model_call: 将活动日志索引和检索规则注入系统提示词
|
- awrap_model_call: 将固定的活动日志检索规则注入系统提示词
|
||||||
- aafter_agent: 从本次对话中提取摘要并追加到当日日志文件
|
- aafter_agent: 从本次对话中提取摘要并追加到当日日志文件
|
||||||
|
|
||||||
参数:
|
参数:
|
||||||
@@ -516,31 +512,9 @@ class ActivityLogMiddleware(AgentMiddleware[ActivityLogState, ContextT, Response
|
|||||||
"""获取指定日期的日志文件路径。"""
|
"""获取指定日期的日志文件路径。"""
|
||||||
return AsyncPath(self.activity_dir) / f"{date_str}.md"
|
return AsyncPath(self.activity_dir) / f"{date_str}.md"
|
||||||
|
|
||||||
def _format_activity_log(self, contents: dict[str, str]) -> str:
|
def _format_activity_log(self, _contents: dict[str, str]) -> str:
|
||||||
"""格式化活动日志索引用于系统提示词注入。"""
|
"""生成不受活动日志内容变化影响的系统提示词。"""
|
||||||
if not contents:
|
|
||||||
return ACTIVITY_LOG_SYSTEM_PROMPT.format(
|
|
||||||
activity_log_index="(近期暂无活动日志索引。需要历史上下文时可调用 query_activity_log。)",
|
|
||||||
retention_days=self.retention_days,
|
|
||||||
)
|
|
||||||
|
|
||||||
# 按日期排序(最近的在前)
|
|
||||||
sorted_dates = sorted(contents.keys(), reverse=True)
|
|
||||||
sections = []
|
|
||||||
for date_str in sorted_dates:
|
|
||||||
content = contents[date_str].strip()
|
|
||||||
if content:
|
|
||||||
sections.append(f"### {date_str}\n{content}")
|
|
||||||
|
|
||||||
if not sections:
|
|
||||||
return ACTIVITY_LOG_SYSTEM_PROMPT.format(
|
|
||||||
activity_log_index="(近期暂无活动日志索引。需要历史上下文时可调用 query_activity_log。)",
|
|
||||||
retention_days=self.retention_days,
|
|
||||||
)
|
|
||||||
|
|
||||||
log_body = "\n".join(sections)
|
|
||||||
return ACTIVITY_LOG_SYSTEM_PROMPT.format(
|
return ACTIVITY_LOG_SYSTEM_PROMPT.format(
|
||||||
activity_log_index=log_body,
|
|
||||||
retention_days=self.retention_days,
|
retention_days=self.retention_days,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -51,6 +51,18 @@ class UsageMiddleware(AgentMiddleware):
|
|||||||
|
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _first_int(
|
||||||
|
cls,
|
||||||
|
candidates: tuple[tuple[Any, tuple[str, ...]], ...],
|
||||||
|
) -> int | None:
|
||||||
|
"""按优先级返回首个可用的 usage 整数值。"""
|
||||||
|
for container, keys in candidates:
|
||||||
|
value = cls._lookup_int(container, *keys)
|
||||||
|
if value is not None:
|
||||||
|
return value
|
||||||
|
return None
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def _extract_model_name(cls, model: Any) -> str | None:
|
def _extract_model_name(cls, model: Any) -> str | None:
|
||||||
return (
|
return (
|
||||||
@@ -82,6 +94,131 @@ class UsageMiddleware(AgentMiddleware):
|
|||||||
or {}
|
or {}
|
||||||
)
|
)
|
||||||
|
|
||||||
|
input_token_details = None
|
||||||
|
if usage_metadata:
|
||||||
|
getter = getattr(usage_metadata, "get", None)
|
||||||
|
input_token_details = (
|
||||||
|
getter("input_token_details")
|
||||||
|
if callable(getter)
|
||||||
|
else getattr(usage_metadata, "input_token_details", None)
|
||||||
|
)
|
||||||
|
|
||||||
|
cache_read_tokens = cls._first_int(
|
||||||
|
(
|
||||||
|
(
|
||||||
|
input_token_details,
|
||||||
|
(
|
||||||
|
"cache_read",
|
||||||
|
"cached_tokens",
|
||||||
|
"cache_read_input_tokens",
|
||||||
|
"cacheReadInputTokens",
|
||||||
|
),
|
||||||
|
),
|
||||||
|
(
|
||||||
|
token_usage,
|
||||||
|
(
|
||||||
|
"prompt_cache_hit_tokens",
|
||||||
|
"cache_read_input_tokens",
|
||||||
|
"cacheReadInputTokens",
|
||||||
|
),
|
||||||
|
),
|
||||||
|
(
|
||||||
|
response_metadata,
|
||||||
|
(
|
||||||
|
"prompt_cache_hit_tokens",
|
||||||
|
"cache_read_input_tokens",
|
||||||
|
"cacheReadInputTokens",
|
||||||
|
"cached_tokens",
|
||||||
|
),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
if cache_read_tokens is None:
|
||||||
|
cache_read_tokens = cls._first_int(
|
||||||
|
(
|
||||||
|
(
|
||||||
|
token_usage.get("prompt_tokens_details", {}),
|
||||||
|
("cached_tokens", "cache_read"),
|
||||||
|
),
|
||||||
|
(
|
||||||
|
token_usage.get("input_tokens_details", {}),
|
||||||
|
("cached_tokens", "cache_read"),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
cache_write_tokens = cls._first_int(
|
||||||
|
(
|
||||||
|
(
|
||||||
|
input_token_details,
|
||||||
|
(
|
||||||
|
"cache_creation",
|
||||||
|
"cache_write",
|
||||||
|
"cache_write_tokens",
|
||||||
|
"cache_write_input_tokens",
|
||||||
|
"cacheWriteInputTokens",
|
||||||
|
),
|
||||||
|
),
|
||||||
|
(
|
||||||
|
token_usage,
|
||||||
|
(
|
||||||
|
"cache_creation_input_tokens",
|
||||||
|
"cache_write_tokens",
|
||||||
|
"cache_write_input_tokens",
|
||||||
|
"cacheWriteInputTokens",
|
||||||
|
),
|
||||||
|
),
|
||||||
|
(
|
||||||
|
response_metadata,
|
||||||
|
(
|
||||||
|
"cache_creation_input_tokens",
|
||||||
|
"cache_write_tokens",
|
||||||
|
"cache_write_input_tokens",
|
||||||
|
"cacheWriteInputTokens",
|
||||||
|
),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
if cache_write_tokens is None:
|
||||||
|
cache_write_tokens = cls._first_int(
|
||||||
|
(
|
||||||
|
(
|
||||||
|
token_usage.get("prompt_tokens_details", {}),
|
||||||
|
("cache_write_tokens", "cache_creation"),
|
||||||
|
),
|
||||||
|
(
|
||||||
|
token_usage.get("input_tokens_details", {}),
|
||||||
|
("cache_write_tokens", "cache_creation"),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
cache_write_ttl_tokens = sum(
|
||||||
|
cls._lookup_int(
|
||||||
|
input_token_details,
|
||||||
|
ttl_key,
|
||||||
|
)
|
||||||
|
or 0
|
||||||
|
for ttl_key in (
|
||||||
|
"ephemeral_5m_input_tokens",
|
||||||
|
"ephemeral_1h_input_tokens",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
if cache_write_ttl_tokens:
|
||||||
|
cache_write_tokens = cache_write_ttl_tokens
|
||||||
|
|
||||||
|
cache_miss_tokens = cls._first_int(
|
||||||
|
(
|
||||||
|
(
|
||||||
|
token_usage,
|
||||||
|
("prompt_cache_miss_tokens", "cache_miss_input_tokens"),
|
||||||
|
),
|
||||||
|
(
|
||||||
|
response_metadata,
|
||||||
|
("prompt_cache_miss_tokens", "cache_miss_input_tokens"),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
if input_tokens is None:
|
if input_tokens is None:
|
||||||
input_tokens = cls._lookup_int(
|
input_tokens = cls._lookup_int(
|
||||||
token_usage,
|
token_usage,
|
||||||
@@ -94,6 +231,27 @@ class UsageMiddleware(AgentMiddleware):
|
|||||||
"prompt_token_count",
|
"prompt_token_count",
|
||||||
"input_tokens",
|
"input_tokens",
|
||||||
)
|
)
|
||||||
|
if input_tokens is None:
|
||||||
|
bedrock_input_tokens = cls._lookup_int(token_usage, "inputTokens")
|
||||||
|
if bedrock_input_tokens is not None:
|
||||||
|
input_tokens = (
|
||||||
|
bedrock_input_tokens
|
||||||
|
+ (cache_read_tokens or 0)
|
||||||
|
+ (cache_write_tokens or 0)
|
||||||
|
)
|
||||||
|
if input_tokens is None and any(
|
||||||
|
value is not None
|
||||||
|
for value in (
|
||||||
|
cache_read_tokens,
|
||||||
|
cache_write_tokens,
|
||||||
|
cache_miss_tokens,
|
||||||
|
)
|
||||||
|
):
|
||||||
|
input_tokens = (
|
||||||
|
(cache_read_tokens or 0)
|
||||||
|
+ (cache_write_tokens or 0)
|
||||||
|
+ (cache_miss_tokens or 0)
|
||||||
|
)
|
||||||
|
|
||||||
if output_tokens is None:
|
if output_tokens is None:
|
||||||
output_tokens = cls._lookup_int(
|
output_tokens = cls._lookup_int(
|
||||||
@@ -113,8 +271,24 @@ class UsageMiddleware(AgentMiddleware):
|
|||||||
if total_tokens is None:
|
if total_tokens is None:
|
||||||
total_tokens = cls._lookup_int(response_metadata, "total_token_count")
|
total_tokens = cls._lookup_int(response_metadata, "total_token_count")
|
||||||
|
|
||||||
|
has_cache_usage = any(
|
||||||
|
value is not None
|
||||||
|
for value in (
|
||||||
|
cache_read_tokens,
|
||||||
|
cache_write_tokens,
|
||||||
|
cache_miss_tokens,
|
||||||
|
)
|
||||||
|
)
|
||||||
has_usage = any(
|
has_usage = any(
|
||||||
value is not None for value in (input_tokens, output_tokens, total_tokens)
|
value is not None
|
||||||
|
for value in (
|
||||||
|
input_tokens,
|
||||||
|
output_tokens,
|
||||||
|
total_tokens,
|
||||||
|
cache_read_tokens,
|
||||||
|
cache_write_tokens,
|
||||||
|
cache_miss_tokens,
|
||||||
|
)
|
||||||
)
|
)
|
||||||
resolved_input = input_tokens or 0
|
resolved_input = input_tokens or 0
|
||||||
resolved_output = output_tokens or 0
|
resolved_output = output_tokens or 0
|
||||||
@@ -123,12 +297,32 @@ class UsageMiddleware(AgentMiddleware):
|
|||||||
if total_tokens is not None
|
if total_tokens is not None
|
||||||
else resolved_input + resolved_output
|
else resolved_input + resolved_output
|
||||||
)
|
)
|
||||||
|
resolved_cache_read = cache_read_tokens or 0
|
||||||
|
resolved_cache_write = cache_write_tokens or 0
|
||||||
|
uncached_input_tokens = (
|
||||||
|
cache_miss_tokens
|
||||||
|
if cache_miss_tokens is not None
|
||||||
|
else max(
|
||||||
|
resolved_input - resolved_cache_read - resolved_cache_write,
|
||||||
|
0,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
cache_hit_ratio = (
|
||||||
|
resolved_cache_read / resolved_input
|
||||||
|
if has_cache_usage and resolved_input
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"has_usage": has_usage,
|
"has_usage": has_usage,
|
||||||
|
"cache_usage_available": has_cache_usage,
|
||||||
"input_tokens": resolved_input,
|
"input_tokens": resolved_input,
|
||||||
"output_tokens": resolved_output,
|
"output_tokens": resolved_output,
|
||||||
"total_tokens": resolved_total,
|
"total_tokens": resolved_total,
|
||||||
|
"cache_read_input_tokens": resolved_cache_read,
|
||||||
|
"cache_write_input_tokens": resolved_cache_write,
|
||||||
|
"uncached_input_tokens": uncached_input_tokens,
|
||||||
|
"cache_hit_ratio": cache_hit_ratio,
|
||||||
}
|
}
|
||||||
|
|
||||||
async def awrap_model_call(
|
async def awrap_model_call(
|
||||||
@@ -157,9 +351,14 @@ class UsageMiddleware(AgentMiddleware):
|
|||||||
if ai_message
|
if ai_message
|
||||||
else {
|
else {
|
||||||
"has_usage": False,
|
"has_usage": False,
|
||||||
|
"cache_usage_available": False,
|
||||||
"input_tokens": 0,
|
"input_tokens": 0,
|
||||||
"output_tokens": 0,
|
"output_tokens": 0,
|
||||||
"total_tokens": 0,
|
"total_tokens": 0,
|
||||||
|
"cache_read_input_tokens": 0,
|
||||||
|
"cache_write_input_tokens": 0,
|
||||||
|
"uncached_input_tokens": 0,
|
||||||
|
"cache_hit_ratio": None,
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
context_window_tokens = self._extract_context_window_tokens(request.model)
|
context_window_tokens = self._extract_context_window_tokens(request.model)
|
||||||
|
|||||||
@@ -1349,6 +1349,33 @@ class MessageChain(ChainBase):
|
|||||||
f"排队消息数: {status.get('pending_messages', 0)}",
|
f"排队消息数: {status.get('pending_messages', 0)}",
|
||||||
f"最后更新: {status.get('last_updated_at') or '暂无'}",
|
f"最后更新: {status.get('last_updated_at') or '暂无'}",
|
||||||
]
|
]
|
||||||
|
if status.get("cache_usage_available"):
|
||||||
|
last_cache_ratio = status.get("last_cache_hit_ratio")
|
||||||
|
total_cache_ratio = status.get("total_cache_hit_ratio")
|
||||||
|
lines.insert(
|
||||||
|
6,
|
||||||
|
"最近一次缓存: "
|
||||||
|
f"命中 {cls._format_token_count(status.get('last_cache_read_input_tokens'))} / "
|
||||||
|
f"写入 {cls._format_token_count(status.get('last_cache_write_input_tokens'))} / "
|
||||||
|
f"未命中 {cls._format_token_count(status.get('last_uncached_input_tokens'))}"
|
||||||
|
+ (
|
||||||
|
f" ({last_cache_ratio * 100:.2f}%)"
|
||||||
|
if last_cache_ratio is not None
|
||||||
|
else ""
|
||||||
|
),
|
||||||
|
)
|
||||||
|
lines.insert(
|
||||||
|
8,
|
||||||
|
"当前会话累计缓存: "
|
||||||
|
f"命中 {cls._format_token_count(status.get('total_cache_read_input_tokens'))} / "
|
||||||
|
f"写入 {cls._format_token_count(status.get('total_cache_write_input_tokens'))} / "
|
||||||
|
f"未命中 {cls._format_token_count(status.get('total_uncached_input_tokens'))}"
|
||||||
|
+ (
|
||||||
|
f" ({total_cache_ratio * 100:.2f}%)"
|
||||||
|
if total_cache_ratio is not None
|
||||||
|
else ""
|
||||||
|
),
|
||||||
|
)
|
||||||
return "\n".join(lines)
|
return "\n".join(lines)
|
||||||
|
|
||||||
def remote_session_status(
|
def remote_session_status(
|
||||||
|
|||||||
@@ -119,6 +119,11 @@ class AgentTokensUsageEventData(BaseEventData):
|
|||||||
input_tokens: int = Field(default=0, description="输入 tokens")
|
input_tokens: int = Field(default=0, description="输入 tokens")
|
||||||
output_tokens: int = Field(default=0, description="输出 tokens")
|
output_tokens: int = Field(default=0, description="输出 tokens")
|
||||||
total_tokens: int = Field(default=0, description="总 tokens")
|
total_tokens: int = Field(default=0, description="总 tokens")
|
||||||
|
cache_read_input_tokens: int = Field(default=0, description="从提示词缓存读取的输入 tokens")
|
||||||
|
cache_write_input_tokens: int = Field(default=0, description="写入提示词缓存的输入 tokens")
|
||||||
|
uncached_input_tokens: int = Field(default=0, description="未命中缓存的输入 tokens")
|
||||||
|
cache_hit_ratio: Optional[float] = Field(default=None, description="提示词缓存命中率")
|
||||||
|
cache_usage_available: bool = Field(default=False, description="供应商是否返回缓存用量明细")
|
||||||
model_call_count: int = Field(default=0, description="模型调用次数")
|
model_call_count: int = Field(default=0, description="模型调用次数")
|
||||||
success: bool = Field(default=False, description="Agent 执行是否成功")
|
success: bool = Field(default=False, description="Agent 执行是否成功")
|
||||||
error: Optional[str] = Field(default=None, description="失败原因")
|
error: Optional[str] = Field(default=None, description="失败原因")
|
||||||
|
|||||||
@@ -53,8 +53,8 @@ def test_activity_log_index_counts_entries_without_body(tmp_path):
|
|||||||
assert "整理了电影文件" not in json.dumps(index, ensure_ascii=False)
|
assert "整理了电影文件" not in json.dumps(index, ensure_ascii=False)
|
||||||
|
|
||||||
|
|
||||||
def test_activity_log_prompt_injects_index_not_full_log(tmp_path):
|
def test_activity_log_prompt_is_stable_and_excludes_log_index(tmp_path):
|
||||||
"""ActivityLogMiddleware 注入系统提示词时不应携带完整活动日志正文。"""
|
"""ActivityLogMiddleware 的系统提示词不应随活动日志索引变化。"""
|
||||||
date_str = datetime.now().strftime("%Y-%m-%d")
|
date_str = datetime.now().strftime("%Y-%m-%d")
|
||||||
_write_activity_log(
|
_write_activity_log(
|
||||||
tmp_path,
|
tmp_path,
|
||||||
@@ -74,10 +74,15 @@ def test_activity_log_prompt_injects_index_not_full_log(tmp_path):
|
|||||||
|
|
||||||
modified = middleware.modify_request(request)
|
modified = middleware.modify_request(request)
|
||||||
system_text = str(modified.system_message.content)
|
system_text = str(modified.system_message.content)
|
||||||
|
stable_prompt = middleware._format_activity_log(
|
||||||
|
{"2099-12-31": "999 条活动记录"}
|
||||||
|
)
|
||||||
|
|
||||||
assert "1 条活动记录" in system_text
|
assert "1 条活动记录" not in system_text
|
||||||
|
assert date_str not in system_text
|
||||||
assert "这是一条不应默认进入上下文的活动正文" not in system_text
|
assert "这是一条不应默认进入上下文的活动正文" not in system_text
|
||||||
assert "query_activity_log" in system_text
|
assert "query_activity_log" in system_text
|
||||||
|
assert middleware._format_activity_log(state_update["activity_log_contents"]) == stable_prompt
|
||||||
|
|
||||||
|
|
||||||
def test_activity_log_abefore_agent_refreshes_existing_state(tmp_path):
|
def test_activity_log_abefore_agent_refreshes_existing_state(tmp_path):
|
||||||
|
|||||||
@@ -0,0 +1,225 @@
|
|||||||
|
import pytest
|
||||||
|
from langchain_core.messages import AIMessage
|
||||||
|
|
||||||
|
from app.agent import MoviePilotAgent
|
||||||
|
from app.agent.llm.helper import LLMHelper
|
||||||
|
from app.agent.llm.provider import LLMProviderManager
|
||||||
|
from app.agent.middleware.usage import UsageMiddleware
|
||||||
|
from app.chain.message import MessageChain
|
||||||
|
|
||||||
|
|
||||||
|
def test_usage_extracts_normalized_cache_details() -> None:
|
||||||
|
"""标准 usage_metadata 应解析缓存读取、写入和未命中 tokens。"""
|
||||||
|
usage = UsageMiddleware._extract_usage(
|
||||||
|
AIMessage(
|
||||||
|
content="ok",
|
||||||
|
usage_metadata={
|
||||||
|
"input_tokens": 1200,
|
||||||
|
"output_tokens": 100,
|
||||||
|
"total_tokens": 1300,
|
||||||
|
"input_token_details": {
|
||||||
|
"cache_read": 700,
|
||||||
|
"cache_creation": 0,
|
||||||
|
"ephemeral_5m_input_tokens": 300,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
assert usage["cache_usage_available"]
|
||||||
|
assert usage["cache_read_input_tokens"] == 700
|
||||||
|
assert usage["cache_write_input_tokens"] == 300
|
||||||
|
assert usage["uncached_input_tokens"] == 200
|
||||||
|
assert usage["cache_hit_ratio"] == pytest.approx(700 / 1200)
|
||||||
|
|
||||||
|
|
||||||
|
def test_usage_extracts_deepseek_cache_hit_and_miss_tokens() -> None:
|
||||||
|
"""DeepSeek 原始 usage 应保留其显式缓存命中与未命中字段。"""
|
||||||
|
usage = UsageMiddleware._extract_usage(
|
||||||
|
AIMessage(
|
||||||
|
content="ok",
|
||||||
|
response_metadata={
|
||||||
|
"token_usage": {
|
||||||
|
"prompt_tokens": 1000,
|
||||||
|
"completion_tokens": 50,
|
||||||
|
"total_tokens": 1050,
|
||||||
|
"prompt_cache_hit_tokens": 800,
|
||||||
|
"prompt_cache_miss_tokens": 200,
|
||||||
|
}
|
||||||
|
},
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
assert usage["cache_usage_available"]
|
||||||
|
assert usage["cache_read_input_tokens"] == 800
|
||||||
|
assert usage["cache_write_input_tokens"] == 0
|
||||||
|
assert usage["uncached_input_tokens"] == 200
|
||||||
|
assert usage["cache_hit_ratio"] == pytest.approx(0.8)
|
||||||
|
|
||||||
|
|
||||||
|
def test_session_usage_aggregates_and_formats_cache_statistics() -> None:
|
||||||
|
"""会话状态应聚合缓存统计并在状态文本中展示。"""
|
||||||
|
agent = MoviePilotAgent(session_id="cache-session", user_id="user-1")
|
||||||
|
agent._record_usage(
|
||||||
|
{
|
||||||
|
"has_usage": True,
|
||||||
|
"cache_usage_available": True,
|
||||||
|
"input_tokens": 100,
|
||||||
|
"output_tokens": 10,
|
||||||
|
"total_tokens": 110,
|
||||||
|
"cache_read_input_tokens": 60,
|
||||||
|
"cache_write_input_tokens": 20,
|
||||||
|
"uncached_input_tokens": 20,
|
||||||
|
"cache_hit_ratio": 0.6,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
agent._record_usage(
|
||||||
|
{
|
||||||
|
"has_usage": True,
|
||||||
|
"cache_usage_available": True,
|
||||||
|
"input_tokens": 50,
|
||||||
|
"output_tokens": 5,
|
||||||
|
"total_tokens": 55,
|
||||||
|
"cache_read_input_tokens": 20,
|
||||||
|
"cache_write_input_tokens": 0,
|
||||||
|
"uncached_input_tokens": 30,
|
||||||
|
"cache_hit_ratio": 0.4,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
status = agent.get_session_status()
|
||||||
|
status.update({"is_processing": False, "pending_messages": 0})
|
||||||
|
status_text = MessageChain._format_session_status_text(status)
|
||||||
|
|
||||||
|
assert status["total_cache_read_input_tokens"] == 80
|
||||||
|
assert status["total_cache_write_input_tokens"] == 20
|
||||||
|
assert status["total_uncached_input_tokens"] == 50
|
||||||
|
assert status["total_cache_hit_ratio"] == pytest.approx(80 / 150)
|
||||||
|
assert "当前会话累计缓存: 命中 80 / 写入 20 / 未命中 50 (53.33%)" in status_text
|
||||||
|
|
||||||
|
|
||||||
|
def test_prompt_cache_key_is_stable_and_private() -> None:
|
||||||
|
"""提示词缓存键应在同一会话内稳定且不暴露原始标识。"""
|
||||||
|
first = MoviePilotAgent(session_id="private-session", user_id="private-user")
|
||||||
|
second = MoviePilotAgent(session_id="private-session", user_id="private-user")
|
||||||
|
other = MoviePilotAgent(session_id="other-session", user_id="private-user")
|
||||||
|
|
||||||
|
cache_key = first._build_prompt_cache_key()
|
||||||
|
|
||||||
|
assert cache_key == second._build_prompt_cache_key()
|
||||||
|
assert cache_key != other._build_prompt_cache_key()
|
||||||
|
assert "private-session" not in cache_key
|
||||||
|
assert "private-user" not in cache_key
|
||||||
|
|
||||||
|
|
||||||
|
def test_openai_prompt_cache_options_only_target_official_endpoints() -> None:
|
||||||
|
"""OpenAI 专属缓存参数不得泄露到第三方兼容端点。"""
|
||||||
|
headers, kwargs = LLMHelper._build_openai_prompt_cache_options(
|
||||||
|
provider="openai",
|
||||||
|
base_url="https://api.openai.com/v1",
|
||||||
|
use_responses_api=False,
|
||||||
|
prompt_cache_key="cache-key",
|
||||||
|
default_headers={"User-Agent": "MoviePilot"},
|
||||||
|
model_kwargs={"extra_body": {"existing": True}},
|
||||||
|
)
|
||||||
|
_, compatible_kwargs = LLMHelper._build_openai_prompt_cache_options(
|
||||||
|
provider="openai",
|
||||||
|
base_url="https://api.openai.com.example/v1",
|
||||||
|
use_responses_api=False,
|
||||||
|
prompt_cache_key="cache-key",
|
||||||
|
default_headers=None,
|
||||||
|
model_kwargs={},
|
||||||
|
)
|
||||||
|
|
||||||
|
assert headers == {"User-Agent": "MoviePilot"}
|
||||||
|
assert kwargs["extra_body"] == {
|
||||||
|
"existing": True,
|
||||||
|
"prompt_cache_key": "cache-key",
|
||||||
|
}
|
||||||
|
assert compatible_kwargs == {}
|
||||||
|
|
||||||
|
|
||||||
|
def test_xai_chat_completions_uses_stable_conversation_header() -> None:
|
||||||
|
"""xAI Chat Completions 应使用官方会话缓存路由请求头。"""
|
||||||
|
headers, kwargs = LLMHelper._build_openai_prompt_cache_options(
|
||||||
|
provider="xai",
|
||||||
|
base_url="https://api.x.ai/v1",
|
||||||
|
use_responses_api=None,
|
||||||
|
prompt_cache_key="cache-key",
|
||||||
|
default_headers=None,
|
||||||
|
model_kwargs={},
|
||||||
|
)
|
||||||
|
|
||||||
|
assert headers == {"x-grok-conv-id": "cache-key"}
|
||||||
|
assert kwargs == {}
|
||||||
|
|
||||||
|
|
||||||
|
def test_prompt_cache_adapter_preserves_control_after_tool_binding() -> None:
|
||||||
|
"""Provider 缓存参数应在 Agent 最终绑定工具时仍然存在。"""
|
||||||
|
|
||||||
|
class FakeModel:
|
||||||
|
"""模拟通过 bind_tools 再调用 bind 的 LangChain 模型。"""
|
||||||
|
|
||||||
|
def bind(self, **kwargs):
|
||||||
|
"""返回最终绑定参数。"""
|
||||||
|
return kwargs
|
||||||
|
|
||||||
|
def bind_tools(self, tools, **kwargs):
|
||||||
|
"""模拟模型的工具绑定流程。"""
|
||||||
|
return self.bind(tools=tools, **kwargs)
|
||||||
|
|
||||||
|
cached_model_cls = LLMHelper._with_prompt_cache_control(
|
||||||
|
FakeModel,
|
||||||
|
{"type": "default"},
|
||||||
|
)
|
||||||
|
|
||||||
|
result = cached_model_cls().bind_tools([{"name": "tool"}])
|
||||||
|
|
||||||
|
assert result["tools"] == [{"name": "tool"}]
|
||||||
|
assert result["cache_control"] == {"type": "default"}
|
||||||
|
|
||||||
|
|
||||||
|
def test_anthropic_cache_control_only_targets_official_endpoint() -> None:
|
||||||
|
"""Anthropic 原生缓存控制不得发送到第三方兼容端点。"""
|
||||||
|
official = LLMHelper._use_anthropic_prompt_cache(
|
||||||
|
provider="anthropic",
|
||||||
|
runtime={
|
||||||
|
"runtime": "anthropic_compatible",
|
||||||
|
"base_url": "https://api.anthropic.com/v1",
|
||||||
|
},
|
||||||
|
prompt_cache_key="cache-key",
|
||||||
|
)
|
||||||
|
compatible = LLMHelper._use_anthropic_prompt_cache(
|
||||||
|
provider="minimax",
|
||||||
|
runtime={
|
||||||
|
"runtime": "anthropic_compatible",
|
||||||
|
"base_url": "https://api.minimax.io/anthropic/v1",
|
||||||
|
},
|
||||||
|
prompt_cache_key="cache-key",
|
||||||
|
)
|
||||||
|
|
||||||
|
assert official
|
||||||
|
assert not compatible
|
||||||
|
|
||||||
|
|
||||||
|
def test_provider_metadata_declares_prompt_cache_without_model_allowlist() -> None:
|
||||||
|
"""Provider 应通过模型能力元数据判断缓存支持,不依赖模型 ID 白名单。"""
|
||||||
|
assert LLMProviderManager._metadata_supports_prompt_cache(
|
||||||
|
{"cost": {"input": 1, "cache_read": 0.1}}
|
||||||
|
)
|
||||||
|
assert LLMProviderManager._metadata_supports_prompt_cache(
|
||||||
|
{"capabilities": {"prompt_cache": True}}
|
||||||
|
)
|
||||||
|
assert not LLMProviderManager._metadata_supports_prompt_cache(
|
||||||
|
{"cost": {"input": 1, "output": 2}}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_bedrock_model_metadata_candidates_remove_region_prefix() -> None:
|
||||||
|
"""Bedrock 跨区域模型应自动回落到基础模型的能力元数据。"""
|
||||||
|
candidates = LLMProviderManager._models_dev_model_candidates(
|
||||||
|
"amazon-bedrock",
|
||||||
|
"us.vendor.model-version",
|
||||||
|
)
|
||||||
|
|
||||||
|
assert candidates == ("us.vendor.model-version", "vendor.model-version")
|
||||||
@@ -89,6 +89,7 @@ def test_initialize_llm_uses_chain_event_selection(monkeypatch) -> None:
|
|||||||
thinking_level="xhigh",
|
thinking_level="xhigh",
|
||||||
api_protocol="auto",
|
api_protocol="auto",
|
||||||
web_search_mode="local",
|
web_search_mode="local",
|
||||||
|
prompt_cache_key=agent._build_prompt_cache_key(),
|
||||||
)
|
)
|
||||||
assert agent._llm_provider_selection["selected_provider_id"] == "provider-1"
|
assert agent._llm_provider_selection["selected_provider_id"] == "provider-1"
|
||||||
|
|
||||||
@@ -119,6 +120,11 @@ def test_execute_agent_broadcasts_usage_on_success() -> None:
|
|||||||
"input_tokens": 12,
|
"input_tokens": 12,
|
||||||
"output_tokens": 8,
|
"output_tokens": 8,
|
||||||
"total_tokens": 20,
|
"total_tokens": 20,
|
||||||
|
"cache_usage_available": True,
|
||||||
|
"cache_read_input_tokens": 8,
|
||||||
|
"cache_write_input_tokens": 2,
|
||||||
|
"uncached_input_tokens": 2,
|
||||||
|
"cache_hit_ratio": 8 / 12,
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
return _FakeAgent([AIMessage(content="ok")])
|
return _FakeAgent([AIMessage(content="ok")])
|
||||||
@@ -138,6 +144,11 @@ def test_execute_agent_broadcasts_usage_on_success() -> None:
|
|||||||
assert usage.input_tokens == 12
|
assert usage.input_tokens == 12
|
||||||
assert usage.output_tokens == 8
|
assert usage.output_tokens == 8
|
||||||
assert usage.total_tokens == 20
|
assert usage.total_tokens == 20
|
||||||
|
assert usage.cache_read_input_tokens == 8
|
||||||
|
assert usage.cache_write_input_tokens == 2
|
||||||
|
assert usage.uncached_input_tokens == 2
|
||||||
|
assert usage.cache_hit_ratio == 8 / 12
|
||||||
|
assert usage.cache_usage_available
|
||||||
|
|
||||||
|
|
||||||
def test_execute_agent_broadcasts_usage_on_failure() -> None:
|
def test_execute_agent_broadcasts_usage_on_failure() -> None:
|
||||||
|
|||||||
Reference in New Issue
Block a user