refactor: update method signature

This commit is contained in:
Krrish Dholakia 2025-03-12 15:23:38 -07:00
parent 738c0b873d
commit 88e9edf7db
23 changed files with 35 additions and 5 deletions

View File

@ -29,6 +29,7 @@ class AiohttpOpenAIChatConfig(OpenAILikeChatConfig):
api_base: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
) -> str:
"""

View File

@ -30,6 +30,7 @@ class BaseAudioTranscriptionConfig(BaseConfig, ABC):
api_base: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
) -> str:
"""

View File

@ -31,6 +31,7 @@ class BaseTextCompletionConfig(BaseConfig, ABC):
api_base: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
) -> str:
"""

View File

@ -45,6 +45,7 @@ class BaseEmbeddingConfig(BaseConfig, ABC):
api_base: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
) -> str:
"""

View File

@ -36,6 +36,7 @@ class BaseImageVariationConfig(BaseConfig, ABC):
api_base: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
) -> str:
"""

View File

@ -76,6 +76,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
api_base: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
) -> str:
"""

View File

@ -79,6 +79,7 @@ class CloudflareChatConfig(BaseConfig):
api_base: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
) -> str:
if api_base is None:

View File

@ -234,6 +234,7 @@ class BaseLLMAIOHTTPHandler:
api_base=api_base,
model=model,
optional_params=optional_params,
litellm_params=litellm_params,
stream=stream,
)
@ -483,6 +484,7 @@ class BaseLLMAIOHTTPHandler:
api_base=api_base,
model=model,
optional_params=optional_params,
litellm_params=litellm_params,
stream=False,
)

View File

@ -605,6 +605,7 @@ class BaseLLMHTTPHandler:
api_base=api_base,
model=model,
optional_params=optional_params,
litellm_params=litellm_params,
)
data = provider_config.transform_embedding_request(
@ -900,6 +901,7 @@ class BaseLLMHTTPHandler:
client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None,
atranscription: bool = False,
headers: dict = {},
litellm_params: dict = {},
) -> TranscriptionResponse:
provider_config = ProviderConfigManager.get_provider_audio_transcription_config(
model=model, provider=litellm.LlmProviders(custom_llm_provider)
@ -923,6 +925,7 @@ class BaseLLMHTTPHandler:
api_base=api_base,
model=model,
optional_params=optional_params,
litellm_params=litellm_params,
)
# Handle the audio file based on type

View File

@ -103,6 +103,7 @@ class DeepgramAudioTranscriptionConfig(BaseAudioTranscriptionConfig):
api_base: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
) -> str:
if api_base is None:

View File

@ -40,6 +40,7 @@ class DeepSeekChatConfig(OpenAIGPTConfig):
api_base: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
) -> str:
"""

View File

@ -356,6 +356,7 @@ class OllamaConfig(BaseConfig):
api_base: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
) -> str:
"""

View File

@ -291,6 +291,7 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
api_base: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
) -> str:
"""

View File

@ -230,7 +230,7 @@ class OpenAILikeChatHandler(OpenAILikeBase):
logging_obj,
optional_params: dict,
acompletion=None,
litellm_params=None,
litellm_params: dict = {},
logger_fn=None,
headers: Optional[dict] = None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
@ -337,7 +337,7 @@ class OpenAILikeChatHandler(OpenAILikeBase):
timeout=timeout,
base_model=base_model,
client=client,
json_mode=json_mode
json_mode=json_mode,
)
else:
## COMPLETION CALL

View File

@ -169,7 +169,10 @@ def completion(
) # for pricing this must remain right before calling api
prediction_url = replicate_config.get_complete_url(
api_base=api_base, model=model, optional_params=optional_params
api_base=api_base,
model=model,
optional_params=optional_params,
litellm_params=litellm_params,
)
## COMPLETION CALL
@ -243,7 +246,10 @@ async def async_completion(
) -> Union[ModelResponse, CustomStreamWrapper]:
prediction_url = replicate_config.get_complete_url(
api_base=api_base, model=model, optional_params=optional_params
api_base=api_base,
model=model,
optional_params=optional_params,
litellm_params=litellm_params,
)
async_handler = get_async_httpx_client(
llm_provider=litellm.LlmProviders.REPLICATE,

View File

@ -141,6 +141,7 @@ class ReplicateConfig(BaseConfig):
api_base: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
) -> str:
version_id = self.model_to_version_id(model)

View File

@ -55,6 +55,7 @@ class TopazImageVariationConfig(BaseImageVariationConfig):
api_base: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
) -> str:
api_base = api_base or "https://api.topazlabs.com"

View File

@ -72,6 +72,7 @@ class TritonConfig(BaseConfig):
api_base: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
) -> str:
if api_base is None:

View File

@ -43,6 +43,7 @@ class VoyageEmbeddingConfig(BaseEmbeddingConfig):
api_base: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
) -> str:
if api_base:

View File

@ -31,7 +31,7 @@ class WatsonXChatHandler(OpenAILikeChatHandler):
logging_obj,
optional_params: dict,
acompletion=None,
litellm_params=None,
litellm_params: dict = {},
headers: Optional[dict] = None,
logger_fn=None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
@ -63,6 +63,7 @@ class WatsonXChatHandler(OpenAILikeChatHandler):
api_base=api_base,
model=model,
optional_params=optional_params,
litellm_params=litellm_params,
stream=optional_params.get("stream", False),
)

View File

@ -83,6 +83,7 @@ class IBMWatsonXChatConfig(IBMWatsonXMixin, OpenAIGPTConfig):
api_base: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
) -> str:
url = self._get_base_url(api_base=api_base)

View File

@ -318,6 +318,7 @@ class IBMWatsonXAIConfig(IBMWatsonXMixin, BaseConfig):
api_base: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
) -> str:
url = self._get_base_url(api_base=api_base)

View File

@ -54,6 +54,7 @@ class IBMWatsonXEmbeddingConfig(IBMWatsonXMixin, BaseEmbeddingConfig):
api_base: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
) -> str:
url = self._get_base_url(api_base=api_base)