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
168 changes: 168 additions & 0 deletions tests/lib/test_azure_responses.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,168 @@
from __future__ import annotations

import json
from typing import Iterator, AsyncIterator

import httpx
import pytest
from respx import MockRouter

from openai.lib.azure import AzureOpenAI, AsyncAzureOpenAI

AZURE_ENDPOINT = "https://example-resource.azure.openai.com"
AZURE_API_VERSION = "2024-02-01"
AZURE_RESPONSES_URL = f"{AZURE_ENDPOINT}/openai/responses?api-version={AZURE_API_VERSION}"
AZURE_CHAT_COMPLETIONS_URL = (
f"{AZURE_ENDPOINT}/openai/deployments/gpt-4/chat/completions?api-version={AZURE_API_VERSION}"
)
AZURE_DEPLOYMENT_MODEL = "gpt-5-nano"
AZURE_SERVED_MODEL = "gpt-5-nano-2025-08-07"


def make_sync_client() -> AzureOpenAI:
return AzureOpenAI(
api_version=AZURE_API_VERSION,
api_key="example API key",
azure_endpoint=AZURE_ENDPOINT,
)


def make_async_client() -> AsyncAzureOpenAI:
return AsyncAzureOpenAI(
api_version=AZURE_API_VERSION,
api_key="example API key",
azure_endpoint=AZURE_ENDPOINT,
)


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 response_created_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_response_created_stream_body() -> AsyncIterator[bytes]:
for chunk in response_created_stream_body():
yield chunk


def mock_responses_create(
respx_mock: MockRouter,
*,
served_model_header: str | None,
stream: bool = False,
) -> None:
headers = {"x-ms-served-model": served_model_header} if served_model_header is not None else {}
if stream:
headers["content-type"] = "text/event-stream"
respx_mock.post(AZURE_RESPONSES_URL).mock(
return_value=httpx.Response(
200,
headers=headers,
content=response_created_stream_body(),
)
)
else:
respx_mock.post(AZURE_RESPONSES_URL).mock(
return_value=httpx.Response(
200,
headers=headers,
json=azure_response_payload(),
)
)


@pytest.mark.respx()
def test_azure_responses_uses_served_model_header(respx_mock: MockRouter) -> None:
mock_responses_create(respx_mock, served_model_header=f" {AZURE_SERVED_MODEL} ")

response = make_sync_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:
mock_responses_create(respx_mock, served_model_header=AZURE_SERVED_MODEL)

response = await make_async_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:
mock_responses_create(respx_mock, served_model_header=AZURE_SERVED_MODEL, stream=True)

stream = make_sync_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_response_created_stream_body(),
)
)

stream = await make_async_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("served_model_header", [None, " "])
@pytest.mark.respx()
def test_azure_responses_preserves_body_model_without_served_model_header(
respx_mock: MockRouter,
served_model_header: str | None,
) -> None:
mock_responses_create(respx_mock, served_model_header=served_model_header)

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

assert response.model == AZURE_DEPLOYMENT_MODEL


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

response = make_sync_client().chat.completions.create(messages=[], model="gpt-4")

assert response.model == "gpt-4"