diff --git a/.circleci/config.yml b/.circleci/config.yml
index f8393be9dff..a5abab254e9 100644
--- a/.circleci/config.yml
+++ b/.circleci/config.yml
@@ -51,6 +51,7 @@ jobs:
pip install prisma==0.11.0
pip install "detect_secrets==1.5.0"
pip install "httpx==0.24.1"
+ pip install "respx==0.21.1"
pip install fastapi
pip install "gunicorn==21.2.0"
pip install "anyio==3.7.1"
diff --git a/docs/my-website/docs/providers/azure.md b/docs/my-website/docs/providers/azure.md
index be3401fd2e9..dc64bffc1cc 100644
--- a/docs/my-website/docs/providers/azure.md
+++ b/docs/my-website/docs/providers/azure.md
@@ -1,3 +1,8 @@
+
+import Image from '@theme/IdealImage';
+import Tabs from '@theme/Tabs';
+import TabItem from '@theme/TabItem';
+
# Azure OpenAI
## API Keys, Params
api_key, api_base, api_version etc can be passed directly to `litellm.completion` - see here or set as `litellm.api_key` params see here
@@ -12,7 +17,7 @@ os.environ["AZURE_AD_TOKEN"] = ""
os.environ["AZURE_API_TYPE"] = ""
```
-## Usage
+## **Usage - LiteLLM Python SDK**
@@ -64,6 +69,125 @@ response = litellm.completion(
)
```
+
+## **Usage - LiteLLM Proxy Server**
+
+Here's how to call Azure OpenAI models with the LiteLLM Proxy Server
+
+### 1. Save key in your environment
+
+```bash
+export AZURE_API_KEY=""
+```
+
+### 2. Start the proxy
+
+
+
+
+```yaml
+model_list:
+ - model_name: gpt-3.5-turbo
+ litellm_params:
+ 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 # The `os.environ/` prefix tells litellm to read this from the env.
+```
+
+
+
+
+```yaml
+model_list:
+ - model_name: gpt-3.5-turbo
+ litellm_params:
+ model: azure/chatgpt-v-2
+ api_base: https://openai-gpt-4-test-v-1.openai.azure.com/
+ api_version: "2023-05-15"
+ tenant_id: os.environ/AZURE_TENANT_ID
+ client_id: os.environ/AZURE_CLIENT_ID
+ client_secret: os.environ/AZURE_CLIENT_SECRET
+```
+
+
+
+
+### 3. Test it
+
+
+
+
+
+```shell
+curl --location 'http://0.0.0.0:4000/chat/completions' \
+--header 'Content-Type: application/json' \
+--data ' {
+ "model": "gpt-3.5-turbo",
+ "messages": [
+ {
+ "role": "user",
+ "content": "what llm are you"
+ }
+ ]
+ }
+'
+```
+
+
+
+```python
+import openai
+client = openai.OpenAI(
+ api_key="anything",
+ base_url="http://0.0.0.0:4000"
+)
+
+response = client.chat.completions.create(model="gpt-3.5-turbo", messages = [
+ {
+ "role": "user",
+ "content": "this is a test request, write a short poem"
+ }
+])
+
+print(response)
+
+```
+
+
+
+```python
+from langchain.chat_models import ChatOpenAI
+from langchain.prompts.chat import (
+ ChatPromptTemplate,
+ HumanMessagePromptTemplate,
+ SystemMessagePromptTemplate,
+)
+from langchain.schema import HumanMessage, SystemMessage
+
+chat = ChatOpenAI(
+ openai_api_base="http://0.0.0.0:4000", # set openai_api_base to the LiteLLM Proxy
+ model = "gpt-3.5-turbo",
+ temperature=0.1
+)
+
+messages = [
+ SystemMessage(
+ content="You are a helpful assistant that im using to make a test request to."
+ ),
+ HumanMessage(
+ content="test from litellm. tell me why it's amazing in 1 sentence"
+ ),
+]
+response = chat(messages)
+
+print(response)
+```
+
+
+
+
+
## Azure OpenAI Chat Completion Models
:::tip
diff --git a/litellm/main.py b/litellm/main.py
index 28054537cf4..16a4f89ed29 100644
--- a/litellm/main.py
+++ b/litellm/main.py
@@ -943,7 +943,9 @@ def completion(
output_cost_per_token=output_cost_per_token,
cooldown_time=cooldown_time,
text_completion=kwargs.get("text_completion"),
+ azure_ad_token_provider=kwargs.get("azure_ad_token_provider"),
user_continue_message=kwargs.get("user_continue_message"),
+
)
logging.update_environment_variables(
model=model,
@@ -3247,6 +3249,10 @@ def embedding(
"model_config",
"cooldown_time",
"tags",
+ "azure_ad_token_provider",
+ "tenant_id",
+ "client_id",
+ "client_secret",
"extra_headers",
]
default_params = openai_params + litellm_params
diff --git a/litellm/proxy/proxy_config.yaml b/litellm/proxy/proxy_config.yaml
index 5f4072434fd..320216a79b9 100644
--- a/litellm/proxy/proxy_config.yaml
+++ b/litellm/proxy/proxy_config.yaml
@@ -1,9 +1,12 @@
model_list:
- model_name: fake-openai-endpoint
litellm_params:
- model: openai/fake
- api_key: fake-key
- api_base: https://exampleopenaiendpoint-production.up.railway.app/
+ model: azure/chatgpt-v-2
+ api_base: https://openai-gpt-4-test-v-1.openai.azure.com/
+ api_version: "2023-05-15"
+ tenant_id: os.environ/AZURE_TENANT_ID
+ client_id: os.environ/AZURE_CLIENT_ID
+ client_secret: os.environ/AZURE_CLIENT_SECRET
guardrails:
- guardrail_name: "bedrock-pre-guard"
diff --git a/litellm/router_utils/client_initalization_utils.py b/litellm/router_utils/client_initalization_utils.py
index f396defb51f..e98b8b4ddc9 100644
--- a/litellm/router_utils/client_initalization_utils.py
+++ b/litellm/router_utils/client_initalization_utils.py
@@ -1,7 +1,7 @@
import asyncio
import os
import traceback
-from typing import TYPE_CHECKING, Any
+from typing import TYPE_CHECKING, Any, Callable
import httpx
import openai
@@ -172,6 +172,14 @@ def set_client(litellm_router_instance: LitellmRouter, model: dict):
organization_env_name = organization.replace("os.environ/", "")
organization = litellm.get_secret(organization_env_name)
litellm_params["organization"] = organization
+ azure_ad_token_provider = None
+ if litellm_params.get("tenant_id"):
+ verbose_router_logger.debug("Using Azure AD Token Provider for Azure Auth")
+ azure_ad_token_provider = get_azure_ad_token_from_entrata_id(
+ tenant_id=litellm_params.get("tenant_id"),
+ client_id=litellm_params.get("client_id"),
+ client_secret=litellm_params.get("client_secret"),
+ )
if custom_llm_provider == "azure" or custom_llm_provider == "azure_text":
if api_base is None or not isinstance(api_base, str):
@@ -190,7 +198,9 @@ def set_client(litellm_router_instance: LitellmRouter, model: dict):
if azure_ad_token.startswith("oidc/"):
azure_ad_token = get_azure_ad_token_from_oidc(azure_ad_token)
if api_version is None:
- api_version = os.getenv("AZURE_API_VERSION", litellm.AZURE_DEFAULT_API_VERSION)
+ api_version = os.getenv(
+ "AZURE_API_VERSION", litellm.AZURE_DEFAULT_API_VERSION
+ )
if "gateway.ai.cloudflare.com" in api_base:
if not api_base.endswith("/"):
@@ -304,6 +314,11 @@ def set_client(litellm_router_instance: LitellmRouter, model: dict):
"api_version": api_version,
"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
+ )
from litellm.llms.azure import select_azure_base_url_or_endpoint
# this decides if we should set azure_endpoint or base_url on Azure OpenAI Client
@@ -493,3 +508,41 @@ def set_client(litellm_router_instance: LitellmRouter, model: dict):
ttl=client_ttl,
local_only=True,
) # cache for 1 hr
+
+
+def get_azure_ad_token_from_entrata_id(
+ tenant_id: str, client_id: str, client_secret: str
+) -> Callable[[], str]:
+ from azure.identity import (
+ ClientSecretCredential,
+ DefaultAzureCredential,
+ get_bearer_token_provider,
+ )
+
+ verbose_router_logger.debug("Getting Azure AD Token from Entrata ID")
+
+ if tenant_id.startswith("os.environ/"):
+ tenant_id = litellm.get_secret(tenant_id)
+
+ if client_id.startswith("os.environ/"):
+ client_id = litellm.get_secret(client_id)
+
+ if client_secret.startswith("os.environ/"):
+ client_secret = litellm.get_secret(client_secret)
+ verbose_router_logger.debug(
+ "tenant_id %s, client_id %s, client_secret %s",
+ tenant_id,
+ client_id,
+ client_secret,
+ )
+ credential = ClientSecretCredential(tenant_id, client_id, client_secret)
+
+ verbose_router_logger.debug("credential %s", credential)
+
+ token_provider = get_bearer_token_provider(
+ credential, "https://cognitiveservices.azure.com/.default"
+ )
+
+ verbose_router_logger.debug("token_provider %s", token_provider)
+
+ return token_provider
diff --git a/litellm/tests/test_azure_openai.py b/litellm/tests/test_azure_openai.py
new file mode 100644
index 00000000000..9972f2833d0
--- /dev/null
+++ b/litellm/tests/test_azure_openai.py
@@ -0,0 +1,98 @@
+import json
+import os
+import sys
+import traceback
+
+from dotenv import load_dotenv
+
+load_dotenv()
+import io
+import os
+
+sys.path.insert(
+ 0, os.path.abspath("../..")
+) # Adds the parent directory to the system path
+
+import os
+from datetime import datetime
+from unittest.mock import AsyncMock, MagicMock, patch
+
+import httpx
+import pytest
+from openai import OpenAI
+from openai.types.chat import ChatCompletionMessage
+from openai.types.chat.chat_completion import ChatCompletion, Choice
+from respx import MockRouter
+
+import litellm
+from litellm import RateLimitError, Timeout, completion, completion_cost, embedding
+from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
+from litellm.llms.prompt_templates.factory import anthropic_messages_pt
+from litellm.router import Router
+
+
+@pytest.mark.asyncio()
+@pytest.mark.respx()
+async def test_azure_tenant_id_auth(respx_mock: MockRouter):
+ """
+
+ Tests when we set tenant_id, client_id, client_secret they don't get sent with the request
+
+ PROD Test
+ """
+
+ router = Router(
+ model_list=[
+ {
+ "model_name": "gpt-3.5-turbo",
+ "litellm_params": { # params for litellm completion/embedding call
+ "model": "azure/chatgpt-v-2",
+ "api_base": os.getenv("AZURE_API_BASE"),
+ "tenant_id": os.getenv("AZURE_TENANT_ID"),
+ "client_id": os.getenv("AZURE_CLIENT_ID"),
+ "client_secret": os.getenv("AZURE_CLIENT_SECRET"),
+ },
+ },
+ ],
+ )
+
+ mock_response = AsyncMock()
+ obj = ChatCompletion(
+ id="foo",
+ model="gpt-4",
+ object="chat.completion",
+ choices=[
+ Choice(
+ finish_reason="stop",
+ index=0,
+ message=ChatCompletionMessage(
+ content="Hello world!",
+ role="assistant",
+ ),
+ )
+ ],
+ created=int(datetime.now().timestamp()),
+ )
+ mock_request = respx_mock.post(url__regex=r".*/chat/completions.*").mock(
+ return_value=httpx.Response(200, json=obj.model_dump(mode="json"))
+ )
+
+ await router.acompletion(
+ model="gpt-3.5-turbo", messages=[{"role": "user", "content": "Hello world!"}]
+ )
+
+ # Ensure all mocks were called
+ respx_mock.assert_all_called()
+
+ for call in mock_request.calls:
+ print(call)
+ print(call.request.content)
+
+ json_body = json.loads(call.request.content)
+ print(json_body)
+
+ assert json_body == {
+ "messages": [{"role": "user", "content": "Hello world!"}],
+ "model": "chatgpt-v-2",
+ "stream": False,
+ }
diff --git a/litellm/types/utils.py b/litellm/types/utils.py
index 6b278efa1b4..14d5cd1b8d5 100644
--- a/litellm/types/utils.py
+++ b/litellm/types/utils.py
@@ -1116,6 +1116,10 @@ all_litellm_params = [
"cooldown_time",
"cache_key",
"max_retries",
+ "azure_ad_token_provider",
+ "tenant_id",
+ "client_id",
+ "client_secret",
"user_continue_message",
]
diff --git a/litellm/utils.py b/litellm/utils.py
index 7596de81d21..d5aefa80ecf 100644
--- a/litellm/utils.py
+++ b/litellm/utils.py
@@ -2323,6 +2323,7 @@ def get_litellm_params(
output_cost_per_second=None,
cooldown_time=None,
text_completion=None,
+ azure_ad_token_provider=None,
user_continue_message=None,
):
litellm_params = {
@@ -2348,6 +2349,7 @@ def get_litellm_params(
"output_cost_per_second": output_cost_per_second,
"cooldown_time": cooldown_time,
"text_completion": text_completion,
+ "azure_ad_token_provider": azure_ad_token_provider,
"user_continue_message": user_continue_message,
}