Fix mypy issues

This commit is contained in:
Sameer Kankute 2026-02-26 10:42:01 +05:30
parent 61e2fdf463
commit 5ed564aeca
7 changed files with 766 additions and 209 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
@ -92,7 +92,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)
@ -428,11 +428,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._send_to_backend(

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

@ -139,53 +139,20 @@ class DBSpendUpdateWriter:
payload["startTime"] = payload["startTime"].isoformat()
if isinstance(payload["endTime"], datetime):
payload["endTime"] = payload["endTime"].isoformat()
if org_id is not None and org_id != "":
payload["organization_id"] = org_id
if team_id is not None and team_id != "":
payload["team_id"] = team_id
asyncio.create_task(
self._update_user_db(
response_cost=response_cost,
user_id=user_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
litellm_proxy_budget_name=litellm_proxy_budget_name,
end_user_id=end_user_id,
)
)
asyncio.create_task(
self._update_key_db(
response_cost=response_cost,
hashed_token=hashed_token,
prisma_client=prisma_client,
)
)
asyncio.create_task(
self._update_team_db(
response_cost=response_cost,
team_id=team_id,
user_id=user_id,
prisma_client=prisma_client,
)
)
asyncio.create_task(
self._update_org_db(
response_cost=response_cost,
org_id=org_id,
prisma_client=prisma_client,
)
)
asyncio.create_task(
self._update_tag_db(
response_cost=response_cost,
request_tags=copy.deepcopy(payload.get("request_tags")),
prisma_client=prisma_client,
)
)
# One deepcopy shared by all 6 daily spend helpers (was 5, fixes agent bug)
payload_copy = copy.deepcopy(payload)
# Deepcopy request_tags for _update_tag_db
request_tags = copy.deepcopy(payload.get("request_tags"))
# Keep _insert_spend_log_to_db awaited inline (not a task, preserve current behavior)
if disable_spend_logs is False:
await self._insert_spend_log_to_db(
payload=copy.deepcopy(payload),
@ -196,44 +163,20 @@ class DBSpendUpdateWriter:
"disable_spend_logs=True. Skipping writing spend logs to db. Other spend updates - Key/User/Team table will still occur."
)
# Single task replaces 11 create_task() calls
asyncio.create_task(
self.add_spend_log_transaction_to_daily_user_transaction(
payload=copy.deepcopy(payload),
prisma_client=prisma_client,
)
)
asyncio.create_task(
self.add_spend_log_transaction_to_daily_end_user_transaction(
payload=copy.deepcopy(payload),
prisma_client=prisma_client,
)
)
asyncio.create_task(
self.add_spend_log_transaction_to_daily_agent_transaction(
payload=payload,
prisma_client=prisma_client,
)
)
asyncio.create_task(
self.add_spend_log_transaction_to_daily_team_transaction(
payload=copy.deepcopy(payload),
prisma_client=prisma_client,
)
)
asyncio.create_task(
self.add_spend_log_transaction_to_daily_org_transaction(
payload=copy.deepcopy(payload),
self._batch_database_updates(
response_cost=response_cost,
user_id=user_id,
hashed_token=hashed_token,
team_id=team_id,
org_id=org_id,
end_user_id=end_user_id,
prisma_client=prisma_client,
)
)
asyncio.create_task(
self.add_spend_log_transaction_to_daily_tag_transaction(
payload=copy.deepcopy(payload),
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
litellm_proxy_budget_name=litellm_proxy_budget_name,
payload_copy=payload_copy,
request_tags=request_tags,
)
)
@ -357,6 +300,157 @@ class DBSpendUpdateWriter:
"_enqueue_tool_registry_upsert error (non-blocking): %s", e
)
async def _batch_database_updates(
self,
*,
response_cost: Optional[float],
user_id: Optional[str],
hashed_token: Optional[str],
team_id: Optional[str],
org_id: Optional[str],
end_user_id: Optional[str],
prisma_client: Optional[PrismaClient],
user_api_key_cache: DualCache,
litellm_proxy_budget_name: Optional[str],
payload_copy: SpendLogsPayload,
request_tags: Optional[Any],
):
"""
Runs all 11 spend-update helpers sequentially inside a single asyncio task.
Each helper is wrapped in try/except so one failure doesn't prevent the others.
"""
try:
await self._update_user_db(
response_cost=response_cost,
user_id=user_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
litellm_proxy_budget_name=litellm_proxy_budget_name,
end_user_id=end_user_id,
)
except Exception:
verbose_proxy_logger.debug(
"_batch_database_updates: _update_user_db failed: %s",
traceback.format_exc(),
)
try:
await self._update_key_db(
response_cost=response_cost,
hashed_token=hashed_token,
prisma_client=prisma_client,
)
except Exception:
verbose_proxy_logger.debug(
"_batch_database_updates: _update_key_db failed: %s",
traceback.format_exc(),
)
try:
await self._update_team_db(
response_cost=response_cost,
team_id=team_id,
user_id=user_id,
prisma_client=prisma_client,
)
except Exception:
verbose_proxy_logger.debug(
"_batch_database_updates: _update_team_db failed: %s",
traceback.format_exc(),
)
try:
await self._update_org_db(
response_cost=response_cost,
org_id=org_id,
prisma_client=prisma_client,
)
except Exception:
verbose_proxy_logger.debug(
"_batch_database_updates: _update_org_db failed: %s",
traceback.format_exc(),
)
try:
await self._update_tag_db(
response_cost=response_cost,
request_tags=request_tags,
prisma_client=prisma_client,
)
except Exception:
verbose_proxy_logger.debug(
"_batch_database_updates: _update_tag_db failed: %s",
traceback.format_exc(),
)
try:
await self.add_spend_log_transaction_to_daily_user_transaction(
payload=payload_copy,
prisma_client=prisma_client,
)
except Exception:
verbose_proxy_logger.debug(
"_batch_database_updates: add_spend_log_transaction_to_daily_user_transaction failed: %s",
traceback.format_exc(),
)
try:
await self.add_spend_log_transaction_to_daily_end_user_transaction(
payload=payload_copy,
prisma_client=prisma_client,
)
except Exception:
verbose_proxy_logger.debug(
"_batch_database_updates: add_spend_log_transaction_to_daily_end_user_transaction failed: %s",
traceback.format_exc(),
)
try:
await self.add_spend_log_transaction_to_daily_agent_transaction(
payload=payload_copy,
prisma_client=prisma_client,
)
except Exception:
verbose_proxy_logger.debug(
"_batch_database_updates: add_spend_log_transaction_to_daily_agent_transaction failed: %s",
traceback.format_exc(),
)
try:
await self.add_spend_log_transaction_to_daily_team_transaction(
payload=payload_copy,
prisma_client=prisma_client,
)
except Exception:
verbose_proxy_logger.debug(
"_batch_database_updates: add_spend_log_transaction_to_daily_team_transaction failed: %s",
traceback.format_exc(),
)
try:
await self.add_spend_log_transaction_to_daily_org_transaction(
payload=payload_copy,
org_id=org_id,
prisma_client=prisma_client,
)
except Exception:
verbose_proxy_logger.debug(
"_batch_database_updates: add_spend_log_transaction_to_daily_org_transaction failed: %s",
traceback.format_exc(),
)
try:
await self.add_spend_log_transaction_to_daily_tag_transaction(
payload=payload_copy,
prisma_client=prisma_client,
)
except Exception:
verbose_proxy_logger.debug(
"_batch_database_updates: add_spend_log_transaction_to_daily_tag_transaction failed: %s",
traceback.format_exc(),
)
async def _update_key_db(
self,
response_cost: Optional[float],
@ -1061,7 +1155,7 @@ class DBSpendUpdateWriter:
team_id = key.split("::")[1]
user_id = key.split("::")[3]
team_memberships_to_invalidate.append((user_id, team_id))
for i in range(n_retry_times + 1):
start_time = time.time()
try:
@ -1098,11 +1192,13 @@ class DBSpendUpdateWriter:
_raise_failed_update_spend_exception(
e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj
)
# Invalidate cache for updated team memberships
# This ensures budget checks read fresh spend data from the database
if team_memberships_to_invalidate and proxy_logging_obj is not None:
user_api_key_cache = proxy_logging_obj.call_details.get("user_api_key_cache")
user_api_key_cache = proxy_logging_obj.call_details.get(
"user_api_key_cache"
)
if user_api_key_cache is not None:
for user_id, team_id in team_memberships_to_invalidate:
cache_key = "team_membership:{}:{}".format(user_id, team_id)
@ -1414,7 +1510,9 @@ class DBSpendUpdateWriter:
),
"endpoint": transaction.get("endpoint") or "",
"prompt_tokens": transaction["prompt_tokens"],
"completion_tokens": transaction["completion_tokens"],
"completion_tokens": transaction[
"completion_tokens"
],
"spend": transaction["spend"],
"api_requests": transaction["api_requests"],
"successful_requests": transaction[
@ -1425,12 +1523,14 @@ class DBSpendUpdateWriter:
# Add cache-related fields if they exist
if "cache_read_input_tokens" in transaction:
common_data["cache_read_input_tokens"] = (
transaction.get("cache_read_input_tokens", 0)
)
common_data[
"cache_read_input_tokens"
] = transaction.get("cache_read_input_tokens", 0)
if "cache_creation_input_tokens" in transaction:
common_data["cache_creation_input_tokens"] = (
transaction.get("cache_creation_input_tokens", 0)
common_data[
"cache_creation_input_tokens"
] = transaction.get(
"cache_creation_input_tokens", 0
)
if entity_type == "tag" and "request_id" in transaction:
@ -1473,10 +1573,14 @@ class DBSpendUpdateWriter:
}
if entity_type == "tag" and "request_id" in transaction:
update_data["request_id"] = transaction.get("request_id")
update_data["request_id"] = transaction.get(
"request_id"
)
# Add endpoint to update_data so existing rows get their endpoint field updated
update_data["endpoint"] = transaction.get("endpoint") or ""
update_data["endpoint"] = (
transaction.get("endpoint") or ""
)
table.upsert(
where=where_clause,
@ -1660,7 +1764,9 @@ class DBSpendUpdateWriter:
self,
payload: Union[dict, SpendLogsPayload],
prisma_client: PrismaClient,
type: Literal["user", "team", "org", "request_tags", "end_user", "agent"] = "user",
type: Literal[
"user", "team", "org", "request_tags", "end_user", "agent"
] = "user",
) -> Optional[BaseDailySpendTransaction]:
common_expected_keys = ["startTime", "api_key"]
if type == "user":
@ -1719,7 +1825,7 @@ class DBSpendUpdateWriter:
endpoint = None
if call_type:
endpoint = ROUTE_ENDPOINT_MAPPING.get(call_type, None)
daily_transaction = BaseDailySpendTransaction(
date=date,
api_key=payload["api_key"],
@ -1931,7 +2037,7 @@ class DBSpendUpdateWriter:
endpoint_str = base_daily_transaction.get("endpoint") or ""
daily_transaction_key = f"{payload['agent_id']}_{base_daily_transaction['date']}_{payload_with_agent_id['api_key']}_{payload_with_agent_id['model']}_{payload_with_agent_id['custom_llm_provider']}_{endpoint_str}"
daily_transaction = DailyAgentSpendTransaction(
agent_id=payload['agent_id'], **base_daily_transaction
agent_id=payload["agent_id"], **base_daily_transaction
)
await self.daily_agent_spend_update_queue.add_update(
update={daily_transaction_key: daily_transaction}

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}

File diff suppressed because it is too large Load diff