refactor(usage): 统一 cache token 提取逻辑,新增请求缓存指纹记录

- 新增 extract_cache_read_tokens() 兼容 OpenAI/Claude/Gemini 多种字段命名
- parsers/stream_processor/cli_event_mixin 统一使用提取函数替换内联逻辑
- 同时兼容 prompt_tokens/completion_tokens (OpenAI) 和 input_tokens/output_tokens (Claude)
- 新增 cache_fingerprint 模块,在 telemetry 记录时自动计算并附带请求缓存指纹
- 新增对应单元测试
This commit is contained in:
fawney19
2026-03-17 02:34:32 +08:00
parent d2f1431269
commit b6cc0bc3a7
9 changed files with 567 additions and 28 deletions

View File

@@ -11,6 +11,7 @@ from src.api.handlers.base.parsers import get_parser_for_format
from src.api.handlers.base.stream_context import StreamContext
from src.api.handlers.base.utils import get_format_converter_registry
from src.core.logger import logger
from src.core.usage_tokens import extract_cache_creation_tokens, extract_cache_read_tokens
from src.services.provider.behavior import get_provider_behavior
from src.utils.sse_parser import SSEEventParser
@@ -261,12 +262,10 @@ class CliEventMixin:
}
if usage and isinstance(usage, dict):
new_input = usage.get("input_tokens", 0) or 0
new_output = usage.get("output_tokens", 0) or 0
new_cached = usage.get("cache_read_tokens") or usage.get("cache_read_input_tokens") or 0
new_cache_creation = (
usage.get("cache_creation_tokens") or usage.get("cache_creation_input_tokens") or 0
)
new_input = usage.get("input_tokens") or usage.get("prompt_tokens") or 0
new_output = usage.get("output_tokens") or usage.get("completion_tokens") or 0
new_cached = extract_cache_read_tokens(usage)
new_cache_creation = extract_cache_creation_tokens(usage)
# 取最大值更新(与 _process_event_data 相同的策略)
if new_input > ctx.input_tokens:

View File

@@ -17,7 +17,7 @@ from src.api.handlers.base.response_parser import (
# is_cli_format 权威定义在 core 层
from src.core.api_format import is_cli_format
from src.core.usage_tokens import extract_cache_creation_tokens
from src.core.usage_tokens import extract_cache_creation_tokens, extract_cache_read_tokens
def _check_nested_error(response: dict[str, Any]) -> tuple[bool, dict[str, Any] | None]:
@@ -208,6 +208,7 @@ class OpenAIResponseParser(ResponseParser):
usage = response.get("usage") or {}
result.input_tokens = usage.get("prompt_tokens", 0)
result.output_tokens = usage.get("completion_tokens", 0)
result.cache_read_tokens = extract_cache_read_tokens(usage)
# 检查错误(支持嵌套错误格式)
is_error, error_info = _check_nested_error(response)
@@ -225,7 +226,7 @@ class OpenAIResponseParser(ResponseParser):
"input_tokens": usage.get("prompt_tokens", 0),
"output_tokens": usage.get("completion_tokens", 0),
"cache_creation_tokens": 0,
"cache_read_tokens": 0,
"cache_read_tokens": extract_cache_read_tokens(usage),
}
def extract_text_content(self, response: dict[str, Any]) -> str:
@@ -331,12 +332,8 @@ class OpenAICliResponseParser(OpenAIResponseParser):
return {
"input_tokens": int(input_tokens),
"output_tokens": int(output_tokens),
"cache_creation_tokens": int(
usage.get("cache_creation_input_tokens") or usage.get("cache_creation_tokens") or 0
),
"cache_read_tokens": int(
usage.get("cache_read_input_tokens") or usage.get("cache_read_tokens") or 0
),
"cache_creation_tokens": extract_cache_creation_tokens(usage),
"cache_read_tokens": extract_cache_read_tokens(usage),
}
@staticmethod

View File

@@ -44,6 +44,7 @@ from src.core.exceptions import (
ProviderTimeoutException,
)
from src.core.logger import logger
from src.core.usage_tokens import extract_cache_creation_tokens, extract_cache_read_tokens
from src.models.database import Provider, ProviderEndpoint
from src.services.provider.behavior import get_provider_behavior
from src.utils.perf import PerfRecorder
@@ -327,12 +328,10 @@ class StreamProcessor:
}
if usage and isinstance(usage, dict):
new_input = usage.get("input_tokens", 0) or 0
new_output = usage.get("output_tokens", 0) or 0
new_cached = usage.get("cache_read_tokens") or usage.get("cache_read_input_tokens") or 0
new_cache_creation = (
usage.get("cache_creation_tokens") or usage.get("cache_creation_input_tokens") or 0
)
new_input = usage.get("input_tokens") or usage.get("prompt_tokens") or 0
new_output = usage.get("output_tokens") or usage.get("completion_tokens") or 0
new_cached = extract_cache_read_tokens(usage)
new_cache_creation = extract_cache_creation_tokens(usage)
if new_input > ctx.input_tokens:
ctx.input_tokens = new_input

View File

@@ -107,7 +107,39 @@ def extract_cache_creation_tokens_detail(usage: dict[str, Any]) -> tuple[int, in
return old, 0, 0
def extract_cache_read_tokens(usage: dict[str, Any]) -> int:
"""
提取缓存读取 tokens兼容多种 OpenAI / Claude / Gemini 字段命名)。
优先级:
1. 直接字段cache_read_input_tokens / cache_read_tokens
2. OpenAI Responses: input_tokens_details.cached_tokens
3. OpenAI Chat: prompt_tokens_details.cached_tokens
4. 通用回退cached_tokens
说明:
- 只要检测到更高优先级字段存在,即便值为 0 也不继续回退,
避免被较低优先级字段覆盖。
"""
if "cache_read_input_tokens" in usage:
return int(usage.get("cache_read_input_tokens", 0) or 0)
if "cache_read_tokens" in usage:
return int(usage.get("cache_read_tokens", 0) or 0)
input_details = usage.get("input_tokens_details")
if isinstance(input_details, dict) and "cached_tokens" in input_details:
return int(input_details.get("cached_tokens", 0) or 0)
prompt_details = usage.get("prompt_tokens_details")
if isinstance(prompt_details, dict) and "cached_tokens" in prompt_details:
return int(prompt_details.get("cached_tokens", 0) or 0)
return int(usage.get("cached_tokens", 0) or 0)
__all__ = [
"extract_cache_creation_tokens",
"extract_cache_creation_tokens_detail",
"extract_cache_read_tokens",
]

View File

@@ -0,0 +1,143 @@
"""Helpers for stable outbound request cache fingerprints."""
from __future__ import annotations
import hashlib
import json
from collections.abc import Mapping
from enum import Enum
from typing import Any
# model 和 prompt_cache_key 统一包含在所有格式中,无需运行时动态添加
_CACHE_RELEVANT_FIELDS_BY_FORMAT: dict[str, frozenset[str]] = {
"openai:chat": frozenset({"model", "messages", "tools", "tool_choice", "prompt_cache_key"}),
"openai:cli": frozenset(
{"model", "input", "instructions", "tools", "tool_choice", "prompt_cache_key"}
),
"openai:compact": frozenset(
{"model", "input", "instructions", "tools", "tool_choice", "prompt_cache_key"}
),
"claude:chat": frozenset(
{"model", "system", "messages", "tools", "tool_choice", "prompt_cache_key"}
),
"gemini:chat": frozenset(
{
"model",
"contents",
"system_instruction",
"systemInstruction",
"tools",
"tool_config",
"toolConfig",
"generation_config",
"generationConfig",
"prompt_cache_key",
}
),
}
def _normalize_for_hash(value: Any) -> Any:
if value is None or isinstance(value, (str, int, float, bool)):
return value
if isinstance(value, bytes):
return value.decode("utf-8", errors="replace")
if isinstance(value, Enum):
return _normalize_for_hash(value.value)
if isinstance(value, Mapping):
return {str(key): _normalize_for_hash(item) for key, item in value.items()}
if isinstance(value, (list, tuple)):
return [_normalize_for_hash(item) for item in value]
if isinstance(value, (set, frozenset)):
normalized = [_normalize_for_hash(item) for item in value]
return sorted(normalized, key=_stable_json_dumps)
return str(value)
def _stable_json_dumps(value: Any) -> str:
return json.dumps(
_normalize_for_hash(value),
ensure_ascii=False,
sort_keys=True,
separators=(",", ":"),
)
def _hash_json_payload(value: Any) -> tuple[str, int]:
"""Return (sha256_hex, json_byte_length) for the canonicalized JSON."""
payload = _stable_json_dumps(value).encode("utf-8")
return hashlib.sha256(payload).hexdigest(), len(payload)
def _normalize_provider_api_format(provider_api_format: str | None) -> str | None:
normalized = str(provider_api_format or "").strip().lower()
return normalized or None
def _get_prompt_cache_key(payload: Any) -> str | None:
if not isinstance(payload, Mapping):
return None
prompt_cache_key = str(payload.get("prompt_cache_key") or "").strip()
return prompt_cache_key or None
def _extract_cache_relevant_payload(
payload: Any, provider_api_format: str | None
) -> tuple[Any, list[str]]:
if not isinstance(payload, Mapping):
return payload, []
fields = _CACHE_RELEVANT_FIELDS_BY_FORMAT.get(provider_api_format or "")
if not fields:
# 未知格式:整个 payload 参与哈希
top_level_keys = sorted(str(key) for key in payload.keys())
return dict(payload), top_level_keys
subset = {field: payload[field] for field in fields if field in payload}
if not subset:
top_level_keys = sorted(str(key) for key in payload.keys())
return dict(payload), top_level_keys
return subset, sorted(subset.keys())
def build_request_cache_fingerprint(
provider_request_body: Any,
*,
provider_api_format: str | None = None,
) -> dict[str, Any] | None:
"""Build stable hashes for the final outbound payload and its cache-relevant subset."""
if provider_request_body is None:
return None
normalized_format = _normalize_provider_api_format(provider_api_format)
payload_sha256, payload_bytes = _hash_json_payload(provider_request_body)
cache_relevant_payload, cache_relevant_keys = _extract_cache_relevant_payload(
provider_request_body,
normalized_format,
)
cache_relevant_sha256, cache_relevant_bytes = _hash_json_payload(cache_relevant_payload)
top_level_keys = []
if isinstance(provider_request_body, Mapping):
top_level_keys = sorted(str(key) for key in provider_request_body.keys())
fingerprint: dict[str, Any] = {
"version": 1,
"provider_api_format": normalized_format,
"payload_sha256": payload_sha256,
"payload_bytes": payload_bytes,
"cache_relevant_sha256": cache_relevant_sha256,
"cache_relevant_bytes": cache_relevant_bytes,
"top_level_keys": top_level_keys,
"cache_relevant_keys": cache_relevant_keys,
}
prompt_cache_key = _get_prompt_cache_key(provider_request_body)
if prompt_cache_key:
fingerprint["prompt_cache_key"] = prompt_cache_key
return fingerprint
__all__ = ["build_request_cache_fingerprint"]

View File

@@ -38,6 +38,7 @@ METADATA_KEEP_KEYS: frozenset[str] = frozenset(
{
"billing_snapshot",
"billing_updated_at",
"cache_fingerprint",
"perf",
"pool_summary",
"scheduling_audit",

View File

@@ -12,6 +12,7 @@ from typing import Any
from sqlalchemy.orm import Session
from src.core.logger import logger
from src.services.provider.cache_fingerprint import build_request_cache_fingerprint
from src.services.system.audit import audit_service
from src.services.usage.service import UsageService
@@ -30,6 +31,42 @@ class MessageTelemetry:
self.request_id = request_id
self.client_ip = client_ip
def _build_usage_metadata(
self,
*,
request_metadata: dict[str, Any] | None = None,
response_metadata: dict[str, Any] | None = None,
provider_request_body: Any | None = None,
provider_api_format: str | None = None,
) -> dict[str, Any] | None:
metadata: dict[str, Any] | None = None
if request_metadata:
metadata = dict(request_metadata)
if response_metadata:
metadata.setdefault("response", response_metadata)
elif response_metadata:
metadata = dict(response_metadata)
fingerprint = build_request_cache_fingerprint(
provider_request_body,
provider_api_format=provider_api_format,
)
if fingerprint:
if metadata is None:
metadata = {}
metadata["cache_fingerprint"] = fingerprint
logger.debug(
"[Telemetry] cache fingerprint: request_id={}, format={}, payload_sha256={}, cache_sha256={}, prompt_cache_key_present={}",
self.request_id,
fingerprint.get("provider_api_format"),
str(fingerprint.get("payload_sha256") or "")[:12],
str(fingerprint.get("cache_relevant_sha256") or "")[:12],
bool(fingerprint.get("prompt_cache_key")),
)
return metadata
async def calculate_cost(
self,
provider: str,
@@ -96,12 +133,12 @@ class MessageTelemetry:
# 请求元数据(用于性能与调试记录)
request_metadata: dict[str, Any] | None = None,
) -> float:
metadata = response_metadata
if request_metadata:
merged = dict(request_metadata)
if response_metadata:
merged.setdefault("response", response_metadata)
metadata = merged
metadata = self._build_usage_metadata(
request_metadata=request_metadata,
response_metadata=response_metadata,
provider_request_body=provider_request_body,
provider_api_format=endpoint_api_format or api_format,
)
usage = await UsageService.record_usage(
db=self.db,
@@ -223,6 +260,12 @@ class MessageTelemetry:
self.request_id,
)
metadata = self._build_usage_metadata(
request_metadata=request_metadata,
provider_request_body=provider_request_body,
provider_api_format=endpoint_api_format or api_format,
)
await UsageService.record_usage(
db=self.db,
user=self.user,
@@ -261,7 +304,7 @@ class MessageTelemetry:
# 模型映射信息
target_model=target_model,
# 请求元数据
metadata=request_metadata,
metadata=metadata,
)
async def record_cancelled(
@@ -307,6 +350,11 @@ class MessageTelemetry:
客户端主动断开连接不算系统失败,使用 cancelled 状态。
"""
provider_name = provider or "unknown"
metadata = self._build_usage_metadata(
request_metadata=request_metadata,
provider_request_body=provider_request_body,
provider_api_format=endpoint_api_format or api_format,
)
await UsageService.record_usage(
db=self.db,
@@ -345,5 +393,5 @@ class MessageTelemetry:
provider_endpoint_id=provider_endpoint_id,
provider_api_key_id=provider_api_key_id,
target_model=target_model,
metadata=request_metadata,
metadata=metadata,
)