Files
video-gen/video-gen-api/app/services/model_pricing/usage_normalizer.py
T
2026-07-11 12:48:29 +08:00

276 lines
11 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
from __future__ import annotations
import hashlib
import json
import re
from copy import deepcopy
from typing import Any, Mapping
from urllib.parse import urlsplit, urlunsplit
def safe_json_dict(value: Any) -> dict[str, Any]:
"""Best-effort JSON object conversion without leaking parse failures into billing."""
if isinstance(value, Mapping):
return deepcopy(dict(value))
if isinstance(value, str) and value.strip():
try:
parsed = json.loads(value)
return deepcopy(dict(parsed)) if isinstance(parsed, Mapping) else {}
except Exception:
return {}
return {}
def safe_int(value: Any, default: int = 0) -> int:
try:
if value in (None, ""):
return default
return int(float(value))
except Exception:
return default
def safe_float(value: Any, default: float = 0.0) -> float:
try:
if value in (None, ""):
return default
return float(value)
except Exception:
return default
def safe_bool(value: Any, default: bool = False) -> bool:
if isinstance(value, bool):
return value
if value in (None, ""):
return default
if isinstance(value, (int, float)):
return value != 0
text = str(value).strip().lower()
if text in {"1", "true", "yes", "on", "enabled"}:
return True
if text in {"0", "false", "no", "off", "disabled"}:
return False
return default
def parse_size(value: Any) -> tuple[int, int]:
text = str(value or "").lower().replace("×", "x")
match = re.search(r"(\d{2,5})\s*x\s*(\d{2,5})", text)
if not match:
return 0, 0
return int(match.group(1)), int(match.group(2))
def _extract_usage(data: Mapping[str, Any]) -> dict[str, Any]:
candidates = [
data.get("usage"),
(data.get("data") or {}).get("usage") if isinstance(data.get("data"), Mapping) else None,
(data.get("result") or {}).get("usage") if isinstance(data.get("result"), Mapping) else None,
]
for value in candidates:
if isinstance(value, Mapping):
return deepcopy(dict(value))
return {}
def normalize_text_pricing_usage(
raw_usage: Mapping[str, Any] | None,
*,
base: Mapping[str, Any] | None = None,
) -> dict[str, Any]:
"""Normalize Ark/OpenAI text usage and retain cache/audio dimensions."""
raw = deepcopy(dict(raw_usage or {}))
result = deepcopy(dict(base or {}))
input_tokens = safe_int(result.get("input_tokens"), safe_int(raw.get("input_tokens"), safe_int(raw.get("prompt_tokens"))))
output_tokens = safe_int(result.get("output_tokens"), safe_int(raw.get("output_tokens"), safe_int(raw.get("completion_tokens"))))
total_tokens = safe_int(result.get("total_tokens"), safe_int(raw.get("total_tokens"), input_tokens + output_tokens))
details: dict[str, Any] = {}
for key in ("prompt_tokens_details", "input_tokens_details"):
value = raw.get(key)
if isinstance(value, Mapping):
details.update(deepcopy(dict(value)))
cached_input_tokens = safe_int(
result.get("cached_input_tokens"),
safe_int(raw.get("cached_input_tokens"), safe_int(details.get("cached_tokens"), safe_int(details.get("cache_read_tokens")))),
)
audio_input_tokens = safe_int(
result.get("audio_input_tokens"),
safe_int(raw.get("audio_input_tokens"), safe_int(details.get("audio_tokens"))),
)
cached_audio_input_tokens = safe_int(
result.get("cached_audio_input_tokens"),
safe_int(raw.get("cached_audio_input_tokens"), safe_int(details.get("cached_audio_tokens"))),
)
result.update(
{
"input_tokens": max(0, input_tokens),
"output_tokens": max(0, output_tokens),
"total_tokens": max(0, total_tokens),
"context_tokens": max(0, safe_int(result.get("context_tokens"), input_tokens)),
"cached_input_tokens": max(0, min(input_tokens, cached_input_tokens)),
"audio_input_tokens": max(0, min(input_tokens, audio_input_tokens)),
"cached_audio_input_tokens": max(0, min(input_tokens, audio_input_tokens, cached_audio_input_tokens)),
"cache_storage_tokens": max(0, safe_int(result.get("cache_storage_tokens"), safe_int(raw.get("cache_storage_tokens")))),
"cache_storage_duration_hours": max(
0.0,
safe_float(result.get("cache_storage_duration_hours"), safe_float(raw.get("cache_storage_duration_hours"))),
),
"provider_usage_primary": safe_bool(result.get("provider_usage_primary"), True),
"usage_source": result.get("usage_source") or "provider",
}
)
if details:
result["provider_input_token_details"] = details
return result
def extract_image_output_items(provider_response: Any) -> list[dict[str, Any]]:
"""Only parse explicit synchronous image output items; never recurse through arbitrary URLs."""
data = safe_json_dict(provider_response)
raw_items = data.get("data")
if isinstance(raw_items, Mapping):
raw_items = raw_items.get("items") or raw_items.get("data")
if not isinstance(raw_items, list):
raw_items = (data.get("result") or {}).get("data") if isinstance(data.get("result"), Mapping) else None
if not isinstance(raw_items, list):
return []
items: list[dict[str, Any]] = []
for index, raw in enumerate(raw_items):
if not isinstance(raw, Mapping):
continue
url = raw.get("url") or raw.get("image_url")
width = safe_int(raw.get("width"))
height = safe_int(raw.get("height"))
if width <= 0 or height <= 0:
width, height = parse_size(raw.get("size"))
item = {
"index": index,
"url": str(url) if url else None,
"width": width or None,
"height": height or None,
"pixels": width * height if width > 0 and height > 0 else None,
"size": raw.get("size"),
"size_source": "provider_response" if width > 0 and height > 0 else "unavailable",
}
# A valid provider output may omit a URL in rare response formats, but it must
# still be represented for generated-count and pixel-tier accounting.
items.append({key: value for key, value in item.items() if value is not None})
return items
def sanitize_output_items(items: list[Mapping[str, Any]]) -> list[dict[str, Any]]:
"""Remove volatile signed URLs before persisting pricing/generation snapshots."""
sanitized: list[dict[str, Any]] = []
for raw in items:
item = deepcopy(dict(raw))
url = str(item.pop("url", "") or "").strip()
if url:
try:
parts = urlsplit(url)
normalized = (
urlunsplit((parts.scheme.lower(), parts.netloc.lower(), parts.path, "", ""))
if parts.scheme and parts.netloc
else url.split("?", 1)[0].split("#", 1)[0]
)
except Exception:
normalized = url.split("?", 1)[0].split("#", 1)[0]
item["url_sha256"] = hashlib.sha256(normalized.encode("utf-8")).hexdigest()
sanitized.append(item)
return sanitized
def normalize_provider_media_usage(
provider_response: Any,
*,
gen_type: str,
fallback_total_tokens: int = 0,
request_image_px: str | None = None,
requested_output_count: int = 1,
provider_input_image_count: int = 0,
) -> dict[str, Any]:
"""Normalize Volcengine synchronous-image or asynchronous-video response usage."""
data = safe_json_dict(provider_response)
raw_usage = _extract_usage(data)
input_tokens = safe_int(raw_usage.get("input_tokens"), safe_int(raw_usage.get("prompt_tokens")))
output_tokens = safe_int(
raw_usage.get("output_tokens"),
safe_int(raw_usage.get("completion_tokens"), safe_int(raw_usage.get("generated_tokens"))),
)
total_tokens = safe_int(raw_usage.get("total_tokens"), fallback_total_tokens)
if total_tokens <= 0:
total_tokens = input_tokens + output_tokens
if output_tokens <= 0 and total_tokens > input_tokens:
output_tokens = total_tokens - input_tokens
result: dict[str, Any] = deepcopy(dict(raw_usage))
result.update(
{
"input_tokens": max(0, input_tokens),
"output_tokens": max(0, output_tokens),
"total_tokens": max(0, total_tokens),
"requested_output_count": max(1, requested_output_count),
"provider_usage_primary": True,
}
)
pricing_meta = data.get("pricing_meta") if isinstance(data.get("pricing_meta"), Mapping) else {}
if gen_type == "image":
output_items = extract_image_output_items(data)
fallback_width, fallback_height = parse_size(request_image_px)
for item in output_items:
if not item.get("width") and fallback_width > 0 and fallback_height > 0:
item.update(
{
"width": fallback_width,
"height": fallback_height,
"pixels": fallback_width * fallback_height,
"size_source": "request_explicit",
}
)
output_items = sanitize_output_items(output_items)
provider_count = safe_int(
pricing_meta.get("provider_input_image_count"),
safe_int(data.get("provider_input_image_count"), provider_input_image_count),
)
generated = len(output_items) or safe_int(raw_usage.get("generated_images"))
provider_billed = safe_int(
raw_usage.get("billed_images"),
safe_int(pricing_meta.get("provider_billed_count"), max(0, generated)),
)
result.update(
{
"provider_input_image_count": max(0, provider_count),
"input_image_count": max(0, provider_count),
"output_items": output_items,
"successful_output_count": max(0, generated),
"provider_billed_count": max(0, provider_billed),
"usage_source": "provider_response",
}
)
return result
width = safe_int(raw_usage.get("width"), safe_int(data.get("width")))
height = safe_int(raw_usage.get("height"), safe_int(data.get("height")))
if width <= 0 or height <= 0:
width, height = parse_size(raw_usage.get("size") or data.get("size"))
result.update(
{
"output_width": width,
"output_height": height,
"dimension_source": "provider_response" if width > 0 and height > 0 else "unavailable",
"fps": safe_float(raw_usage.get("fps"), safe_float(data.get("fps"))),
"resolution": str(raw_usage.get("resolution") or data.get("resolution") or "").lower(),
"aspect_ratio": str(raw_usage.get("aspect_ratio") or data.get("aspect_ratio") or data.get("ratio") or ""),
"generate_audio": safe_bool(raw_usage.get("generate_audio"), safe_bool(data.get("generate_audio"))),
"inference_mode": str(raw_usage.get("inference_mode") or data.get("service_tier") or "online").lower(),
"usage_source": "provider" if total_tokens > 0 else "provider_response",
}
)
return result