Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
Capture LangChain chat token usage from generation metadata and LLM output fallbacks.
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,7 @@
_message_name,
_normalize_role,
extract_token_details,
extract_usage_tokens,
is_stream_end_marker,
make_input_message,
make_last_output_message,
Expand All @@ -55,6 +56,8 @@
from opentelemetry.util.genai.types import (
InputMessage,
MessagePart,
Modality,
ModalityTokens,
OutputMessage,
RetrievalDocument,
Role,
Expand Down Expand Up @@ -83,6 +86,58 @@ def _conversation_id(metadata: dict[str, Any] | None) -> str | None:
return None


def _usage_metadata_candidates(
chat_generation: Any,
llm_output: Any,
) -> list[Mapping[str, Any]]:
candidates: list[Mapping[str, Any]] = []
message = getattr(chat_generation, "message", None)
usage_metadata = _usage_mapping(getattr(message, "usage_metadata", None))
if usage_metadata is not None:
candidates.append(usage_metadata)

generation_info = _usage_mapping(
getattr(chat_generation, "generation_info", None)
)
if generation_info is not None:
usage_metadata = _usage_mapping(generation_info.get("usage_metadata"))
if usage_metadata is not None:
candidates.append(usage_metadata)

llm_output_mapping = _usage_mapping(llm_output)
if llm_output_mapping is not None:
for key in ("token_usage", "usage"):
Comment thread
amarrtech marked this conversation as resolved.
usage_metadata = _usage_mapping(llm_output_mapping.get(key))
if usage_metadata is not None:
candidates.append(usage_metadata)

return candidates


def _usage_mapping(value: Any) -> Mapping[str, Any] | None:
if isinstance(value, Mapping):
return cast(Mapping[str, Any], value)
return None


def _contains_supported_modality_tokens(
entries: ModalityTokens | None,
) -> bool:
supported_modalities = {
Modality.TEXT.value,
Modality.IMAGE.value,
Modality.AUDIO.value,
}
return entries is not None and any(
str(getattr(modality, "value", modality)).lower()
in supported_modalities
and isinstance(token_count, int)
and not isinstance(token_count, bool)
and token_count >= 0
for modality, token_count in entries
)


def _extract_document_score(doc: Any) -> float | int | None:
"""Extract relevance score polymorphically from a Document or Mapping.

Expand Down Expand Up @@ -504,6 +559,7 @@ def on_llm_end(
served_model: str | None = None
generation_model: str | None = None
generation_response_id: str | None = None
llm_output = getattr(response, "llm_output", None)
for generation in getattr(response, "generations", []):
for chat_generation in generation:
message = chat_generation.message
Expand Down Expand Up @@ -624,25 +680,26 @@ def on_llm_end(
finish_reasons.append(finish_reason)

# Get token usage if available
if chat_generation.message.usage_metadata:
usage_metadata = chat_generation.message.usage_metadata
input_tokens = usage_metadata.get("input_tokens", 0)
if not isinstance(input_tokens, int) or isinstance(
input_tokens, bool
):
input_tokens = 0
llm_invocation.input_tokens = input_tokens

output_tokens = usage_metadata.get("output_tokens", 0)
if not isinstance(output_tokens, int) or isinstance(
output_tokens, bool
):
output_tokens = 0
has_input_tokens = False
has_output_tokens = False
input_modality_tokens = None
output_modality_tokens = None
for usage_metadata in _usage_metadata_candidates(
chat_generation,
llm_output,
):
input_tokens, output_tokens = extract_usage_tokens(
usage_metadata
)
if input_tokens is not None and not has_input_tokens:
llm_invocation.input_tokens = input_tokens
has_input_tokens = True
if output_tokens is not None and not has_output_tokens:
llm_invocation.output_tokens = output_tokens
has_output_tokens = True

# Cache, reasoning, and modality token break-downs
token_details = extract_token_details(
cast(dict[str, Any], usage_metadata)
)
token_details = extract_token_details(usage_metadata)
if (
cache_write := token_details.get(
"cache_write_input_tokens"
Expand All @@ -664,25 +721,39 @@ def on_llm_end(
) is not None:
llm_invocation.thinking_tokens = reasoning_tokens

llm_invocation.set_input_tokens(
modality_tokens(
usage_metadata, "input_token_details"
)
candidate_input_modalities = modality_tokens(
usage_metadata, "input_token_details"
)
llm_invocation.set_output_tokens(
modality_tokens(
usage_metadata, "output_token_details"
if (
input_modality_tokens is None
and _contains_supported_modality_tokens(
candidate_input_modalities
)
):
input_modality_tokens = candidate_input_modalities

candidate_output_modalities = modality_tokens(
usage_metadata, "output_token_details"
)
if (
output_modality_tokens is None
and _contains_supported_modality_tokens(
candidate_output_modalities
)
):
output_modality_tokens = (
candidate_output_modalities
)

llm_invocation.output_tokens = output_tokens
llm_invocation.set_input_tokens(input_modality_tokens)
llm_invocation.set_output_tokens(output_modality_tokens)

llm_invocation.output_messages = output_messages
if finish_reasons:
llm_invocation.finish_reasons = finish_reasons

response_model, response_id = resolve_response_model_and_id(
llm_output=getattr(response, "llm_output", None),
llm_output=llm_output,
served_model=served_model,
generation_model=generation_model,
generation_response_id=generation_response_id,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -618,7 +618,14 @@ def resolve_response_model_and_id(
return response_model, response_id


def extract_token_details(usage_metadata: dict[str, Any]) -> dict[str, int]:
def _get_positive_int(values: Mapping[str, Any], key: str) -> int | None:
value = values.get(key)
if isinstance(value, int) and not isinstance(value, bool) and value > 0:
return value
return None


def extract_token_details(usage_metadata: Mapping[str, Any]) -> dict[str, int]:
"""Extract cache and reasoning token break-downs from LangChain usage metadata."""

token_details: dict[str, int] = {}
Expand All @@ -635,21 +642,22 @@ def extract_token_details(usage_metadata: dict[str, Any]) -> dict[str, int]:
else {}
)

def _get_positive_int(d: dict[str, Any], key: str) -> int | None:
val = d.get(key)
if isinstance(val, int) and not isinstance(val, bool) and val > 0:
return val
return None

cache_write = _get_positive_int(input_details, "cache_write")
if cache_write is None:
cache_write = _get_positive_int(input_details, "cache_creation")
if cache_write is None:
cache_write = _get_positive_int(
usage_metadata, "cache_creation_input_tokens"
)
if cache_write is not None:
token_details["cache_write_input_tokens"] = cache_write

if (
cache_read := _get_positive_int(input_details, "cache_read")
) is not None:
cache_read = _get_positive_int(input_details, "cache_read")
if cache_read is None:
cache_read = _get_positive_int(
usage_metadata, "cache_read_input_tokens"
)
if cache_read is not None:
token_details["cache_read_input_tokens"] = cache_read

if (
Expand All @@ -660,6 +668,44 @@ def _get_positive_int(d: dict[str, Any], key: str) -> int | None:
return token_details


def extract_usage_tokens(
usage_metadata: Mapping[str, Any],
) -> tuple[
int | None,
int | None,
]:
"""Extract input and output token counts from LangChain usage mappings."""

input_tokens = _first_int_value(
usage_metadata.get("input_tokens"),
usage_metadata.get("prompt_tokens"),

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Could we also handle prompt_token_count and candidates_token_count here? Those are the keys in VertexAI’s raw generation_info["usage_metadata"]. With the current helper, both values are ignored, so the VertexAI fallback described in #583 still emits no token usage. An exact VertexAI-shaped test would help cover this.

usage_metadata.get("prompt_token_count"),
)
if input_tokens is not None:
input_tokens += sum(
cache_tokens
for key in (
"cache_creation_input_tokens",
"cache_read_input_tokens",
)
if (cache_tokens := _get_positive_int(usage_metadata, key))
is not None
)
output_tokens = _first_int_value(
usage_metadata.get("output_tokens"),
usage_metadata.get("completion_tokens"),
usage_metadata.get("candidates_token_count"),
)
return input_tokens, output_tokens


def _first_int_value(*values: Any) -> int | None:
for value in values:
if isinstance(value, int) and not isinstance(value, bool):
return value
return None


def modality_tokens(
usage_metadata: Mapping[str, Any], key: str
) -> ModalityTokens | None:
Expand Down
Loading