fix: fix linting errors

This commit is contained in:
Krrish Dholakia 2025-09-27 14:02:27 -07:00
parent 7be9a32934
commit c9f29dd0c5
8 changed files with 106 additions and 72 deletions

View file

@ -3,14 +3,15 @@ Transformation for Bedrock Invoke Agent
https://docs.aws.amazon.com/bedrock/latest/APIReference/API_agent-runtime_InvokeAgent.html
"""
import base64
import json
from litellm._uuid import uuid
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union
import httpx
from litellm._logging import verbose_logger
from litellm._uuid import uuid
from litellm.litellm_core_utils.prompt_templates.common_utils import (
convert_content_list_to_str,
)
@ -22,6 +23,11 @@ from litellm.types.llms.bedrock_invoke_agents import (
InvokeAgentEvent,
InvokeAgentEventHeaders,
InvokeAgentEventList,
InvokeAgentMetadata,
InvokeAgentModelInvocationInput,
InvokeAgentModelInvocationOutput,
InvokeAgentOrchestrationTrace,
InvokeAgentPreProcessingTrace,
InvokeAgentTrace,
InvokeAgentTracePayload,
InvokeAgentUsage,
@ -389,15 +395,19 @@ class AmazonInvokeAgentConfig(BaseConfig, BaseAWSLLM):
self, trace_data: InvokeAgentTrace, usage_info: InvokeAgentUsage
) -> None:
"""Extract usage information from preprocessing trace."""
pre_processing = trace_data.get("preProcessingTrace", {})
pre_processing: Optional[InvokeAgentPreProcessingTrace] = trace_data.get(
"preProcessingTrace"
)
if not pre_processing:
return
model_output = pre_processing.get("modelInvocationOutput", {})
model_output: Optional[InvokeAgentModelInvocationOutput] = pre_processing.get(
"modelInvocationOutput", {}
)
if not model_output:
return
metadata = model_output.get("metadata", {})
metadata: Optional[InvokeAgentMetadata] = model_output.get("metadata", {})
if not metadata:
return
@ -412,11 +422,15 @@ class AmazonInvokeAgentConfig(BaseConfig, BaseAWSLLM):
self, trace_data: InvokeAgentTrace
) -> Optional[str]:
"""Extract model information from orchestration trace."""
orchestration_trace = trace_data.get("orchestrationTrace", {})
orchestration_trace: Optional[InvokeAgentOrchestrationTrace] = trace_data.get(
"orchestrationTrace"
)
if not orchestration_trace:
return None
model_invocation = orchestration_trace.get("modelInvocationInput", {})
model_invocation: Optional[InvokeAgentModelInvocationInput] = (
orchestration_trace.get("modelInvocationInput", {})
)
if not model_invocation:
return None

View file

@ -1,6 +1,7 @@
"""
Transformation for Calling Google models in their native format.
"""
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Tuple, Union, cast
import httpx
@ -25,27 +26,29 @@ else:
GenerateContentContentListUnionDict = Any
GenerateContentResponse = Any
ToolConfigDict = Any
from ..common_utils import get_api_key_from_env
class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM):
"""
Configuration for calling Google models in their native format.
"""
##############################
# Constants
##############################
XGOOGLE_API_KEY = "x-goog-api-key"
##############################
@property
def custom_llm_provider(self) -> Literal["gemini", "vertex_ai"]:
return "gemini"
def __init__(self):
super().__init__()
VertexLLM.__init__(self)
def get_supported_generate_content_optional_params(self, model: str) -> List[str]:
"""
Get the list of supported Google GenAI parameters for the model.
@ -58,7 +61,7 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM):
"""
return [
"http_options",
"system_instruction",
"system_instruction",
"temperature",
"top_p",
"top_k",
@ -84,10 +87,9 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM):
"speech_config",
"audio_timestamp",
"automatic_function_calling",
"thinking_config"
"thinking_config",
]
def map_generate_content_optional_params(
self,
generate_content_config_dict: GenerateContentConfigDict,
@ -103,26 +105,29 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM):
Returns:
Mapped parameters for the provider
"""
from litellm.types.google_genai.main import GenerateContentConfigDict
_generate_content_config_dict = GenerateContentConfigDict()
supported_google_genai_params = self.get_supported_generate_content_optional_params(model)
_generate_content_config_dict: Dict[str, Any] = {}
supported_google_genai_params = (
self.get_supported_generate_content_optional_params(model)
)
for param, value in generate_content_config_dict.items():
if param in supported_google_genai_params:
_generate_content_config_dict[param] = value
return dict(_generate_content_config_dict)
return _generate_content_config_dict
def validate_environment(
self,
self,
api_key: Optional[str],
headers: Optional[dict],
model: str,
litellm_params: Optional[Union[GenericLiteLLMParams, dict]]
litellm_params: Optional[Union[GenericLiteLLMParams, dict]],
) -> dict:
default_headers = {
"Content-Type": "application/json",
}
# Use the passed api_key first, then fall back to litellm_params and environment
gemini_api_key = api_key or self._get_google_ai_studio_api_key(dict(litellm_params or {}))
gemini_api_key = api_key or self._get_google_ai_studio_api_key(
dict(litellm_params or {})
)
if gemini_api_key is not None:
default_headers[self.XGOOGLE_API_KEY] = gemini_api_key
if headers is not None:
@ -137,14 +142,14 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM):
or get_api_key_from_env()
or litellm.api_key
)
def _get_common_auth_components(
self,
litellm_params: dict,
) -> Tuple[Any, Optional[str], Optional[str]]:
"""
Get common authentication components used by both sync and async methods.
Returns:
Tuple of (vertex_credentials, vertex_project, vertex_location)
"""
@ -152,7 +157,7 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM):
vertex_project = self.get_vertex_ai_project(litellm_params)
vertex_location = self.get_vertex_ai_location(litellm_params)
return vertex_credentials, vertex_project, vertex_location
def _build_final_headers_and_url(
self,
model: str,
@ -168,7 +173,7 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM):
Build final headers and API URL from auth components.
"""
gemini_api_key = self._get_google_ai_studio_api_key(litellm_params)
auth_header, api_base = self._get_token_and_url(
model=model,
gemini_api_key=gemini_api_key,
@ -201,7 +206,9 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM):
"""
Sync version of get_auth_token_and_url.
"""
vertex_credentials, vertex_project, vertex_location = self._get_common_auth_components(litellm_params)
vertex_credentials, vertex_project, vertex_location = (
self._get_common_auth_components(litellm_params)
)
_auth_header, vertex_project = self._ensure_access_token(
credentials=vertex_credentials,
@ -238,7 +245,9 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM):
Returns:
Tuple of headers and API base
"""
vertex_credentials, vertex_project, vertex_location = self._get_common_auth_components(litellm_params)
vertex_credentials, vertex_project, vertex_location = (
self._get_common_auth_components(litellm_params)
)
_auth_header, vertex_project = await self._ensure_access_token_async(
credentials=vertex_credentials,
@ -256,7 +265,6 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM):
api_base=api_base,
litellm_params=litellm_params,
)
def transform_generate_content_request(
self,
@ -269,6 +277,7 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM):
GenerateContentConfigDict,
GenerateContentRequestDict,
)
typed_generate_content_request = GenerateContentRequestDict(
model=model,
contents=contents,
@ -279,7 +288,7 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM):
request_dict = cast(dict, typed_generate_content_request)
return request_dict
def transform_generate_content_response(
self,
model: str,
@ -297,6 +306,7 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM):
Transformed response data
"""
from litellm.types.google_genai.main import GenerateContentResponse
try:
response = raw_response.json()
except Exception as e:
@ -305,7 +315,7 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM):
status_code=raw_response.status_code,
headers=raw_response.headers,
)
logging_obj.model_call_details["httpx_response"] = raw_response
return GenerateContentResponse(**response)
return GenerateContentResponse(**response)

View file

@ -1,7 +1,8 @@
"""
Transformation for Calling Google models in their native format.
"""
from typing import Dict, Literal, Optional, Union
from typing import Any, Dict, Literal, Optional, Union
from litellm.llms.gemini.google_genai.transformation import GoogleGenAIConfig
from litellm.types.router import GenericLiteLLMParams
@ -58,22 +59,21 @@ class VertexAIGoogleGenAIConfig(GoogleGenAIConfig):
Returns:
Mapped parameters for the provider
"""
from litellm.types.google_genai.main import GenerateContentConfigDict
_generate_content_config_dict = GenerateContentConfigDict()
_generate_content_config_dict: Dict = {}
for param, value in generate_content_config_dict.items():
camel_case_key = self._camel_to_snake(param)
_generate_content_config_dict[camel_case_key] = value
return dict(_generate_content_config_dict)
return _generate_content_config_dict
def transform_generate_content_request(
self,
model: str,
contents: any,
tools: Optional[any],
contents: Any,
tools: Optional[Any],
generate_content_config_dict: Dict,
system_instruction: Optional[any] = None,
system_instruction: Optional[Any] = None,
) -> dict:
"""
Transform the generate content request for Vertex AI.

View file

@ -3,7 +3,6 @@ import asyncio
import copy
import json
import traceback
from litellm._uuid import uuid
from base64 import b64encode
from datetime import datetime
from typing import Dict, List, Optional, Tuple, Union
@ -25,6 +24,7 @@ from starlette.datastructures import UploadFile as StarletteUploadFile
import litellm
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.constants import MAXIMUM_TRACEBACK_LINES_TO_LOG
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
@ -424,10 +424,10 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
for field_name, field_value in form_data.items():
if isinstance(field_value, (StarletteUploadFile, UploadFile)):
files[
field_name
] = await HttpPassThroughEndpointHelpers._build_request_files_from_upload_file(
upload_file=field_value
files[field_name] = (
await HttpPassThroughEndpointHelpers._build_request_files_from_upload_file(
upload_file=field_value
)
)
else:
form_data_dict[field_name] = field_value
@ -476,7 +476,11 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
user_api_key_request_route=user_api_key_dict.request_route,
user_api_key_spend=user_api_key_dict.spend,
user_api_key_max_budget=user_api_key_dict.max_budget,
user_api_key_budget_reset_at=user_api_key_dict.budget_reset_at.isoformat() if user_api_key_dict.budget_reset_at else None,
user_api_key_budget_reset_at=(
user_api_key_dict.budget_reset_at.isoformat()
if user_api_key_dict.budget_reset_at
else None
),
)
)
@ -496,7 +500,7 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
kwargs = {
"litellm_params": {
**litellm_params_in_body,
**litellm_params_in_body, # type: ignore
"metadata": _metadata,
"proxy_server_request": {
"url": str(request.url),
@ -509,9 +513,9 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
"passthrough_logging_payload": passthrough_logging_payload,
}
logging_obj.model_call_details[
"passthrough_logging_payload"
] = passthrough_logging_payload
logging_obj.model_call_details["passthrough_logging_payload"] = (
passthrough_logging_payload
)
return kwargs
@ -923,7 +927,6 @@ def create_pass_through_route(
):
# check if target is an adapter.py or a url
from litellm._uuid import uuid
from litellm.proxy.types_utils.utils import get_instance_fn
try:
@ -1367,7 +1370,6 @@ async def create_pass_through_endpoints(
Create new pass-through endpoint
"""
from litellm._uuid import uuid
from litellm.proxy.proxy_server import (
get_config_general_settings,
update_config_general_settings,

View file

@ -10,7 +10,7 @@ from pydantic import BaseModel
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.constants import REDACTED_BY_LITELM_STRING, MAX_STRING_LENGTH_PROMPT_IN_DB
from litellm.constants import MAX_STRING_LENGTH_PROMPT_IN_DB, REDACTED_BY_LITELM_STRING
from litellm.litellm_core_utils.core_helpers import get_litellm_metadata_from_kwargs
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.proxy._types import SpendLogsMetadata, SpendLogsPayload
@ -21,6 +21,7 @@ from litellm.types.utils import (
StandardLoggingModelInformation,
StandardLoggingPayload,
StandardLoggingVectorStoreRequest,
VectorStoreSearchResponse,
)
from litellm.utils import get_end_user_id_for_cost_tracking
@ -297,7 +298,9 @@ def get_logging_payload( # noqa: PLR0915
id = f"{id}_cache_hit{time.time()}" # SpendLogs does not allow duplicate request_id
mcp_namespaced_tool_name = None
mcp_tool_call_metadata = clean_metadata.get("mcp_tool_call_metadata", {})
mcp_tool_call_metadata: Optional[StandardLoggingMCPToolCall] = clean_metadata.get(
"mcp_tool_call_metadata"
)
if mcp_tool_call_metadata is not None:
mcp_namespaced_tool_name = mcp_tool_call_metadata.get(
"namespaced_tool_name", None
@ -505,23 +508,23 @@ def _sanitize_request_body_for_spend_logs_payload(
# This split ensures we keep more context from the end of conversations
start_ratio = 0.35
end_ratio = 0.65
# Calculate character distribution
start_chars = int(MAX_STRING_LENGTH_PROMPT_IN_DB * start_ratio)
end_chars = int(MAX_STRING_LENGTH_PROMPT_IN_DB * end_ratio)
# Ensure we don't exceed the total limit
total_keep = start_chars + end_chars
if total_keep > MAX_STRING_LENGTH_PROMPT_IN_DB:
end_chars = MAX_STRING_LENGTH_PROMPT_IN_DB - start_chars
# If the string length is less than what we want to keep, just truncate normally
if len(value) <= MAX_STRING_LENGTH_PROMPT_IN_DB:
return value
# Calculate how many characters are being skipped
skipped_chars = len(value) - total_keep
# Build the truncated string: beginning + truncation marker + end
truncated_value = (
f"{value[:start_chars]}"
@ -567,8 +570,9 @@ def _get_vector_store_request_for_spend_logs_payload(
if vector_store_request_metadata is None:
return None
for vector_store_request in vector_store_request_metadata:
vector_store_search_response = (
vector_store_request.get("vector_store_search_response", {}) or {}
vector_store_search_response: VectorStoreSearchResponse = (
vector_store_request.get("vector_store_search_response")
or VectorStoreSearchResponse()
)
response_data = vector_store_search_response.get("data", []) or []
for response_item in response_data:

View file

@ -3442,7 +3442,7 @@ class Router:
*[try_retrieve_batch(model) for model in filtered_model_list]
)
final_results = {
final_results: Dict = {
"object": "list",
"data": [],
"first_id": None,

View file

@ -1308,7 +1308,7 @@ class MCPListToolsFailedEvent(BaseLiteLLMOpenAIResponseObject):
item_id: str
# MCP Call Events
# MCP Call Events
class MCPCallInProgressEvent(BaseLiteLLMOpenAIResponseObject):
type: Literal[ResponsesAPIStreamEvents.MCP_CALL_IN_PROGRESS]
sequence_number: int

View file

@ -1,10 +1,14 @@
main.py:1503: error: Argument "api_key" to "completion" of "AzureChatCompletion" has incompatible type "str | None"; expected "str" [arg-type]
main.py:1508: error: Argument "azure_ad_token" to "completion" of "AzureChatCompletion" has incompatible type "Any | str | None"; expected "str" [arg-type]
main.py:1509: error: Argument "azure_ad_token_provider" to "completion" of "AzureChatCompletion" has incompatible type "Any | None"; expected "Callable[..., Any]" [arg-type]
main.py:1579: error: Argument "api_key" to "completion" of "AzureTextCompletion" has incompatible type "str | None"; expected "str" [arg-type]
main.py:1580: error: Argument "api_base" to "completion" of "AzureTextCompletion" has incompatible type "str | None"; expected "str" [arg-type]
main.py:1583: error: Argument "azure_ad_token" to "completion" of "AzureTextCompletion" has incompatible type "Any | str | None"; expected "str" [arg-type]
proxy/hooks/parallel_request_limiter_v3.py:383: error: Item "None" of "RateLimitDescriptorRateLimitObject | None" has no attribute "get" [union-attr]
proxy/hooks/parallel_request_limiter_v3.py:384: error: Item "None" of "RateLimitDescriptorRateLimitObject | None" has no attribute "get" [union-attr]
proxy/hooks/parallel_request_limiter_v3.py:385: error: Item "None" of "RateLimitDescriptorRateLimitObject | None" has no attribute "get" [union-attr]
proxy/hooks/parallel_request_limiter_v3.py:386: error: Item "None" of "RateLimitDescriptorRateLimitObject | None" has no attribute "get" [union-attr]
types/llms/openai.py:46: error: Module "openai.types.responses.response_create_params" has no attribute "Text" [attr-defined]
router.py:3445: error: Need type annotation for "final_results" [var-annotated]
main.py:1581: error: Argument "api_base" to "completion" of "AzureTextCompletion" has incompatible type "str | None"; expected "str" [arg-type]
llms/gemini/google_genai/transformation.py:111: error: TypedDict key must be a string literal; expected one of ("http_options", "system_instruction", "temperature", "top_p", "top_k", ...) [literal-required]
llms/bedrock/chat/invoke_agent/transformation.py:392: error: Need type annotation for "pre_processing" [var-annotated]
llms/bedrock/chat/invoke_agent/transformation.py:396: error: Need type annotation for "model_output" [var-annotated]
llms/bedrock/chat/invoke_agent/transformation.py:400: error: Need type annotation for "metadata" [var-annotated]
llms/bedrock/chat/invoke_agent/transformation.py:415: error: Need type annotation for "orchestration_trace" [var-annotated]
llms/bedrock/chat/invoke_agent/transformation.py:419: error: Need type annotation for "model_invocation" [var-annotated]
llms/vertex_ai/google_genai/transformation.py:67: error: TypedDict key must be a string literal; expected one of ("http_options", "system_instruction", "temperature", "top_p", "top_k", ...) [literal-required]
proxy/spend_tracking/spend_tracking_utils.py:300: error: Need type annotation for "mcp_tool_call_metadata" [var-annotated]
proxy/spend_tracking/spend_tracking_utils.py:571: error: Need type annotation for "vector_store_search_response" [var-annotated]
proxy/pass_through_endpoints/pass_through_endpoints.py:499: error: Unsupported type "dict[str, Any]" for ** expansion in TypedDict [typeddict-item]
Found 13 errors in 8 files (checked 1114 source files)