Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
48 changes: 47 additions & 1 deletion src/openai/lib/azure.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@
import httpx

from ..auth import WorkloadIdentity
from .._types import NOT_GIVEN, Omit, Query, Headers, Timeout, NotGiven
from .._types import NOT_GIVEN, Omit, Query, Headers, Timeout, NotGiven, ResponseT
from .._utils import is_given, is_mapping
from .._client import OpenAI, AsyncOpenAI
from .._compat import model_copy
Expand Down Expand Up @@ -43,6 +43,7 @@
# as we don't want to make the `api_key` in the main client Optional
# and Azure AD tokens may be retrieved on a per-request basis
API_KEY_SENTINEL = "".join(["<", "missing API key", ">"])
_AZURE_RESPONSES_SERVED_MODEL_HEADER = "x-ms-served-model"


def _has_header(headers: Headers, header: str) -> bool:
Expand All @@ -54,6 +55,37 @@ def _has_auth_header(headers: Headers) -> bool:
return _has_header(headers, "Authorization") or _has_header(headers, "api-key")


def _is_responses_request(response: httpx.Response) -> bool:
path = response.request.url.path.rstrip("/")
return path.endswith("/responses") or "/responses/" in path


def _served_model_from_response(response: httpx.Response) -> str | None:
if not _is_responses_request(response):
return None

served_model = response.headers.get(_AZURE_RESPONSES_SERVED_MODEL_HEADER)
if served_model is None:
return None

served_model = served_model.strip()
return served_model or None


def _replace_response_model(data: object, served_model: str | None) -> object:
if served_model is None or not is_mapping(data):
return data

nested_response = data.get("response")
if is_mapping(nested_response):
return {**data, "response": {**nested_response, "model": served_model}}

if "model" in data:
return {**data, "model": served_model}

return data


class MutuallyExclusiveAuthError(OpenAIError):
def __init__(self) -> None:
super().__init__(
Expand All @@ -65,6 +97,20 @@ class BaseAzureClient(BaseClient[_HttpxClientT, _DefaultStreamT]):
_azure_endpoint: httpx.URL | None
_azure_deployment: str | None

@override
def _process_response_data(
self,
*,
data: object,
cast_to: type[ResponseT],
response: httpx.Response,
) -> ResponseT:
return super()._process_response_data(
data=_replace_response_model(data, _served_model_from_response(response)),
cast_to=cast_to,
response=response,
)

@override
def _build_request(
self,
Expand Down
131 changes: 130 additions & 1 deletion tests/lib/test_azure.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,8 @@
from __future__ import annotations

import json
import logging
from typing import Union, cast
from typing import Union, Iterator, AsyncIterator, cast
from typing_extensions import Literal, Protocol

import httpx
Expand Down Expand Up @@ -30,11 +31,139 @@
azure_endpoint="https://example-resource.azure.openai.com",
)

AZURE_RESPONSES_URL = "https://example-resource.azure.openai.com/openai/responses?api-version=2024-02-01"
AZURE_DEPLOYMENT_MODEL = "gpt-5-nano"
AZURE_SERVED_MODEL = "gpt-5-nano-2025-08-07"


def _azure_response_payload(*, model: str = AZURE_DEPLOYMENT_MODEL) -> dict[str, object]:
return {
"id": "resp_123",
"object": "response",
"created_at": 0,
"model": model,
"output": [],
"parallel_tool_calls": True,
"tool_choice": "auto",
"tools": [],
}


def _azure_response_stream_body() -> Iterator[bytes]:
yield b"event: response.created\n"
yield (
b'data: {"type":"response.created","sequence_number":0,"response":'
+ json.dumps(_azure_response_payload(), separators=(",", ":")).encode()
+ b"}\n\n"
)
yield b"data: [DONE]\n\n"


async def _async_azure_response_stream_body() -> AsyncIterator[bytes]:
for chunk in _azure_response_stream_body():
yield chunk


class MockRequestCall(Protocol):
request: httpx.Request


@pytest.mark.respx()
def test_azure_responses_uses_served_model_header(respx_mock: MockRouter) -> None:
respx_mock.post(AZURE_RESPONSES_URL).mock(
return_value=httpx.Response(
200,
headers={"x-ms-served-model": f" {AZURE_SERVED_MODEL} "},
json=_azure_response_payload(),
)
)

client = AzureOpenAI(
api_version="2024-02-01",
api_key="example API key",
azure_endpoint="https://example-resource.azure.openai.com",
)

response = client.responses.create(model=AZURE_DEPLOYMENT_MODEL, input="ping")

assert response.model == AZURE_SERVED_MODEL


@pytest.mark.asyncio
@pytest.mark.respx()
async def test_async_azure_responses_uses_served_model_header(respx_mock: MockRouter) -> None:
respx_mock.post(AZURE_RESPONSES_URL).mock(
return_value=httpx.Response(
200,
headers={"x-ms-served-model": AZURE_SERVED_MODEL},
json=_azure_response_payload(),
)
)

client = AsyncAzureOpenAI(
api_version="2024-02-01",
api_key="example API key",
azure_endpoint="https://example-resource.azure.openai.com",
)

response = await client.responses.create(model=AZURE_DEPLOYMENT_MODEL, input="ping")

assert response.model == AZURE_SERVED_MODEL


@pytest.mark.respx()
def test_azure_responses_stream_uses_served_model_header(respx_mock: MockRouter) -> None:
respx_mock.post(AZURE_RESPONSES_URL).mock(
return_value=httpx.Response(
200,
headers={
"content-type": "text/event-stream",
"x-ms-served-model": AZURE_SERVED_MODEL,
},
content=_azure_response_stream_body(),
)
)

client = AzureOpenAI(
api_version="2024-02-01",
api_key="example API key",
azure_endpoint="https://example-resource.azure.openai.com",
)

stream = client.responses.create(model=AZURE_DEPLOYMENT_MODEL, input="ping", stream=True)
event = next(stream)

assert event.type == "response.created"
assert event.response.model == AZURE_SERVED_MODEL


@pytest.mark.asyncio
@pytest.mark.respx()
async def test_async_azure_responses_stream_uses_served_model_header(respx_mock: MockRouter) -> None:
respx_mock.post(AZURE_RESPONSES_URL).mock(
return_value=httpx.Response(
200,
headers={
"content-type": "text/event-stream",
"x-ms-served-model": AZURE_SERVED_MODEL,
},
content=_async_azure_response_stream_body(),
)
)

client = AsyncAzureOpenAI(
api_version="2024-02-01",
api_key="example API key",
azure_endpoint="https://example-resource.azure.openai.com",
)

stream = await client.responses.create(model=AZURE_DEPLOYMENT_MODEL, input="ping", stream=True)
event = await stream.__anext__()

assert event.type == "response.created"
assert event.response.model == AZURE_SERVED_MODEL


@pytest.mark.parametrize("client", [sync_client, async_client])
def test_implicit_deployment_path(client: Client) -> None:
req = client._build_request(
Expand Down