Skip to content

Commit accd05b

Browse files
committed
fix(azure): keep provider validation errors value-free
1 parent 06ef57c commit accd05b

2 files changed

Lines changed: 140 additions & 11 deletions

File tree

src/openai/lib/azure.py

Lines changed: 7 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -443,11 +443,9 @@ def _get_azure_ad_token(self) -> str | None:
443443

444444
provider = self._azure_ad_token_provider
445445
if provider is not None:
446-
token = provider()
447-
if not token or not isinstance(token, str): # pyright: ignore[reportUnnecessaryIsInstance]
448-
raise ValueError(
449-
f"Expected `azure_ad_token_provider` argument to return a string but it returned {token}",
450-
)
446+
token = cast(object, provider())
447+
if not isinstance(token, str) or not token:
448+
raise ValueError("Expected `azure_ad_token_provider` argument to return a non-empty string.")
451449
return token
452450

453451
return None
@@ -794,14 +792,12 @@ async def _get_azure_ad_token(self) -> str | None:
794792

795793
provider = self._azure_ad_token_provider
796794
if provider is not None:
797-
token = provider()
795+
token = cast(object, provider())
798796
if inspect.isawaitable(token):
799797
token = await token
800-
if not token or not isinstance(cast(Any, token), str):
801-
raise ValueError(
802-
f"Expected `azure_ad_token_provider` argument to return a string but it returned {token}",
803-
)
804-
return str(token)
798+
if not isinstance(token, str) or not token:
799+
raise ValueError("Expected `azure_ad_token_provider` argument to return a non-empty string.")
800+
return token
805801

806802
return None
807803

Lines changed: 133 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,133 @@
1+
from __future__ import annotations
2+
3+
import asyncio
4+
import inspect
5+
import traceback
6+
from typing import Any, Callable, NoReturn, NamedTuple
7+
from typing_extensions import override
8+
9+
import httpx2
10+
import pytest
11+
12+
from openai import AzureOpenAI, AsyncAzureOpenAI
13+
14+
FAKE_TOKEN = "fake-azure-provider-token"
15+
ERROR_MESSAGE = "Expected `azure_ad_token_provider` argument to return a non-empty string."
16+
PROVIDER_MODES = ["sync", "async-direct", "async-coroutine", "async-awaitable"]
17+
18+
19+
class FakeAccessToken(NamedTuple):
20+
token: str
21+
expires_on: int
22+
23+
24+
class UninspectableToken:
25+
def __bool__(self) -> NoReturn:
26+
raise AssertionError(FAKE_TOKEN)
27+
28+
@override
29+
def __str__(self) -> NoReturn:
30+
raise AssertionError(FAKE_TOKEN)
31+
32+
@override
33+
def __repr__(self) -> NoReturn:
34+
raise AssertionError(FAKE_TOKEN)
35+
36+
37+
@pytest.fixture(autouse=True)
38+
def offline_environment(monkeypatch: pytest.MonkeyPatch) -> None:
39+
for name in ("AZURE_OPENAI_API_KEY", "AZURE_OPENAI_AD_TOKEN", "OPENAI_API_KEY"):
40+
monkeypatch.delenv(name, raising=False)
41+
42+
def unexpected_connection(*_args: Any, **_kwargs: Any) -> NoReturn:
43+
pytest.fail("Azure provider diagnostics must not open a network connection")
44+
45+
monkeypatch.setattr("socket.socket.connect", unexpected_connection)
46+
monkeypatch.setattr("socket.socket.connect_ex", unexpected_connection)
47+
48+
49+
def make_provider(mode: str, value: object) -> Callable[[], Any]:
50+
async def coroutine() -> object:
51+
return value
52+
53+
def awaitable() -> asyncio.Future[object]:
54+
future: asyncio.Future[object] = asyncio.get_running_loop().create_future()
55+
future.set_result(value)
56+
return future
57+
58+
if mode == "async-coroutine":
59+
return coroutine
60+
if mode == "async-awaitable":
61+
return awaitable
62+
return lambda: value
63+
64+
65+
def make_client(mode: str, value: object, requests: list[httpx2.Request]) -> Any:
66+
def send(request: httpx2.Request) -> httpx2.Response:
67+
requests.append(request)
68+
return httpx2.Response(200, json={"data": []})
69+
70+
transport = httpx2.MockTransport(send)
71+
http_client = httpx2.Client(transport=transport) if mode == "sync" else httpx2.AsyncClient(transport=transport)
72+
cls: Any = AzureOpenAI if mode == "sync" else AsyncAzureOpenAI
73+
return cls(
74+
azure_endpoint="https://azure.test",
75+
api_version="2024-02-01",
76+
azure_ad_token_provider=make_provider(mode, value),
77+
http_client=http_client,
78+
max_retries=0,
79+
)
80+
81+
82+
async def resolve(value: Any) -> Any:
83+
return await value if inspect.isawaitable(value) else value
84+
85+
86+
@pytest.mark.parametrize("mode", PROVIDER_MODES)
87+
@pytest.mark.parametrize("entrypoint", ["http", "realtime-config", "realtime", "beta-realtime"])
88+
@pytest.mark.parametrize(
89+
"value",
90+
[
91+
pytest.param({"access_token": FAKE_TOKEN}, id="dict"),
92+
pytest.param(FakeAccessToken(FAKE_TOKEN, 0), id="access-token"),
93+
pytest.param(UninspectableToken(), id="uninspectable-object"),
94+
pytest.param("", id="empty-string"),
95+
pytest.param(None, id="none"),
96+
],
97+
)
98+
async def test_invalid_provider_result_is_value_free(mode: str, entrypoint: str, value: object) -> None:
99+
requests: list[httpx2.Request] = []
100+
client = make_client(mode, value, requests)
101+
try:
102+
with pytest.raises(ValueError) as exc_info:
103+
if entrypoint == "http":
104+
await resolve(client.models.list())
105+
elif entrypoint == "realtime-config":
106+
await resolve(client._configure_realtime("test-model", {}))
107+
else:
108+
resource = client.realtime if entrypoint == "realtime" else client.beta.realtime
109+
await resolve(resource.connect(model="test-model").enter())
110+
111+
error = exc_info.value
112+
assert str(error) == ERROR_MESSAGE
113+
assert FAKE_TOKEN not in repr(error)
114+
assert FAKE_TOKEN not in "".join(traceback.format_exception(type(error), error, error.__traceback__))
115+
assert error.__cause__ is None
116+
assert error.__context__ is None
117+
assert requests == []
118+
finally:
119+
await resolve(client.close())
120+
121+
122+
@pytest.mark.parametrize("mode", PROVIDER_MODES)
123+
async def test_nonempty_provider_result_remains_usable(mode: str) -> None:
124+
requests: list[httpx2.Request] = []
125+
client = make_client(mode, FAKE_TOKEN, requests)
126+
try:
127+
await resolve(client.models.list())
128+
assert len(requests) == 1
129+
assert requests[0].headers["Authorization"] == f"Bearer {FAKE_TOKEN}"
130+
_, headers = await resolve(client._configure_realtime("test-model", {}))
131+
assert headers == {"Authorization": f"Bearer {FAKE_TOKEN}"}
132+
finally:
133+
await resolve(client.close())

0 commit comments

Comments
 (0)