fix: resolve UP045 lint violations (Optional[X] -> X | None)

Convert Optional[X] type annotations to X | None syntax across rerank
transformations, spend tracking, and other modules to satisfy ruff strict gate.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
Sameer Kankute 2026-06-24 17:21:33 +05:30
parent e0c8a6b483
commit 7d89450c8e
No known key found for this signature in database
19 changed files with 313 additions and 338 deletions

View file

@ -73,17 +73,17 @@ if TYPE_CHECKING:
async def make_call(
client: Optional[AsyncHTTPHandler],
client: AsyncHTTPHandler | None,
api_base: str,
headers: dict,
data: str,
model: str,
messages: list,
logging_obj,
timeout: Optional[Union[float, httpx.Timeout]],
timeout: Union[float, httpx.Timeout] | None,
json_mode: bool,
speed: Optional[str] = None,
tool_name_reverse_map: Optional[Dict[str, str]] = None,
speed: str | None = None,
tool_name_reverse_map: Dict[str, str] | None = None,
) -> Tuple[Any, httpx.Headers]:
if client is None:
client = litellm.module_level_aclient
@ -133,17 +133,17 @@ async def make_call(
def make_sync_call(
client: Optional[HTTPHandler],
client: HTTPHandler | None,
api_base: str,
headers: dict,
data: str,
model: str,
messages: list,
logging_obj,
timeout: Optional[Union[float, httpx.Timeout]],
timeout: Union[float, httpx.Timeout] | None,
json_mode: bool,
speed: Optional[str] = None,
tool_name_reverse_map: Optional[Dict[str, str]] = None,
speed: str | None = None,
tool_name_reverse_map: Dict[str, str] | None = None,
) -> Tuple[Any, httpx.Headers]:
if client is None:
client = litellm.module_level_client # re-use a module level client
@ -213,7 +213,7 @@ class AnthropicChatCompletion(BaseLLM):
model_response: ModelResponse,
print_verbose: Callable,
timeout: Union[float, httpx.Timeout],
client: Optional[AsyncHTTPHandler],
client: AsyncHTTPHandler | None,
encoding,
api_key,
logging_obj,
@ -277,7 +277,7 @@ class AnthropicChatCompletion(BaseLLM):
provider_config: "BaseConfig",
logger_fn=None,
headers={},
client: Optional[AsyncHTTPHandler] = None,
client: AsyncHTTPHandler | None = None,
) -> Union[ModelResponse, "CustomStreamWrapper"]:
async_handler = client or get_async_httpx_client(
llm_provider=litellm.LlmProviders.ANTHROPIC
@ -538,9 +538,9 @@ class ModelResponseIterator:
self,
streaming_response,
sync_stream: bool,
json_mode: Optional[bool] = False,
speed: Optional[str] = None,
tool_name_reverse_map: Optional[Dict[str, str]] = None,
json_mode: bool | None = False,
speed: str | None = None,
tool_name_reverse_map: Dict[str, str] | None = None,
):
self.streaming_response = streaming_response
self.response_iterator = self.streaming_response
@ -570,7 +570,7 @@ class ModelResponseIterator:
# Track current content block type to avoid emitting tool calls for non-tool blocks
# See: https://github.com/BerriAI/litellm/issues/17254
self.current_content_block_type: Optional[str] = None
self.current_content_block_type: str | None = None
# Accumulate web_search_tool_result blocks for multi-turn reconstruction
# See: https://github.com/BerriAI/litellm/issues/17737
@ -586,8 +586,8 @@ class ModelResponseIterator:
# Track server tool use inputs and results for code_interpreter_results
self._server_tool_inputs: Dict[str, Any] = {}
self.tool_results: List[Dict[str, Any]] = []
self._current_server_tool_id: Optional[str] = None
self._container_id: Optional[str] = None
self._current_server_tool_id: str | None = None
self._container_id: str | None = None
def check_empty_tool_call_args(self) -> bool:
"""
@ -626,7 +626,7 @@ class ModelResponseIterator:
def _content_block_delta_helper(self, chunk: dict) -> Tuple[
str,
Optional[ChatCompletionToolCallChunk],
ChatCompletionToolCallChunk | None,
List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]],
Dict[str, Any],
]:
@ -634,7 +634,7 @@ class ModelResponseIterator:
Helper function to handle the content block delta
"""
text = ""
tool_use: Optional[ChatCompletionToolCallChunk] = None
tool_use: ChatCompletionToolCallChunk | None = None
provider_specific_fields = {}
content_block = ContentBlockDelta(**chunk) # type: ignore
thinking_blocks: List[
@ -695,13 +695,13 @@ class ModelResponseIterator:
thinking_blocks: List[
Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]
],
) -> Optional[str]:
) -> str | None:
"""
Handle the reasoning content
"""
reasoning_content = None
for block in thinking_blocks:
thinking_content = cast(Optional[str], block.get("thinking"))
thinking_content = cast(str | None, block.get("thinking"))
if reasoning_content is None:
reasoning_content = ""
if thinking_content is not None:
@ -777,18 +777,12 @@ class ModelResponseIterator:
type_chunk = chunk.get("type", "") or ""
text = ""
tool_use: Optional[ChatCompletionToolCallChunk] = None
tool_use: ChatCompletionToolCallChunk | None = None
finish_reason = ""
usage: Optional[Usage] = None
usage: Usage | None = None
provider_specific_fields: Dict[str, Any] = {}
reasoning_content: Optional[str] = None
thinking_blocks: Optional[
List[
Union[
ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock
]
]
] = None
reasoning_content: str | None = None
thinking_blocks: List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] | None = None
# Always use index=0 for OpenAI choice format (fixes multi-choice errors)
index = 0
@ -1058,8 +1052,8 @@ class ModelResponseIterator:
raise ValueError(f"Failed to decode JSON from chunk: {chunk}")
def _handle_json_mode_chunk(
self, text: str, tool_use: Optional[ChatCompletionToolCallChunk]
) -> Tuple[str, Optional[ChatCompletionToolCallChunk]]:
self, text: str, tool_use: ChatCompletionToolCallChunk | None
) -> Tuple[str, ChatCompletionToolCallChunk | None]:
"""
If JSON mode is enabled, convert the tool call to a message.
@ -1107,7 +1101,7 @@ class ModelResponseIterator:
def _handle_message_delta(
self, chunk: dict
) -> Tuple[str, Optional[Usage], Optional[Dict[str, Any]]]:
) -> Tuple[str, Usage | None, Dict[str, Any] | None]:
"""
Handle message_delta event for finish_reason, usage, and container.
@ -1131,7 +1125,7 @@ class ModelResponseIterator:
def _handle_accumulated_json_chunk(
self, data_str: str
) -> Optional[ModelResponseStream]:
) -> ModelResponseStream | None:
"""
Handle partial JSON chunks by accumulating them until valid JSON is received.
@ -1156,7 +1150,7 @@ class ModelResponseIterator:
# If it's not valid JSON yet, continue to the next chunk
return None
def _parse_sse_data(self, str_line: str) -> Optional[ModelResponseStream]:
def _parse_sse_data(self, str_line: str) -> ModelResponseStream | None:
"""
Parse SSE data line, handling both complete and partial JSON chunks.

View file

@ -22,8 +22,8 @@ class BaseRerankConfig(ABC):
self,
headers: dict,
model: str,
api_key: Optional[str] = None,
optional_params: Optional[dict] = None,
api_key: str | None = None,
optional_params: dict | None = None,
) -> dict:
pass
@ -33,7 +33,7 @@ class BaseRerankConfig(ABC):
model: str,
optional_rerank_params: Dict,
headers: dict,
litellm_params: Optional[dict] = None,
litellm_params: dict | None = None,
) -> dict:
return {}
@ -44,7 +44,7 @@ class BaseRerankConfig(ABC):
raw_response: httpx.Response,
model_response: RerankResponse,
logging_obj: LiteLLMLoggingObj,
api_key: Optional[str] = None,
api_key: str | None = None,
request_data: dict = {},
optional_params: dict = {},
litellm_params: dict = {},
@ -54,9 +54,9 @@ class BaseRerankConfig(ABC):
@abstractmethod
def get_complete_url(
self,
api_base: Optional[str],
api_base: str | None,
model: str,
optional_params: Optional[dict] = None,
optional_params: dict | None = None,
) -> str:
"""
OPTIONAL
@ -79,12 +79,12 @@ class BaseRerankConfig(ABC):
drop_params: bool,
query: str,
documents: List[Union[str, Dict[str, Any]]],
custom_llm_provider: Optional[str] = None,
top_n: Optional[int] = None,
rank_fields: Optional[List[str]] = None,
return_documents: Optional[bool] = True,
max_chunks_per_doc: Optional[int] = None,
max_tokens_per_doc: Optional[int] = None,
custom_llm_provider: str | None = None,
top_n: int | None = None,
rank_fields: List[str] | None = None,
return_documents: bool | None = True,
max_chunks_per_doc: int | None = None,
max_tokens_per_doc: int | None = None,
) -> Dict:
pass
@ -100,9 +100,9 @@ class BaseRerankConfig(ABC):
def calculate_rerank_cost(
self,
model: str,
custom_llm_provider: Optional[str] = None,
billed_units: Optional[RerankBilledUnits] = None,
model_info: Optional[ModelInfo] = None,
custom_llm_provider: str | None = None,
billed_units: RerankBilledUnits | None = None,
model_info: ModelInfo | None = None,
) -> Tuple[float, float]:
"""
Calculates the cost per query for a given rerank model.

View file

@ -22,9 +22,9 @@ class CohereRerankConfig(BaseRerankConfig):
def get_complete_url(
self,
api_base: Optional[str],
api_base: str | None,
model: str,
optional_params: Optional[dict] = None,
optional_params: dict | None = None,
) -> str:
if api_base:
# Remove trailing slashes and ensure clean base URL
@ -46,17 +46,17 @@ class CohereRerankConfig(BaseRerankConfig):
def map_cohere_rerank_params(
self,
non_default_params: Optional[dict],
non_default_params: dict | None,
model: str,
drop_params: bool,
query: str,
documents: List[Union[str, Dict[str, Any]]],
custom_llm_provider: Optional[str] = None,
top_n: Optional[int] = None,
rank_fields: Optional[List[str]] = None,
return_documents: Optional[bool] = True,
max_chunks_per_doc: Optional[int] = None,
max_tokens_per_doc: Optional[int] = None,
custom_llm_provider: str | None = None,
top_n: int | None = None,
rank_fields: List[str] | None = None,
return_documents: bool | None = True,
max_chunks_per_doc: int | None = None,
max_tokens_per_doc: int | None = None,
) -> Dict:
"""
Map Cohere rerank params
@ -78,8 +78,8 @@ class CohereRerankConfig(BaseRerankConfig):
self,
headers: dict,
model: str,
api_key: Optional[str] = None,
optional_params: Optional[dict] = None,
api_key: str | None = None,
optional_params: dict | None = None,
) -> dict:
if api_key is None:
api_key = (
@ -111,7 +111,7 @@ class CohereRerankConfig(BaseRerankConfig):
model: str,
optional_rerank_params: Dict,
headers: dict,
litellm_params: Optional[dict] = None,
litellm_params: dict | None = None,
) -> dict:
if "query" not in optional_rerank_params:
raise ValueError("query is required for Cohere rerank")
@ -134,7 +134,7 @@ class CohereRerankConfig(BaseRerankConfig):
raw_response: httpx.Response,
model_response: RerankResponse,
logging_obj: LiteLLMLoggingObj,
api_key: Optional[str] = None,
api_key: str | None = None,
request_data: dict = {},
optional_params: dict = {},
litellm_params: dict = {},

View file

@ -14,9 +14,9 @@ class CohereRerankV2Config(CohereRerankConfig):
def get_complete_url(
self,
api_base: Optional[str],
api_base: str | None,
model: str,
optional_params: Optional[dict] = None,
optional_params: dict | None = None,
) -> str:
if api_base:
# Remove trailing slashes and ensure clean base URL
@ -38,17 +38,17 @@ class CohereRerankV2Config(CohereRerankConfig):
def map_cohere_rerank_params(
self,
non_default_params: Optional[dict],
non_default_params: dict | None,
model: str,
drop_params: bool,
query: str,
documents: List[Union[str, Dict[str, Any]]],
custom_llm_provider: Optional[str] = None,
top_n: Optional[int] = None,
rank_fields: Optional[List[str]] = None,
return_documents: Optional[bool] = True,
max_chunks_per_doc: Optional[int] = None,
max_tokens_per_doc: Optional[int] = None,
custom_llm_provider: str | None = None,
top_n: int | None = None,
rank_fields: List[str] | None = None,
return_documents: bool | None = True,
max_chunks_per_doc: int | None = None,
max_tokens_per_doc: int | None = None,
) -> Dict:
"""
Map Cohere rerank params
@ -71,7 +71,7 @@ class CohereRerankV2Config(CohereRerankConfig):
model: str,
optional_rerank_params: Dict,
headers: dict,
litellm_params: Optional[dict] = None,
litellm_params: dict | None = None,
) -> dict:
if "query" not in optional_rerank_params:
raise ValueError("query is required for Cohere rerank")

View file

@ -59,9 +59,9 @@ class DashScopeRerankConfig(BaseRerankConfig):
def get_complete_url(
self,
api_base: Optional[str],
api_base: str | None,
model: str,
optional_params: Optional[dict] = None,
optional_params: dict | None = None,
) -> str:
if api_base is None:
api_base = get_secret_str("DASHSCOPE_API_BASE_RERANK") or DEFAULT_RERANK_URL
@ -83,8 +83,8 @@ class DashScopeRerankConfig(BaseRerankConfig):
self,
headers: dict,
model: str,
api_key: Optional[str] = None,
optional_params: Optional[dict] = None,
api_key: str | None = None,
optional_params: dict | None = None,
) -> dict:
if api_key is None:
api_key = get_secret_str("DASHSCOPE_API_KEY")
@ -105,17 +105,17 @@ class DashScopeRerankConfig(BaseRerankConfig):
def map_cohere_rerank_params(
self,
non_default_params: Optional[dict],
non_default_params: dict | None,
model: str,
drop_params: bool,
query: str,
documents: List[Union[str, Dict[str, Any]]],
custom_llm_provider: Optional[str] = None,
top_n: Optional[int] = None,
rank_fields: Optional[List[str]] = None,
return_documents: Optional[bool] = True,
max_chunks_per_doc: Optional[int] = None,
max_tokens_per_doc: Optional[int] = None,
custom_llm_provider: str | None = None,
top_n: int | None = None,
rank_fields: List[str] | None = None,
return_documents: bool | None = True,
max_chunks_per_doc: int | None = None,
max_tokens_per_doc: int | None = None,
) -> Dict:
# qwen3-rerank accepts query/documents/top_n/return_documents. The
# rest (rank_fields, max_*_per_doc) are silently dropped.
@ -134,7 +134,7 @@ class DashScopeRerankConfig(BaseRerankConfig):
model: str,
optional_rerank_params: Dict,
headers: dict,
litellm_params: Optional[dict] = None,
litellm_params: dict | None = None,
) -> dict:
if "query" not in optional_rerank_params:
raise ValueError("query is required for DashScope rerank")
@ -158,10 +158,10 @@ class DashScopeRerankConfig(BaseRerankConfig):
raw_response: httpx.Response,
model_response: RerankResponse,
logging_obj: LiteLLMLoggingObj,
api_key: Optional[str] = None,
request_data: Optional[dict] = None,
optional_params: Optional[dict] = None,
litellm_params: Optional[dict] = None,
api_key: str | None = None,
request_data: dict | None = None,
optional_params: dict | None = None,
litellm_params: dict | None = None,
) -> RerankResponse:
request_data = request_data or {}
optional_params = optional_params or {}

View file

@ -30,9 +30,9 @@ class DeepinfraRerankConfig(BaseRerankConfig):
def get_complete_url(
self,
api_base: Optional[str],
api_base: str | None,
model: str,
optional_params: Optional[dict] = None,
optional_params: dict | None = None,
) -> str:
"""
Constructs the complete DeepInfra inference endpoint URL for rerank.
@ -67,8 +67,8 @@ class DeepinfraRerankConfig(BaseRerankConfig):
self,
headers: dict,
model: str,
api_key: Optional[str] = None,
optional_params: Optional[dict] = None,
api_key: str | None = None,
optional_params: dict | None = None,
) -> dict:
if api_key is None:
api_key = get_secret_str("DEEPINFRA_API_KEY")
@ -98,12 +98,12 @@ class DeepinfraRerankConfig(BaseRerankConfig):
drop_params: bool,
query: str,
documents: List[Union[str, Dict[str, Any]]],
custom_llm_provider: Optional[str] = None,
top_n: Optional[int] = None,
rank_fields: Optional[List[str]] = None,
return_documents: Optional[bool] = True,
max_chunks_per_doc: Optional[int] = None,
max_tokens_per_doc: Optional[int] = None,
custom_llm_provider: str | None = None,
top_n: int | None = None,
rank_fields: List[str] | None = None,
return_documents: bool | None = True,
max_chunks_per_doc: int | None = None,
max_tokens_per_doc: int | None = None,
) -> Dict:
# Start with the basic parameters
optional_rerank_params = {}
@ -132,7 +132,7 @@ class DeepinfraRerankConfig(BaseRerankConfig):
model: str,
optional_rerank_params: Dict,
headers: dict,
litellm_params: Optional[dict] = None,
litellm_params: dict | None = None,
) -> dict:
# Convert OptionalRerankParams to dict as expected by parent class
if optional_rerank_params is None:
@ -145,7 +145,7 @@ class DeepinfraRerankConfig(BaseRerankConfig):
raw_response: httpx.Response,
model_response: RerankResponse,
logging_obj: LiteLLMLoggingObj,
api_key: Optional[str] = None,
api_key: str | None = None,
request_data: dict = {},
optional_params: dict = {},
litellm_params: dict = {},

View file

@ -29,9 +29,9 @@ class FireworksAIRerankConfig(FireworksAIMixin, BaseRerankConfig):
def get_complete_url(
self,
api_base: Optional[str],
api_base: str | None,
model: str,
optional_params: Optional[dict] = None,
optional_params: dict | None = None,
) -> str:
if api_base:
# Remove trailing slashes and ensure clean base URL
@ -56,17 +56,17 @@ class FireworksAIRerankConfig(FireworksAIMixin, BaseRerankConfig):
def map_cohere_rerank_params(
self,
non_default_params: Optional[dict],
non_default_params: dict | None,
model: str,
drop_params: bool,
query: str,
documents: List[Union[str, Dict[str, Any]]],
custom_llm_provider: Optional[str] = None,
top_n: Optional[int] = None,
rank_fields: Optional[List[str]] = None,
return_documents: Optional[bool] = True,
max_chunks_per_doc: Optional[int] = None,
max_tokens_per_doc: Optional[int] = None,
custom_llm_provider: str | None = None,
top_n: int | None = None,
rank_fields: List[str] | None = None,
return_documents: bool | None = True,
max_chunks_per_doc: int | None = None,
max_tokens_per_doc: int | None = None,
) -> Dict[str, Any]:
"""
Map Cohere rerank params to Fireworks AI rerank params
@ -101,8 +101,8 @@ class FireworksAIRerankConfig(FireworksAIMixin, BaseRerankConfig):
self,
headers: dict,
model: str,
api_key: Optional[str] = None,
optional_params: Optional[dict] = None,
api_key: str | None = None,
optional_params: dict | None = None,
) -> dict:
api_key = self._get_api_key(api_key)
if api_key is None:
@ -127,7 +127,7 @@ class FireworksAIRerankConfig(FireworksAIMixin, BaseRerankConfig):
model: str,
optional_rerank_params: Dict,
headers: dict,
litellm_params: Optional[dict] = None,
litellm_params: dict | None = None,
) -> dict:
"""
Transform request to Fireworks AI rerank format
@ -175,7 +175,7 @@ class FireworksAIRerankConfig(FireworksAIMixin, BaseRerankConfig):
raw_response: httpx.Response,
model_response: RerankResponse,
logging_obj: LiteLLMLoggingObj,
api_key: Optional[str] = None,
api_key: str | None = None,
request_data: dict = {},
optional_params: dict = {},
litellm_params: dict = {},
@ -220,7 +220,7 @@ class FireworksAIRerankConfig(FireworksAIMixin, BaseRerankConfig):
rerank_meta = RerankResponseMeta(billed_units=_billed_units, tokens=_tokens)
# Extract results - Fireworks AI uses "data" instead of "results"
_results: Optional[List[dict]] = raw_response_json.get(
_results: List[dict] | None = raw_response_json.get(
"data"
) or raw_response_json.get("results")

View file

@ -22,8 +22,8 @@ from ..common_utils import (
class GithubCopilotConfig(OpenAIConfig):
def __init__(
self,
api_key: Optional[str] = None,
api_base: Optional[str] = None,
api_key: str | None = None,
api_base: str | None = None,
custom_llm_provider: str = "openai",
) -> None:
super().__init__()
@ -32,10 +32,10 @@ class GithubCopilotConfig(OpenAIConfig):
def _get_openai_compatible_provider_info(
self,
model: str,
api_base: Optional[str],
api_key: Optional[str],
api_base: str | None,
api_key: str | None,
custom_llm_provider: str,
) -> Tuple[Optional[str], Optional[str], str]:
) -> Tuple[str | None, str | None, str]:
dynamic_api_base = (
api_base
or self.authenticator.get_api_base()
@ -85,8 +85,8 @@ class GithubCopilotConfig(OpenAIConfig):
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:
# Get base headers from parent
validated_headers = super().validate_environment(
@ -173,7 +173,7 @@ class GithubCopilotConfig(OpenAIConfig):
@staticmethod
def _parse_anthropic_native_content(
content_blocks: List[Any],
) -> Tuple[str, List[ChatCompletionToolCallChunk], Optional[List[Any]]]:
) -> Tuple[str, List[ChatCompletionToolCallChunk], List[Any] | None]:
"""
Parse Anthropic-native content blocks into OpenAI-compatible fields.
@ -205,8 +205,8 @@ class GithubCopilotConfig(OpenAIConfig):
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":
"""
Handle newer Copilot models (e.g. claude-opus-4.7, claude-opus-4.8) that
@ -237,7 +237,7 @@ class GithubCopilotConfig(OpenAIConfig):
if not response_json.get("choices"):
content = ""
tool_calls: List[ChatCompletionToolCallChunk] = []
thinking_blocks: Optional[List[Any]] = None
thinking_blocks: List[Any] | None = None
if "content" in response_json and isinstance(
response_json["content"], list
):

View file

@ -28,7 +28,7 @@ class HostedVLLMRerankError(BaseLLMException):
self,
status_code: int,
message: str,
headers: Optional[Union[dict, httpx.Headers]] = None,
headers: Union[dict, httpx.Headers] | None = None,
):
super().__init__(status_code=status_code, message=message, headers=headers)
@ -39,9 +39,9 @@ class HostedVLLMRerankConfig(BaseRerankConfig):
def get_complete_url(
self,
api_base: Optional[str],
api_base: str | None,
model: str,
optional_params: Optional[dict] = None,
optional_params: dict | None = None,
) -> str:
if api_base:
# Remove trailing slashes and ensure clean base URL
@ -65,17 +65,17 @@ class HostedVLLMRerankConfig(BaseRerankConfig):
def map_cohere_rerank_params(
self,
non_default_params: Optional[dict],
non_default_params: dict | None,
model: str,
drop_params: bool,
query: str,
documents: List[Union[str, Dict[str, Any]]],
custom_llm_provider: Optional[str] = None,
top_n: Optional[int] = None,
rank_fields: Optional[List[str]] = None,
return_documents: Optional[bool] = True,
max_chunks_per_doc: Optional[int] = None,
max_tokens_per_doc: Optional[int] = None,
custom_llm_provider: str | None = None,
top_n: int | None = None,
rank_fields: List[str] | None = None,
return_documents: bool | None = True,
max_chunks_per_doc: int | None = None,
max_tokens_per_doc: int | None = None,
) -> Dict:
"""
Map parameters for Hosted VLLM rerank
@ -97,8 +97,8 @@ class HostedVLLMRerankConfig(BaseRerankConfig):
self,
headers: dict,
model: str,
api_key: Optional[str] = None,
optional_params: Optional[dict] = None,
api_key: str | None = None,
optional_params: dict | None = None,
) -> dict:
if api_key is None:
api_key = get_secret_str("HOSTED_VLLM_API_KEY") or "fake-api-key"
@ -121,7 +121,7 @@ class HostedVLLMRerankConfig(BaseRerankConfig):
model: str,
optional_rerank_params: Dict,
headers: dict,
litellm_params: Optional[dict] = None,
litellm_params: dict | None = None,
) -> dict:
if "query" not in optional_rerank_params:
raise ValueError("query is required for Hosted VLLM rerank")
@ -144,7 +144,7 @@ class HostedVLLMRerankConfig(BaseRerankConfig):
raw_response: httpx.Response,
model_response: RerankResponse,
logging_obj: LiteLLMLoggingObj,
api_key: Optional[str] = None,
api_key: str | None = None,
request_data: dict = {},
optional_params: dict = {},
litellm_params: dict = {},
@ -178,7 +178,7 @@ class HostedVLLMRerankConfig(BaseRerankConfig):
rerank_meta = RerankResponseMeta(billed_units=_billed_units, tokens=_tokens)
# Extract results
_results: Optional[List[dict]] = response.get("results")
_results: List[dict] | None = response.get("results")
if _results is None:
raise ValueError(f"No results found in the response={response}")

View file

@ -35,7 +35,7 @@ class HuggingFaceRerankResponseItem(TypedDict):
index: int
score: float
text: Optional[str] # Optional, included when return_text=True
text: str | None # Optional, included when return_text=True
class HuggingFaceRerankResponse(TypedDict):
@ -50,7 +50,7 @@ HuggingFaceRerankResponseList = List[HuggingFaceRerankResponseItem]
class HuggingFaceRerankConfig(BaseRerankConfig):
def get_api_base(self, model: str, api_base: Optional[str]) -> str:
def get_api_base(self, model: str, api_base: str | None) -> str:
if api_base is not None:
return api_base
elif os.getenv("HF_API_BASE") is not None:
@ -62,9 +62,9 @@ class HuggingFaceRerankConfig(BaseRerankConfig):
def get_complete_url(
self,
api_base: Optional[str],
api_base: str | None,
model: str,
optional_params: Optional[dict] = None,
optional_params: dict | None = None,
) -> str:
"""
Get the complete URL for the API call, including the /rerank suffix if necessary.
@ -89,17 +89,17 @@ class HuggingFaceRerankConfig(BaseRerankConfig):
def map_cohere_rerank_params(
self,
non_default_params: Optional[dict],
non_default_params: dict | None,
model: str,
drop_params: bool,
query: str,
documents: List[Union[str, Dict[str, Any]]],
custom_llm_provider: Optional[str] = None,
top_n: Optional[int] = None,
rank_fields: Optional[List[str]] = None,
return_documents: Optional[bool] = True,
max_chunks_per_doc: Optional[int] = None,
max_tokens_per_doc: Optional[int] = None,
custom_llm_provider: str | None = None,
top_n: int | None = None,
rank_fields: List[str] | None = None,
return_documents: bool | None = True,
max_chunks_per_doc: int | None = None,
max_tokens_per_doc: int | None = None,
) -> Dict:
optional_rerank_params = {}
if non_default_params is not None:
@ -121,9 +121,9 @@ class HuggingFaceRerankConfig(BaseRerankConfig):
self,
headers: dict,
model: str,
api_key: Optional[str] = None,
optional_params: Optional[dict] = None,
api_base: Optional[str] = None,
api_key: str | None = None,
optional_params: dict | None = None,
api_base: str | None = None,
) -> dict:
# Get API credentials
api_key, api_base = self.get_api_credentials(api_key=api_key, api_base=api_base)
@ -146,7 +146,7 @@ class HuggingFaceRerankConfig(BaseRerankConfig):
model: str,
optional_rerank_params: Union[OptionalRerankParams, dict],
headers: dict,
litellm_params: Optional[dict] = None,
litellm_params: dict | None = None,
) -> dict:
if "query" not in optional_rerank_params:
raise ValueError("query is required for HuggingFace rerank")
@ -172,7 +172,7 @@ class HuggingFaceRerankConfig(BaseRerankConfig):
raw_response: httpx.Response,
model_response: RerankResponse,
logging_obj: LoggingClass,
api_key: Optional[str] = None,
api_key: str | None = None,
request_data: dict = {},
optional_params: dict = {},
litellm_params: dict = {},
@ -275,9 +275,9 @@ class HuggingFaceRerankConfig(BaseRerankConfig):
def get_api_credentials(
self,
api_key: Optional[str] = None,
api_base: Optional[str] = None,
) -> Tuple[Optional[str], Optional[str]]:
api_key: str | None = None,
api_base: str | None = None,
) -> Tuple[str | None, str | None]:
"""
Get API key and base URL from multiple sources.
Returns tuple of (api_key, api_base).

View file

@ -39,12 +39,12 @@ class JinaAIRerankConfig(BaseRerankConfig):
drop_params: bool,
query: str,
documents: List[Union[str, Dict[str, Any]]],
custom_llm_provider: Optional[str] = None,
top_n: Optional[int] = None,
rank_fields: Optional[List[str]] = None,
return_documents: Optional[bool] = True,
max_chunks_per_doc: Optional[int] = None,
max_tokens_per_doc: Optional[int] = None,
custom_llm_provider: str | None = None,
top_n: int | None = None,
rank_fields: List[str] | None = None,
return_documents: bool | None = True,
max_chunks_per_doc: int | None = None,
max_tokens_per_doc: int | None = None,
) -> Dict:
optional_params = {}
supported_params = self.get_supported_cohere_rerank_params(model)
@ -59,9 +59,9 @@ class JinaAIRerankConfig(BaseRerankConfig):
def get_complete_url(
self,
api_base: Optional[str],
api_base: str | None,
model: str,
optional_params: Optional[dict] = None,
optional_params: dict | None = None,
) -> str:
base_path = "/v1/rerank"
@ -78,7 +78,7 @@ class JinaAIRerankConfig(BaseRerankConfig):
model: str,
optional_rerank_params: Dict,
headers: Dict,
litellm_params: Optional[dict] = None,
litellm_params: dict | None = None,
) -> Dict:
return {"model": model, **optional_rerank_params}
@ -88,7 +88,7 @@ class JinaAIRerankConfig(BaseRerankConfig):
raw_response: Response,
model_response: RerankResponse,
logging_obj: LiteLLMLoggingObj,
api_key: Optional[str] = None,
api_key: str | None = None,
request_data: Dict = {},
optional_params: Dict = {},
litellm_params: Dict = {},
@ -104,7 +104,7 @@ class JinaAIRerankConfig(BaseRerankConfig):
_tokens = RerankTokens(**_json_response.get("usage", {}))
rerank_meta = RerankResponseMeta(billed_units=_billed_units, tokens=_tokens)
_results: Optional[List[dict]] = _json_response.get("results")
_results: List[dict] | None = _json_response.get("results")
if _results is None:
raise ValueError(f"No results found in the response={_json_response}")
@ -136,8 +136,8 @@ class JinaAIRerankConfig(BaseRerankConfig):
self,
headers: Dict,
model: str,
api_key: Optional[str] = None,
optional_params: Optional[dict] = None,
api_key: str | None = None,
optional_params: dict | None = None,
) -> Dict:
if api_key is None:
raise ValueError(
@ -152,9 +152,9 @@ class JinaAIRerankConfig(BaseRerankConfig):
def calculate_rerank_cost(
self,
model: str,
custom_llm_provider: Optional[str] = None,
billed_units: Optional[RerankBilledUnits] = None,
model_info: Optional[ModelInfo] = None,
custom_llm_provider: str | None = None,
billed_units: RerankBilledUnits | None = None,
model_info: ModelInfo | None = None,
) -> Tuple[float, float]:
"""
Jina AI reranker is priced at $0.000000018 per token.

View file

@ -64,9 +64,9 @@ class NvidiaNimRerankConfig(BaseRerankConfig):
def get_complete_url(
self,
api_base: Optional[str],
api_base: str | None,
model: str,
optional_params: Optional[dict] = None,
optional_params: dict | None = None,
) -> str:
"""
Construct the Nvidia NIM rerank URL.
@ -106,17 +106,17 @@ class NvidiaNimRerankConfig(BaseRerankConfig):
def map_cohere_rerank_params(
self,
non_default_params: Optional[dict],
non_default_params: dict | None,
model: str,
drop_params: bool,
query: str,
documents: List[Union[str, Dict[str, Any]]],
custom_llm_provider: Optional[str] = None,
top_n: Optional[int] = None,
rank_fields: Optional[List[str]] = None,
return_documents: Optional[bool] = True,
max_chunks_per_doc: Optional[int] = None,
max_tokens_per_doc: Optional[int] = None,
custom_llm_provider: str | None = None,
top_n: int | None = None,
rank_fields: List[str] | None = None,
return_documents: bool | None = True,
max_chunks_per_doc: int | None = None,
max_tokens_per_doc: int | None = None,
) -> Dict:
"""
Map Cohere/OpenAI rerank params to Nvidia NIM format.
@ -145,8 +145,8 @@ class NvidiaNimRerankConfig(BaseRerankConfig):
self,
headers: dict,
model: str,
api_key: Optional[str] = None,
optional_params: Optional[dict] = None,
api_key: str | None = None,
optional_params: dict | None = None,
) -> dict:
"""
Validate that the Nvidia NIM API key is present.
@ -177,7 +177,7 @@ class NvidiaNimRerankConfig(BaseRerankConfig):
model: str,
optional_rerank_params: Dict,
headers: dict,
litellm_params: Optional[dict] = None,
litellm_params: dict | None = None,
) -> dict:
"""
Transform request to Nvidia NIM format.
@ -252,7 +252,7 @@ class NvidiaNimRerankConfig(BaseRerankConfig):
raw_response: httpx.Response,
model_response: RerankResponse,
logging_obj: LiteLLMLoggingObj,
api_key: Optional[str] = None,
api_key: str | None = None,
request_data: dict = {},
optional_params: dict = {},
litellm_params: dict = {},

View file

@ -36,9 +36,9 @@ class VertexAIRerankConfig(BaseRerankConfig, VertexBase):
def get_complete_url(
self,
api_base: Optional[str],
api_base: str | None,
model: str,
optional_params: Optional[Dict] = None,
optional_params: Dict | None = None,
) -> str:
"""
Get the complete URL for the Vertex AI Discovery Engine ranking API
@ -76,8 +76,8 @@ class VertexAIRerankConfig(BaseRerankConfig, VertexBase):
self,
headers: dict,
model: str,
api_key: Optional[str] = None,
optional_params: Optional[Dict] = None,
api_key: str | None = None,
optional_params: Dict | None = None,
) -> dict:
"""
Validate and set up authentication for Vertex AI Discovery Engine API
@ -112,7 +112,7 @@ class VertexAIRerankConfig(BaseRerankConfig, VertexBase):
model: str,
optional_rerank_params: Dict,
headers: dict,
litellm_params: Optional[dict] = None,
litellm_params: dict | None = None,
) -> dict:
"""
Transform the request from Cohere format to Vertex AI Discovery Engine format
@ -161,7 +161,7 @@ class VertexAIRerankConfig(BaseRerankConfig, VertexBase):
raw_response: httpx.Response,
model_response: RerankResponse,
logging_obj: LiteLLMLoggingObj,
api_key: Optional[str] = None,
api_key: str | None = None,
request_data: dict = {},
optional_params: dict = {},
litellm_params: dict = {},
@ -236,12 +236,12 @@ class VertexAIRerankConfig(BaseRerankConfig, VertexBase):
drop_params: bool,
query: str,
documents: List[Union[str, Dict[str, Any]]],
custom_llm_provider: Optional[str] = None,
top_n: Optional[int] = None,
rank_fields: Optional[List[str]] = None,
return_documents: Optional[bool] = True,
max_chunks_per_doc: Optional[int] = None,
max_tokens_per_doc: Optional[int] = None,
custom_llm_provider: str | None = None,
top_n: int | None = None,
rank_fields: List[str] | None = None,
return_documents: bool | None = True,
max_chunks_per_doc: int | None = None,
max_tokens_per_doc: int | None = None,
) -> Dict:
"""
Map Cohere rerank params to Vertex AI format

View file

@ -33,12 +33,12 @@ class VoyageRerankConfig(BaseRerankConfig):
drop_params: bool,
query: str,
documents: List[Union[str, Dict[str, Any]]],
custom_llm_provider: Optional[str] = None,
top_n: Optional[int] = None,
rank_fields: Optional[List[str]] = None,
return_documents: Optional[bool] = True,
max_chunks_per_doc: Optional[int] = None,
max_tokens_per_doc: Optional[int] = None,
custom_llm_provider: str | None = None,
top_n: int | None = None,
rank_fields: List[str] | None = None,
return_documents: bool | None = True,
max_chunks_per_doc: int | None = None,
max_tokens_per_doc: int | None = None,
) -> Dict:
# Voyage AI uses 'top_k' instead of 'top_n'
optional_params: Dict[str, Any] = {"query": query, "documents": documents}
@ -52,9 +52,9 @@ class VoyageRerankConfig(BaseRerankConfig):
def get_complete_url(
self,
api_base: Optional[str],
api_base: str | None,
model: str,
optional_params: Optional[dict] = None,
optional_params: dict | None = None,
) -> str:
if api_base is None:
return "https://api.voyageai.com/v1/rerank"
@ -71,7 +71,7 @@ class VoyageRerankConfig(BaseRerankConfig):
model: str,
optional_rerank_params: Dict,
headers: Dict,
litellm_params: Optional[dict] = None,
litellm_params: dict | None = None,
) -> Dict:
return {"model": model, **optional_rerank_params}
@ -81,7 +81,7 @@ class VoyageRerankConfig(BaseRerankConfig):
raw_response: httpx.Response,
model_response: RerankResponse,
logging_obj: LiteLLMLoggingObj,
api_key: Optional[str] = None,
api_key: str | None = None,
request_data: Dict = {},
optional_params: Dict = {},
litellm_params: Dict = {},
@ -102,7 +102,7 @@ class VoyageRerankConfig(BaseRerankConfig):
)
# Voyage AI returns results in "data" key, not "results"
_results: Optional[List[dict]] = _json_response.get("data")
_results: List[dict] | None = _json_response.get("data")
if _results is None:
raise ValueError(f"No results found in the response={_json_response}")
@ -136,8 +136,8 @@ class VoyageRerankConfig(BaseRerankConfig):
self,
headers: Dict,
model: str,
api_key: Optional[str] = None,
optional_params: Optional[dict] = None,
api_key: str | None = None,
optional_params: dict | None = None,
) -> Dict:
if api_key is None:
api_key = get_secret_str("VOYAGE_API_KEY") or get_secret_str(
@ -155,9 +155,9 @@ class VoyageRerankConfig(BaseRerankConfig):
def calculate_rerank_cost(
self,
model: str,
custom_llm_provider: Optional[str] = None,
billed_units: Optional[RerankBilledUnits] = None,
model_info: Optional[ModelInfo] = None,
custom_llm_provider: str | None = None,
billed_units: RerankBilledUnits | None = None,
model_info: ModelInfo | None = None,
) -> Tuple[float, float]:
if (
model_info is None

View file

@ -31,9 +31,9 @@ class IBMWatsonXRerankConfig(IBMWatsonXMixin, BaseRerankConfig):
def get_complete_url(
self,
api_base: Optional[str],
api_base: str | None,
model: str,
optional_params: Optional[dict] = None,
optional_params: dict | None = None,
) -> str:
base_url = self._get_base_url(api_base=api_base)
endpoint = WatsonXAIEndpoint.RERANK.value
@ -60,8 +60,8 @@ class IBMWatsonXRerankConfig(IBMWatsonXMixin, BaseRerankConfig):
self,
headers: dict,
model: str,
api_key: Optional[str] = None,
optional_params: Optional[dict] = None,
api_key: str | None = None,
optional_params: dict | None = None,
) -> Dict:
optional_params = optional_params or {}
@ -73,11 +73,11 @@ class IBMWatsonXRerankConfig(IBMWatsonXMixin, BaseRerankConfig):
if "Authorization" in headers:
return {**default_headers, **headers}
token = cast(
Optional[str],
str | None,
optional_params.pop("token", None) or get_secret_str("WATSONX_TOKEN"),
)
zen_api_key = cast(
Optional[str],
str | None,
optional_params.pop("zen_api_key", None)
or get_secret_str("WATSONX_ZENAPIKEY"),
)
@ -93,17 +93,17 @@ class IBMWatsonXRerankConfig(IBMWatsonXMixin, BaseRerankConfig):
def map_cohere_rerank_params(
self,
non_default_params: Optional[dict],
non_default_params: dict | None,
model: str,
drop_params: bool,
query: str,
documents: List[Union[str, Dict[str, Any]]],
custom_llm_provider: Optional[str] = None,
top_n: Optional[int] = None,
rank_fields: Optional[List[str]] = None,
return_documents: Optional[bool] = True,
max_chunks_per_doc: Optional[int] = None,
max_tokens_per_doc: Optional[int] = None,
custom_llm_provider: str | None = None,
top_n: int | None = None,
rank_fields: List[str] | None = None,
return_documents: bool | None = True,
max_chunks_per_doc: int | None = None,
max_tokens_per_doc: int | None = None,
) -> Dict:
"""
Map Cohere rerank params to IBM watsonx.ai rerank params
@ -143,7 +143,7 @@ class IBMWatsonXRerankConfig(IBMWatsonXMixin, BaseRerankConfig):
model: str,
optional_rerank_params: Dict,
headers: dict,
litellm_params: Optional[dict] = None,
litellm_params: dict | None = None,
) -> dict:
"""
Transform request to IBM watsonx.ai rerank format
@ -162,7 +162,7 @@ class IBMWatsonXRerankConfig(IBMWatsonXMixin, BaseRerankConfig):
raw_response: httpx.Response,
model_response: RerankResponse,
logging_obj: LiteLLMLoggingObj,
api_key: Optional[str] = None,
api_key: str | None = None,
request_data: dict = {},
optional_params: dict = {},
litellm_params: dict = {},
@ -179,7 +179,7 @@ class IBMWatsonXRerankConfig(IBMWatsonXMixin, BaseRerankConfig):
headers=raw_response.headers,
)
_results: Optional[List[dict]] = raw_response_json.get("results")
_results: List[dict] | None = raw_response_json.get("results")
if _results is None:
raise ValueError(f"No results found in the response={raw_response_json}")

View file

@ -21,7 +21,7 @@ class ColdStorageHandler:
async def get_proxy_server_request_from_cold_storage_with_object_key(
self,
object_key: str,
) -> Optional[dict]:
) -> dict | None:
"""
Get the proxy server request from cold storage using the object key directly.
@ -33,7 +33,7 @@ class ColdStorageHandler:
"""
# select the custom logger to use for cold storage
custom_logger_name: Optional[_custom_logger_compatible_callbacks_literal] = (
custom_logger_name: _custom_logger_compatible_callbacks_literal | None = (
self._select_custom_logger_for_cold_storage()
)
@ -42,7 +42,7 @@ class ColdStorageHandler:
return None
# get the active/initialized custom logger
custom_logger: Optional[CustomLogger] = (
custom_logger: CustomLogger | None = (
litellm.logging_callback_manager.get_active_custom_logger_for_callback_name(
custom_logger_name
)
@ -60,9 +60,7 @@ class ColdStorageHandler:
def _select_custom_logger_for_cold_storage(
self,
) -> Optional[_custom_logger_compatible_callbacks_literal]:
cold_storage_custom_logger: Optional[
_custom_logger_compatible_callbacks_literal
] = litellm.cold_storage_custom_logger
) -> _custom_logger_compatible_callbacks_literal | None:
cold_storage_custom_logger: _custom_logger_compatible_callbacks_literal | None = litellm.cold_storage_custom_logger
return cold_storage_custom_logger

View file

@ -104,7 +104,7 @@ def _strip_password_from_users(users) -> None:
include_in_schema=False,
)
async def spend_user_fn(
user_id: Optional[str] = fastapi.Query(
user_id: str | None = fastapi.Query(
default=None,
description="Get User Table row for user_id",
),
@ -185,11 +185,11 @@ async def spend_user_fn(
},
)
async def view_spend_tags(
start_date: Optional[str] = fastapi.Query(
start_date: str | None = fastapi.Query(
default=None,
description="Time from which to start viewing key spend",
),
end_date: Optional[str] = fastapi.Query(
end_date: str | None = fastapi.Query(
default=None,
description="Time till which to view key spend",
),
@ -290,11 +290,11 @@ async def get_global_activity_internal_user(
include_in_schema=False,
)
async def get_global_activity(
start_date: Optional[str] = fastapi.Query(
start_date: str | None = fastapi.Query(
default=None,
description="Time from which to start viewing spend",
),
end_date: Optional[str] = fastapi.Query(
end_date: str | None = fastapi.Query(
default=None,
description="Time till which to view spend",
),
@ -437,11 +437,11 @@ async def get_global_activity_model_internal_user(
include_in_schema=False,
)
async def get_global_activity_model(
start_date: Optional[str] = fastapi.Query(
start_date: str | None = fastapi.Query(
default=None,
description="Time from which to start viewing spend",
),
end_date: Optional[str] = fastapi.Query(
end_date: str | None = fastapi.Query(
default=None,
description="Time till which to view spend",
),
@ -599,11 +599,11 @@ async def get_global_activity_exceptions_per_deployment(
model_group: str = fastapi.Query(
description="Filter by model group",
),
start_date: Optional[str] = fastapi.Query(
start_date: str | None = fastapi.Query(
default=None,
description="Time from which to start viewing spend",
),
end_date: Optional[str] = fastapi.Query(
end_date: str | None = fastapi.Query(
default=None,
description="Time till which to view spend",
),
@ -754,11 +754,11 @@ async def get_global_activity_exceptions(
model_group: str = fastapi.Query(
description="Filter by model group",
),
start_date: Optional[str] = fastapi.Query(
start_date: str | None = fastapi.Query(
default=None,
description="Time from which to start viewing spend",
),
end_date: Optional[str] = fastapi.Query(
end_date: str | None = fastapi.Query(
default=None,
description="Time till which to view spend",
),
@ -861,11 +861,11 @@ async def get_global_activity_exceptions(
},
)
async def get_global_spend_provider(
start_date: Optional[str] = fastapi.Query(
start_date: str | None = fastapi.Query(
default=None,
description="Time from which to start viewing spend",
),
end_date: Optional[str] = fastapi.Query(
end_date: str | None = fastapi.Query(
default=None,
description="Time till which to view spend",
),
@ -996,31 +996,31 @@ async def get_global_spend_provider(
},
)
async def get_global_spend_report(
start_date: Optional[str] = fastapi.Query(
start_date: str | None = fastapi.Query(
default=None,
description="Time from which to start viewing spend",
),
end_date: Optional[str] = fastapi.Query(
end_date: str | None = fastapi.Query(
default=None,
description="Time till which to view spend",
),
group_by: Optional[Literal["team", "customer", "api_key"]] = fastapi.Query(
group_by: Literal["team", "customer", "api_key"] | None = fastapi.Query(
default="team",
description="Group spend by internal team or customer or api_key",
),
api_key: Optional[str] = fastapi.Query(
api_key: str | None = fastapi.Query(
default=None,
description="View spend for a specific api_key. Example api_key='sk-1234",
),
internal_user_id: Optional[str] = fastapi.Query(
internal_user_id: str | None = fastapi.Query(
default=None,
description="View spend for a specific internal_user_id. Example internal_user_id='1234",
),
team_id: Optional[str] = fastapi.Query(
team_id: str | None = fastapi.Query(
default=None,
description="View spend for a specific team_id. Example team_id='1234",
),
customer_id: Optional[str] = fastapi.Query(
customer_id: str | None = fastapi.Query(
default=None,
description="View spend for a specific customer_id. Example customer_id='1234. Can be used in conjunction with team_id as well.",
),
@ -1416,15 +1416,15 @@ async def global_get_all_tag_names():
},
)
async def global_view_spend_tags(
start_date: Optional[str] = fastapi.Query(
start_date: str | None = fastapi.Query(
default=None,
description="Time from which to start viewing key spend",
),
end_date: Optional[str] = fastapi.Query(
end_date: str | None = fastapi.Query(
default=None,
description="Time till which to view key spend",
),
tags: Optional[str] = fastapi.Query(
tags: str | None = fastapi.Query(
default=None,
description="comman separated tags to filter on",
),
@ -1643,7 +1643,7 @@ async def calculate_spend(request: SpendCalculateRequest):
# check if model in llm_router
_model_in_llm_router = None
cost_per_token: Optional[CostPerToken] = None
cost_per_token: CostPerToken | None = None
if llm_router is not None:
if (
llm_router.model_group_alias is not None
@ -1736,35 +1736,35 @@ async def calculate_spend(request: SpendCalculateRequest):
)
async def ui_view_spend_logs(
request: Request,
api_key: Optional[str] = fastapi.Query(
api_key: str | None = fastapi.Query(
default=None,
description="Get spend logs based on api key",
),
user_id: Optional[str] = fastapi.Query(
user_id: str | None = fastapi.Query(
default=None,
description="Get spend logs based on user_id",
),
request_id: Optional[str] = fastapi.Query(
request_id: str | None = fastapi.Query(
default=None,
description="request_id to get spend logs for specific request_id",
),
team_id: Optional[str] = fastapi.Query(
team_id: str | None = fastapi.Query(
default=None,
description="Filter spend logs by team_id",
),
min_spend: Optional[float] = fastapi.Query(
min_spend: float | None = fastapi.Query(
default=None,
description="Filter logs with spend greater than or equal to this value",
),
max_spend: Optional[float] = fastapi.Query(
max_spend: float | None = fastapi.Query(
default=None,
description="Filter logs with spend less than or equal to this value",
),
start_date: Optional[str] = fastapi.Query(
start_date: str | None = fastapi.Query(
default=None,
description="Time from which to start viewing key spend",
),
end_date: Optional[str] = fastapi.Query(
end_date: str | None = fastapi.Query(
default=None,
description="Time till which to view key spend",
),
@ -1775,36 +1775,36 @@ async def ui_view_spend_logs(
default=50, description="Number of items per page", ge=1, le=100
),
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
status_filter: Optional[str] = fastapi.Query(
status_filter: str | None = fastapi.Query(
default=None, description="Filter logs by status (e.g., success, failure)"
),
model: Optional[str] = fastapi.Query(
model: str | None = fastapi.Query(
default=None, description="Filter logs by model"
),
model_id: Optional[str] = fastapi.Query(
model_id: str | None = fastapi.Query(
default=None,
description="Filter logs by model ID (litellm model deployment id)",
),
model_group: Optional[str] = fastapi.Query(
model_group: str | None = fastapi.Query(
default=None, description="Filter logs by model group"
),
key_alias: Optional[str] = fastapi.Query(
key_alias: str | None = fastapi.Query(
default=None, description="Filter logs by key alias"
),
end_user: Optional[str] = fastapi.Query(
end_user: str | None = fastapi.Query(
default=None, description="Filter logs by end user"
),
error_code: Optional[str] = fastapi.Query(
error_code: str | None = fastapi.Query(
default=None, description="Filter logs by error code (e.g., '404', '500')"
),
error_message: Optional[str] = fastapi.Query(
error_message: str | None = fastapi.Query(
default=None, description="Filter logs by error message (partial string match)"
),
sort_by: str = fastapi.Query(
default="startTime",
description="Sort logs by field: spend, total_tokens, startTime, endTime, request_duration_ms, model, or ttft_ms",
),
sort_order: Optional[str] = fastapi.Query(
sort_order: str | None = fastapi.Query(
default="desc",
description="Sort order: asc or desc",
),
@ -1968,7 +1968,7 @@ async def ui_view_spend_logs(
if max_spend is not None:
where_conditions["spend"]["lte"] = max_spend
is_admin_view = _is_admin_view_safe(user_api_key_dict=user_api_key_dict)
permitted_team_ids: Optional[List[str]] = None
permitted_team_ids: List[str] | None = None
if not is_admin_view:
if team_id is not None:
can_view_team = await _can_team_member_view_log(
@ -2183,11 +2183,11 @@ async def ui_view_spend_logs(
)
async def ui_view_request_response_for_request_id(
request_id: str,
start_date: Optional[str] = fastapi.Query(
start_date: str | None = fastapi.Query(
default=None,
description="Time from which to start viewing key spend",
),
end_date: Optional[str] = fastapi.Query(
end_date: str | None = fastapi.Query(
default=None,
description="Time till which to view key spend",
),
@ -2221,8 +2221,8 @@ async def ui_view_request_response_for_request_id(
custom_loggers = (
litellm.logging_callback_manager.get_active_additional_logging_utils_from_custom_logger()
)
start_date_obj: Optional[datetime] = None
end_date_obj: Optional[datetime] = None
start_date_obj: datetime | None = None
end_date_obj: datetime | None = None
if start_date is not None:
start_date_obj = datetime.strptime(start_date, "%Y-%m-%d %H:%M:%S").replace(
tzinfo=timezone.utc
@ -2274,23 +2274,23 @@ async def ui_view_request_response_for_request_id(
},
)
async def view_spend_logs(
api_key: Optional[str] = fastapi.Query(
api_key: str | None = fastapi.Query(
default=None,
description="Get spend logs based on api key",
),
user_id: Optional[str] = fastapi.Query(
user_id: str | None = fastapi.Query(
default=None,
description="Get spend logs based on user_id",
),
request_id: Optional[str] = fastapi.Query(
request_id: str | None = fastapi.Query(
default=None,
description="request_id to get spend logs for specific request_id. If none passed then pass spend logs for all requests",
),
start_date: Optional[str] = fastapi.Query(
start_date: str | None = fastapi.Query(
default=None,
description="Time from which to start viewing key spend",
),
end_date: Optional[str] = fastapi.Query(
end_date: str | None = fastapi.Query(
default=None,
description="Time till which to view key spend",
),
@ -2626,7 +2626,7 @@ async def global_spend_refresh():
async def global_spend_for_internal_user(
api_key: Optional[str] = None,
api_key: str | None = None,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
from litellm.proxy.proxy_server import prisma_client
@ -2670,7 +2670,7 @@ async def global_spend_for_internal_user(
include_in_schema=False,
)
async def global_spend_logs(
api_key: Optional[str] = fastapi.Query(
api_key: str | None = fastapi.Query(
default=None,
description="API Key to get global spend (spend per day for last 30d). Admin-only endpoint",
),
@ -3029,7 +3029,7 @@ async def global_view_all_end_users():
dependencies=[Depends(user_api_key_auth)],
include_in_schema=False,
)
async def global_spend_end_users(data: Optional[GlobalEndUsersSpend] = None):
async def global_spend_end_users(data: GlobalEndUsersSpend | None = None):
"""
[BETA] This is a beta endpoint. It will change.
@ -3262,8 +3262,8 @@ async def get_spend_by_tags(
async def ui_get_spend_by_tags(
start_date: str,
end_date: str,
prisma_client: Optional[PrismaClient] = None,
tags_str: Optional[str] = None,
prisma_client: PrismaClient | None = None,
tags_str: str | None = None,
):
"""
Should cover 2 cases:
@ -3274,7 +3274,7 @@ async def ui_get_spend_by_tags(
# tags_str is a list of strings csv of tags
# tags_str = tag1,tag2,tag3
# convert to list if it's not None
tags_list: Optional[List[str]] = None
tags_list: List[str] | None = None
if tags_str is not None and len(tags_str) > 0:
tags_list = tags_str.split(",")
@ -3538,7 +3538,7 @@ async def _build_ui_spend_logs_response(
}
def _build_status_filter_condition(status_filter: Optional[str]) -> Dict[str, Any]:
def _build_status_filter_condition(status_filter: str | None) -> Dict[str, Any]:
"""
Helper function to build the status filter condition for database queries.
@ -3577,7 +3577,7 @@ def _is_admin_view_safe(user_api_key_dict: UserAPIKeyAuth) -> bool:
async def _can_team_member_view_log(
prisma_client,
user_api_key_dict: UserAPIKeyAuth,
team_id: Optional[str],
team_id: str | None,
) -> bool:
"""
Check if the requesting user can view spend logs for the given team.

View file

@ -30,15 +30,11 @@ async def arerank(
model: str,
query: str,
documents: List[Union[str, Dict[str, Any]]],
custom_llm_provider: Optional[
Literal[
"cohere", "together_ai", "deepinfra", "fireworks_ai", "voyage", "watsonx"
]
] = None,
top_n: Optional[int] = None,
rank_fields: Optional[List[str]] = None,
return_documents: Optional[bool] = None,
max_chunks_per_doc: Optional[int] = None,
custom_llm_provider: Literal["cohere", "together_ai", "deepinfra", "fireworks_ai", "voyage", "watsonx"] | None = None,
top_n: int | None = None,
rank_fields: List[str] | None = None,
return_documents: bool | None = None,
max_chunks_per_doc: int | None = None,
**kwargs,
) -> Union[RerankResponse, Coroutine[Any, Any, RerankResponse]]:
"""
@ -79,33 +75,20 @@ def rerank(
model: str,
query: str,
documents: List[Union[str, Dict[str, Any]]],
custom_llm_provider: Optional[
Literal[
"cohere",
"together_ai",
"azure_ai",
"infinity",
"litellm_proxy",
"hosted_vllm",
"deepinfra",
"fireworks_ai",
"voyage",
"watsonx",
]
] = None,
top_n: Optional[int] = None,
rank_fields: Optional[List[str]] = None,
return_documents: Optional[bool] = True,
max_chunks_per_doc: Optional[int] = None,
max_tokens_per_doc: Optional[int] = None,
custom_llm_provider: Literal["cohere", "together_ai", "azure_ai", "infinity", "litellm_proxy", "hosted_vllm", "deepinfra", "fireworks_ai", "voyage", "watsonx"] | None = None,
top_n: int | None = None,
rank_fields: List[str] | None = None,
return_documents: bool | None = True,
max_chunks_per_doc: int | None = None,
max_tokens_per_doc: int | None = None,
**kwargs,
) -> Union[RerankResponse, Coroutine[Any, Any, RerankResponse]]:
"""
Reranks a list of documents based on their relevance to the query
"""
headers: Optional[dict] = kwargs.get("headers") # type: ignore
headers: dict | None = kwargs.get("headers") # type: ignore
litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore
litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None)
litellm_call_id: str | None = kwargs.get("litellm_call_id", None)
proxy_server_request = kwargs.get("proxy_server_request", None)
model_info = kwargs.get("model_info", None)
user = kwargs.get("user", None)
@ -187,11 +170,11 @@ def rerank(
or _custom_llm_provider == litellm.LlmProviders.LITELLM_PROXY
):
# Implement Cohere rerank logic
api_key: Optional[str] = (
api_key: str | None = (
dynamic_api_key or optional_params.api_key or litellm.api_key
)
api_base: Optional[str] = (
api_base: str | None = (
dynamic_api_base
or optional_params.api_base
or litellm.api_base

View file

@ -9,13 +9,13 @@ def get_optional_rerank_params(
drop_params: bool,
query: str,
documents: List[Union[str, Dict[str, Any]]],
custom_llm_provider: Optional[str] = None,
top_n: Optional[int] = None,
rank_fields: Optional[List[str]] = None,
return_documents: Optional[bool] = True,
max_chunks_per_doc: Optional[int] = None,
max_tokens_per_doc: Optional[int] = None,
non_default_params: Optional[dict] = None,
custom_llm_provider: str | None = None,
top_n: int | None = None,
rank_fields: List[str] | None = None,
return_documents: bool | None = True,
max_chunks_per_doc: int | None = None,
max_tokens_per_doc: int | None = None,
non_default_params: dict | None = None,
) -> Dict:
all_non_default_params = non_default_params or {}
if query is not None: