mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-08 22:21:35 +00:00
fix: fix minor linting errors
This commit is contained in:
parent
305fff8ffb
commit
7ab737e4a8
4 changed files with 41 additions and 28 deletions
|
|
@ -1,5 +1,5 @@
|
|||
"""
|
||||
Support for gpt model family
|
||||
Support for gpt model family
|
||||
"""
|
||||
|
||||
from typing import List, Optional, Union
|
||||
|
|
@ -87,7 +87,7 @@ class OpenAITextCompletionConfig(BaseTextCompletionConfig, OpenAIGPTConfig):
|
|||
## RESPONSE OBJECT
|
||||
if response_object is None or model_response_object is None:
|
||||
raise ValueError("Error in response object format")
|
||||
choice_list = []
|
||||
choice_list: List[Choices] = []
|
||||
for idx, choice in enumerate(response_object["choices"]):
|
||||
message = Message(
|
||||
content=choice["text"],
|
||||
|
|
@ -100,7 +100,7 @@ class OpenAITextCompletionConfig(BaseTextCompletionConfig, OpenAIGPTConfig):
|
|||
logprobs=choice.get("logprobs", None),
|
||||
)
|
||||
choice_list.append(choice)
|
||||
model_response_object.choices = choice_list
|
||||
model_response_object.choices = choice_list # type: ignore
|
||||
|
||||
if "usage" in response_object:
|
||||
setattr(model_response_object, "usage", response_object["usage"])
|
||||
|
|
@ -111,9 +111,9 @@ class OpenAITextCompletionConfig(BaseTextCompletionConfig, OpenAIGPTConfig):
|
|||
if "model" in response_object:
|
||||
model_response_object.model = response_object["model"]
|
||||
|
||||
model_response_object._hidden_params[
|
||||
"original_response"
|
||||
] = response_object # track original response, if users make a litellm.text_completion() request, we can return the original response
|
||||
model_response_object._hidden_params["original_response"] = (
|
||||
response_object # track original response, if users make a litellm.text_completion() request, we can return the original response
|
||||
)
|
||||
return model_response_object
|
||||
except Exception as e:
|
||||
raise e
|
||||
|
|
|
|||
|
|
@ -64,9 +64,9 @@ class VertexFineTuningAPI(VertexLLM):
|
|||
)
|
||||
|
||||
if create_fine_tuning_job_data.validation_file:
|
||||
supervised_tuning_spec[
|
||||
"validation_dataset"
|
||||
] = create_fine_tuning_job_data.validation_file
|
||||
supervised_tuning_spec["validation_dataset"] = (
|
||||
create_fine_tuning_job_data.validation_file
|
||||
)
|
||||
|
||||
_vertex_hyperparameters = (
|
||||
self._transform_openai_hyperparameters_to_vertex_hyperparameters(
|
||||
|
|
@ -140,7 +140,9 @@ class VertexFineTuningAPI(VertexLLM):
|
|||
fine_tuned_model=response.get("tunedModelDisplayName", ""),
|
||||
finished_at=None,
|
||||
hyperparameters=self._translate_vertex_response_hyperparameters(
|
||||
vertex_hyper_parameters=_supervisedTuningSpec.get("hyperParameters", {})
|
||||
vertex_hyper_parameters=_supervisedTuningSpec.get(
|
||||
"hyperParameters", FineTuneHyperparameters()
|
||||
)
|
||||
or {}
|
||||
),
|
||||
model=response.get("baseModel", "") or "",
|
||||
|
|
@ -343,9 +345,9 @@ class VertexFineTuningAPI(VertexLLM):
|
|||
elif "cachedContents" in request_route:
|
||||
_model = request_data.get("model")
|
||||
if _model is not None and "/publishers/google/models/" not in _model:
|
||||
request_data[
|
||||
"model"
|
||||
] = f"projects/{vertex_project}/locations/{vertex_location}/publishers/google/models/{_model}"
|
||||
request_data["model"] = (
|
||||
f"projects/{vertex_project}/locations/{vertex_location}/publishers/google/models/{_model}"
|
||||
)
|
||||
|
||||
url = f"https://{vertex_location}-aiplatform.googleapis.com/v1beta1/projects/{vertex_project}/locations/{vertex_location}{request_route}"
|
||||
else:
|
||||
|
|
|
|||
|
|
@ -70,7 +70,7 @@ class WatsonXChatHandler(OpenAILikeChatHandler):
|
|||
)
|
||||
|
||||
return super().completion(
|
||||
model=watsonx_auth_payload.get("model_id", None),
|
||||
model=watsonx_auth_payload.get("model_id") or "",
|
||||
messages=messages,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
|
|
|
|||
|
|
@ -17,7 +17,6 @@ import logging
|
|||
import threading
|
||||
import time
|
||||
import traceback
|
||||
from litellm._uuid import uuid
|
||||
from collections import defaultdict
|
||||
from functools import lru_cache
|
||||
from typing import (
|
||||
|
|
@ -45,6 +44,7 @@ import litellm.litellm_core_utils
|
|||
import litellm.litellm_core_utils.exception_mapping_utils
|
||||
from litellm import get_secret_str
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.caching.caching import (
|
||||
DualCache,
|
||||
InMemoryCache,
|
||||
|
|
@ -2005,11 +2005,17 @@ class Router:
|
|||
|
||||
# Filter out prompt management specific parameters from data before merging
|
||||
prompt_management_params = {
|
||||
"bitbucket_config", "dotprompt_config", "prompt_id",
|
||||
"prompt_variables", "prompt_label", "prompt_version"
|
||||
"bitbucket_config",
|
||||
"dotprompt_config",
|
||||
"prompt_id",
|
||||
"prompt_variables",
|
||||
"prompt_label",
|
||||
"prompt_version",
|
||||
}
|
||||
filtered_data = {k: v for k, v in data.items() if k not in prompt_management_params}
|
||||
|
||||
filtered_data = {
|
||||
k: v for k, v in data.items() if k not in prompt_management_params
|
||||
}
|
||||
|
||||
kwargs = {**filtered_data, **kwargs, **optional_params}
|
||||
kwargs["model"] = model
|
||||
kwargs["messages"] = messages
|
||||
|
|
@ -4108,7 +4114,9 @@ class Router:
|
|||
"""
|
||||
model_group = kwargs.get("model")
|
||||
response = original_function(*args, **kwargs)
|
||||
if coroutine_checker.is_async_callable(response) or inspect.isawaitable(response):
|
||||
if coroutine_checker.is_async_callable(response) or inspect.isawaitable(
|
||||
response
|
||||
):
|
||||
response = await response
|
||||
## PROCESS RESPONSE HEADERS
|
||||
response = await self.set_response_headers(
|
||||
|
|
@ -4517,7 +4525,9 @@ class Router:
|
|||
_time_to_cooldown = self.cooldown_time
|
||||
|
||||
if isinstance(_model_info, dict):
|
||||
deployment_id = _model_info.get("id", None)
|
||||
deployment_id: Optional[str] = _model_info.get("id")
|
||||
if deployment_id is None:
|
||||
return False
|
||||
increment_deployment_failures_for_current_minute(
|
||||
litellm_router_instance=self,
|
||||
deployment_id=deployment_id,
|
||||
|
|
@ -5134,12 +5144,12 @@ class Router:
|
|||
# Check if this is a prompt management model before validating as LLM provider
|
||||
litellm_model = deployment.litellm_params.model
|
||||
is_prompt_management_model = False
|
||||
|
||||
|
||||
if "/" in litellm_model:
|
||||
split_litellm_model = litellm_model.split("/")[0]
|
||||
if split_litellm_model in litellm._known_custom_logger_compatible_callbacks:
|
||||
is_prompt_management_model = True
|
||||
|
||||
|
||||
if is_prompt_management_model:
|
||||
# For prompt management models, skip LLM provider validation
|
||||
# The actual model will be resolved at runtime from the prompt file
|
||||
|
|
@ -5229,11 +5239,12 @@ class Router:
|
|||
# litellm_router_instance=self, model=deployment.to_json(exclude_none=True)
|
||||
# )
|
||||
|
||||
self._initialize_deployment_for_pass_through(
|
||||
deployment=deployment,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
model=deployment.litellm_params.model,
|
||||
)
|
||||
if custom_llm_provider is not None:
|
||||
self._initialize_deployment_for_pass_through(
|
||||
deployment=deployment,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
model=deployment.litellm_params.model,
|
||||
)
|
||||
|
||||
#########################################################
|
||||
# Check if this is an auto-router deployment
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue