chore: Add model tracking to LiteLlm and introduce a LiteLLM with fallbacks demo

Related: #2292

Co-authored-by: Shan Cao <caoshan@google.com>
PiperOrigin-RevId: 828024955
This commit is contained in:
Shan Cao
2025-11-04 10:10:09 -08:00
committed by Copybara-Service
parent e25beb4bce
commit d4c63fc562
7 changed files with 180 additions and 12 deletions
+13 -5
View File
@@ -558,7 +558,9 @@ def _model_response_to_generate_content_response(
if not message:
raise ValueError("No message in response")
llm_response = _message_to_generate_content_response(message)
llm_response = _message_to_generate_content_response(
message, model_version=response.model
)
if finish_reason:
# If LiteLLM already provides a FinishReason enum (e.g., for Gemini), use
# it directly. Otherwise, map the finish_reason string to the enum.
@@ -579,13 +581,14 @@ def _model_response_to_generate_content_response(
def _message_to_generate_content_response(
message: Message, is_partial: bool = False
message: Message, *, is_partial: bool = False, model_version: str = None
) -> LlmResponse:
"""Converts a litellm message to LlmResponse.
Args:
message: The message to convert.
is_partial: Whether the message is partial.
model_version: The model version used to generate the response.
Returns:
The LlmResponse.
@@ -606,7 +609,9 @@ def _message_to_generate_content_response(
parts.append(part)
return LlmResponse(
content=types.Content(role="model", parts=parts), partial=is_partial
content=types.Content(role="model", parts=parts),
partial=is_partial,
model_version=model_version,
)
@@ -950,6 +955,7 @@ class LiteLlm(BaseLlm):
content=chunk.text,
),
is_partial=True,
model_version=part.model,
)
elif isinstance(chunk, UsageMetadataChunk):
usage_metadata = types.GenerateContentResponseUsageMetadata(
@@ -981,14 +987,16 @@ class LiteLlm(BaseLlm):
role="assistant",
content=text,
tool_calls=tool_calls,
)
),
model_version=part.model,
)
)
text = ""
function_calls.clear()
elif finish_reason == "stop" and text:
aggregated_llm_response = _message_to_generate_content_response(
ChatCompletionAssistantMessage(role="assistant", content=text)
ChatCompletionAssistantMessage(role="assistant", content=text),
model_version=part.model,
)
text = ""
+7
View File
@@ -55,6 +55,9 @@ class LlmResponse(BaseModel):
)
"""The pydantic model config."""
model_version: Optional[str] = None
"""Output only. The model version used to generate the response."""
content: Optional[types.Content] = None
"""The generative content of the response.
@@ -159,6 +162,7 @@ class LlmResponse(BaseModel):
citation_metadata=candidate.citation_metadata,
avg_logprobs=candidate.avg_logprobs,
logprobs_result=candidate.logprobs_result,
model_version=generate_content_response.model_version,
)
else:
return LlmResponse(
@@ -169,6 +173,7 @@ class LlmResponse(BaseModel):
finish_reason=candidate.finish_reason,
avg_logprobs=candidate.avg_logprobs,
logprobs_result=candidate.logprobs_result,
model_version=generate_content_response.model_version,
)
else:
if generate_content_response.prompt_feedback:
@@ -177,10 +182,12 @@ class LlmResponse(BaseModel):
error_code=prompt_feedback.block_reason,
error_message=prompt_feedback.block_reason_message,
usage_metadata=usage_metadata,
model_version=generate_content_response.model_version,
)
else:
return LlmResponse(
error_code='UNKNOWN_ERROR',
error_message='Unknown error.',
usage_metadata=usage_metadata,
model_version=generate_content_response.model_version,
)