diff --git a/src/neo4j_graphrag/llm/types.py b/src/neo4j_graphrag/llm/types.py index d199919ec..84eb116fa 100644 --- a/src/neo4j_graphrag/llm/types.py +++ b/src/neo4j_graphrag/llm/types.py @@ -27,11 +27,15 @@ class LLMUsage(BaseModel): ``None`` when not reported by the provider. total_tokens (Optional[int]): Total tokens consumed by the call. ``None`` when not reported by the provider. + cached_tokens (Optional[int]): Number of prompt tokens served from the + provider's context cache (e.g. Vertex AI implicit caching). + ``None`` when not reported by the provider. """ request_tokens: Optional[int] = None response_tokens: Optional[int] = None total_tokens: Optional[int] = None + cached_tokens: Optional[int] = None class LLMResponse(BaseModel): diff --git a/src/neo4j_graphrag/llm/vertexai_llm.py b/src/neo4j_graphrag/llm/vertexai_llm.py index f9ee804be..313899e7e 100644 --- a/src/neo4j_graphrag/llm/vertexai_llm.py +++ b/src/neo4j_graphrag/llm/vertexai_llm.py @@ -593,9 +593,19 @@ def _parse_content_response(self, response: GenerationResponse) -> LLMResponse: usage = None metadata = response.usage_metadata if metadata: + cached = metadata.cached_content_token_count or None usage = LLMUsage( request_tokens=metadata.prompt_token_count, response_tokens=metadata.candidates_token_count, total_tokens=metadata.total_token_count, + cached_tokens=cached, ) + if cached: + from openinference.semconv.trace import SpanAttributes + from opentelemetry import trace as otel_trace + + otel_trace.get_current_span().set_attribute( + SpanAttributes.LLM_TOKEN_COUNT_PROMPT_DETAILS_CACHE_READ, + cached, + ) return LLMResponse(content=response.text, usage=usage)