Skip to content
Merged
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
16 changes: 15 additions & 1 deletion src/openai/_base_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -530,7 +530,7 @@ def _custom_auth(

def _build_headers(self, options: FinalRequestOptions, *, retries_taken: int = 0) -> httpx2.Headers:
custom_headers = options.headers or {}
headers_dict = _merge_mappings({**self._auth_headers(options.security), **self.default_headers}, custom_headers)
headers_dict = _merge_headers(self._auth_headers(options.security), self.default_headers, custom_headers)
self._validate_headers(headers_dict, custom_headers)

# headers are case-insensitive while dictionaries are not.
Expand Down Expand Up @@ -2342,3 +2342,17 @@ def _merge_mappings(
"""
merged = {**obj1, **obj2}
return {key: value for key, value in merged.items() if not isinstance(value, Omit)}


def _merge_headers(*mappings: Headers) -> dict[str, str]:
"""Merge headers case-insensitively, with later mappings taking precedence."""
merged: dict[str, tuple[str, str]] = {}
for mapping in mappings:
for name, value in mapping.items():
normalized_name = name.lower()
if isinstance(value, Omit):
merged.pop(normalized_name, None)
else:
merged[normalized_name] = (name, value)

return {name: value for name, value in merged.values()}
34 changes: 34 additions & 0 deletions tests/test_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -156,6 +156,40 @@ def _get_open_connections(client: OpenAI | AsyncOpenAI) -> int:
return len(cast(Any, transport)._pool._requests)


@pytest.mark.parametrize("is_async", [False, True])
@pytest.mark.parametrize(
"extra_headers,expected",
[
({}, ["Bearer fake-default"]),
({"AUTHORIZATION": "Bearer fake-request"}, ["Bearer fake-request"]),
({"AUTHORIZATION": Omit()}, []),
],
ids=["default", "override", "omit"],
)
async def test_case_insensitive_auth_headers(
is_async: bool, extra_headers: dict[str, str | Omit], expected: list[str]
) -> None:
def handler(request: httpx2.Request) -> httpx2.Response:
assert request.headers.get_list("authorization") == expected
return httpx2.Response(200, json={"object": "list", "data": []})

transport = httpx2.MockTransport(handler)
if is_async:
async with AsyncOpenAI(
api_key="fake-original",
default_headers={"authorization": "Bearer fake-default"},
http_client=httpx2.AsyncClient(transport=transport),
) as async_client:
await async_client.models.list(extra_headers=extra_headers)
else:
with OpenAI(
api_key="fake-original",
default_headers={"authorization": "Bearer fake-default"},
http_client=httpx2.Client(transport=transport),
) as client:
client.models.list(extra_headers=extra_headers)


class TestOpenAI:
@pytest.mark.parametrize(
"code_fields,expected_code",
Expand Down
Loading