Fix mypy issues

This commit is contained in:
Sameer Kankute 2026-02-26 10:42:01 +05:30
parent f78104d34c
commit 3326893420
7 changed files with 30 additions and 24 deletions

View file

@ -1,7 +1,7 @@
import asyncio
import concurrent.futures
import json
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union, cast
import litellm
from litellm._logging import verbose_logger
@ -88,7 +88,7 @@ class RealTimeStreaming:
message_obj = message
else:
message_obj = json.loads(message)
self._collect_tool_calls_from_response_done(message_obj)
self._collect_tool_calls_from_response_done(cast(dict, message_obj))
try:
if (
not isinstance(message, dict)
@ -355,11 +355,11 @@ class RealTimeStreaming:
== "conversation.item.input_audio_transcription.completed"
):
transcript = event.get("transcript", "")
self._collect_user_input_from_backend_event(event)
self._collect_user_input_from_backend_event(cast(dict, event))
self.store_message(event_str)
await self.websocket.send_text(event_str)
blocked = await self.run_realtime_guardrails(
transcript, item_id=event.get("item_id")
cast(str, transcript), item_id=cast(Optional[str], event.get("item_id"))
)
if not blocked:
await self.backend_ws.send(

View file

@ -11,7 +11,6 @@ from typing import Any, List, Optional
import httpx
from litellm.types.utils import Usage
from litellm.llms.bedrock.chat.invoke_transformations.amazon_qwen3_transformation import (
AmazonQwen3Config,
)
@ -19,7 +18,7 @@ from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation
LiteLLMLoggingObj,
)
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import ModelResponse
from litellm.types.utils import ModelResponse, Usage
class AmazonQwen2Config(AmazonQwen3Config):
@ -80,10 +79,14 @@ class AmazonQwen2Config(AmazonQwen3Config):
# Set usage information if available in response
if "usage" in response_data:
usage_data = response_data["usage"]
model_response.usage = Usage(
prompt_tokens=usage_data.get("prompt_tokens", 0),
completion_tokens=usage_data.get("completion_tokens", 0),
total_tokens=usage_data.get("total_tokens", 0),
setattr(
model_response,
"usage",
Usage(
prompt_tokens=usage_data.get("prompt_tokens", 0),
completion_tokens=usage_data.get("completion_tokens", 0),
total_tokens=usage_data.get("total_tokens", 0),
),
)
return model_response

View file

@ -10,14 +10,13 @@ from typing import Any, List, Optional
import httpx
from litellm.types.utils import Usage
from litellm.llms.base_llm.chat.transformation import BaseConfig
from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation import (
AmazonInvokeConfig,
LiteLLMLoggingObj,
)
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import ModelResponse
from litellm.types.utils import ModelResponse, Usage
class AmazonQwen3Config(AmazonInvokeConfig, BaseConfig):
@ -202,10 +201,14 @@ class AmazonQwen3Config(AmazonInvokeConfig, BaseConfig):
# Set usage information if available in response
if "usage" in response_data:
usage_data = response_data["usage"]
model_response.usage = Usage(
prompt_tokens=usage_data.get("prompt_tokens", 0),
completion_tokens=usage_data.get("completion_tokens", 0),
total_tokens=usage_data.get("total_tokens", 0),
setattr(
model_response,
"usage",
Usage(
prompt_tokens=usage_data.get("prompt_tokens", 0),
completion_tokens=usage_data.get("completion_tokens", 0),
total_tokens=usage_data.get("total_tokens", 0),
),
)
return model_response

View file

@ -1,4 +1,4 @@
from typing import Dict, List, Optional, Set, Tuple
from typing import Dict, List, Optional, Set, Tuple, cast
from fastapi import HTTPException
from starlette.datastructures import Headers
@ -539,7 +539,7 @@ class MCPRequestHandler:
allowed_tools = team_tools
else:
# No team restrictions → use key restrictions
allowed_tools = key_tools
allowed_tools = cast(List[str], key_tools)
# Intersect with agent's tool permissions if agent_id is set
if user_api_key_auth.agent_id:

View file

@ -312,7 +312,7 @@ class DBSpendUpdateWriter:
prisma_client: Optional[PrismaClient],
user_api_key_cache: DualCache,
litellm_proxy_budget_name: Optional[str],
payload_copy: dict,
payload_copy: SpendLogsPayload,
request_tags: Optional[Any],
):
"""

View file

@ -5,7 +5,7 @@ usage/spend data by querying the aggregated daily activity endpoints.
import json
from datetime import date
from typing import Any, AsyncIterator, Callable, Dict, List, Literal, Optional
from typing import Any, AsyncIterator, Callable, Dict, List, Literal, Optional, cast
import litellm
from litellm._logging import verbose_proxy_logger
@ -492,17 +492,17 @@ async def _process_tool_call(
"tool_label": handler["label"],
"arguments": fn_args,
}
yield _sse({**tool_event_base, "status": "running"})
yield _sse(cast(SSEToolCallEvent, {**tool_event_base, "status": "running"}))
try:
tool_result = await _execute_tool_call(
handler, fn_name, fn_args, user_id, is_admin
)
yield _sse({**tool_event_base, "status": "complete"})
yield _sse(cast(SSEToolCallEvent, {**tool_event_base, "status": "complete"}))
except Exception as e:
verbose_proxy_logger.error("Tool %s failed: %s", fn_name, e)
tool_result = f"Error fetching {handler['label']}. Please try again."
yield _sse({**tool_event_base, "status": "error"})
yield _sse(cast(SSEToolCallEvent, {**tool_event_base, "status": "error"}))
chat_messages.append(
{"role": "tool", "tool_call_id": tc.id, "content": tool_result}

View file

@ -7053,7 +7053,7 @@ class Router:
user_model_info = deployment.get("model_info") or {}
if model_info is not None:
model_info.update(user_model_info)
model_info.update(cast(ModelInfo, user_model_info))
return model_info