mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix some lint issues
This commit is contained in:
parent
d185b8fa49
commit
93d272f40c
11 changed files with 199 additions and 194 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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).
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -732,7 +732,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
|
||||
@staticmethod
|
||||
def _merge_passthrough_streaming_headers(
|
||||
response_headers: Optional[Any],
|
||||
response_headers: Any | None,
|
||||
custom_headers: dict,
|
||||
) -> dict:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue