fix(vertex_httpx.py): fix assumptions on usagemetadata

This commit is contained in:
Krrish Dholakia 2024-07-05 11:01:33 -07:00
parent 6cd7631e2e
commit 015a398713
2 changed files with 26 additions and 18 deletions

View file

@ -603,15 +603,15 @@ class VertexLLM(BaseLLM):
## GET USAGE ##
usage = litellm.Usage(
prompt_tokens=completion_response["usageMetadata"][
"promptTokenCount"
],
prompt_tokens=completion_response["usageMetadata"].get(
"promptTokenCount", 0
),
completion_tokens=completion_response["usageMetadata"].get(
"candidatesTokenCount", 0
),
total_tokens=completion_response["usageMetadata"][
"totalTokenCount"
],
total_tokens=completion_response["usageMetadata"].get(
"totalTokenCount", 0
),
)
setattr(model_response, "usage", usage)
@ -647,15 +647,15 @@ class VertexLLM(BaseLLM):
## GET USAGE ##
usage = litellm.Usage(
prompt_tokens=completion_response["usageMetadata"][
"promptTokenCount"
],
prompt_tokens=completion_response["usageMetadata"].get(
"promptTokenCount", 0
),
completion_tokens=completion_response["usageMetadata"].get(
"candidatesTokenCount", 0
),
total_tokens=completion_response["usageMetadata"][
"totalTokenCount"
],
total_tokens=completion_response["usageMetadata"].get(
"totalTokenCount", 0
),
)
setattr(model_response, "usage", usage)
@ -705,11 +705,15 @@ class VertexLLM(BaseLLM):
## GET USAGE ##
usage = litellm.Usage(
prompt_tokens=completion_response["usageMetadata"]["promptTokenCount"],
prompt_tokens=completion_response["usageMetadata"].get(
"promptTokenCount", 0
),
completion_tokens=completion_response["usageMetadata"].get(
"candidatesTokenCount", 0
),
total_tokens=completion_response["usageMetadata"]["totalTokenCount"],
total_tokens=completion_response["usageMetadata"].get(
"totalTokenCount", 0
),
)
setattr(model_response, "usage", usage)
@ -1340,11 +1344,15 @@ class ModelResponseIterator:
if "usageMetadata" in processed_chunk:
usage = ChatCompletionUsageBlock(
prompt_tokens=processed_chunk["usageMetadata"]["promptTokenCount"],
prompt_tokens=processed_chunk["usageMetadata"].get(
"promptTokenCount", 0
),
completion_tokens=processed_chunk["usageMetadata"].get(
"candidatesTokenCount", 0
),
total_tokens=processed_chunk["usageMetadata"]["totalTokenCount"],
total_tokens=processed_chunk["usageMetadata"].get(
"totalTokenCount", 0
),
)
returned_chunk = GenericStreamingChunk(

View file

@ -239,8 +239,8 @@ class PromptFeedback(TypedDict):
class UsageMetadata(TypedDict, total=False):
promptTokenCount: Required[int]
totalTokenCount: Required[int]
promptTokenCount: int
totalTokenCount: int
candidatesTokenCount: int