Skip to content
Draft
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
4 changes: 4 additions & 0 deletions src/neo4j_graphrag/llm/types.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
10 changes: 10 additions & 0 deletions src/neo4j_graphrag/llm/vertexai_llm.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Loading