mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Merge pull request #9341 from BerriAI/litellm_fix_ssl_verify
[Bug Fix] - Azure OpenAI - ensure SSL verification runs
This commit is contained in:
commit
23a09f1359
14 changed files with 914 additions and 306 deletions
|
|
@ -43,6 +43,7 @@ class AzureAssistantsAPI(BaseAzureLLM):
|
|||
api_base=api_base,
|
||||
model_name="",
|
||||
api_version=api_version,
|
||||
is_async=False,
|
||||
)
|
||||
azure_openai_client = AzureOpenAI(**azure_client_params) # type: ignore
|
||||
else:
|
||||
|
|
@ -68,6 +69,7 @@ class AzureAssistantsAPI(BaseAzureLLM):
|
|||
api_base=api_base,
|
||||
model_name="",
|
||||
api_version=api_version,
|
||||
is_async=True,
|
||||
)
|
||||
|
||||
azure_openai_client = AsyncAzureOpenAI(**azure_client_params)
|
||||
|
|
|
|||
|
|
@ -1,10 +1,9 @@
|
|||
import uuid
|
||||
from typing import Any, Optional
|
||||
from typing import Any, Coroutine, Optional, Union
|
||||
|
||||
from openai import AsyncAzureOpenAI, AzureOpenAI
|
||||
from pydantic import BaseModel
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.audio_utils.utils import get_audio_file_name
|
||||
from litellm.types.utils import FileTypes
|
||||
from litellm.utils import (
|
||||
|
|
@ -14,6 +13,7 @@ from litellm.utils import (
|
|||
)
|
||||
|
||||
from .azure import AzureChatCompletion
|
||||
from .common_utils import AzureOpenAIError
|
||||
|
||||
|
||||
class AzureAudioTranscription(AzureChatCompletion):
|
||||
|
|
@ -33,20 +33,11 @@ class AzureAudioTranscription(AzureChatCompletion):
|
|||
azure_ad_token: Optional[str] = None,
|
||||
atranscription: bool = False,
|
||||
litellm_params: Optional[dict] = None,
|
||||
) -> TranscriptionResponse:
|
||||
) -> Union[TranscriptionResponse, Coroutine[Any, Any, TranscriptionResponse]]:
|
||||
data = {"model": model, "file": audio_file, **optional_params}
|
||||
|
||||
# init AzureOpenAI Client
|
||||
azure_client_params = self.initialize_azure_sdk_client(
|
||||
litellm_params=litellm_params or {},
|
||||
api_key=api_key,
|
||||
model_name=model,
|
||||
api_version=api_version,
|
||||
api_base=api_base,
|
||||
)
|
||||
|
||||
if atranscription is True:
|
||||
return self.async_audio_transcriptions( # type: ignore
|
||||
return self.async_audio_transcriptions(
|
||||
audio_file=audio_file,
|
||||
data=data,
|
||||
model_response=model_response,
|
||||
|
|
@ -54,14 +45,26 @@ class AzureAudioTranscription(AzureChatCompletion):
|
|||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
client=client,
|
||||
azure_client_params=azure_client_params,
|
||||
max_retries=max_retries,
|
||||
logging_obj=logging_obj,
|
||||
model=model,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
azure_client = self.get_azure_openai_client(
|
||||
api_version=api_version,
|
||||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
model=model,
|
||||
_is_async=False,
|
||||
client=client,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
if not isinstance(azure_client, AzureOpenAI):
|
||||
raise AzureOpenAIError(
|
||||
status_code=500,
|
||||
message="azure_client is not an instance of AzureOpenAI",
|
||||
)
|
||||
if client is None:
|
||||
azure_client = AzureOpenAI(http_client=litellm.client_session, **azure_client_params) # type: ignore
|
||||
else:
|
||||
azure_client = client
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
|
|
@ -98,24 +101,34 @@ class AzureAudioTranscription(AzureChatCompletion):
|
|||
async def async_audio_transcriptions(
|
||||
self,
|
||||
audio_file: FileTypes,
|
||||
model: str,
|
||||
data: dict,
|
||||
model_response: TranscriptionResponse,
|
||||
timeout: float,
|
||||
azure_client_params: dict,
|
||||
logging_obj: Any,
|
||||
api_version: Optional[str] = None,
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
client=None,
|
||||
max_retries=None,
|
||||
):
|
||||
litellm_params: Optional[dict] = None,
|
||||
) -> TranscriptionResponse:
|
||||
response = None
|
||||
try:
|
||||
if client is None:
|
||||
async_azure_client = AsyncAzureOpenAI(
|
||||
**azure_client_params,
|
||||
async_azure_client = self.get_azure_openai_client(
|
||||
api_version=api_version,
|
||||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
model=model,
|
||||
_is_async=True,
|
||||
client=client,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
if not isinstance(async_azure_client, AsyncAzureOpenAI):
|
||||
raise AzureOpenAIError(
|
||||
status_code=500,
|
||||
message="async_azure_client is not an instance of AsyncAzureOpenAI",
|
||||
)
|
||||
else:
|
||||
async_azure_client = client
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
|
|
@ -168,7 +181,12 @@ class AzureAudioTranscription(AzureChatCompletion):
|
|||
model_response_object=model_response,
|
||||
hidden_params=hidden_params,
|
||||
response_type="audio_transcription",
|
||||
) # type: ignore
|
||||
)
|
||||
if not isinstance(response, TranscriptionResponse):
|
||||
raise AzureOpenAIError(
|
||||
status_code=500,
|
||||
message="response is not an instance of TranscriptionResponse",
|
||||
)
|
||||
return response
|
||||
except Exception as e:
|
||||
## LOGGING
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
import asyncio
|
||||
import json
|
||||
import time
|
||||
from typing import Any, Callable, Dict, List, Literal, Optional, Union
|
||||
from typing import Any, Callable, Coroutine, Dict, List, Optional, Union
|
||||
|
||||
import httpx # type: ignore
|
||||
from openai import APITimeoutError, AsyncAzureOpenAI, AzureOpenAI
|
||||
|
|
@ -141,41 +141,6 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
|
||||
return headers
|
||||
|
||||
def _get_sync_azure_client(
|
||||
self,
|
||||
api_version: Optional[str],
|
||||
api_base: Optional[str],
|
||||
api_key: Optional[str],
|
||||
azure_ad_token: Optional[str],
|
||||
azure_ad_token_provider: Optional[Callable],
|
||||
model: str,
|
||||
max_retries: int,
|
||||
timeout: Union[float, httpx.Timeout],
|
||||
client: Optional[Any],
|
||||
client_type: Literal["sync", "async"],
|
||||
litellm_params: Optional[dict] = None,
|
||||
):
|
||||
# init AzureOpenAI Client
|
||||
azure_client_params: Dict[str, Any] = self.initialize_azure_sdk_client(
|
||||
litellm_params=litellm_params or {},
|
||||
api_key=api_key,
|
||||
model_name=model,
|
||||
api_version=api_version,
|
||||
api_base=api_base,
|
||||
)
|
||||
if client is None:
|
||||
if client_type == "sync":
|
||||
azure_client = AzureOpenAI(**azure_client_params) # type: ignore
|
||||
elif client_type == "async":
|
||||
azure_client = AsyncAzureOpenAI(**azure_client_params) # type: ignore
|
||||
else:
|
||||
azure_client = client
|
||||
if api_version is not None and isinstance(azure_client._custom_query, dict):
|
||||
# set api_version to version passed by user
|
||||
azure_client._custom_query.setdefault("api-version", api_version)
|
||||
|
||||
return azure_client
|
||||
|
||||
def make_sync_azure_openai_chat_completion_request(
|
||||
self,
|
||||
azure_client: AzureOpenAI,
|
||||
|
|
@ -263,47 +228,21 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
max_retries = DEFAULT_MAX_RETRIES
|
||||
json_mode: Optional[bool] = optional_params.pop("json_mode", False)
|
||||
|
||||
azure_client_params = self.initialize_azure_sdk_client(
|
||||
litellm_params=litellm_params or {},
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
model_name=model,
|
||||
api_version=api_version,
|
||||
)
|
||||
### CHECK IF CLOUDFLARE AI GATEWAY ###
|
||||
### if so - set the model as part of the base url
|
||||
if "gateway.ai.cloudflare.com" in api_base:
|
||||
## build base url - assume api base includes resource name
|
||||
if client is None:
|
||||
if not api_base.endswith("/"):
|
||||
api_base += "/"
|
||||
api_base += f"{model}"
|
||||
|
||||
azure_client_params = {
|
||||
"api_version": api_version,
|
||||
"base_url": f"{api_base}",
|
||||
"http_client": litellm.client_session,
|
||||
"max_retries": max_retries,
|
||||
"timeout": timeout,
|
||||
}
|
||||
if api_key is not None:
|
||||
azure_client_params["api_key"] = api_key
|
||||
elif azure_ad_token is not None:
|
||||
if azure_ad_token.startswith("oidc/"):
|
||||
azure_ad_token = get_azure_ad_token_from_oidc(
|
||||
azure_ad_token
|
||||
)
|
||||
|
||||
azure_client_params["azure_ad_token"] = azure_ad_token
|
||||
elif azure_ad_token_provider is not None:
|
||||
azure_client_params["azure_ad_token_provider"] = (
|
||||
azure_ad_token_provider
|
||||
)
|
||||
|
||||
if acompletion is True:
|
||||
client = AsyncAzureOpenAI(**azure_client_params)
|
||||
else:
|
||||
client = AzureOpenAI(**azure_client_params)
|
||||
client = self._init_azure_client_for_cloudflare_ai_gateway(
|
||||
api_base=api_base,
|
||||
model=model,
|
||||
api_version=api_version,
|
||||
max_retries=max_retries,
|
||||
timeout=timeout,
|
||||
api_key=api_key,
|
||||
azure_ad_token=azure_ad_token,
|
||||
azure_ad_token_provider=azure_ad_token_provider,
|
||||
acompletion=acompletion,
|
||||
client=client,
|
||||
)
|
||||
|
||||
data = {"model": None, "messages": messages, **optional_params}
|
||||
else:
|
||||
|
|
@ -330,7 +269,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
timeout=timeout,
|
||||
client=client,
|
||||
max_retries=max_retries,
|
||||
azure_client_params=azure_client_params,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
else:
|
||||
return self.acompletion(
|
||||
|
|
@ -348,7 +287,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
logging_obj=logging_obj,
|
||||
max_retries=max_retries,
|
||||
convert_tool_call_to_json_mode=json_mode,
|
||||
azure_client_params=azure_client_params,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
elif "stream" in optional_params and optional_params["stream"] is True:
|
||||
return self.streaming(
|
||||
|
|
@ -364,6 +303,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
timeout=timeout,
|
||||
client=client,
|
||||
max_retries=max_retries,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
else:
|
||||
## LOGGING
|
||||
|
|
@ -385,21 +325,15 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
status_code=422, message="max retries must be an int"
|
||||
)
|
||||
# init AzureOpenAI Client
|
||||
if (
|
||||
client is None
|
||||
or not isinstance(client, AzureOpenAI)
|
||||
or dynamic_params
|
||||
):
|
||||
azure_client = AzureOpenAI(**azure_client_params)
|
||||
else:
|
||||
azure_client = client
|
||||
if api_version is not None and isinstance(
|
||||
azure_client._custom_query, dict
|
||||
):
|
||||
# set api_version to version passed by user
|
||||
azure_client._custom_query.setdefault(
|
||||
"api-version", api_version
|
||||
)
|
||||
azure_client = self.get_azure_openai_client(
|
||||
api_version=api_version,
|
||||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
model=model,
|
||||
client=client,
|
||||
_is_async=False,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
if not isinstance(azure_client, AzureOpenAI):
|
||||
raise AzureOpenAIError(
|
||||
status_code=500,
|
||||
|
|
@ -459,16 +393,22 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
azure_ad_token_provider: Optional[Callable] = None,
|
||||
convert_tool_call_to_json_mode: Optional[bool] = None,
|
||||
client=None, # this is the AsyncAzureOpenAI
|
||||
azure_client_params: dict = {},
|
||||
litellm_params: Optional[dict] = {},
|
||||
):
|
||||
response = None
|
||||
try:
|
||||
# setting Azure client
|
||||
if client is None or dynamic_params:
|
||||
azure_client = AsyncAzureOpenAI(**azure_client_params)
|
||||
else:
|
||||
azure_client = client
|
||||
|
||||
azure_client = self.get_azure_openai_client(
|
||||
api_version=api_version,
|
||||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
model=model,
|
||||
client=client,
|
||||
_is_async=True,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
if not isinstance(azure_client, AsyncAzureOpenAI):
|
||||
raise ValueError("Azure client is not an instance of AsyncAzureOpenAI")
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
input=data["messages"],
|
||||
|
|
@ -554,6 +494,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
azure_ad_token: Optional[str] = None,
|
||||
azure_ad_token_provider: Optional[Callable] = None,
|
||||
client=None,
|
||||
litellm_params: Optional[dict] = {},
|
||||
):
|
||||
# init AzureOpenAI Client
|
||||
azure_client_params = {
|
||||
|
|
@ -576,10 +517,20 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
elif azure_ad_token_provider is not None:
|
||||
azure_client_params["azure_ad_token_provider"] = azure_ad_token_provider
|
||||
|
||||
if client is None or dynamic_params:
|
||||
azure_client = AzureOpenAI(**azure_client_params)
|
||||
else:
|
||||
azure_client = client
|
||||
azure_client = self.get_azure_openai_client(
|
||||
api_version=api_version,
|
||||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
model=model,
|
||||
client=client,
|
||||
_is_async=False,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
if not isinstance(azure_client, AzureOpenAI):
|
||||
raise AzureOpenAIError(
|
||||
status_code=500,
|
||||
message="azure_client is not an instance of AzureOpenAI",
|
||||
)
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
input=data["messages"],
|
||||
|
|
@ -621,13 +572,21 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
azure_ad_token: Optional[str] = None,
|
||||
azure_ad_token_provider: Optional[Callable] = None,
|
||||
client=None,
|
||||
azure_client_params: dict = {},
|
||||
litellm_params: Optional[dict] = {},
|
||||
):
|
||||
try:
|
||||
if client is None or dynamic_params:
|
||||
azure_client = AsyncAzureOpenAI(**azure_client_params)
|
||||
else:
|
||||
azure_client = client
|
||||
azure_client = self.get_azure_openai_client(
|
||||
api_version=api_version,
|
||||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
model=model,
|
||||
client=client,
|
||||
_is_async=True,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
if not isinstance(azure_client, AsyncAzureOpenAI):
|
||||
raise ValueError("Azure client is not an instance of AsyncAzureOpenAI")
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
input=data["messages"],
|
||||
|
|
@ -678,22 +637,36 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
|
||||
async def aembedding(
|
||||
self,
|
||||
model: str,
|
||||
data: dict,
|
||||
model_response: EmbeddingResponse,
|
||||
azure_client_params: dict,
|
||||
input: list,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
api_base: str,
|
||||
api_key: Optional[str] = None,
|
||||
api_version: Optional[str] = None,
|
||||
client: Optional[AsyncAzureOpenAI] = None,
|
||||
timeout=None,
|
||||
):
|
||||
timeout: Optional[Union[float, httpx.Timeout]] = None,
|
||||
max_retries: Optional[int] = None,
|
||||
azure_ad_token: Optional[str] = None,
|
||||
azure_ad_token_provider: Optional[Callable] = None,
|
||||
litellm_params: Optional[dict] = {},
|
||||
) -> EmbeddingResponse:
|
||||
response = None
|
||||
try:
|
||||
|
||||
if client is None:
|
||||
openai_aclient = AsyncAzureOpenAI(**azure_client_params)
|
||||
else:
|
||||
openai_aclient = client
|
||||
openai_aclient = self.get_azure_openai_client(
|
||||
api_version=api_version,
|
||||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
model=model,
|
||||
_is_async=True,
|
||||
client=client,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
if not isinstance(openai_aclient, AsyncAzureOpenAI):
|
||||
raise ValueError("Azure client is not an instance of AsyncAzureOpenAI")
|
||||
|
||||
raw_response = await openai_aclient.embeddings.with_raw_response.create(
|
||||
**data, timeout=timeout
|
||||
)
|
||||
|
|
@ -707,13 +680,19 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
additional_args={"complete_input_dict": data},
|
||||
original_response=stringified_response,
|
||||
)
|
||||
return convert_to_model_response_object(
|
||||
embedding_response = convert_to_model_response_object(
|
||||
response_object=stringified_response,
|
||||
model_response_object=model_response,
|
||||
hidden_params={"headers": headers},
|
||||
_response_headers=process_azure_headers(headers),
|
||||
response_type="embedding",
|
||||
)
|
||||
if not isinstance(embedding_response, EmbeddingResponse):
|
||||
raise AzureOpenAIError(
|
||||
status_code=500,
|
||||
message="embedding_response is not an instance of EmbeddingResponse",
|
||||
)
|
||||
return embedding_response
|
||||
except Exception as e:
|
||||
## LOGGING
|
||||
logging_obj.post_call(
|
||||
|
|
@ -742,7 +721,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
aembedding=None,
|
||||
headers: Optional[dict] = None,
|
||||
litellm_params: Optional[dict] = None,
|
||||
) -> EmbeddingResponse:
|
||||
) -> Union[EmbeddingResponse, Coroutine[Any, Any, EmbeddingResponse]]:
|
||||
if headers:
|
||||
optional_params["extra_headers"] = headers
|
||||
if self._client_session is None:
|
||||
|
|
@ -751,20 +730,6 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
data = {"model": model, "input": input, **optional_params}
|
||||
if max_retries is None:
|
||||
max_retries = litellm.DEFAULT_MAX_RETRIES
|
||||
if not isinstance(max_retries, int):
|
||||
raise AzureOpenAIError(
|
||||
status_code=422, message="max retries must be an int"
|
||||
)
|
||||
|
||||
# init AzureOpenAI Client
|
||||
|
||||
azure_client_params = self.initialize_azure_sdk_client(
|
||||
litellm_params=litellm_params or {},
|
||||
api_key=api_key,
|
||||
model_name=model,
|
||||
api_version=api_version,
|
||||
api_base=api_base,
|
||||
)
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
input=input,
|
||||
|
|
@ -776,20 +741,33 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
)
|
||||
|
||||
if aembedding is True:
|
||||
return self.aembedding( # type: ignore
|
||||
return self.aembedding(
|
||||
data=data,
|
||||
input=input,
|
||||
model=model,
|
||||
logging_obj=logging_obj,
|
||||
api_key=api_key,
|
||||
model_response=model_response,
|
||||
azure_client_params=azure_client_params,
|
||||
timeout=timeout,
|
||||
client=client,
|
||||
litellm_params=litellm_params,
|
||||
api_base=api_base,
|
||||
)
|
||||
if client is None:
|
||||
azure_client = AzureOpenAI(**azure_client_params) # type: ignore
|
||||
else:
|
||||
azure_client = client
|
||||
azure_client = self.get_azure_openai_client(
|
||||
api_version=api_version,
|
||||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
model=model,
|
||||
_is_async=False,
|
||||
client=client,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
if not isinstance(azure_client, AzureOpenAI):
|
||||
raise AzureOpenAIError(
|
||||
status_code=500,
|
||||
message="azure_client is not an instance of AzureOpenAI",
|
||||
)
|
||||
|
||||
## COMPLETION CALL
|
||||
raw_response = azure_client.embeddings.with_raw_response.create(**data, timeout=timeout) # type: ignore
|
||||
headers = dict(raw_response.headers)
|
||||
|
|
@ -1155,6 +1133,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
model_name=model or "",
|
||||
api_version=api_version,
|
||||
api_base=api_base,
|
||||
is_async=False,
|
||||
)
|
||||
if aimg_generation is True:
|
||||
return self.aimage_generation(data=data, input=input, logging_obj=logging_obj, model_response=model_response, api_key=api_key, client=client, azure_client_params=azure_client_params, timeout=timeout, headers=headers) # type: ignore
|
||||
|
|
@ -1240,17 +1219,13 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
litellm_params=litellm_params,
|
||||
) # type: ignore
|
||||
|
||||
azure_client: AzureOpenAI = self._get_sync_azure_client(
|
||||
azure_client: AzureOpenAI = self.get_azure_openai_client(
|
||||
api_base=api_base,
|
||||
api_version=api_version,
|
||||
api_key=api_key,
|
||||
azure_ad_token=azure_ad_token,
|
||||
azure_ad_token_provider=azure_ad_token_provider,
|
||||
model=model,
|
||||
max_retries=max_retries,
|
||||
timeout=timeout,
|
||||
_is_async=False,
|
||||
client=client,
|
||||
client_type="sync",
|
||||
litellm_params=litellm_params,
|
||||
) # type: ignore
|
||||
|
||||
|
|
@ -1279,17 +1254,13 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
litellm_params: Optional[dict] = None,
|
||||
) -> HttpxBinaryResponseContent:
|
||||
|
||||
azure_client: AsyncAzureOpenAI = self._get_sync_azure_client(
|
||||
azure_client: AsyncAzureOpenAI = self.get_azure_openai_client(
|
||||
api_base=api_base,
|
||||
api_version=api_version,
|
||||
api_key=api_key,
|
||||
azure_ad_token=azure_ad_token,
|
||||
azure_ad_token_provider=azure_ad_token_provider,
|
||||
model=model,
|
||||
max_retries=max_retries,
|
||||
timeout=timeout,
|
||||
_is_async=True,
|
||||
client=client,
|
||||
client_type="async",
|
||||
litellm_params=litellm_params,
|
||||
) # type: ignore
|
||||
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
import json
|
||||
import os
|
||||
from typing import Callable, Optional, Union
|
||||
from typing import Any, Callable, Dict, Optional, Union
|
||||
|
||||
import httpx
|
||||
from openai import AsyncAzureOpenAI, AzureOpenAI
|
||||
|
|
@ -9,6 +9,7 @@ import litellm
|
|||
from litellm._logging import verbose_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.llms.openai.common_utils import BaseOpenAILLM
|
||||
from litellm.secret_managers.get_azure_ad_token_provider import (
|
||||
get_azure_ad_token_provider,
|
||||
)
|
||||
|
|
@ -244,24 +245,37 @@ def select_azure_base_url_or_endpoint(azure_client_params: dict):
|
|||
return azure_client_params
|
||||
|
||||
|
||||
class BaseAzureLLM:
|
||||
class BaseAzureLLM(BaseOpenAILLM):
|
||||
def get_azure_openai_client(
|
||||
self,
|
||||
litellm_params: dict,
|
||||
api_key: Optional[str],
|
||||
api_base: Optional[str],
|
||||
api_version: Optional[str] = None,
|
||||
client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI]] = None,
|
||||
litellm_params: Optional[dict] = None,
|
||||
_is_async: bool = False,
|
||||
model: Optional[str] = None,
|
||||
) -> Optional[Union[AzureOpenAI, AsyncAzureOpenAI]]:
|
||||
openai_client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI]] = None
|
||||
client_initialization_params: dict = locals()
|
||||
if client is None:
|
||||
cached_client = self.get_cached_openai_client(
|
||||
client_initialization_params=client_initialization_params,
|
||||
client_type="azure",
|
||||
)
|
||||
if cached_client:
|
||||
if isinstance(cached_client, AzureOpenAI) or isinstance(
|
||||
cached_client, AsyncAzureOpenAI
|
||||
):
|
||||
return cached_client
|
||||
|
||||
azure_client_params = self.initialize_azure_sdk_client(
|
||||
litellm_params=litellm_params,
|
||||
litellm_params=litellm_params or {},
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
model_name="",
|
||||
model_name=model,
|
||||
api_version=api_version,
|
||||
is_async=_is_async,
|
||||
)
|
||||
if _is_async is True:
|
||||
openai_client = AsyncAzureOpenAI(**azure_client_params)
|
||||
|
|
@ -269,7 +283,18 @@ class BaseAzureLLM:
|
|||
openai_client = AzureOpenAI(**azure_client_params) # type: ignore
|
||||
else:
|
||||
openai_client = client
|
||||
if api_version is not None and isinstance(
|
||||
openai_client._custom_query, dict
|
||||
):
|
||||
# set api_version to version passed by user
|
||||
openai_client._custom_query.setdefault("api-version", api_version)
|
||||
|
||||
# save client in-memory cache
|
||||
self.set_cached_openai_client(
|
||||
openai_client=openai_client,
|
||||
client_initialization_params=client_initialization_params,
|
||||
client_type="azure",
|
||||
)
|
||||
return openai_client
|
||||
|
||||
def initialize_azure_sdk_client(
|
||||
|
|
@ -277,8 +302,9 @@ class BaseAzureLLM:
|
|||
litellm_params: dict,
|
||||
api_key: Optional[str],
|
||||
api_base: Optional[str],
|
||||
model_name: str,
|
||||
model_name: Optional[str],
|
||||
api_version: Optional[str],
|
||||
is_async: bool,
|
||||
) -> dict:
|
||||
|
||||
azure_ad_token_provider: Optional[Callable[[], str]] = None
|
||||
|
|
@ -334,8 +360,13 @@ class BaseAzureLLM:
|
|||
"api_version": api_version,
|
||||
"azure_ad_token": azure_ad_token,
|
||||
"azure_ad_token_provider": azure_ad_token_provider,
|
||||
"http_client": litellm.client_session,
|
||||
}
|
||||
# init http client + SSL Verification settings
|
||||
if is_async is True:
|
||||
azure_client_params["http_client"] = self._get_async_http_client()
|
||||
else:
|
||||
azure_client_params["http_client"] = self._get_sync_http_client()
|
||||
|
||||
if max_retries is not None:
|
||||
azure_client_params["max_retries"] = max_retries
|
||||
if timeout is not None:
|
||||
|
|
@ -351,3 +382,45 @@ class BaseAzureLLM:
|
|||
)
|
||||
|
||||
return azure_client_params
|
||||
|
||||
def _init_azure_client_for_cloudflare_ai_gateway(
|
||||
self,
|
||||
api_base: str,
|
||||
model: str,
|
||||
api_version: str,
|
||||
max_retries: int,
|
||||
timeout: Union[float, httpx.Timeout],
|
||||
api_key: Optional[str],
|
||||
azure_ad_token: Optional[str],
|
||||
azure_ad_token_provider: Optional[Callable[[], str]],
|
||||
acompletion: bool,
|
||||
client: Optional[Union[AzureOpenAI, AsyncAzureOpenAI]] = None,
|
||||
) -> Union[AzureOpenAI, AsyncAzureOpenAI]:
|
||||
## build base url - assume api base includes resource name
|
||||
if client is None:
|
||||
if not api_base.endswith("/"):
|
||||
api_base += "/"
|
||||
api_base += f"{model}"
|
||||
|
||||
azure_client_params: Dict[str, Any] = {
|
||||
"api_version": api_version,
|
||||
"base_url": f"{api_base}",
|
||||
"http_client": litellm.client_session,
|
||||
"max_retries": max_retries,
|
||||
"timeout": timeout,
|
||||
}
|
||||
if api_key is not None:
|
||||
azure_client_params["api_key"] = api_key
|
||||
elif azure_ad_token is not None:
|
||||
if azure_ad_token.startswith("oidc/"):
|
||||
azure_ad_token = get_azure_ad_token_from_oidc(azure_ad_token)
|
||||
|
||||
azure_client_params["azure_ad_token"] = azure_ad_token
|
||||
if azure_ad_token_provider is not None:
|
||||
azure_client_params["azure_ad_token_provider"] = azure_ad_token_provider
|
||||
|
||||
if acompletion is True:
|
||||
client = AsyncAzureOpenAI(**azure_client_params) # type: ignore
|
||||
else:
|
||||
client = AzureOpenAI(**azure_client_params) # type: ignore
|
||||
return client
|
||||
|
|
|
|||
|
|
@ -2,7 +2,6 @@ from typing import Any, Callable, Optional
|
|||
|
||||
from openai import AsyncAzureOpenAI, AzureOpenAI
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.prompt_templates.factory import prompt_factory
|
||||
from litellm.utils import CustomStreamWrapper, ModelResponse, TextCompletionResponse
|
||||
|
||||
|
|
@ -12,18 +11,6 @@ from ..common_utils import AzureOpenAIError, BaseAzureLLM
|
|||
openai_text_completion_config = OpenAITextCompletionConfig()
|
||||
|
||||
|
||||
def select_azure_base_url_or_endpoint(azure_client_params: dict):
|
||||
azure_endpoint = azure_client_params.get("azure_endpoint", None)
|
||||
if azure_endpoint is not None:
|
||||
# see : https://github.com/openai/openai-python/blob/3d61ed42aba652b547029095a7eb269ad4e1e957/src/openai/lib/azure.py#L192
|
||||
if "/openai/deployments" in azure_endpoint:
|
||||
# this is base_url, not an azure_endpoint
|
||||
azure_client_params["base_url"] = azure_endpoint
|
||||
azure_client_params.pop("azure_endpoint")
|
||||
|
||||
return azure_client_params
|
||||
|
||||
|
||||
class AzureTextCompletion(BaseAzureLLM):
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
|
|
@ -70,39 +57,22 @@ class AzureTextCompletion(BaseAzureLLM):
|
|||
messages=messages, model=model, custom_llm_provider="azure_text"
|
||||
)
|
||||
|
||||
azure_client_params = self.initialize_azure_sdk_client(
|
||||
litellm_params=litellm_params or {},
|
||||
api_key=api_key,
|
||||
model_name=model,
|
||||
api_version=api_version,
|
||||
api_base=api_base,
|
||||
)
|
||||
|
||||
### CHECK IF CLOUDFLARE AI GATEWAY ###
|
||||
### if so - set the model as part of the base url
|
||||
if "gateway.ai.cloudflare.com" in api_base:
|
||||
## build base url - assume api base includes resource name
|
||||
if client is None:
|
||||
if not api_base.endswith("/"):
|
||||
api_base += "/"
|
||||
api_base += f"{model}"
|
||||
|
||||
azure_client_params = {
|
||||
"api_version": api_version,
|
||||
"base_url": f"{api_base}",
|
||||
"http_client": litellm.client_session,
|
||||
"max_retries": max_retries,
|
||||
"timeout": timeout,
|
||||
}
|
||||
if api_key is not None:
|
||||
azure_client_params["api_key"] = api_key
|
||||
elif azure_ad_token is not None:
|
||||
azure_client_params["azure_ad_token"] = azure_ad_token
|
||||
|
||||
if acompletion is True:
|
||||
client = AsyncAzureOpenAI(**azure_client_params)
|
||||
else:
|
||||
client = AzureOpenAI(**azure_client_params)
|
||||
client = self._init_azure_client_for_cloudflare_ai_gateway(
|
||||
api_key=api_key,
|
||||
api_version=api_version,
|
||||
api_base=api_base,
|
||||
model=model,
|
||||
client=client,
|
||||
max_retries=max_retries,
|
||||
timeout=timeout,
|
||||
azure_ad_token=azure_ad_token,
|
||||
azure_ad_token_provider=azure_ad_token_provider,
|
||||
acompletion=acompletion,
|
||||
)
|
||||
|
||||
data = {"model": None, "prompt": prompt, **optional_params}
|
||||
else:
|
||||
|
|
@ -124,7 +94,7 @@ class AzureTextCompletion(BaseAzureLLM):
|
|||
azure_ad_token=azure_ad_token,
|
||||
timeout=timeout,
|
||||
client=client,
|
||||
azure_client_params=azure_client_params,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
else:
|
||||
return self.acompletion(
|
||||
|
|
@ -139,7 +109,7 @@ class AzureTextCompletion(BaseAzureLLM):
|
|||
client=client,
|
||||
logging_obj=logging_obj,
|
||||
max_retries=max_retries,
|
||||
azure_client_params=azure_client_params,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
elif "stream" in optional_params and optional_params["stream"] is True:
|
||||
return self.streaming(
|
||||
|
|
@ -152,7 +122,6 @@ class AzureTextCompletion(BaseAzureLLM):
|
|||
azure_ad_token=azure_ad_token,
|
||||
timeout=timeout,
|
||||
client=client,
|
||||
azure_client_params=azure_client_params,
|
||||
)
|
||||
else:
|
||||
## LOGGING
|
||||
|
|
@ -174,17 +143,21 @@ class AzureTextCompletion(BaseAzureLLM):
|
|||
status_code=422, message="max retries must be an int"
|
||||
)
|
||||
# init AzureOpenAI Client
|
||||
if client is None:
|
||||
azure_client = AzureOpenAI(**azure_client_params)
|
||||
else:
|
||||
azure_client = client
|
||||
if api_version is not None and isinstance(
|
||||
azure_client._custom_query, dict
|
||||
):
|
||||
# set api_version to version passed by user
|
||||
azure_client._custom_query.setdefault(
|
||||
"api-version", api_version
|
||||
)
|
||||
azure_client = self.get_azure_openai_client(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
api_version=api_version,
|
||||
client=client,
|
||||
litellm_params=litellm_params,
|
||||
_is_async=False,
|
||||
model=model,
|
||||
)
|
||||
|
||||
if not isinstance(azure_client, AzureOpenAI):
|
||||
raise AzureOpenAIError(
|
||||
status_code=500,
|
||||
message="azure_client is not an instance of AzureOpenAI",
|
||||
)
|
||||
|
||||
raw_response = azure_client.completions.with_raw_response.create(
|
||||
**data, timeout=timeout
|
||||
|
|
@ -233,21 +206,27 @@ class AzureTextCompletion(BaseAzureLLM):
|
|||
max_retries: int,
|
||||
azure_ad_token: Optional[str] = None,
|
||||
client=None, # this is the AsyncAzureOpenAI
|
||||
azure_client_params: dict = {},
|
||||
litellm_params: dict = {},
|
||||
):
|
||||
response = None
|
||||
try:
|
||||
# init AzureOpenAI Client
|
||||
# setting Azure client
|
||||
if client is None:
|
||||
azure_client = AsyncAzureOpenAI(**azure_client_params)
|
||||
else:
|
||||
azure_client = client
|
||||
if api_version is not None and isinstance(
|
||||
azure_client._custom_query, dict
|
||||
):
|
||||
# set api_version to version passed by user
|
||||
azure_client._custom_query.setdefault("api-version", api_version)
|
||||
azure_client = self.get_azure_openai_client(
|
||||
api_version=api_version,
|
||||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
model=model,
|
||||
_is_async=True,
|
||||
client=client,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
if not isinstance(azure_client, AsyncAzureOpenAI):
|
||||
raise AzureOpenAIError(
|
||||
status_code=500,
|
||||
message="azure_client is not an instance of AsyncAzureOpenAI",
|
||||
)
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
input=data["prompt"],
|
||||
|
|
@ -290,7 +269,7 @@ class AzureTextCompletion(BaseAzureLLM):
|
|||
timeout: Any,
|
||||
azure_ad_token: Optional[str] = None,
|
||||
client=None,
|
||||
azure_client_params: dict = {},
|
||||
litellm_params: dict = {},
|
||||
):
|
||||
max_retries = data.pop("max_retries", 2)
|
||||
if not isinstance(max_retries, int):
|
||||
|
|
@ -298,13 +277,21 @@ class AzureTextCompletion(BaseAzureLLM):
|
|||
status_code=422, message="max retries must be an int"
|
||||
)
|
||||
# init AzureOpenAI Client
|
||||
if client is None:
|
||||
azure_client = AzureOpenAI(**azure_client_params)
|
||||
else:
|
||||
azure_client = client
|
||||
if api_version is not None and isinstance(azure_client._custom_query, dict):
|
||||
# set api_version to version passed by user
|
||||
azure_client._custom_query.setdefault("api-version", api_version)
|
||||
azure_client = self.get_azure_openai_client(
|
||||
api_version=api_version,
|
||||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
model=model,
|
||||
_is_async=False,
|
||||
client=client,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
if not isinstance(azure_client, AzureOpenAI):
|
||||
raise AzureOpenAIError(
|
||||
status_code=500,
|
||||
message="azure_client is not an instance of AzureOpenAI",
|
||||
)
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
input=data["prompt"],
|
||||
|
|
@ -339,19 +326,24 @@ class AzureTextCompletion(BaseAzureLLM):
|
|||
timeout: Any,
|
||||
azure_ad_token: Optional[str] = None,
|
||||
client=None,
|
||||
azure_client_params: dict = {},
|
||||
litellm_params: dict = {},
|
||||
):
|
||||
try:
|
||||
# init AzureOpenAI Client
|
||||
if client is None:
|
||||
azure_client = AsyncAzureOpenAI(**azure_client_params)
|
||||
else:
|
||||
azure_client = client
|
||||
if api_version is not None and isinstance(
|
||||
azure_client._custom_query, dict
|
||||
):
|
||||
# set api_version to version passed by user
|
||||
azure_client._custom_query.setdefault("api-version", api_version)
|
||||
azure_client = self.get_azure_openai_client(
|
||||
api_version=api_version,
|
||||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
model=model,
|
||||
_is_async=True,
|
||||
client=client,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
if not isinstance(azure_client, AsyncAzureOpenAI):
|
||||
raise AzureOpenAIError(
|
||||
status_code=500,
|
||||
message="azure_client is not an instance of AsyncAzureOpenAI",
|
||||
)
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
input=data["prompt"],
|
||||
|
|
|
|||
|
|
@ -2,13 +2,17 @@
|
|||
Common helpers / utils across al OpenAI endpoints
|
||||
"""
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
from typing import Any, Dict, List, Literal, Optional, Union
|
||||
|
||||
import httpx
|
||||
import openai
|
||||
from openai import AsyncAzureOpenAI, AsyncOpenAI, AzureOpenAI, OpenAI
|
||||
|
||||
import litellm
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.llms.custom_httpx.http_handler import _DEFAULT_TTL_FOR_HTTPX_CLIENTS
|
||||
|
||||
|
||||
class OpenAIError(BaseLLMException):
|
||||
|
|
@ -92,3 +96,113 @@ def drop_params_from_unprocessable_entity_error(
|
|||
new_data = {k: v for k, v in data.items() if k not in invalid_params}
|
||||
|
||||
return new_data
|
||||
|
||||
|
||||
class BaseOpenAILLM:
|
||||
"""
|
||||
Base class for OpenAI LLMs for getting their httpx clients and SSL verification settings
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def get_cached_openai_client(
|
||||
client_initialization_params: dict, client_type: Literal["openai", "azure"]
|
||||
) -> Optional[Union[OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI]]:
|
||||
"""Retrieves the OpenAI client from the in-memory cache based on the client initialization parameters"""
|
||||
_cache_key = BaseOpenAILLM.get_openai_client_cache_key(
|
||||
client_initialization_params=client_initialization_params,
|
||||
client_type=client_type,
|
||||
)
|
||||
_cached_client = litellm.in_memory_llm_clients_cache.get_cache(_cache_key)
|
||||
return _cached_client
|
||||
|
||||
@staticmethod
|
||||
def set_cached_openai_client(
|
||||
openai_client: Union[OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI],
|
||||
client_type: Literal["openai", "azure"],
|
||||
client_initialization_params: dict,
|
||||
):
|
||||
"""Stores the OpenAI client in the in-memory cache for _DEFAULT_TTL_FOR_HTTPX_CLIENTS SECONDS"""
|
||||
_cache_key = BaseOpenAILLM.get_openai_client_cache_key(
|
||||
client_initialization_params=client_initialization_params,
|
||||
client_type=client_type,
|
||||
)
|
||||
litellm.in_memory_llm_clients_cache.set_cache(
|
||||
key=_cache_key,
|
||||
value=openai_client,
|
||||
ttl=_DEFAULT_TTL_FOR_HTTPX_CLIENTS,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def get_openai_client_cache_key(
|
||||
client_initialization_params: dict, client_type: Literal["openai", "azure"]
|
||||
) -> str:
|
||||
"""Creates a cache key for the OpenAI client based on the client initialization parameters"""
|
||||
hashed_api_key = None
|
||||
if client_initialization_params.get("api_key") is not None:
|
||||
hash_object = hashlib.sha256(
|
||||
client_initialization_params.get("api_key", "").encode()
|
||||
)
|
||||
# Hexadecimal representation of the hash
|
||||
hashed_api_key = hash_object.hexdigest()
|
||||
|
||||
# Create a more readable cache key using a list of key-value pairs
|
||||
key_parts = [
|
||||
f"hashed_api_key={hashed_api_key}",
|
||||
f"is_async={client_initialization_params.get('is_async')}",
|
||||
]
|
||||
|
||||
LITELLM_CLIENT_SPECIFIC_PARAMS = [
|
||||
"timeout",
|
||||
"max_retries",
|
||||
"organization",
|
||||
"api_base",
|
||||
]
|
||||
openai_client_fields = (
|
||||
BaseOpenAILLM.get_openai_client_initialization_param_fields(
|
||||
client_type=client_type
|
||||
)
|
||||
+ LITELLM_CLIENT_SPECIFIC_PARAMS
|
||||
)
|
||||
|
||||
for param in openai_client_fields:
|
||||
key_parts.append(f"{param}={client_initialization_params.get(param)}")
|
||||
|
||||
_cache_key = ",".join(key_parts)
|
||||
return _cache_key
|
||||
|
||||
@staticmethod
|
||||
def get_openai_client_initialization_param_fields(
|
||||
client_type: Literal["openai", "azure"]
|
||||
) -> List[str]:
|
||||
"""Returns a list of fields that are used to initialize the OpenAI client"""
|
||||
import inspect
|
||||
|
||||
from openai import AzureOpenAI, OpenAI
|
||||
|
||||
if client_type == "openai":
|
||||
signature = inspect.signature(OpenAI.__init__)
|
||||
else:
|
||||
signature = inspect.signature(AzureOpenAI.__init__)
|
||||
|
||||
# Extract parameter names, excluding 'self'
|
||||
param_names = [param for param in signature.parameters if param != "self"]
|
||||
return param_names
|
||||
|
||||
@staticmethod
|
||||
def _get_async_http_client() -> Optional[httpx.AsyncClient]:
|
||||
if litellm.aclient_session is not None:
|
||||
return litellm.aclient_session
|
||||
|
||||
return httpx.AsyncClient(
|
||||
limits=httpx.Limits(max_connections=1000, max_keepalive_connections=100),
|
||||
verify=litellm.ssl_verify,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _get_sync_http_client() -> Optional[httpx.Client]:
|
||||
if litellm.client_session is not None:
|
||||
return litellm.client_session
|
||||
return httpx.Client(
|
||||
limits=httpx.Limits(max_connections=1000, max_keepalive_connections=100),
|
||||
verify=litellm.ssl_verify,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -33,7 +33,6 @@ from litellm.litellm_core_utils.logging_utils import track_llm_api_timing
|
|||
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
|
||||
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
|
||||
from litellm.llms.bedrock.chat.invoke_handler import MockResponseIterator
|
||||
from litellm.llms.custom_httpx.http_handler import _DEFAULT_TTL_FOR_HTTPX_CLIENTS
|
||||
from litellm.types.utils import (
|
||||
EmbeddingResponse,
|
||||
ImageResponse,
|
||||
|
|
@ -50,7 +49,11 @@ from litellm.utils import (
|
|||
from ...types.llms.openai import *
|
||||
from ..base import BaseLLM
|
||||
from .chat.o_series_transformation import OpenAIOSeriesConfig
|
||||
from .common_utils import OpenAIError, drop_params_from_unprocessable_entity_error
|
||||
from .common_utils import (
|
||||
BaseOpenAILLM,
|
||||
OpenAIError,
|
||||
drop_params_from_unprocessable_entity_error,
|
||||
)
|
||||
|
||||
openaiOSeriesConfig = OpenAIOSeriesConfig()
|
||||
|
||||
|
|
@ -317,7 +320,7 @@ class OpenAIChatCompletionResponseIterator(BaseModelResponseIterator):
|
|||
raise e
|
||||
|
||||
|
||||
class OpenAIChatCompletion(BaseLLM):
|
||||
class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
|
|
@ -343,7 +346,8 @@ class OpenAIChatCompletion(BaseLLM):
|
|||
max_retries: Optional[int] = DEFAULT_MAX_RETRIES,
|
||||
organization: Optional[str] = None,
|
||||
client: Optional[Union[OpenAI, AsyncOpenAI]] = None,
|
||||
):
|
||||
) -> Optional[Union[OpenAI, AsyncOpenAI]]:
|
||||
client_initialization_params: Dict = locals()
|
||||
if client is None:
|
||||
if not isinstance(max_retries, int):
|
||||
raise OpenAIError(
|
||||
|
|
@ -352,25 +356,21 @@ class OpenAIChatCompletion(BaseLLM):
|
|||
max_retries
|
||||
),
|
||||
)
|
||||
# Creating a new OpenAI Client
|
||||
# check in memory cache before creating a new one
|
||||
# Convert the API key to bytes
|
||||
hashed_api_key = None
|
||||
if api_key is not None:
|
||||
hash_object = hashlib.sha256(api_key.encode())
|
||||
# Hexadecimal representation of the hash
|
||||
hashed_api_key = hash_object.hexdigest()
|
||||
cached_client = self.get_cached_openai_client(
|
||||
client_initialization_params=client_initialization_params,
|
||||
client_type="openai",
|
||||
)
|
||||
|
||||
_cache_key = f"hashed_api_key={hashed_api_key},api_base={api_base},timeout={timeout},max_retries={max_retries},organization={organization},is_async={is_async}"
|
||||
|
||||
_cached_client = litellm.in_memory_llm_clients_cache.get_cache(_cache_key)
|
||||
if _cached_client:
|
||||
return _cached_client
|
||||
if cached_client:
|
||||
if isinstance(cached_client, OpenAI) or isinstance(
|
||||
cached_client, AsyncOpenAI
|
||||
):
|
||||
return cached_client
|
||||
if is_async:
|
||||
_new_client: Union[OpenAI, AsyncOpenAI] = AsyncOpenAI(
|
||||
api_key=api_key,
|
||||
base_url=api_base,
|
||||
http_client=litellm.aclient_session,
|
||||
http_client=OpenAIChatCompletion._get_async_http_client(),
|
||||
timeout=timeout,
|
||||
max_retries=max_retries,
|
||||
organization=organization,
|
||||
|
|
@ -379,17 +379,17 @@ class OpenAIChatCompletion(BaseLLM):
|
|||
_new_client = OpenAI(
|
||||
api_key=api_key,
|
||||
base_url=api_base,
|
||||
http_client=litellm.client_session,
|
||||
http_client=OpenAIChatCompletion._get_sync_http_client(),
|
||||
timeout=timeout,
|
||||
max_retries=max_retries,
|
||||
organization=organization,
|
||||
)
|
||||
|
||||
## SAVE CACHE KEY
|
||||
litellm.in_memory_llm_clients_cache.set_cache(
|
||||
key=_cache_key,
|
||||
value=_new_client,
|
||||
ttl=_DEFAULT_TTL_FOR_HTTPX_CLIENTS,
|
||||
self.set_cached_openai_client(
|
||||
openai_client=_new_client,
|
||||
client_initialization_params=client_initialization_params,
|
||||
client_type="openai",
|
||||
)
|
||||
return _new_client
|
||||
|
||||
|
|
|
|||
|
|
@ -25,6 +25,7 @@ from functools import partial
|
|||
from typing import (
|
||||
Any,
|
||||
Callable,
|
||||
Coroutine,
|
||||
Dict,
|
||||
List,
|
||||
Literal,
|
||||
|
|
@ -3288,7 +3289,7 @@ def embedding( # noqa: PLR0915
|
|||
litellm_call_id=None,
|
||||
logger_fn=None,
|
||||
**kwargs,
|
||||
) -> EmbeddingResponse:
|
||||
) -> Union[EmbeddingResponse, Coroutine[Any, Any, EmbeddingResponse]]:
|
||||
"""
|
||||
Embedding function that calls an API to generate embeddings for the given input.
|
||||
|
||||
|
|
@ -3409,7 +3410,9 @@ def embedding( # noqa: PLR0915
|
|||
if mock_response is not None:
|
||||
return mock_embedding(model=model, mock_response=mock_response)
|
||||
try:
|
||||
response: Optional[EmbeddingResponse] = None
|
||||
response: Optional[
|
||||
Union[EmbeddingResponse, Coroutine[Any, Any, EmbeddingResponse]]
|
||||
] = None
|
||||
|
||||
if azure is True or custom_llm_provider == "azure":
|
||||
# azure configs
|
||||
|
|
@ -3901,7 +3904,11 @@ def embedding( # noqa: PLR0915
|
|||
raise LiteLLMUnknownProvider(
|
||||
model=model, custom_llm_provider=custom_llm_provider
|
||||
)
|
||||
if response is not None and hasattr(response, "_hidden_params"):
|
||||
if (
|
||||
response is not None
|
||||
and hasattr(response, "_hidden_params")
|
||||
and isinstance(response, EmbeddingResponse)
|
||||
):
|
||||
response._hidden_params["custom_llm_provider"] = custom_llm_provider
|
||||
|
||||
if response is None:
|
||||
|
|
@ -4944,6 +4951,10 @@ async def atranscription(*args, **kwargs) -> TranscriptionResponse:
|
|||
else:
|
||||
# Call the synchronous function using run_in_executor
|
||||
response = await loop.run_in_executor(None, func_with_context)
|
||||
if not isinstance(response, TranscriptionResponse):
|
||||
raise ValueError(
|
||||
f"Invalid response from transcription provider, expected TranscriptionResponse, but got {type(response)}"
|
||||
)
|
||||
return response
|
||||
except Exception as e:
|
||||
custom_llm_provider = custom_llm_provider or "openai"
|
||||
|
|
@ -4977,7 +4988,7 @@ def transcription(
|
|||
max_retries: Optional[int] = None,
|
||||
custom_llm_provider=None,
|
||||
**kwargs,
|
||||
) -> TranscriptionResponse:
|
||||
) -> Union[TranscriptionResponse, Coroutine[Any, Any, TranscriptionResponse]]:
|
||||
"""
|
||||
Calls openai + azure whisper endpoints.
|
||||
|
||||
|
|
@ -5046,7 +5057,9 @@ def transcription(
|
|||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
response: Optional[TranscriptionResponse] = None
|
||||
response: Optional[
|
||||
Union[TranscriptionResponse, Coroutine[Any, Any, TranscriptionResponse]]
|
||||
] = None
|
||||
if custom_llm_provider == "azure":
|
||||
# azure configs
|
||||
api_base = api_base or litellm.api_base or get_secret_str("AZURE_API_BASE")
|
||||
|
|
|
|||
|
|
@ -1,6 +1,9 @@
|
|||
model_list:
|
||||
- model_name: gpt-4o
|
||||
- model_name: gpt-3.5-turbo-end-user-test
|
||||
litellm_params:
|
||||
model: gpt-4o
|
||||
model: azure/chatgpt-v-2
|
||||
api_base: https://openai-gpt-4-test-v-1.openai.azure.com/
|
||||
api_version: "2023-05-15"
|
||||
api_key: os.environ/AZURE_API_KEY
|
||||
|
||||
|
||||
|
|
|
|||
108
tests/code_coverage_tests/azure_client_usage_test.py
Normal file
108
tests/code_coverage_tests/azure_client_usage_test.py
Normal file
|
|
@ -0,0 +1,108 @@
|
|||
import ast
|
||||
import os
|
||||
import re
|
||||
|
||||
|
||||
def find_azure_files(base_dir):
|
||||
"""
|
||||
Find all Python files in the Azure directory.
|
||||
"""
|
||||
azure_files = []
|
||||
for root, _, files in os.walk(base_dir):
|
||||
for file in files:
|
||||
if file.endswith(".py"):
|
||||
azure_files.append(os.path.join(root, file))
|
||||
return azure_files
|
||||
|
||||
|
||||
def check_direct_instantiation(file_path):
|
||||
"""
|
||||
Check if a file directly instantiates AzureOpenAI or AsyncAzureOpenAI
|
||||
outside of the BaseAzureLLM class methods.
|
||||
"""
|
||||
with open(file_path, "r") as file:
|
||||
content = file.read()
|
||||
|
||||
# Parse the file
|
||||
tree = ast.parse(content)
|
||||
|
||||
# Track issues found
|
||||
issues = []
|
||||
|
||||
# Find all class definitions
|
||||
for node in ast.walk(tree):
|
||||
if isinstance(node, ast.ClassDef):
|
||||
class_name = node.name
|
||||
|
||||
# Skip BaseAzureLLM class since it's allowed to define the client creation methods
|
||||
if class_name == "BaseAzureLLM":
|
||||
continue
|
||||
|
||||
# Check method bodies for direct instantiation
|
||||
for method in node.body:
|
||||
if isinstance(method, ast.FunctionDef) or isinstance(
|
||||
method, ast.AsyncFunctionDef
|
||||
):
|
||||
method_name = method.name
|
||||
|
||||
# Skip methods that are specifically for client creation
|
||||
if method_name in [
|
||||
"get_azure_openai_client",
|
||||
"initialize_azure_sdk_client",
|
||||
]:
|
||||
continue
|
||||
|
||||
# Look for direct instantiation in the method body
|
||||
for subnode in ast.walk(method):
|
||||
if isinstance(subnode, ast.Call):
|
||||
if hasattr(subnode, "func") and hasattr(subnode.func, "id"):
|
||||
if subnode.func.id in [
|
||||
"AzureOpenAI",
|
||||
"AsyncAzureOpenAI",
|
||||
]:
|
||||
issues.append(
|
||||
f"Direct instantiation of {subnode.func.id} in {class_name}.{method_name}"
|
||||
)
|
||||
elif hasattr(subnode, "func") and hasattr(
|
||||
subnode.func, "attr"
|
||||
):
|
||||
if subnode.func.attr in [
|
||||
"AzureOpenAI",
|
||||
"AsyncAzureOpenAI",
|
||||
]:
|
||||
issues.append(
|
||||
f"Direct instantiation of {subnode.func.attr} in {class_name}.{method_name}"
|
||||
)
|
||||
|
||||
return issues
|
||||
|
||||
|
||||
def main():
|
||||
"""
|
||||
Main function to run the test.
|
||||
"""
|
||||
# local
|
||||
base_dir = "../../litellm/llms/azure"
|
||||
azure_files = find_azure_files(base_dir)
|
||||
print(f"Found {len(azure_files)} Azure Python files to check")
|
||||
|
||||
all_issues = []
|
||||
|
||||
for file_path in azure_files:
|
||||
issues = check_direct_instantiation(file_path)
|
||||
if issues:
|
||||
all_issues.extend([f"{file_path}: {issue}" for issue in issues])
|
||||
|
||||
if all_issues:
|
||||
print("Found direct instantiations of AzureOpenAI or AsyncAzureOpenAI:")
|
||||
for issue in all_issues:
|
||||
print(f" - {issue}")
|
||||
raise Exception(
|
||||
f"Found {len(all_issues)} direct instantiations of AzureOpenAI or AsyncAzureOpenAI classes. Use get_azure_openai_client instead."
|
||||
)
|
||||
else:
|
||||
print("All Azure modules are correctly using get_azure_openai_client!")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
|
@ -10,6 +10,7 @@ ALLOWED_FILES = [
|
|||
"../../litellm/llms/huggingface_restapi.py",
|
||||
"../../litellm/llms/base.py",
|
||||
"../../litellm/llms/custom_httpx/httpx_handler.py",
|
||||
"../../litellm/llms/openai/common_utils.py",
|
||||
# when running on ci/cd
|
||||
"./litellm/__init__.py",
|
||||
"./litellm/llms/custom_httpx/http_handler.py",
|
||||
|
|
@ -18,6 +19,7 @@ ALLOWED_FILES = [
|
|||
"./litellm/llms/huggingface_restapi.py",
|
||||
"./litellm/llms/base.py",
|
||||
"./litellm/llms/custom_httpx/httpx_handler.py",
|
||||
"./litellm/llms/openai/common_utils.py",
|
||||
]
|
||||
|
||||
warning_msg = "this is a serious violation that can impact latency. Creating Async clients per request can add +500ms per request"
|
||||
|
|
|
|||
|
|
@ -68,6 +68,7 @@ def test_initialize_with_api_key(setup_mocks):
|
|||
api_base="https://test.openai.azure.com",
|
||||
model_name="gpt-4",
|
||||
api_version="2023-06-01",
|
||||
is_async=False,
|
||||
)
|
||||
|
||||
# Verify expected result
|
||||
|
|
@ -90,6 +91,7 @@ def test_initialize_with_tenant_credentials(setup_mocks):
|
|||
api_base="https://test.openai.azure.com",
|
||||
model_name="gpt-4",
|
||||
api_version=None,
|
||||
is_async=False,
|
||||
)
|
||||
|
||||
# Verify that get_azure_ad_token_from_entrata_id was called
|
||||
|
|
@ -117,6 +119,7 @@ def test_initialize_with_username_password(setup_mocks):
|
|||
api_base="https://test.openai.azure.com",
|
||||
model_name="gpt-4",
|
||||
api_version=None,
|
||||
is_async=False,
|
||||
)
|
||||
|
||||
# Verify that get_azure_ad_token_from_username_password was called
|
||||
|
|
@ -138,6 +141,7 @@ def test_initialize_with_oidc_token(setup_mocks):
|
|||
api_base="https://test.openai.azure.com",
|
||||
model_name="gpt-4",
|
||||
api_version=None,
|
||||
is_async=False,
|
||||
)
|
||||
|
||||
# Verify that get_azure_ad_token_from_oidc was called
|
||||
|
|
@ -158,6 +162,7 @@ def test_initialize_with_enable_token_refresh(setup_mocks):
|
|||
api_base="https://test.openai.azure.com",
|
||||
model_name="gpt-4",
|
||||
api_version=None,
|
||||
is_async=False,
|
||||
)
|
||||
|
||||
# Verify that get_azure_ad_token_provider was called
|
||||
|
|
@ -179,6 +184,7 @@ def test_initialize_with_token_refresh_error(setup_mocks):
|
|||
api_base="https://test.openai.azure.com",
|
||||
model_name="gpt-4",
|
||||
api_version=None,
|
||||
is_async=False,
|
||||
)
|
||||
|
||||
# Verify error was logged
|
||||
|
|
@ -196,6 +202,7 @@ def test_api_version_from_env_var(setup_mocks):
|
|||
api_base="https://test.openai.azure.com",
|
||||
model_name="gpt-4",
|
||||
api_version=None,
|
||||
is_async=False,
|
||||
)
|
||||
|
||||
# Verify expected result
|
||||
|
|
@ -210,6 +217,7 @@ def test_select_azure_base_url_called(setup_mocks):
|
|||
api_base="https://test.openai.azure.com",
|
||||
model_name="gpt-4",
|
||||
api_version="2023-06-01",
|
||||
is_async=False,
|
||||
)
|
||||
|
||||
# Verify that select_azure_base_url_or_endpoint was called
|
||||
|
|
@ -300,12 +308,10 @@ async def test_ensure_initialize_azure_sdk_client_always_used(call_type):
|
|||
# Get appropriate input for this call type
|
||||
input_kwarg = test_inputs.get(call_type.value, {})
|
||||
|
||||
patch_target = "litellm.main.azure_chat_completions.initialize_azure_sdk_client"
|
||||
if call_type == CallTypes.atranscription:
|
||||
patch_target = (
|
||||
"litellm.main.azure_audio_transcriptions.initialize_azure_sdk_client"
|
||||
)
|
||||
elif call_type == CallTypes.arerank:
|
||||
patch_target = (
|
||||
"litellm.llms.azure.common_utils.BaseAzureLLM.initialize_azure_sdk_client"
|
||||
)
|
||||
if call_type == CallTypes.arerank:
|
||||
patch_target = (
|
||||
"litellm.rerank_api.main.azure_rerank.initialize_azure_sdk_client"
|
||||
)
|
||||
|
|
@ -455,3 +461,177 @@ async def test_ensure_initialize_azure_sdk_client_always_used_azure_text(call_ty
|
|||
for call in azure_calls:
|
||||
assert "api_key" in call.kwargs, "api_key not found in parameters"
|
||||
assert "api_base" in call.kwargs, "api_base not found in parameters"
|
||||
|
||||
|
||||
# Test parameters for different API functions with Azure models
|
||||
AZURE_API_FUNCTION_PARAMS = [
|
||||
# (function_name, is_async, args)
|
||||
(
|
||||
"completion",
|
||||
False,
|
||||
{
|
||||
"model": "azure/gpt-4",
|
||||
"messages": [{"role": "user", "content": "Hello"}],
|
||||
"max_tokens": 10,
|
||||
"api_key": "test-api-key",
|
||||
"api_base": "https://test.openai.azure.com",
|
||||
"api_version": "2023-05-15",
|
||||
},
|
||||
),
|
||||
(
|
||||
"completion",
|
||||
True,
|
||||
{
|
||||
"model": "azure/gpt-4",
|
||||
"messages": [{"role": "user", "content": "Hello"}],
|
||||
"max_tokens": 10,
|
||||
"stream": True,
|
||||
"api_key": "test-api-key",
|
||||
"api_base": "https://test.openai.azure.com",
|
||||
"api_version": "2023-05-15",
|
||||
},
|
||||
),
|
||||
(
|
||||
"embedding",
|
||||
False,
|
||||
{
|
||||
"model": "azure/text-embedding-ada-002",
|
||||
"input": "Hello world",
|
||||
"api_key": "test-api-key",
|
||||
"api_base": "https://test.openai.azure.com",
|
||||
"api_version": "2023-05-15",
|
||||
},
|
||||
),
|
||||
(
|
||||
"embedding",
|
||||
True,
|
||||
{
|
||||
"model": "azure/text-embedding-ada-002",
|
||||
"input": "Hello world",
|
||||
"api_key": "test-api-key",
|
||||
"api_base": "https://test.openai.azure.com",
|
||||
"api_version": "2023-05-15",
|
||||
},
|
||||
),
|
||||
(
|
||||
"speech",
|
||||
False,
|
||||
{
|
||||
"model": "azure/tts-1",
|
||||
"input": "Hello, this is a test of text to speech",
|
||||
"voice": "alloy",
|
||||
"api_key": "test-api-key",
|
||||
"api_base": "https://test.openai.azure.com",
|
||||
"api_version": "2023-05-15",
|
||||
},
|
||||
),
|
||||
(
|
||||
"speech",
|
||||
True,
|
||||
{
|
||||
"model": "azure/tts-1",
|
||||
"input": "Hello, this is a test of text to speech",
|
||||
"voice": "alloy",
|
||||
"api_key": "test-api-key",
|
||||
"api_base": "https://test.openai.azure.com",
|
||||
"api_version": "2023-05-15",
|
||||
},
|
||||
),
|
||||
(
|
||||
"transcription",
|
||||
False,
|
||||
{
|
||||
"model": "azure/whisper-1",
|
||||
"file": MagicMock(),
|
||||
"api_key": "test-api-key",
|
||||
"api_base": "https://test.openai.azure.com",
|
||||
"api_version": "2023-05-15",
|
||||
},
|
||||
),
|
||||
(
|
||||
"transcription",
|
||||
True,
|
||||
{
|
||||
"model": "azure/whisper-1",
|
||||
"file": MagicMock(),
|
||||
"api_key": "test-api-key",
|
||||
"api_base": "https://test.openai.azure.com",
|
||||
"api_version": "2023-05-15",
|
||||
},
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("function_name,is_async,args", AZURE_API_FUNCTION_PARAMS)
|
||||
@pytest.mark.asyncio
|
||||
async def test_azure_client_reuse(function_name, is_async, args):
|
||||
"""
|
||||
Test that multiple Azure API calls reuse the same Azure OpenAI client
|
||||
"""
|
||||
litellm.set_verbose = True
|
||||
|
||||
# Determine which client class to mock based on whether the test is async
|
||||
client_path = (
|
||||
"litellm.llms.azure.common_utils.AsyncAzureOpenAI"
|
||||
if is_async
|
||||
else "litellm.llms.azure.common_utils.AzureOpenAI"
|
||||
)
|
||||
|
||||
# Create a proper mock class that can pass isinstance checks
|
||||
mock_client = MagicMock()
|
||||
|
||||
# Create the appropriate patches
|
||||
with patch(client_path) as mock_client_class, patch.object(
|
||||
BaseAzureLLM, "set_cached_openai_client"
|
||||
) as mock_set_cache, patch.object(
|
||||
BaseAzureLLM, "get_cached_openai_client"
|
||||
) as mock_get_cache, patch.object(
|
||||
BaseAzureLLM, "initialize_azure_sdk_client"
|
||||
) as mock_init_azure:
|
||||
# Configure the mock client class to return our mock instance
|
||||
mock_client_class.return_value = mock_client
|
||||
|
||||
# Setup the mock to return None first time (cache miss) then a client for subsequent calls
|
||||
mock_get_cache.side_effect = [None] + [
|
||||
mock_client
|
||||
] * 9 # First call returns None, rest return the mock client
|
||||
|
||||
# Mock the initialize_azure_sdk_client to return a dict with the necessary params
|
||||
mock_init_azure.return_value = {
|
||||
"api_key": args.get("api_key"),
|
||||
"azure_endpoint": args.get("api_base"),
|
||||
"api_version": args.get("api_version"),
|
||||
"azure_ad_token": None,
|
||||
"azure_ad_token_provider": None,
|
||||
}
|
||||
|
||||
# Make 10 API calls
|
||||
for _ in range(10):
|
||||
try:
|
||||
# Call the appropriate function based on parameters
|
||||
if is_async:
|
||||
# Add 'a' prefix for async functions
|
||||
func = getattr(litellm, f"a{function_name}")
|
||||
await func(**args)
|
||||
else:
|
||||
func = getattr(litellm, function_name)
|
||||
func(**args)
|
||||
except Exception:
|
||||
# We expect exceptions since we're mocking the client
|
||||
pass
|
||||
|
||||
# Verify client was created only once
|
||||
assert (
|
||||
mock_client_class.call_count == 1
|
||||
), f"{'Async' if is_async else ''}AzureOpenAI client should be created only once"
|
||||
|
||||
# Verify initialize_azure_sdk_client was called once
|
||||
assert (
|
||||
mock_init_azure.call_count == 1
|
||||
), "initialize_azure_sdk_client should be called once"
|
||||
|
||||
# Verify the client was cached
|
||||
assert mock_set_cache.call_count == 1, "Client should be cached once"
|
||||
|
||||
# Verify we tried to get from cache 10 times (once per request)
|
||||
assert mock_get_cache.call_count == 10, "Should check cache for each request"
|
||||
|
|
|
|||
132
tests/litellm/llms/openai/test_openai_common_utils.py
Normal file
132
tests/litellm/llms/openai/test_openai_common_utils.py
Normal file
|
|
@ -0,0 +1,132 @@
|
|||
import os
|
||||
import sys
|
||||
from unittest.mock import MagicMock, call, patch
|
||||
|
||||
import pytest
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../../..")
|
||||
) # Adds the parent directory to the system path
|
||||
|
||||
import litellm
|
||||
from litellm.llms.openai.common_utils import BaseOpenAILLM
|
||||
|
||||
# Test parameters for different API functions
|
||||
API_FUNCTION_PARAMS = [
|
||||
# (function_name, is_async, args)
|
||||
(
|
||||
"completion",
|
||||
False,
|
||||
{
|
||||
"model": "gpt-4o-mini",
|
||||
"messages": [{"role": "user", "content": "Hello"}],
|
||||
"max_tokens": 10,
|
||||
},
|
||||
),
|
||||
(
|
||||
"completion",
|
||||
True,
|
||||
{
|
||||
"model": "gpt-4o-mini",
|
||||
"messages": [{"role": "user", "content": "Hello"}],
|
||||
"max_tokens": 10,
|
||||
},
|
||||
),
|
||||
(
|
||||
"completion",
|
||||
True,
|
||||
{
|
||||
"model": "gpt-4o-mini",
|
||||
"messages": [{"role": "user", "content": "Hello"}],
|
||||
"max_tokens": 10,
|
||||
"stream": True,
|
||||
},
|
||||
),
|
||||
("embedding", False, {"model": "text-embedding-ada-002", "input": "Hello world"}),
|
||||
("embedding", True, {"model": "text-embedding-ada-002", "input": "Hello world"}),
|
||||
(
|
||||
"image_generation",
|
||||
False,
|
||||
{"model": "dall-e-3", "prompt": "A beautiful sunset over mountains"},
|
||||
),
|
||||
(
|
||||
"image_generation",
|
||||
True,
|
||||
{"model": "dall-e-3", "prompt": "A beautiful sunset over mountains"},
|
||||
),
|
||||
(
|
||||
"speech",
|
||||
False,
|
||||
{
|
||||
"model": "tts-1",
|
||||
"input": "Hello, this is a test of text to speech",
|
||||
"voice": "alloy",
|
||||
},
|
||||
),
|
||||
(
|
||||
"speech",
|
||||
True,
|
||||
{
|
||||
"model": "tts-1",
|
||||
"input": "Hello, this is a test of text to speech",
|
||||
"voice": "alloy",
|
||||
},
|
||||
),
|
||||
("transcription", False, {"model": "whisper-1", "file": MagicMock()}),
|
||||
("transcription", True, {"model": "whisper-1", "file": MagicMock()}),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("function_name,is_async,args", API_FUNCTION_PARAMS)
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_client_reuse(function_name, is_async, args):
|
||||
"""
|
||||
Test that multiple API calls reuse the same OpenAI client
|
||||
"""
|
||||
litellm.set_verbose = True
|
||||
|
||||
# Determine which client class to mock based on whether the test is async
|
||||
client_path = (
|
||||
"litellm.llms.openai.openai.AsyncOpenAI"
|
||||
if is_async
|
||||
else "litellm.llms.openai.openai.OpenAI"
|
||||
)
|
||||
|
||||
# Create the appropriate patches
|
||||
with patch(client_path) as mock_client_class, patch.object(
|
||||
BaseOpenAILLM, "set_cached_openai_client"
|
||||
) as mock_set_cache, patch.object(
|
||||
BaseOpenAILLM, "get_cached_openai_client"
|
||||
) as mock_get_cache:
|
||||
|
||||
# Setup the mock to return None first time (cache miss) then a client for subsequent calls
|
||||
mock_client = MagicMock()
|
||||
mock_get_cache.side_effect = [None] + [
|
||||
mock_client
|
||||
] * 9 # First call returns None, rest return the mock client
|
||||
|
||||
# Make 10 API calls
|
||||
for _ in range(10):
|
||||
try:
|
||||
# Call the appropriate function based on parameters
|
||||
if is_async:
|
||||
# Add 'a' prefix for async functions
|
||||
func = getattr(litellm, f"a{function_name}")
|
||||
await func(**args)
|
||||
else:
|
||||
func = getattr(litellm, function_name)
|
||||
func(**args)
|
||||
except Exception:
|
||||
# We expect exceptions since we're mocking the client
|
||||
pass
|
||||
|
||||
# Verify client was created only once
|
||||
assert (
|
||||
mock_client_class.call_count == 1
|
||||
), f"{'Async' if is_async else ''}OpenAI client should be created only once"
|
||||
|
||||
# Verify the client was cached
|
||||
assert mock_set_cache.call_count == 1, "Client should be cached once"
|
||||
|
||||
# Verify we tried to get from cache 10 times (once per request)
|
||||
assert mock_get_cache.call_count == 10, "Should check cache for each request"
|
||||
|
|
@ -4364,14 +4364,14 @@ async def test_dynamic_azure_params(stream, sync_mode):
|
|||
|
||||
## recreate mock client
|
||||
if sync_mode:
|
||||
mock_client = MagicMock(return_value="Hello world!")
|
||||
new_mock_client = MagicMock(return_value="Hello world!")
|
||||
else:
|
||||
mock_client = AsyncMock(return_value="Hello world!")
|
||||
new_mock_client = AsyncMock(return_value="Hello world!")
|
||||
|
||||
## CHECK IF NEW CLIENT IS USED (PARAM CHANGE)
|
||||
with patch.object(
|
||||
client.chat.completions.with_raw_response, "create", new=mock_client
|
||||
) as mock_client:
|
||||
client.chat.completions.with_raw_response, "create", new=new_mock_client
|
||||
) as new_mock_client:
|
||||
try:
|
||||
if sync_mode:
|
||||
_ = completion(
|
||||
|
|
@ -4393,7 +4393,7 @@ async def test_dynamic_azure_params(stream, sync_mode):
|
|||
pass
|
||||
|
||||
try:
|
||||
mock_client.assert_not_called()
|
||||
new_mock_client.assert_called()
|
||||
except Exception as e:
|
||||
raise e
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue