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
Original file line number Diff line number Diff line change
@@ -0,0 +1,7 @@
---
default: patch
---

# Fixed generated endpoints breaking when a schema is named `Client`, `AuthenticatedClient`, `Response` or `HTTPStatus`

Generated endpoint modules imported the SDK's `Client`, `AuthenticatedClient` and `Response` classes and the standard library's `HTTPStatus` under their bare names, so a model with one of those names shadowed them. Depending on the name this produced wrong type annotations, a failed response build, or a `Response.status_code` that held a model instance instead of an `HTTPStatus`. Endpoint modules now bind these generator-owned symbols under private aliases, leaving the public model names unchanged.
Original file line number Diff line number Diff line change
@@ -0,0 +1,141 @@
import asyncio
import json
from collections.abc import Iterator
from http import HTTPStatus
from typing import get_type_hints

import httpx
import pytest

from end_to_end_tests.generated_client import GeneratedClientContext, generate_client_from_inline_spec

ISSUE_SPEC = """
openapi: "3.0.3"
info:
title: Minimal repro
version: "1.0"
paths:
/api/clients/{id}:
patch:
operationId: update_client
tags:
- client
parameters:
- name: id
in: path
required: true
schema:
type: string
requestBody:
content:
application/json:
schema:
type: object
properties:
name:
type: string
responses:
"200":
description: OK
content:
application/json:
schema:
$ref: "#/components/schemas/Client"
components:
schemas:
Client:
type: object
properties:
name:
type: string
"""


@pytest.fixture(params=["Client", "AuthenticatedClient", "Response", "HTTPStatus", "Customer"])
def generated_client(request: pytest.FixtureRequest) -> Iterator[GeneratedClientContext]:
# Client is the original issue document; Customer is a non-collision control.
spec = ISSUE_SPEC.replace("Client", request.param)
with generate_client_from_inline_spec(
spec, base_module="minimal_repro_client", add_missing_sections=False
) as generated:
yield generated


def test_endpoint_types_and_requests(generated_client: GeneratedClientContext) -> None:
endpoint = generated_client.import_module(".api.client.update_client")
transport = generated_client.import_module(".client")
model = get_type_hints(endpoint._parse_response)["return"]
models = generated_client.import_module(".models")
response_type = generated_client.import_symbol(".types", "Response")
body_type = generated_client.import_symbol(".models", "UpdateClientBody")
package = generated_client.import_module("")

# Compare with public exports, independently of endpoint import aliases.
schema_model = next(
getattr(models, name)
for name in ("Client", "AuthenticatedClient", "Response", "HTTPStatus", "Customer")
if hasattr(models, name)
)
assert model == schema_model | None
assert package.Client is transport.Client
assert package.AuthenticatedClient is transport.AuthenticatedClient
assert schema_model is not transport.Client
assert schema_model is not transport.AuthenticatedClient
assert schema_model is not response_type

for function_name in ("_parse_response", "_build_response", "sync", "sync_detailed", "asyncio", "asyncio_detailed"):
hints = get_type_hints(getattr(endpoint, function_name))
assert hints["client"] == transport.AuthenticatedClient | transport.Client
detailed = function_name == "_build_response" or function_name.endswith("_detailed")
assert hints["return"] == (response_type[schema_model] if detailed else schema_model | None)

def handle_request(request: httpx.Request) -> httpx.Response:
assert request.method == "PATCH"
assert request.url.path == "/api/clients/one"
assert json.loads(request.content) == {"name": "Ada"}
return httpx.Response(200, json={"name": "Ada"})

def check_result(result, detailed: bool) -> None:
if detailed:
assert type(result) is response_type
assert result.status_code is HTTPStatus.OK
result = result.parsed
assert type(result) is schema_model
assert result.name == "Ada"

with httpx.Client(base_url="https://example.test", transport=httpx.MockTransport(handle_request)) as http_client:
client = transport.Client(base_url="https://example.test").set_httpx_client(http_client)
for operation in (endpoint.sync, endpoint.sync_detailed):
check_result(
operation("one", client=client, body=body_type(name="Ada")), operation is endpoint.sync_detailed
)

async def call_async() -> None:
async with httpx.AsyncClient(
base_url="https://example.test", transport=httpx.MockTransport(handle_request)
) as http_client:
client = transport.AuthenticatedClient(
base_url="https://example.test", token="token"
).set_async_httpx_client(http_client)
for operation in (endpoint.asyncio, endpoint.asyncio_detailed):
check_result(
await operation("one", client=client, body=body_type(name="Ada")),
operation is endpoint.asyncio_detailed,
)

asyncio.run(call_async())


def test_authenticated_endpoint_client_type() -> None:
spec = (
ISSUE_SPEC.replace("Client", "AuthenticatedClient")
.replace("operationId: update_client", "operationId: update_client\n security: [{bearer: []}]")
.replace("components:", "components:\n securitySchemes:\n bearer:\n type: http\n scheme: bearer")
)
with generate_client_from_inline_spec(
spec, base_module="minimal_repro_client", add_missing_sections=False
) as generated:
endpoint = generated.import_module(".api.client.update_client")
transport = generated.import_symbol(".client", "AuthenticatedClient")
for function_name in ("sync", "sync_detailed", "asyncio", "asyncio_detailed"):
assert get_type_hints(getattr(endpoint, function_name))["client"] is transport
Original file line number Diff line number Diff line change
@@ -1,11 +1,12 @@
from http import HTTPStatus
from http import HTTPStatus as _HTTPStatus
from typing import Any

import httpx

from ... import errors
from ...client import AuthenticatedClient, Client
from ...types import Response
from ...client import AuthenticatedClient as _AuthenticatedClient
from ...client import Client as _Client
from ...types import Response as _Response


def _get_kwargs() -> dict[str, Any]:
Expand All @@ -18,16 +19,16 @@ def _get_kwargs() -> dict[str, Any]:
return _kwargs


def _parse_response(*, client: AuthenticatedClient | Client, response: httpx.Response) -> Any | None:
def _parse_response(*, client: _AuthenticatedClient | _Client, response: httpx.Response) -> Any | None:
if client.raise_on_unexpected_status:
raise errors.UnexpectedStatus(response.status_code, response.content)
else:
return None


def _build_response(*, client: AuthenticatedClient | Client, response: httpx.Response) -> Response[Any]:
return Response(
status_code=HTTPStatus(response.status_code),
def _build_response(*, client: _AuthenticatedClient | _Client, response: httpx.Response) -> _Response[Any]:
return _Response(
status_code=_HTTPStatus(response.status_code),
content=response.content,
headers=response.headers,
parsed=_parse_response(client=client, response=response),
Expand All @@ -36,8 +37,8 @@ def _build_response(*, client: AuthenticatedClient | Client, response: httpx.Res

def sync_detailed(
*,
client: AuthenticatedClient | Client,
) -> Response[Any]:
client: _AuthenticatedClient | _Client,
) -> _Response[Any]:
"""
Raises:
errors.UnexpectedStatus: If the server returns an undocumented status code and Client.raise_on_unexpected_status is True.
Expand All @@ -58,8 +59,8 @@ def sync_detailed(

async def asyncio_detailed(
*,
client: AuthenticatedClient | Client,
) -> Response[Any]:
client: _AuthenticatedClient | _Client,
) -> _Response[Any]:
"""
Raises:
errors.UnexpectedStatus: If the server returns an undocumented status code and Client.raise_on_unexpected_status is True.
Expand Down
Original file line number Diff line number Diff line change
@@ -1,12 +1,13 @@
from http import HTTPStatus
from http import HTTPStatus as _HTTPStatus
from typing import Any, cast

import httpx

from ... import errors
from ...client import AuthenticatedClient, Client
from ...client import AuthenticatedClient as _AuthenticatedClient
from ...client import Client as _Client
from ...models.misc_metadata_escapes_body import MiscMetadataEscapesBody
from ...types import Response
from ...types import Response as _Response


def _get_kwargs(
Expand All @@ -28,7 +29,7 @@ def _get_kwargs(
return _kwargs


def _parse_response(*, client: AuthenticatedClient | Client, response: httpx.Response) -> str | None:
def _parse_response(*, client: _AuthenticatedClient | _Client, response: httpx.Response) -> str | None:
if response.status_code == 200:
response_200 = cast(str, response.json())
return response_200
Expand All @@ -39,9 +40,9 @@ def _parse_response(*, client: AuthenticatedClient | Client, response: httpx.Res
return None


def _build_response(*, client: AuthenticatedClient | Client, response: httpx.Response) -> Response[str]:
return Response(
status_code=HTTPStatus(response.status_code),
def _build_response(*, client: _AuthenticatedClient | _Client, response: httpx.Response) -> _Response[str]:
return _Response(
status_code=_HTTPStatus(response.status_code),
content=response.content,
headers=response.headers,
parsed=_parse_response(client=client, response=response),
Expand All @@ -50,9 +51,9 @@ def _build_response(*, client: AuthenticatedClient | Client, response: httpx.Res

def sync_detailed(
*,
client: AuthenticatedClient | Client,
client: _AuthenticatedClient | _Client,
body: MiscMetadataEscapesBody,
) -> Response[str]:
) -> _Response[str]:
"""
Args:
body (MiscMetadataEscapesBody):
Expand All @@ -78,7 +79,7 @@ def sync_detailed(

def sync(
*,
client: AuthenticatedClient | Client,
client: _AuthenticatedClient | _Client,
body: MiscMetadataEscapesBody,
) -> str | None:
"""
Expand All @@ -101,9 +102,9 @@ def sync(

async def asyncio_detailed(
*,
client: AuthenticatedClient | Client,
client: _AuthenticatedClient | _Client,
body: MiscMetadataEscapesBody,
) -> Response[str]:
) -> _Response[str]:
"""
Args:
body (MiscMetadataEscapesBody):
Expand All @@ -127,7 +128,7 @@ async def asyncio_detailed(

async def asyncio(
*,
client: AuthenticatedClient | Client,
client: _AuthenticatedClient | _Client,
body: MiscMetadataEscapesBody,
) -> str | None:
"""
Expand Down
Original file line number Diff line number Diff line change
@@ -1,12 +1,13 @@
from http import HTTPStatus
from http import HTTPStatus as _HTTPStatus
from typing import Any

import httpx

from ... import errors
from ...client import AuthenticatedClient, Client
from ...client import AuthenticatedClient as _AuthenticatedClient
from ...client import Client as _Client
from ...models.non_string_example_body import NonStringExampleBody
from ...types import Response
from ...types import Response as _Response


def _get_kwargs(
Expand All @@ -28,7 +29,7 @@ def _get_kwargs(
return _kwargs


def _parse_response(*, client: AuthenticatedClient | Client, response: httpx.Response) -> Any | None:
def _parse_response(*, client: _AuthenticatedClient | _Client, response: httpx.Response) -> Any | None:
if response.status_code == 200:
return None

Expand All @@ -38,9 +39,9 @@ def _parse_response(*, client: AuthenticatedClient | Client, response: httpx.Res
return None


def _build_response(*, client: AuthenticatedClient | Client, response: httpx.Response) -> Response[Any]:
return Response(
status_code=HTTPStatus(response.status_code),
def _build_response(*, client: _AuthenticatedClient | _Client, response: httpx.Response) -> _Response[Any]:
return _Response(
status_code=_HTTPStatus(response.status_code),
content=response.content,
headers=response.headers,
parsed=_parse_response(client=client, response=response),
Expand All @@ -49,9 +50,9 @@ def _build_response(*, client: AuthenticatedClient | Client, response: httpx.Res

def sync_detailed(
*,
client: AuthenticatedClient | Client,
client: _AuthenticatedClient | _Client,
body: NonStringExampleBody,
) -> Response[Any]:
) -> _Response[Any]:
"""
Args:
body (NonStringExampleBody):
Expand All @@ -77,9 +78,9 @@ def sync_detailed(

async def asyncio_detailed(
*,
client: AuthenticatedClient | Client,
client: _AuthenticatedClient | _Client,
body: NonStringExampleBody,
) -> Response[Any]:
) -> _Response[Any]:
"""
Args:
body (NonStringExampleBody):
Expand Down
Loading