mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-01 17:00:21 +08:00
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:
@@ -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.stream_context import StreamContext
|
||||||
from src.api.handlers.base.utils import get_format_converter_registry
|
from src.api.handlers.base.utils import get_format_converter_registry
|
||||||
from src.core.logger import logger
|
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.services.provider.behavior import get_provider_behavior
|
||||||
from src.utils.sse_parser import SSEEventParser
|
from src.utils.sse_parser import SSEEventParser
|
||||||
|
|
||||||
@@ -261,12 +262,10 @@ class CliEventMixin:
|
|||||||
}
|
}
|
||||||
|
|
||||||
if usage and isinstance(usage, dict):
|
if usage and isinstance(usage, dict):
|
||||||
new_input = usage.get("input_tokens", 0) or 0
|
new_input = usage.get("input_tokens") or usage.get("prompt_tokens") or 0
|
||||||
new_output = usage.get("output_tokens", 0) or 0
|
new_output = usage.get("output_tokens") or usage.get("completion_tokens") or 0
|
||||||
new_cached = usage.get("cache_read_tokens") or usage.get("cache_read_input_tokens") or 0
|
new_cached = extract_cache_read_tokens(usage)
|
||||||
new_cache_creation = (
|
new_cache_creation = extract_cache_creation_tokens(usage)
|
||||||
usage.get("cache_creation_tokens") or usage.get("cache_creation_input_tokens") or 0
|
|
||||||
)
|
|
||||||
|
|
||||||
# 取最大值更新(与 _process_event_data 相同的策略)
|
# 取最大值更新(与 _process_event_data 相同的策略)
|
||||||
if new_input > ctx.input_tokens:
|
if new_input > ctx.input_tokens:
|
||||||
|
|||||||
@@ -17,7 +17,7 @@ from src.api.handlers.base.response_parser import (
|
|||||||
|
|
||||||
# is_cli_format 权威定义在 core 层
|
# is_cli_format 权威定义在 core 层
|
||||||
from src.core.api_format import is_cli_format
|
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]:
|
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 {}
|
usage = response.get("usage") or {}
|
||||||
result.input_tokens = usage.get("prompt_tokens", 0)
|
result.input_tokens = usage.get("prompt_tokens", 0)
|
||||||
result.output_tokens = usage.get("completion_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)
|
is_error, error_info = _check_nested_error(response)
|
||||||
@@ -225,7 +226,7 @@ class OpenAIResponseParser(ResponseParser):
|
|||||||
"input_tokens": usage.get("prompt_tokens", 0),
|
"input_tokens": usage.get("prompt_tokens", 0),
|
||||||
"output_tokens": usage.get("completion_tokens", 0),
|
"output_tokens": usage.get("completion_tokens", 0),
|
||||||
"cache_creation_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:
|
def extract_text_content(self, response: dict[str, Any]) -> str:
|
||||||
@@ -331,12 +332,8 @@ class OpenAICliResponseParser(OpenAIResponseParser):
|
|||||||
return {
|
return {
|
||||||
"input_tokens": int(input_tokens),
|
"input_tokens": int(input_tokens),
|
||||||
"output_tokens": int(output_tokens),
|
"output_tokens": int(output_tokens),
|
||||||
"cache_creation_tokens": int(
|
"cache_creation_tokens": extract_cache_creation_tokens(usage),
|
||||||
usage.get("cache_creation_input_tokens") or usage.get("cache_creation_tokens") or 0
|
"cache_read_tokens": extract_cache_read_tokens(usage),
|
||||||
),
|
|
||||||
"cache_read_tokens": int(
|
|
||||||
usage.get("cache_read_input_tokens") or usage.get("cache_read_tokens") or 0
|
|
||||||
),
|
|
||||||
}
|
}
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
|
|||||||
@@ -44,6 +44,7 @@ from src.core.exceptions import (
|
|||||||
ProviderTimeoutException,
|
ProviderTimeoutException,
|
||||||
)
|
)
|
||||||
from src.core.logger import logger
|
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.models.database import Provider, ProviderEndpoint
|
||||||
from src.services.provider.behavior import get_provider_behavior
|
from src.services.provider.behavior import get_provider_behavior
|
||||||
from src.utils.perf import PerfRecorder
|
from src.utils.perf import PerfRecorder
|
||||||
@@ -327,12 +328,10 @@ class StreamProcessor:
|
|||||||
}
|
}
|
||||||
|
|
||||||
if usage and isinstance(usage, dict):
|
if usage and isinstance(usage, dict):
|
||||||
new_input = usage.get("input_tokens", 0) or 0
|
new_input = usage.get("input_tokens") or usage.get("prompt_tokens") or 0
|
||||||
new_output = usage.get("output_tokens", 0) or 0
|
new_output = usage.get("output_tokens") or usage.get("completion_tokens") or 0
|
||||||
new_cached = usage.get("cache_read_tokens") or usage.get("cache_read_input_tokens") or 0
|
new_cached = extract_cache_read_tokens(usage)
|
||||||
new_cache_creation = (
|
new_cache_creation = extract_cache_creation_tokens(usage)
|
||||||
usage.get("cache_creation_tokens") or usage.get("cache_creation_input_tokens") or 0
|
|
||||||
)
|
|
||||||
|
|
||||||
if new_input > ctx.input_tokens:
|
if new_input > ctx.input_tokens:
|
||||||
ctx.input_tokens = new_input
|
ctx.input_tokens = new_input
|
||||||
|
|||||||
@@ -107,7 +107,39 @@ def extract_cache_creation_tokens_detail(usage: dict[str, Any]) -> tuple[int, in
|
|||||||
return old, 0, 0
|
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__ = [
|
__all__ = [
|
||||||
"extract_cache_creation_tokens",
|
"extract_cache_creation_tokens",
|
||||||
"extract_cache_creation_tokens_detail",
|
"extract_cache_creation_tokens_detail",
|
||||||
|
"extract_cache_read_tokens",
|
||||||
]
|
]
|
||||||
|
|||||||
143
src/services/provider/cache_fingerprint.py
Normal file
143
src/services/provider/cache_fingerprint.py
Normal 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"]
|
||||||
@@ -38,6 +38,7 @@ METADATA_KEEP_KEYS: frozenset[str] = frozenset(
|
|||||||
{
|
{
|
||||||
"billing_snapshot",
|
"billing_snapshot",
|
||||||
"billing_updated_at",
|
"billing_updated_at",
|
||||||
|
"cache_fingerprint",
|
||||||
"perf",
|
"perf",
|
||||||
"pool_summary",
|
"pool_summary",
|
||||||
"scheduling_audit",
|
"scheduling_audit",
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ from typing import Any
|
|||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
from src.core.logger import logger
|
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.system.audit import audit_service
|
||||||
from src.services.usage.service import UsageService
|
from src.services.usage.service import UsageService
|
||||||
|
|
||||||
@@ -30,6 +31,42 @@ class MessageTelemetry:
|
|||||||
self.request_id = request_id
|
self.request_id = request_id
|
||||||
self.client_ip = client_ip
|
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(
|
async def calculate_cost(
|
||||||
self,
|
self,
|
||||||
provider: str,
|
provider: str,
|
||||||
@@ -96,12 +133,12 @@ class MessageTelemetry:
|
|||||||
# 请求元数据(用于性能与调试记录)
|
# 请求元数据(用于性能与调试记录)
|
||||||
request_metadata: dict[str, Any] | None = None,
|
request_metadata: dict[str, Any] | None = None,
|
||||||
) -> float:
|
) -> float:
|
||||||
metadata = response_metadata
|
metadata = self._build_usage_metadata(
|
||||||
if request_metadata:
|
request_metadata=request_metadata,
|
||||||
merged = dict(request_metadata)
|
response_metadata=response_metadata,
|
||||||
if response_metadata:
|
provider_request_body=provider_request_body,
|
||||||
merged.setdefault("response", response_metadata)
|
provider_api_format=endpoint_api_format or api_format,
|
||||||
metadata = merged
|
)
|
||||||
|
|
||||||
usage = await UsageService.record_usage(
|
usage = await UsageService.record_usage(
|
||||||
db=self.db,
|
db=self.db,
|
||||||
@@ -223,6 +260,12 @@ class MessageTelemetry:
|
|||||||
self.request_id,
|
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(
|
await UsageService.record_usage(
|
||||||
db=self.db,
|
db=self.db,
|
||||||
user=self.user,
|
user=self.user,
|
||||||
@@ -261,7 +304,7 @@ class MessageTelemetry:
|
|||||||
# 模型映射信息
|
# 模型映射信息
|
||||||
target_model=target_model,
|
target_model=target_model,
|
||||||
# 请求元数据
|
# 请求元数据
|
||||||
metadata=request_metadata,
|
metadata=metadata,
|
||||||
)
|
)
|
||||||
|
|
||||||
async def record_cancelled(
|
async def record_cancelled(
|
||||||
@@ -307,6 +350,11 @@ class MessageTelemetry:
|
|||||||
客户端主动断开连接不算系统失败,使用 cancelled 状态。
|
客户端主动断开连接不算系统失败,使用 cancelled 状态。
|
||||||
"""
|
"""
|
||||||
provider_name = provider or "unknown"
|
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(
|
await UsageService.record_usage(
|
||||||
db=self.db,
|
db=self.db,
|
||||||
@@ -345,5 +393,5 @@ class MessageTelemetry:
|
|||||||
provider_endpoint_id=provider_endpoint_id,
|
provider_endpoint_id=provider_endpoint_id,
|
||||||
provider_api_key_id=provider_api_key_id,
|
provider_api_key_id=provider_api_key_id,
|
||||||
target_model=target_model,
|
target_model=target_model,
|
||||||
metadata=request_metadata,
|
metadata=metadata,
|
||||||
)
|
)
|
||||||
|
|||||||
113
tests/api/handlers/base/test_openai_usage_parsing.py
Normal file
113
tests/api/handlers/base/test_openai_usage_parsing.py
Normal file
@@ -0,0 +1,113 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from src.api.handlers.base.parsers import OpenAICliResponseParser, OpenAIResponseParser
|
||||||
|
from src.api.handlers.base.response_parser import (
|
||||||
|
ParsedChunk,
|
||||||
|
ParsedResponse,
|
||||||
|
ResponseParser,
|
||||||
|
StreamStats,
|
||||||
|
)
|
||||||
|
from src.api.handlers.base.stream_context import StreamContext
|
||||||
|
from src.api.handlers.base.stream_processor import StreamProcessor
|
||||||
|
|
||||||
|
|
||||||
|
class _DummyParser(ResponseParser):
|
||||||
|
def parse_sse_line(self, line: str, stats: StreamStats) -> ParsedChunk | None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
def parse_response(self, response: dict[str, Any], status_code: int) -> ParsedResponse:
|
||||||
|
return ParsedResponse(raw_response=response, status_code=status_code)
|
||||||
|
|
||||||
|
def extract_usage_from_response(self, response: dict[str, Any]) -> dict[str, int]:
|
||||||
|
return {}
|
||||||
|
|
||||||
|
def extract_text_content(self, response: dict[str, Any]) -> str:
|
||||||
|
return ""
|
||||||
|
|
||||||
|
|
||||||
|
def test_openai_response_parser_extracts_cached_tokens_from_prompt_tokens_details() -> None:
|
||||||
|
parser = OpenAIResponseParser()
|
||||||
|
|
||||||
|
usage = parser.extract_usage_from_response(
|
||||||
|
{
|
||||||
|
"usage": {
|
||||||
|
"prompt_tokens": 120,
|
||||||
|
"completion_tokens": 18,
|
||||||
|
"prompt_tokens_details": {"cached_tokens": 96},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
assert usage["input_tokens"] == 120
|
||||||
|
assert usage["output_tokens"] == 18
|
||||||
|
assert usage["cache_read_tokens"] == 96
|
||||||
|
|
||||||
|
|
||||||
|
def test_openai_cli_response_parser_extracts_cached_tokens_from_input_tokens_details() -> None:
|
||||||
|
parser = OpenAICliResponseParser()
|
||||||
|
|
||||||
|
usage = parser.extract_usage_from_response(
|
||||||
|
{
|
||||||
|
"type": "response.completed",
|
||||||
|
"response": {
|
||||||
|
"usage": {
|
||||||
|
"input_tokens": 2048,
|
||||||
|
"output_tokens": 128,
|
||||||
|
"input_tokens_details": {"cached_tokens": 1792},
|
||||||
|
}
|
||||||
|
},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
assert usage["input_tokens"] == 2048
|
||||||
|
assert usage["output_tokens"] == 128
|
||||||
|
assert usage["cache_read_tokens"] == 1792
|
||||||
|
|
||||||
|
|
||||||
|
def test_stream_processor_extracts_cached_tokens_from_openai_cli_converted_event() -> None:
|
||||||
|
processor = StreamProcessor(request_id="req_test", default_parser=_DummyParser())
|
||||||
|
ctx = StreamContext(model="gpt-5", api_format="openai:chat")
|
||||||
|
|
||||||
|
processor._extract_usage_from_converted_event(
|
||||||
|
ctx,
|
||||||
|
{
|
||||||
|
"type": "response.completed",
|
||||||
|
"response": {
|
||||||
|
"usage": {
|
||||||
|
"input_tokens": 4096,
|
||||||
|
"output_tokens": 64,
|
||||||
|
"input_tokens_details": {"cached_tokens": 3584},
|
||||||
|
}
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"response.completed",
|
||||||
|
)
|
||||||
|
|
||||||
|
assert ctx.input_tokens == 4096
|
||||||
|
assert ctx.output_tokens == 64
|
||||||
|
assert ctx.cached_tokens == 3584
|
||||||
|
|
||||||
|
|
||||||
|
def test_stream_processor_extracts_cached_tokens_from_openai_chat_converted_event() -> None:
|
||||||
|
processor = StreamProcessor(request_id="req_test", default_parser=_DummyParser())
|
||||||
|
ctx = StreamContext(model="gpt-5", api_format="openai:chat")
|
||||||
|
|
||||||
|
processor._extract_usage_from_converted_event(
|
||||||
|
ctx,
|
||||||
|
{
|
||||||
|
"object": "chat.completion.chunk",
|
||||||
|
"choices": [],
|
||||||
|
"usage": {
|
||||||
|
"prompt_tokens": 512,
|
||||||
|
"completion_tokens": 21,
|
||||||
|
"prompt_tokens_details": {"cached_tokens": 480},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"chat.completion.chunk",
|
||||||
|
)
|
||||||
|
|
||||||
|
assert ctx.input_tokens == 512
|
||||||
|
assert ctx.output_tokens == 21
|
||||||
|
assert ctx.cached_tokens == 480
|
||||||
207
tests/services/test_request_cache_fingerprint.py
Normal file
207
tests/services/test_request_cache_fingerprint.py
Normal file
@@ -0,0 +1,207 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from src.config.settings import config
|
||||||
|
from src.services.provider.cache_fingerprint import build_request_cache_fingerprint
|
||||||
|
from src.services.usage._recording_helpers import sanitize_request_metadata
|
||||||
|
from src.services.usage.service import UsageService
|
||||||
|
from src.services.usage.telemetry import MessageTelemetry
|
||||||
|
|
||||||
|
|
||||||
|
def _build_openai_cli_body() -> dict[str, Any]:
|
||||||
|
return {
|
||||||
|
"model": "gpt-5.4",
|
||||||
|
"instructions": "You are precise.",
|
||||||
|
"input": [{"role": "user", "content": [{"type": "input_text", "text": "hello"}]}],
|
||||||
|
"tools": [
|
||||||
|
{
|
||||||
|
"type": "function",
|
||||||
|
"name": "lookup_weather",
|
||||||
|
"parameters": {
|
||||||
|
"type": "object",
|
||||||
|
"required": ["city", "country"],
|
||||||
|
"properties": {
|
||||||
|
"country": {"type": "string"},
|
||||||
|
"city": {"type": "string"},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"temperature": 0.2,
|
||||||
|
"prompt_cache_key": "pcache-123",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_request_cache_fingerprint_is_stable_for_dict_key_reordering() -> None:
|
||||||
|
body_a = _build_openai_cli_body()
|
||||||
|
body_b = {
|
||||||
|
"prompt_cache_key": "pcache-123",
|
||||||
|
"temperature": 0.2,
|
||||||
|
"tools": [
|
||||||
|
{
|
||||||
|
"parameters": {
|
||||||
|
"properties": {
|
||||||
|
"city": {"type": "string"},
|
||||||
|
"country": {"type": "string"},
|
||||||
|
},
|
||||||
|
"required": ["city", "country"],
|
||||||
|
"type": "object",
|
||||||
|
},
|
||||||
|
"name": "lookup_weather",
|
||||||
|
"type": "function",
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"input": [{"content": [{"text": "hello", "type": "input_text"}], "role": "user"}],
|
||||||
|
"instructions": "You are precise.",
|
||||||
|
"model": "gpt-5.4",
|
||||||
|
}
|
||||||
|
|
||||||
|
fingerprint_a = build_request_cache_fingerprint(body_a, provider_api_format="openai:cli")
|
||||||
|
fingerprint_b = build_request_cache_fingerprint(body_b, provider_api_format="openai:cli")
|
||||||
|
|
||||||
|
assert fingerprint_a is not None
|
||||||
|
assert fingerprint_b is not None
|
||||||
|
assert fingerprint_a["payload_sha256"] == fingerprint_b["payload_sha256"]
|
||||||
|
assert fingerprint_a["cache_relevant_sha256"] == fingerprint_b["cache_relevant_sha256"]
|
||||||
|
assert fingerprint_a["prompt_cache_key"] == "pcache-123"
|
||||||
|
assert fingerprint_a["cache_relevant_keys"] == [
|
||||||
|
"input",
|
||||||
|
"instructions",
|
||||||
|
"model",
|
||||||
|
"prompt_cache_key",
|
||||||
|
"tools",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_request_cache_fingerprint_ignores_non_prompt_fields_in_cache_hash() -> None:
|
||||||
|
body_a = _build_openai_cli_body()
|
||||||
|
body_b = _build_openai_cli_body()
|
||||||
|
body_b["temperature"] = 0.9
|
||||||
|
|
||||||
|
fingerprint_a = build_request_cache_fingerprint(body_a, provider_api_format="openai:cli")
|
||||||
|
fingerprint_b = build_request_cache_fingerprint(body_b, provider_api_format="openai:cli")
|
||||||
|
|
||||||
|
assert fingerprint_a is not None
|
||||||
|
assert fingerprint_b is not None
|
||||||
|
assert fingerprint_a["payload_sha256"] != fingerprint_b["payload_sha256"]
|
||||||
|
assert fingerprint_a["cache_relevant_sha256"] == fingerprint_b["cache_relevant_sha256"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_request_cache_fingerprint_tracks_prompt_changes() -> None:
|
||||||
|
body_a = _build_openai_cli_body()
|
||||||
|
body_b = _build_openai_cli_body()
|
||||||
|
body_b["instructions"] = "You are terse."
|
||||||
|
|
||||||
|
fingerprint_a = build_request_cache_fingerprint(body_a, provider_api_format="openai:cli")
|
||||||
|
fingerprint_b = build_request_cache_fingerprint(body_b, provider_api_format="openai:cli")
|
||||||
|
|
||||||
|
assert fingerprint_a is not None
|
||||||
|
assert fingerprint_b is not None
|
||||||
|
assert fingerprint_a["cache_relevant_sha256"] != fingerprint_b["cache_relevant_sha256"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_sanitize_request_metadata_preserves_cache_fingerprint(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
monkeypatch.setattr(config, "usage_metadata_max_bytes", 120, raising=False)
|
||||||
|
|
||||||
|
metadata = {
|
||||||
|
"trace": {"payload": "x" * 400},
|
||||||
|
"debug": {"payload": "y" * 400},
|
||||||
|
"cache_fingerprint": {
|
||||||
|
"payload_sha256": "a" * 64,
|
||||||
|
"cache_relevant_sha256": "b" * 64,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
sanitized = sanitize_request_metadata(metadata)
|
||||||
|
|
||||||
|
assert sanitized["_metadata_truncated"] is True
|
||||||
|
assert sanitized["cache_fingerprint"]["payload_sha256"] == "a" * 64
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_message_telemetry_record_success_keeps_response_shape_and_adds_fingerprint(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
captured: dict[str, Any] = {}
|
||||||
|
|
||||||
|
async def _fake_record_usage(**kwargs: Any) -> Any:
|
||||||
|
captured.update(kwargs)
|
||||||
|
return SimpleNamespace(total_cost_usd=0.0, input_tokens=1, output_tokens=2)
|
||||||
|
|
||||||
|
monkeypatch.setattr(UsageService, "record_usage", _fake_record_usage)
|
||||||
|
|
||||||
|
telemetry = MessageTelemetry(
|
||||||
|
db=SimpleNamespace(), # type: ignore[arg-type]
|
||||||
|
user=None,
|
||||||
|
api_key=None,
|
||||||
|
request_id="req-cache-fingerprint",
|
||||||
|
client_ip="127.0.0.1",
|
||||||
|
)
|
||||||
|
|
||||||
|
await telemetry.record_success(
|
||||||
|
provider="openai",
|
||||||
|
model="gpt-5.4",
|
||||||
|
input_tokens=1,
|
||||||
|
output_tokens=2,
|
||||||
|
response_time_ms=10,
|
||||||
|
status_code=200,
|
||||||
|
request_body={"messages": [{"role": "user", "content": "hello"}]},
|
||||||
|
request_headers={"user-agent": "codex desktop"},
|
||||||
|
response_body={"id": "resp-1"},
|
||||||
|
response_headers={"x-test": "1"},
|
||||||
|
provider_request_body=_build_openai_cli_body(),
|
||||||
|
response_metadata={"model_version": "gpt-5.4-2026-03-01"},
|
||||||
|
endpoint_api_format="openai:cli",
|
||||||
|
)
|
||||||
|
|
||||||
|
metadata = captured["metadata"]
|
||||||
|
assert metadata["model_version"] == "gpt-5.4-2026-03-01"
|
||||||
|
assert "response" not in metadata
|
||||||
|
assert metadata["cache_fingerprint"]["provider_api_format"] == "openai:cli"
|
||||||
|
assert metadata["cache_fingerprint"]["prompt_cache_key"] == "pcache-123"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_message_telemetry_record_failure_keeps_request_metadata_and_adds_fingerprint(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
captured: dict[str, Any] = {}
|
||||||
|
|
||||||
|
async def _fake_record_usage(**kwargs: Any) -> Any:
|
||||||
|
captured.update(kwargs)
|
||||||
|
return SimpleNamespace()
|
||||||
|
|
||||||
|
monkeypatch.setattr(UsageService, "record_usage", _fake_record_usage)
|
||||||
|
|
||||||
|
telemetry = MessageTelemetry(
|
||||||
|
db=SimpleNamespace(), # type: ignore[arg-type]
|
||||||
|
user=None,
|
||||||
|
api_key=None,
|
||||||
|
request_id="req-cache-fingerprint-fail",
|
||||||
|
client_ip="127.0.0.1",
|
||||||
|
)
|
||||||
|
|
||||||
|
await telemetry.record_failure(
|
||||||
|
provider="openai",
|
||||||
|
model="gpt-5.4",
|
||||||
|
response_time_ms=10,
|
||||||
|
status_code=502,
|
||||||
|
error_message="upstream failed",
|
||||||
|
request_body={"messages": [{"role": "user", "content": "hello"}]},
|
||||||
|
request_headers={"user-agent": "codex desktop"},
|
||||||
|
is_stream=False,
|
||||||
|
provider_request_body=_build_openai_cli_body(),
|
||||||
|
request_metadata={"perf": {"ttfb_ms": 12}},
|
||||||
|
endpoint_api_format="openai:cli",
|
||||||
|
)
|
||||||
|
|
||||||
|
metadata = captured["metadata"]
|
||||||
|
assert metadata["perf"]["ttfb_ms"] == 12
|
||||||
|
assert metadata["cache_fingerprint"]["provider_api_format"] == "openai:cli"
|
||||||
|
assert metadata["cache_fingerprint"]["prompt_cache_key"] == "pcache-123"
|
||||||
Reference in New Issue
Block a user