修复: 校验真实 Claude usage 元数据

This commit is contained in:
zizi 2026-07-22 22:01:08 +08:00
parent 3d0b35018f
commit 0c7427360b
2 changed files with 231 additions and 2 deletions

View File

@ -54,6 +54,17 @@ MODEL_ID_PATTERN = re.compile(
r"^(?=.{6,128}$)[A-Za-z0-9][A-Za-z0-9._:-]{5,127}(?:\[[A-Za-z0-9._:-]+\])?$" r"^(?=.{6,128}$)[A-Za-z0-9][A-Za-z0-9._:-]{5,127}(?:\[[A-Za-z0-9._:-]+\])?$"
) )
MONEY_QUANTUM = Decimal("0.000001") MONEY_QUANTUM = Decimal("0.000001")
USAGE_REQUIRED_COUNT_FIELDS = frozenset({"input_tokens", "output_tokens"})
USAGE_OPTIONAL_COUNT_FIELDS = frozenset(
{"cache_creation_input_tokens", "cache_read_input_tokens"}
)
USAGE_COUNT_MAP_FIELDS = {
"server_tool_use": frozenset({"web_search_requests", "web_fetch_requests"}),
"cache_creation": frozenset(
{"ephemeral_5m_input_tokens", "ephemeral_1h_input_tokens"}
),
}
USAGE_STRING_FIELDS = frozenset({"service_tier", "speed", "inference_geo"})
class ClaudeRuntimeError(RuntimeError): class ClaudeRuntimeError(RuntimeError):
@ -503,7 +514,7 @@ def _minimal_environment(
def _validate_number_tree(value: Any, path: str) -> None: def _validate_number_tree(value: Any, path: str) -> None:
"""递归校验 usage 中的数值均有限且非负。""" """递归校验严格数字树,用于 modelUsage 的计数核账。"""
if isinstance(value, Mapping): if isinstance(value, Mapping):
if not value: if not value:
@ -520,6 +531,89 @@ def _validate_number_tree(value: Any, path: str) -> None:
raise ValueError(f"{path} 数值非法") raise ValueError(f"{path} 数值非法")
def _validate_non_negative_finite_number(value: Any, path: str) -> None:
"""校验计数字段是非负有限数,不接受 bool、NaN 或 Infinity。"""
if isinstance(value, bool) or not isinstance(value, (int, float, Decimal)):
raise ValueError(f"{path} 必须是非负有限数")
try:
decimal_value = Decimal(str(value))
except (InvalidOperation, ValueError) as exc:
raise ValueError(f"{path} 必须是非负有限数") from exc
if not decimal_value.is_finite() or decimal_value < 0:
raise ValueError(f"{path} 必须是非负有限数")
def _validate_safe_json(value: Any, path: str) -> None:
"""递归校验前向兼容字段仍然只包含安全 JSON 值。"""
if value is None or isinstance(value, (str, bool, int)):
return
if isinstance(value, (float, Decimal)):
try:
decimal_value = Decimal(str(value))
except (InvalidOperation, ValueError) as exc:
raise ValueError(f"{path} 不是安全 JSON") from exc
if not decimal_value.is_finite():
raise ValueError(f"{path} 不是安全 JSON")
return
if isinstance(value, list):
for index, item in enumerate(value):
_validate_safe_json(item, f"{path}[{index}]")
return
if isinstance(value, Mapping):
for key, item in value.items():
if not isinstance(key, str):
raise ValueError(f"{path} 对象键不是字符串")
_validate_safe_json(item, f"{path}.{key}")
return
raise ValueError(f"{path} 不是安全 JSON")
def _validate_usage_count_mapping(
value: Any, path: str, known_count_fields: frozenset[str]
) -> None:
"""校验已知计数映射,并允许未来字段继续使用安全 JSON。"""
if not isinstance(value, Mapping):
raise ValueError(f"{path} 必须是对象")
for key, item in value.items():
if not isinstance(key, str):
raise ValueError(f"{path} 对象键不是字符串")
if key in known_count_fields:
_validate_non_negative_finite_number(item, f"{path}.{key}")
else:
_validate_safe_json(item, f"{path}.{key}")
def _validate_usage(value: Any, path: str = "usage") -> None:
"""精确校验 usage 的权威计数,同时兼容安全的未来 metadata。"""
if not isinstance(value, Mapping) or not value:
raise ValueError(f"{path} 必须是非空对象")
missing_fields = USAGE_REQUIRED_COUNT_FIELDS - set(value)
if missing_fields:
raise ValueError(f"{path} 缺少必需计数字段")
for key, item in value.items():
if not isinstance(key, str):
raise ValueError(f"{path} 对象键不是字符串")
field_path = f"{path}.{key}"
if key in USAGE_REQUIRED_COUNT_FIELDS or key in USAGE_OPTIONAL_COUNT_FIELDS:
_validate_non_negative_finite_number(item, field_path)
elif key in USAGE_COUNT_MAP_FIELDS:
_validate_usage_count_mapping(item, field_path, USAGE_COUNT_MAP_FIELDS[key])
elif key in USAGE_STRING_FIELDS:
if not isinstance(item, str) or (key != "inference_geo" and not item):
raise ValueError(f"{field_path} 必须是字符串")
elif key == "iterations":
if not isinstance(item, list):
raise ValueError(f"{field_path} 必须是数组")
_validate_safe_json(item, field_path)
else:
_validate_safe_json(item, field_path)
def _schema_type_matches(value: Any, expected: str) -> bool: def _schema_type_matches(value: Any, expected: str) -> bool:
"""按 JSON 类型语义判断 Python 值,显式排除 bool 伪装整数。""" """按 JSON 类型语义判断 Python 值,显式排除 bool 伪装整数。"""
@ -775,7 +869,7 @@ def run_claude(
issues.add(f"{prefix}_BUDGET_EXCEEDED") issues.add(f"{prefix}_BUDGET_EXCEEDED")
try: try:
_validate_number_tree(envelope.get("usage"), "usage") _validate_usage(envelope.get("usage"))
except (ValueError, InvalidOperation): except (ValueError, InvalidOperation):
issues.add(f"{prefix}_RECEIPT_INVALID") issues.add(f"{prefix}_RECEIPT_INVALID")

View File

@ -3,6 +3,7 @@
from __future__ import annotations from __future__ import annotations
import copy
import hashlib import hashlib
import json import json
import os import os
@ -20,6 +21,7 @@ from claude_runtime import ( # noqa: E402
ClaudeRuntimeError, ClaudeRuntimeError,
ExecutionProfile, ExecutionProfile,
_minimal_environment, _minimal_environment,
_validate_usage,
build_sandbox_command, build_sandbox_command,
canonical_json, canonical_json,
run_claude, run_claude,
@ -95,6 +97,27 @@ def success_envelope() -> dict[str, object]:
} }
def realistic_usage() -> dict[str, object]:
"""构造包含 Claude metadata、已知计数映射和迭代数组的真实 usage 形状。"""
return {
"input_tokens": 10,
"output_tokens": 2,
"cache_creation_input_tokens": 3,
"cache_read_input_tokens": 4,
"server_tool_use": {"web_search_requests": 1, "web_fetch_requests": 0},
"cache_creation": {"ephemeral_5m_input_tokens": 5, "ephemeral_1h_input_tokens": 6},
"service_tier": "standard",
"speed": "standard",
"inference_geo": "",
"iterations": [
{"type": "tool_use", "duration_ms": 12},
{"type": "metadata", "values": [True, None, "safe"]},
],
"future_metadata": {"trace_id": "safe", "enabled": True},
}
class FakeRunner: class FakeRunner:
"""记录 subprocess 参数并返回预设结果,测试绝不调用真实模型。""" """记录 subprocess 参数并返回预设结果,测试绝不调用真实模型。"""
@ -210,6 +233,15 @@ class ClaudeRuntimeTest(unittest.TestCase):
self.assertEqual(pathlib.Path(str(kwargs["cwd"])), pathlib.Path(kwargs["env"]["HOME"])) self.assertEqual(pathlib.Path(str(kwargs["cwd"])), pathlib.Path(kwargs["env"]["HOME"]))
self.assertFalse(pathlib.Path(str(kwargs["cwd"])).exists()) self.assertFalse(pathlib.Path(str(kwargs["cwd"])).exists())
with tempfile.TemporaryDirectory(dir="/private/tmp") as directory:
real_auth_environment = _minimal_environment(
{"ANTHROPIC_AUTH_TOKEN": auth_token, "ANTHROPIC_BASE_URL": base_url},
pathlib.Path(directory),
require_authentication=True,
)
self.assertEqual(real_auth_environment["ANTHROPIC_AUTH_TOKEN"], auth_token)
self.assertEqual(real_auth_environment["ANTHROPIC_BASE_URL"], base_url)
def test_auth_token_requires_a_safe_base_url(self): def test_auth_token_requires_a_safe_base_url(self):
"""AUTH_TOKEN 必须绑定无凭据、无查询和无片段的 HTTP(S) 网关地址。""" """AUTH_TOKEN 必须绑定无凭据、无查询和无片段的 HTTP(S) 网关地址。"""
@ -242,6 +274,109 @@ class ClaudeRuntimeTest(unittest.TestCase):
self.assertNotIn(base_url, str(raised.exception)) self.assertNotIn(base_url, str(raised.exception))
self.assertNotIn(base_url, canonical_json(raised.exception.details)) self.assertNotIn(base_url, canonical_json(raised.exception.details))
def test_realistic_usage_metadata_is_accepted_without_changing_cost_authority(self):
"""真实 usage metadata 可通过,但模型成本仍只由 modelUsage 对账。"""
envelope = success_envelope()
envelope["usage"] = realistic_usage()
result = run_claude(
execution_profile(),
{"request": "x"},
runner=runner_for(envelope),
binding_verifier=lambda _profile: None,
)
self.assertEqual(result.structured_output, {"value": "ok"})
self.assertEqual(result.receipt.as_dict()["totalCostUsd"], "0.120000")
self.assertEqual(result.receipt.as_dict()["usage"], realistic_usage())
def test_usage_rejects_bad_counts_metadata_and_non_json_iterations(self):
"""usage 的必需计数、已知映射和 metadata 破坏时必须失败关闭。"""
cases: list[tuple[str, dict[str, object]]] = []
missing_required = realistic_usage()
del missing_required["input_tokens"]
cases.append(("missing input_tokens", missing_required))
negative_token = realistic_usage()
negative_token["output_tokens"] = -1
cases.append(("negative output_tokens", negative_token))
negative_server_count = realistic_usage()
negative_server_count["server_tool_use"] = {"web_search_requests": -1}
cases.append(("negative server_tool_use count", negative_server_count))
negative_cache_count = realistic_usage()
negative_cache_count["cache_creation"] = {"ephemeral_5m_input_tokens": -1}
cases.append(("negative cache_creation count", negative_cache_count))
invalid_service_tier = realistic_usage()
invalid_service_tier["service_tier"] = 1
cases.append(("non-string service_tier", invalid_service_tier))
invalid_inference_geo = realistic_usage()
invalid_inference_geo["inference_geo"] = None
cases.append(("non-string inference_geo", invalid_inference_geo))
invalid_unknown_metadata = realistic_usage()
invalid_unknown_metadata["future_metadata"] = {"ratio": float("inf")}
cases.append(("infinite unknown metadata", invalid_unknown_metadata))
invalid_iteration_number = realistic_usage()
invalid_iteration_number["iterations"] = [{"duration_ms": float("nan")}]
cases.append(("non-finite iteration metadata", invalid_iteration_number))
for label, usage in cases:
envelope = success_envelope()
envelope["usage"] = usage
with self.subTest(label=label), self.assertRaises(ClaudeRuntimeError) as raised:
run_claude(
execution_profile(),
{"request": "x"},
runner=runner_for(envelope),
binding_verifier=lambda _profile: None,
)
self.assertEqual(raised.exception.primary_code, "WRITER_RECEIPT_INVALID")
invalid_iteration_json = copy.deepcopy(realistic_usage())
invalid_iteration_json["iterations"] = [object()]
with self.assertRaises(ValueError):
_validate_usage(invalid_iteration_json)
def test_realistic_usage_cannot_bypass_model_cost_reconciliation(self):
"""丰富 usage metadata 不能掩盖 modelUsage 与 total cost 对账失败。"""
cases = []
mismatched_model = success_envelope()
mismatched_model["usage"] = realistic_usage()
mismatched_model["modelUsage"] = {
"claude-sonnet-4-20250514": {
"inputTokens": 10,
"outputTokens": 2,
"costUSD": "0.120000",
}
}
cases.append(mismatched_model)
mismatched_total = success_envelope()
mismatched_total["usage"] = realistic_usage()
mismatched_total["total_cost_usd"] = "0.130000"
cases.append(mismatched_total)
for envelope in cases:
with self.assertRaises(ClaudeRuntimeError) as raised:
run_claude(
execution_profile(),
{"request": "x"},
runner=runner_for(envelope),
binding_verifier=lambda _profile: None,
)
self.assertIn(
raised.exception.primary_code,
{"WRITER_MODEL_MISMATCH", "WRITER_RECEIPT_INVALID"},
)
def test_api_key_and_oauth_remain_supported_and_exclusive(self): def test_api_key_and_oauth_remain_supported_and_exclusive(self):
"""API key/OAuth 仍可进入最小环境,多个认证字段仍然失败关闭。""" """API key/OAuth 仍可进入最小环境,多个认证字段仍然失败关闭。"""