"""Token-usage extraction from LangChain messages. Both the agent run and the grader call report token usage the same way — on the message's ``usage_metadata`` (or, for some providers, nested under ``response_metadata['usage']``). These helpers read it from one message or sum it across many so callers (agent runner, grader) don't each re-implement the lookup. """ from typing import Any def message_usage(message: Any) -> tuple[int, int]: """Return ``(input_tokens, output_tokens)`` for one message; ``(0, 0)`` if absent.""" usage_metadata = getattr(message, "usage_metadata", None) if isinstance(usage_metadata, dict): return ( usage_metadata.get("input_tokens", 0), usage_metadata.get("output_tokens", 0), ) response_metadata = getattr(message, "response_metadata", {}) or {} usage = response_metadata.get("usage", {}) or {} return ( _coalesce(usage, "input_tokens", "prompt_tokens"), _coalesce(usage, "output_tokens", "completion_tokens"), ) def _coalesce(usage: dict[str, Any], primary: str, secondary: str) -> int: """Return ``primary``'s token count, falling back to ``secondary`` only when needed. The fallback happens only when ``primary`` is missing or None, so a legitimate 0 isn't mistaken for a missing value. """ value = usage.get(primary) if value is None: value = usage.get(secondary) return value or 0 def sum_usage(messages: list[Any]) -> tuple[int, int]: """Sum ``(input_tokens, output_tokens)`` across every message in ``messages``.""" total_input = 0 total_output = 0 for message in messages: message_input, message_output = message_usage(message) total_input += message_input total_output += message_output return total_input, total_output