fix some lint issues

This commit is contained in:
KnyazSh 2026-06-18 14:35:46 +00:00
parent d185b8fa49
commit 93d272f40c
11 changed files with 199 additions and 194 deletions

View file

@ -5936,7 +5936,7 @@ def emit_standard_logging_payload(payload: StandardLoggingPayload):
if os.getenv("LITELLM_PRINT_STANDARD_LOGGING_PAYLOAD"):
try:
print(json.dumps(payload, indent=4, default=str)) # noqa: T201
except Exception as e:
except Exception as e: # noqa: BLE001
verbose_logger.exception(
"Error serializing standard logging payload for debug output: {}".format(
str(e)

View file

@ -7,7 +7,6 @@ Based on official GigaChat SDK authentication flow.
import time
import uuid
from typing import Optional, Tuple
import httpx
@ -41,7 +40,7 @@ class GigaChatAuthError(BaseLLMException):
pass
def _get_credentials() -> Optional[str]:
def _get_credentials() -> str | None:
"""Get GigaChat credentials from environment."""
return get_secret_str("GIGACHAT_CREDENTIALS") or get_secret_str("GIGACHAT_API_KEY")
@ -62,10 +61,10 @@ def _get_http_client() -> HTTPHandler:
def get_access_token(
credentials: Optional[str] = None,
scope: Optional[str] = None,
auth_url: Optional[str] = None,
litellm_params: Optional[dict] = None,
credentials: str | None = None,
scope: str | None = None,
auth_url: str | None = None,
litellm_params: dict | None = None,
) -> str:
"""
Get valid access token, using cache if available.
@ -124,10 +123,10 @@ def get_access_token(
async def get_access_token_async(
credentials: Optional[str] = None,
scope: Optional[str] = None,
auth_url: Optional[str] = None,
litellm_params: Optional[dict] = None,
credentials: str | None = None,
scope: str | None = None,
auth_url: str | None = None,
litellm_params: dict | None = None,
) -> str:
"""Async version of get_access_token."""
if not litellm_params:
@ -176,12 +175,12 @@ def _request_token_sync(
credentials: str,
scope: str,
auth_url: str,
) -> Tuple[str, int]:
) -> tuple[str, int]:
"""
Request new access token from GigaChat OAuth endpoint (sync).
Returns:
Tuple of (access_token, expires_at_ms)
tuple of (access_token, expires_at_ms)
"""
headers = {
"Authorization": f"Basic {credentials}",
@ -213,7 +212,7 @@ async def _request_token_async(
credentials: str,
scope: str,
auth_url: str,
) -> Tuple[str, int]:
) -> tuple[str, int]:
"""Async version of _request_token_sync."""
headers = {
"Authorization": f"Basic {credentials}",
@ -244,7 +243,7 @@ async def _request_token_async(
)
def _parse_token_response(response: httpx.Response) -> Tuple[str, int]:
def _parse_token_response(response: httpx.Response) -> tuple[str, int]:
"""Parse OAuth token response."""
data = response.json()

View file

@ -4,7 +4,7 @@ GigaChat Streaming Response Handler
import json
import uuid
from typing import Any, Optional
from typing import Any
from litellm.llms.gigachat.utils import convert_usage
from litellm.types.llms.openai import (
@ -21,7 +21,7 @@ class GigaChatModelResponseIterator:
self,
streaming_response: Any,
sync_stream: bool,
json_mode: Optional[bool] = False,
json_mode: bool | None = False,
):
self.streaming_response = streaming_response
self.response_iterator = self.streaming_response
@ -30,9 +30,9 @@ class GigaChatModelResponseIterator:
def chunk_parser(self, chunk: dict) -> GenericStreamingChunk:
"""Parse a single streaming chunk from GigaChat."""
text = ""
tool_use: Optional[ChatCompletionToolCallChunk] = None
tool_use: ChatCompletionToolCallChunk | None = None
is_finished = False
finish_reason: Optional[str] = None
finish_reason: str | None = None
choices = chunk.get("choices", [])
if not choices:

View file

@ -7,7 +7,7 @@ Transforms OpenAI-format requests to GigaChat format and back.
import json
import time
import uuid
from typing import TYPE_CHECKING, Any, AsyncIterator, Iterator, List, Optional, Union
from typing import TYPE_CHECKING, Any, AsyncIterator, Iterator, Union
import httpx
@ -60,36 +60,36 @@ class GigaChatConfig(BaseConfig):
stream: Enable streaming
"""
temperature: Optional[float] = None
top_p: Optional[float] = None
max_tokens: Optional[int] = None
repetition_penalty: Optional[float] = None
profanity_check: Optional[bool] = None
temperature: float | None = None
top_p: float | None = None
max_tokens: int | None = None
repetition_penalty: float | None = None
profanity_check: bool | None = None
def __init__(
self,
temperature: Optional[float] = None,
top_p: Optional[float] = None,
max_tokens: Optional[int] = None,
repetition_penalty: Optional[float] = None,
profanity_check: Optional[bool] = None,
temperature: float | None = None,
top_p: float | None = None,
max_tokens: int | None = None,
repetition_penalty: float | None = None,
profanity_check: bool | None = None,
) -> None:
locals_ = locals().copy()
for key, value in locals_.items():
if key != "self" and value is not None:
setattr(self.__class__, key, value)
# Instance variables for current request context
self._current_credentials: Optional[str] = None
self._current_api_base: Optional[str] = None
self._current_credentials: str | None = None
self._current_api_base: str | None = None
def get_complete_url(
self,
api_base: Optional[str],
api_key: Optional[str],
api_base: str | None,
api_key: str | None,
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
stream: bool | None = None,
) -> str:
"""Get complete API URL for chat completions."""
base = get_api_base(api_base)
@ -99,11 +99,11 @@ class GigaChatConfig(BaseConfig):
self,
headers: dict,
model: str,
messages: List[AllMessageValues],
messages: list[AllMessageValues],
optional_params: dict,
litellm_params: dict,
api_key: Optional[str] = None,
api_base: Optional[str] = None,
api_key: str | None = None,
api_base: str | None = None,
) -> dict:
"""
Set up headers with OAuth token.
@ -128,7 +128,7 @@ class GigaChatConfig(BaseConfig):
return headers
def get_supported_openai_params(self, model: str) -> List[str]:
def get_supported_openai_params(self, model: str) -> list[str]:
"""Return list of supported OpenAI parameters."""
return [
"stream",
@ -201,7 +201,7 @@ class GigaChatConfig(BaseConfig):
return optional_params
def _convert_tools_to_functions(self, tools: List[dict]) -> List[dict]:
def _convert_tools_to_functions(self, tools: list[dict]) -> list[dict]:
"""Convert OpenAI tools format to GigaChat functions format."""
functions = []
for tool in tools:
@ -218,7 +218,7 @@ class GigaChatConfig(BaseConfig):
def _map_tool_choice(
self, tool_choice: Union[str, dict]
) -> Optional[Union[str, dict]]:
) -> Union[str, dict] | None:
"""
Map OpenAI tool_choice to GigaChat function_call format.
@ -258,7 +258,7 @@ class GigaChatConfig(BaseConfig):
# Default to None (don't set function_call)
return None
def _upload_image(self, image_url: str) -> Optional[str]:
def _upload_image(self, image_url: str) -> str | None:
"""
Upload image to GigaChat and return file_id.
@ -281,7 +281,7 @@ class GigaChatConfig(BaseConfig):
def transform_request(
self,
model: str,
messages: List[AllMessageValues],
messages: list[AllMessageValues],
optional_params: dict,
litellm_params: dict,
headers: dict,
@ -316,7 +316,7 @@ class GigaChatConfig(BaseConfig):
return request_data
def _transform_messages(self, messages: List[AllMessageValues]) -> List[dict]:
def _transform_messages(self, messages: list[AllMessageValues]) -> list[dict]:
"""Transform OpenAI messages to GigaChat format."""
transformed = []
@ -395,12 +395,12 @@ class GigaChatConfig(BaseConfig):
model_response: ModelResponse,
logging_obj: LiteLLMLoggingObj,
request_data: dict,
messages: List[AllMessageValues],
messages: list[AllMessageValues],
optional_params: dict,
litellm_params: dict,
encoding: Any,
api_key: Optional[str] = None,
json_mode: Optional[bool] = None,
api_key: str | None = None,
json_mode: bool | None = None,
) -> ModelResponse:
"""Transform GigaChat response to OpenAI format."""
try:
@ -494,7 +494,7 @@ class GigaChatConfig(BaseConfig):
self,
streaming_response: Union[Iterator[str], AsyncIterator[str], ModelResponse],
sync_stream: bool,
json_mode: Optional[bool] = False,
json_mode: bool | None = False,
):
"""Return streaming response iterator."""
from .streaming import GigaChatModelResponseIterator

View file

@ -6,7 +6,7 @@ API Documentation: https://developers.sber.ru/docs/ru/gigachat/api/reference/res
"""
import types
from typing import List, Optional, Tuple, Union
from typing import Union
import httpx
@ -55,7 +55,7 @@ class GigaChatEmbeddingConfig(BaseEmbeddingConfig):
and v is not None
}
def get_supported_openai_params(self, model: str) -> List[str]:
def get_supported_openai_params(self, model: str) -> list[str]:
"""GigaChat embeddings don't support additional parameters."""
return []
@ -71,26 +71,26 @@ class GigaChatEmbeddingConfig(BaseEmbeddingConfig):
def _get_openai_compatible_provider_info(
self,
api_base: Optional[str],
api_key: Optional[str],
) -> Tuple[str, Optional[str], Optional[str]]:
api_base: str | None,
api_key: str | None,
) -> tuple[str, str | None, str | None]:
"""
Returns provider info for GigaChat.
Returns:
Tuple of (custom_llm_provider, api_base, dynamic_api_key)
tuple of (custom_llm_provider, api_base, dynamic_api_key)
"""
api_base = get_api_base(api_base)
return LlmProviders.GIGACHAT.value, api_base, api_key
def get_complete_url(
self,
api_base: Optional[str],
api_key: Optional[str],
api_base: str | None,
api_key: str | None,
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
stream: bool | None = None,
) -> str:
"""Get the complete URL for embeddings endpoint."""
base = get_api_base(api_base)
@ -135,7 +135,7 @@ class GigaChatEmbeddingConfig(BaseEmbeddingConfig):
raw_response: httpx.Response,
model_response: EmbeddingResponse,
logging_obj: LiteLLMLoggingObj,
api_key: Optional[str],
api_key: str | None,
request_data: dict,
optional_params: dict,
litellm_params: dict,
@ -182,11 +182,11 @@ class GigaChatEmbeddingConfig(BaseEmbeddingConfig):
self,
headers: dict,
model: str,
messages: List[AllMessageValues],
messages: list[AllMessageValues],
optional_params: dict,
litellm_params: dict,
api_key: Optional[str] = None,
api_base: Optional[str] = None,
api_key: str | None = None,
api_base: str | None = None,
) -> dict:
"""
Set up headers with OAuth token for GigaChat.

View file

@ -9,7 +9,6 @@ import base64
import hashlib
import re
import uuid
from typing import Dict, Optional, Tuple
from litellm._logging import verbose_logger
from litellm.llms.custom_httpx.http_handler import (
@ -22,7 +21,7 @@ from litellm.types.utils import LlmProviders
from .authenticator import get_access_token, get_access_token_async
# Simple in-memory cache for file IDs
_file_cache: Dict[str, str] = {}
_file_cache: dict[str, str] = {}
def _get_url_hash(url: str) -> str:
@ -30,7 +29,7 @@ def _get_url_hash(url: str) -> str:
return hashlib.sha256(url.encode()).hexdigest()
def _parse_data_url(data_url: str) -> Optional[Tuple[bytes, str, str]]:
def _parse_data_url(data_url: str) -> tuple[bytes, str, str] | None:
"""
Parse data URL (base64 image).
@ -49,7 +48,7 @@ def _parse_data_url(data_url: str) -> Optional[Tuple[bytes, str, str]]:
return content_bytes, content_type, ext
def _download_image_sync(url: str) -> Tuple[bytes, str, str]:
def _download_image_sync(url: str) -> tuple[bytes, str, str]:
"""Download image from URL synchronously."""
client = _get_httpx_client(params={"ssl_verify": False})
response = client.get(url)
@ -61,7 +60,7 @@ def _download_image_sync(url: str) -> Tuple[bytes, str, str]:
return response.content, content_type, ext
async def _download_image_async(url: str) -> Tuple[bytes, str, str]:
async def _download_image_async(url: str) -> tuple[bytes, str, str]:
"""Download image from URL asynchronously."""
client = get_async_httpx_client(
llm_provider=LlmProviders.GIGACHAT,
@ -78,10 +77,10 @@ async def _download_image_async(url: str) -> Tuple[bytes, str, str]:
def upload_file_sync(
image_url: str,
credentials: Optional[str] = None,
api_base: Optional[str] = None,
litellm_params: Optional[dict] = None,
) -> Optional[str]:
credentials: str | None = None,
api_base: str | None = None,
litellm_params: dict | None = None,
) -> str | None:
"""
Upload file to GigaChat and return file_id (sync).
@ -146,10 +145,10 @@ def upload_file_sync(
async def upload_file_async(
image_url: str,
credentials: Optional[str] = None,
api_base: Optional[str] = None,
litellm_params: Optional[dict] = None,
) -> Optional[str]:
credentials: str | None = None,
api_base: str | None = None,
litellm_params: dict | None = None,
) -> str | None:
"""
Upload file to GigaChat and return file_id (async).

View file

@ -1,5 +1,5 @@
import json
from typing import TYPE_CHECKING, List, Optional, Tuple, cast
from typing import TYPE_CHECKING, cast
import httpx
@ -24,13 +24,13 @@ class GigaChatPassthroughConfig(BasePassthroughConfig):
def get_complete_url(
self,
api_base: Optional[str],
api_key: Optional[str],
api_base: str | None,
api_key: str | None,
model: str,
endpoint: str,
request_query_params: Optional[dict],
request_query_params: dict | None,
litellm_params: dict,
) -> Tuple["URL", str]:
) -> tuple["URL", str]:
"""Get complete API URL for chat completions."""
base_target_url = self.get_api_base(api_base)
@ -48,11 +48,11 @@ class GigaChatPassthroughConfig(BasePassthroughConfig):
self,
headers: dict,
model: str,
messages: List[AllMessageValues],
messages: list[AllMessageValues],
optional_params: dict,
litellm_params: dict,
api_key: Optional[str] = None,
api_base: Optional[str] = None,
api_key: str | None = None,
api_base: str | None = None,
) -> dict:
"""
Set up headers with OAuth token.
@ -76,7 +76,7 @@ class GigaChatPassthroughConfig(BasePassthroughConfig):
request_data: dict,
logging_obj: "LiteLLMLoggingObj",
endpoint: str,
) -> Optional["CostResponseTypes"]:
) -> "CostResponseTypes" | None:
from litellm import encoding
from litellm.types.utils import LlmProviders, ModelResponse
from litellm.utils import ProviderConfigManager
@ -140,12 +140,12 @@ class GigaChatPassthroughConfig(BasePassthroughConfig):
def handle_logging_collected_chunks(
self,
all_chunks: List[str],
all_chunks: list[str],
litellm_logging_obj: "LiteLLMLoggingObj",
model: str,
custom_llm_provider: str,
endpoint: str,
) -> Optional["CostResponseTypes"]:
) -> "CostResponseTypes" | None:
"""
1. Convert all_chunks to a ModelResponseStream
2. combine model_response_stream to model_response
@ -208,20 +208,20 @@ class GigaChatPassthroughConfig(BasePassthroughConfig):
return None
@staticmethod
def get_api_base(api_base: Optional[str] = None) -> Optional[str]:
def get_api_base(api_base: str | None = None) -> str | None:
return api_base or get_secret_str("GIGACHAT_API_BASE") or GIGACHAT_BASE_URL
@staticmethod
def get_api_key(
api_key: Optional[str] = None,
) -> Optional[str]:
api_key: str | None = None,
) -> str | None:
return api_key or get_secret_str("GIGACHAT_API_KEY")
@staticmethod
def get_base_model(model: str) -> Optional[str]:
def get_base_model(model: str) -> str | None:
return model
def get_models(
self, api_key: Optional[str] = None, api_base: Optional[str] = None
) -> List[str]:
self, api_key: str | None = None, api_base: str | None = None
) -> list[str]:
return super().get_models(api_key, api_base)

View file

@ -1,5 +1,3 @@
from typing import Optional
from litellm.secret_managers.main import get_secret_str
from litellm.types.utils import PromptTokensDetailsWrapper, Usage
@ -30,5 +28,5 @@ def convert_usage(usage_data: dict[str, int]) -> Usage:
)
def get_api_base(api_base: Optional[str] = None) -> Optional[str]:
def get_api_base(api_base: str | None = None) -> str | None:
return api_base or get_secret_str("GIGACHAT_API_BASE") or GIGACHAT_BASE_URL

View file

@ -50,7 +50,7 @@ class AsyncPassthroughStreamingResponse(AsyncGenerator[Any, Any]):
self._iterator: AsyncGenerator[bytes, Any]
self._litellm_logging_obj = litellm_logging_obj
self._provider_config = provider_config
self._raw_bytes: List[bytes] = []
self._raw_bytes: list[bytes] = []
self._flush_scheduled = False
self._background_tasks: set[asyncio.Task] = set()
@ -92,10 +92,10 @@ class AsyncPassthroughStreamingResponse(AsyncGenerator[Any, Any]):
self._iterator = cast(
AsyncGenerator[bytes, Any], self._response.aiter_bytes()
)
except Exception:
except Exception: # noqa: BLE001
try:
await self._response.aclose()
except Exception:
except Exception: # noqa: BLE001
pass
raise
return self
@ -120,7 +120,7 @@ class AsyncPassthroughStreamingResponse(AsyncGenerator[Any, Any]):
# Remove the task from the set when it finishes to avoid memory leaks
task.add_done_callback(self._background_tasks.discard)
except Exception as e:
except Exception as e: # noqa: BLE001
verbose_logger.exception(
"Failed to schedule passthrough spend-tracking flush; "
"%d buffered chunks dropped: %s",
@ -138,11 +138,11 @@ class AsyncPassthroughStreamingResponse(AsyncGenerator[Any, Any]):
chunk = await self._iterator.__anext__()
self._raw_bytes.append(chunk)
return chunk
except Exception:
except Exception: # noqa: BLE001
self._start_flush()
try:
await self._response.aclose()
except Exception:
except Exception: # noqa: BLE001
pass
raise
@ -161,7 +161,7 @@ class AsyncPassthroughStreamingResponse(AsyncGenerator[Any, Any]):
try:
if self._initialized:
await self._response.aclose()
except Exception:
except Exception: # noqa: BLE001
pass
@ -196,7 +196,7 @@ class PassthroughStreamingResponse(Generator[Any, Any, Any]):
raw_bytes=self._raw_bytes,
provider_config=self._provider_config,
)
except Exception as e:
except Exception as e: # noqa: BLE001
verbose_logger.exception(
"Failed to schedule passthrough spend-tracking flush; "
"%d buffered chunks dropped: %s",
@ -212,11 +212,11 @@ class PassthroughStreamingResponse(Generator[Any, Any, Any]):
chunk = next(self._iterator)
self._raw_bytes.append(chunk)
return chunk
except Exception:
except Exception: # noqa: BLE001
self._start_flush()
try:
self._response.close()
except Exception:
except Exception: # noqa: BLE001
pass
raise
@ -230,7 +230,7 @@ class PassthroughStreamingResponse(Generator[Any, Any, Any]):
self._start_flush()
try:
self._response.close()
except Exception:
except Exception: # noqa: BLE001
pass
@ -240,17 +240,17 @@ async def allm_passthrough_route(
method: str,
endpoint: str,
model: str,
custom_llm_provider: Optional[str] = None,
api_base: Optional[str] = None,
api_key: Optional[str] = None,
request_query_params: Optional[dict] = None,
request_headers: Optional[dict] = None,
content: Optional[Any] = None,
data: Optional[dict] = None,
files: Optional[RequestFiles] = None,
json: Optional[Any] = None,
params: Optional[QueryParamTypes] = None,
cookies: Optional[CookieTypes] = None,
custom_llm_provider: str | None = None,
api_base: str | None = None,
api_key: str | None = None,
request_query_params: dict | None = None,
request_headers: dict | None = None,
content: Any | None = None,
data: dict | None = None,
files: RequestFiles | None = None,
json: Any | None = None,
params: QueryParamTypes | None = None,
cookies: CookieTypes | None = None,
client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None,
**kwargs,
) -> Union[httpx.Response, AsyncGenerator[Any, Any]]:
@ -365,19 +365,19 @@ def llm_passthrough_route(
method: str,
endpoint: str,
model: str,
custom_llm_provider: Optional[str] = None,
api_base: Optional[str] = None,
api_key: Optional[str] = None,
request_query_params: Optional[dict] = None,
request_headers: Optional[dict] = None,
custom_llm_provider: str | None = None,
api_base: str | None = None,
api_key: str | None = None,
request_query_params: dict | None = None,
request_headers: dict | None = None,
allm_passthrough_route: bool = False,
content: Optional[Any] = None,
data: Optional[dict] = None,
files: Optional[RequestFiles] = None,
json: Optional[Any] = None,
params: Optional[QueryParamTypes] = None,
cookies: Optional[CookieTypes] = None,
client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None,
content: Any | None = None,
data: dict | None = None,
files: RequestFiles | None = None,
json: Any | None = None,
params: QueryParamTypes | None = None,
cookies: CookieTypes | None = None,
client: Union[HTTPHandler, AsyncHTTPHandler] | None = None,
**kwargs,
) -> Union[
httpx.Response,

View file

@ -732,7 +732,7 @@ class ProxyBaseLLMRequestProcessing:
@staticmethod
def _merge_passthrough_streaming_headers(
response_headers: Optional[Any],
response_headers: Any | None,
custom_headers: dict,
) -> dict:
"""

View file

@ -9,7 +9,7 @@ Use litellm with Anthropic SDK, Vertex AI SDK, Cohere SDK, etc.
import json
import os
import re
from typing import Any, Optional, Tuple, Union, cast
from typing import TYPE_CHECKING, Any, Callable, Union, cast
import httpx
from fastapi import APIRouter, Depends, HTTPException, Request, Response, WebSocket
@ -59,6 +59,11 @@ from litellm.utils import ProviderConfigManager
from .passthrough_endpoint_router import PassthroughEndpointRouter
if TYPE_CHECKING:
from litellm.proxy.proxy_server import ProxyConfig
from litellm.proxy.utils import ProxyLogging
vertex_llm_base = VertexBase()
router = APIRouter()
default_vertex_config = None
@ -77,7 +82,7 @@ def create_request_copy(request: Request):
def is_passthrough_request_using_router_model(
request_body: dict, llm_router: Optional[litellm.Router]
request_body: dict, llm_router: litellm.Router | None
) -> bool:
"""
Returns True if the model is in the llm_router model names
@ -225,7 +230,7 @@ async def gemini_proxy_route(
)
# Add or update query parameters
gemini_api_key: Optional[str] = passthrough_endpoint_router.get_credentials(
gemini_api_key: str | None = passthrough_endpoint_router.get_credentials(
custom_llm_provider="gemini",
region_name=None,
)
@ -473,9 +478,9 @@ async def milvus_proxy_route(
request_body = await get_request_body(request)
# check collectionName
collection_name = cast(Optional[str], request_body.get("collectionName"))
collection_name = cast(str | None, request_body.get("collectionName"))
extra_headers = {}
base_target_url: Optional[str] = None
base_target_url: str | None = None
if not collection_name:
raise HTTPException(
status_code=400,
@ -760,12 +765,12 @@ async def handle_bedrock_passthrough_router_model(
general_settings: dict,
proxy_config,
select_data_generator,
user_model: Optional[str],
user_temperature: Optional[float],
user_request_timeout: Optional[float],
user_max_tokens: Optional[int],
user_api_base: Optional[str],
version: Optional[str],
user_model: str | None,
user_temperature: float | None,
user_request_timeout: float | None,
user_max_tokens: int | None,
user_api_base: str | None,
version: str | None,
) -> Union[Response, StreamingResponse]:
"""
Handle Bedrock passthrough for router models (models defined in config.yaml).
@ -1134,12 +1139,12 @@ async def bedrock_proxy_route(
def _resolve_vertex_model_from_router(
model_id: str,
llm_router: Optional[litellm.Router],
llm_router: litellm.Router | None,
encoded_endpoint: str,
endpoint: str,
vertex_project: Optional[str],
vertex_location: Optional[str],
) -> Tuple[str, str, Optional[str], Optional[str]]:
vertex_project: str | None,
vertex_location: str | None,
) -> tuple[str, str, str | None, str | None]:
"""
Resolve Vertex AI model configuration from router.
@ -1152,7 +1157,7 @@ def _resolve_vertex_model_from_router(
vertex_location: Current vertex location (may be from URL)
Returns:
Tuple of (encoded_endpoint, endpoint, vertex_project, vertex_location)
tuple of (encoded_endpoint, endpoint, vertex_project, vertex_location)
with resolved values from router config
"""
if not llm_router:
@ -1501,42 +1506,42 @@ from abc import ABC, abstractmethod
class BaseVertexAIPassThroughHandler(ABC):
@staticmethod
@abstractmethod
def get_default_base_target_url(vertex_location: Optional[str]) -> str:
def get_default_base_target_url(vertex_location: str | None) -> str:
pass
@staticmethod
@abstractmethod
def update_base_target_url_with_credential_location(
base_target_url: str, vertex_location: Optional[str]
base_target_url: str, vertex_location: str | None
) -> str:
pass
class VertexAIDiscoveryPassThroughHandler(BaseVertexAIPassThroughHandler):
@staticmethod
def get_default_base_target_url(vertex_location: Optional[str]) -> str:
def get_default_base_target_url(vertex_location: str | None) -> str:
return "https://discoveryengine.googleapis.com/"
@staticmethod
def update_base_target_url_with_credential_location(
base_target_url: str, vertex_location: Optional[str]
base_target_url: str, vertex_location: str | None
) -> str:
return base_target_url
class VertexAIPassThroughHandler(BaseVertexAIPassThroughHandler):
@staticmethod
def get_default_base_target_url(vertex_location: Optional[str]) -> str:
def get_default_base_target_url(vertex_location: str | None) -> str:
return get_vertex_base_url(vertex_location)
@staticmethod
def update_base_target_url_with_credential_location(
base_target_url: str, vertex_location: Optional[str]
base_target_url: str, vertex_location: str | None
) -> str:
return get_vertex_base_url(vertex_location)
def get_vertex_base_url(vertex_location: Optional[str]) -> str:
def get_vertex_base_url(vertex_location: str | None) -> str:
"""
Base URL for Vertex AI pass-through (trailing slash for URL joining).
@ -1586,10 +1591,10 @@ def get_vertex_pass_through_handler(
def _override_vertex_params_from_router_credentials(
router_credentials: Optional[Any],
vertex_project: Optional[str],
vertex_location: Optional[str],
) -> Tuple[Optional[str], Optional[str]]:
router_credentials: Any | None,
vertex_project: str | None,
vertex_location: str | None,
) -> tuple[str | None, str | None]:
"""
Override vertex_project and vertex_location with values from router_credentials if available.
@ -1599,7 +1604,7 @@ def _override_vertex_params_from_router_credentials(
vertex_location: Current vertex location (from URL)
Returns:
Tuple of (vertex_project, vertex_location) with overridden values if applicable
tuple of (vertex_project, vertex_location) with overridden values if applicable
"""
if router_credentials is None:
return vertex_project, vertex_location
@ -1648,13 +1653,13 @@ def _override_vertex_params_from_router_credentials(
async def _prepare_vertex_auth_headers(
request: Request,
vertex_credentials: Optional[Any],
router_credentials: Optional[Any],
vertex_project: Optional[str],
vertex_location: Optional[str],
base_target_url: Optional[str],
vertex_credentials: Any | None,
router_credentials: Any | None,
vertex_project: str | None,
vertex_location: str | None,
base_target_url: str | None,
get_vertex_pass_through_handler: BaseVertexAIPassThroughHandler,
) -> Tuple[dict, Optional[str], bool, Optional[str], Optional[str]]:
) -> tuple[dict, str | None, bool, str | None, str | None]:
"""
Prepare authentication headers for Vertex AI pass-through requests.
@ -1668,12 +1673,12 @@ async def _prepare_vertex_auth_headers(
get_vertex_pass_through_handler: Handler for the specific Vertex AI service
Returns:
Tuple containing:
tuple containing:
- headers: dict - Authentication headers to use
- base_target_url: Optional[str] - Updated base target URL
- base_target_url: str | None - Updated base target URL
- headers_passed_through: bool - Whether headers were passed through from request
- vertex_project: Optional[str] - Updated vertex project ID
- vertex_location: Optional[str] - Updated vertex location
- vertex_project: str | None - Updated vertex project ID
- vertex_location: str | None - Updated vertex location
"""
vertex_llm_base = VertexBase()
headers_passed_through = False
@ -1746,8 +1751,8 @@ async def _base_vertex_proxy_route(
request: Request,
fastapi_response: Response,
get_vertex_pass_through_handler: BaseVertexAIPassThroughHandler,
user_api_key_dict: Optional[UserAPIKeyAuth] = None,
router_credentials: Optional[Any] = None,
user_api_key_dict: UserAPIKeyAuth | None = None,
router_credentials: Any | None = None,
):
"""
Base function for Vertex AI passthrough routes.
@ -1793,8 +1798,8 @@ async def _base_vertex_proxy_route(
user_api_key_dict=user_api_key_dict,
)
vertex_project: Optional[str] = get_vertex_project_id_from_url(endpoint)
vertex_location: Optional[str] = get_vertex_location_from_url(endpoint)
vertex_project: str | None = get_vertex_project_id_from_url(endpoint)
vertex_location: str | None = get_vertex_location_from_url(endpoint)
# Override with vector store credentials if available
vertex_project, vertex_location = _override_vertex_params_from_router_credentials(
@ -1919,7 +1924,7 @@ async def vertex_discovery_proxy_route(
from litellm.types.vector_stores import LiteLLM_ManagedVectorStore
# Extract vector store ID from endpoint if present (e.g., dataStores/test-litellm-app_1761094730750)
vector_store_credentials: Optional[LiteLLM_ManagedVectorStore] = None
vector_store_credentials: LiteLLM_ManagedVectorStore | None = None
vector_store_id_match = re.search(r"dataStores/([^/]+)", endpoint)
if vector_store_id_match:
@ -2057,9 +2062,9 @@ class BaseOpenAIPassThroughHandler:
fastapi_response: Response,
user_api_key_dict: UserAPIKeyAuth,
base_target_url: str,
api_key: Optional[str],
api_key: str | None,
custom_llm_provider: litellm.LlmProviders,
extra_headers: Optional[dict] = None,
extra_headers: dict | None = None,
):
encoded_endpoint = httpx.URL(endpoint).path
# Ensure endpoint starts with '/' for proper URL construction
@ -2115,7 +2120,7 @@ class BaseOpenAIPassThroughHandler:
@staticmethod
def _assemble_headers(
api_key: Optional[str], request: Request, extra_headers: Optional[dict] = None
api_key: str | None, request: Request, extra_headers: dict | None = None
) -> dict:
base_headers = {}
if api_key is not None:
@ -2251,10 +2256,10 @@ async def cursor_proxy_route(
async def vertex_ai_live_websocket_passthrough(
websocket: WebSocket,
model: Optional[str] = None,
vertex_project: Optional[str] = None,
vertex_location: Optional[str] = None,
user_api_key_dict: Optional[UserAPIKeyAuth] = None,
model: str | None = None,
vertex_project: str | None = None,
vertex_location: str | None = None,
user_api_key_dict: UserAPIKeyAuth | None = None,
):
"""
Vertex AI Live API WebSocket Pass-through Function
@ -2286,8 +2291,8 @@ async def vertex_ai_live_websocket_passthrough(
)
resolved_project = vertex_project
resolved_location: Optional[str] = vertex_location
credentials_value: Optional[str] = None
resolved_location: str | None = vertex_location
credentials_value: str | None = None
if vertex_credentials_config is not None:
resolved_project = resolved_project or vertex_credentials_config.vertex_project
@ -2391,9 +2396,9 @@ def create_vertex_ai_live_websocket_endpoint():
def create_generic_websocket_passthrough_endpoint(
provider: str,
target_url: str,
custom_headers: Optional[dict] = None,
custom_headers: dict | None = None,
forward_headers: bool = False,
cost_per_request: Optional[float] = None,
cost_per_request: float | None = None,
):
"""
Create a generic WebSocket passthrough endpoint for any provider.
@ -2540,7 +2545,7 @@ async def gigachat_proxy_route(
)
return result
except Exception as e:
except Exception as e: # noqa: BLE001
raise await base_llm_response_processor._handle_llm_api_exception(
e=e,
user_api_key_dict=user_api_key_dict,
@ -2556,16 +2561,16 @@ async def handle_gigachat_passthrough_router_model(
fastapi_response: Response,
llm_router: litellm.Router,
user_api_key_dict: UserAPIKeyAuth,
proxy_logging_obj,
proxy_logging_obj: ProxyLogging,
general_settings: dict,
proxy_config,
select_data_generator,
user_model: Optional[str],
user_temperature: Optional[float],
user_request_timeout: Optional[float],
user_max_tokens: Optional[int],
user_api_base: Optional[str],
version: Optional[str],
proxy_config: ProxyConfig,
select_data_generator: Callable,
user_model: str | None,
user_temperature: float | None,
user_request_timeout: float | None,
user_max_tokens: int | None,
user_api_base: str | None,
version: str | None,
) -> Union[Response, StreamingResponse]:
"""
Handle Gigachat passthrough for router models (models defined in config.yaml).
@ -2580,6 +2585,10 @@ async def handle_gigachat_passthrough_router_model(
request_body: The parsed request body
llm_router: The LiteLLM router instance
user_api_key_dict: The user API key authentication dictionary
proxy_logging_obj: Proxy logging
general_settings: Proxy general settings
proxy_config: Proxy config
select_data_generator: Select data generator function
(additional args for common processing)
Returns:
@ -2677,7 +2686,7 @@ async def handle_gigachat_passthrough_router_model(
return result
return result
except Exception as e:
except Exception as e: # noqa: BLE001
# Use common exception handling
raise await base_llm_response_processor._handle_llm_api_exception(
e=e,