Skip to content

Commit a253d2a

Browse files
authored
Merge pull request #610 from koxudaxi/fix-201-response-generation
Fix 201 response generation
2 parents 29a4368 + 3a218b5 commit a253d2a

20 files changed

Lines changed: 224 additions & 68 deletions

File tree

.github/workflows/test.yml

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -302,6 +302,7 @@ jobs:
302302
files: .tox/coverage.xml
303303
flags: unittests
304304
fail_ci_if_error: true
305+
use_pypi: true
305306
env:
306307
CODECOV_TOKEN: ${{ secrets.CODECOV_TOKEN }}
307308
- name: Fail if coverage check failed

fastapi_code_generator/parser.py

Lines changed: 131 additions & 31 deletions
Original file line numberDiff line numberDiff line change
@@ -48,6 +48,12 @@
4848
from pydantic import BaseModel, ConfigDict, ValidationInfo
4949

5050
RE_APPLICATION_JSON_PATTERN: Pattern[str] = re.compile(r'^application/.*json$')
51+
ResponseStatusCode = Union[str, int]
52+
ResponseDefinitions = Mapping[
53+
ResponseStatusCode, Union[ResponseObject, ReferenceObject]
54+
]
55+
ResponseStatusLookup = Mapping[ResponseStatusCode, object]
56+
ParsedResponseDataTypes = Mapping[ResponseStatusCode, Dict[str, DataType]]
5157
RE_SNAKECASE_REPLACE_PATTERN: Pattern[str] = re.compile(r"[\-\.\s]")
5258
RE_UPPERCASE_PATTERN: Pattern[str] = re.compile(r"[A-Z]")
5359
RE_CAMELCASE_STRIP_PATTERN: Pattern[str] = re.compile(r"\w[\s\W]+\w")
@@ -722,43 +728,137 @@ def _is_upload_file_schema(self, schema: Any) -> bool:
722728
and schema.format == 'binary'
723729
)
724730

725-
def parse_responses( # type: ignore[override]
731+
def _parse_success_status_code(self, status_code: object) -> int | None:
732+
if not (status_code_text := str(status_code)).isdigit():
733+
return None
734+
parsed_status_code = int(status_code_text)
735+
if 200 <= parsed_status_code < 300:
736+
return parsed_status_code
737+
return None
738+
739+
def _get_success_status_codes(self, responses: ResponseDefinitions) -> list[int]:
740+
return sorted(
741+
parsed_status_code
742+
for status_code in responses
743+
if (parsed_status_code := self._parse_success_status_code(status_code))
744+
is not None
745+
)
746+
747+
def _find_response_status_code_key(
748+
self, responses: ResponseStatusLookup, status_code: int
749+
) -> ResponseStatusCode | None:
750+
if status_code in responses:
751+
return status_code
752+
if (status_code_text := str(status_code)) in responses:
753+
return status_code_text
754+
return None
755+
756+
def _get_response_data_types(
757+
self, data_types: ParsedResponseDataTypes, status_code: int
758+
) -> Dict[str, DataType] | None:
759+
if (
760+
status_code_key := self._find_response_status_code_key(
761+
data_types, status_code
762+
)
763+
) is None:
764+
return None
765+
return data_types[status_code_key] or None
766+
767+
def _select_primary_response_status_code(
768+
self,
769+
responses: ResponseDefinitions,
770+
data_types: ParsedResponseDataTypes,
771+
success_status_codes: list[int],
772+
) -> ResponseStatusCode | None:
773+
if status_code_key := self._find_response_status_code_key(responses, 200):
774+
return status_code_key
775+
for status_code in success_status_codes:
776+
match self._get_response_data_types(data_types, status_code):
777+
case response_data_types if response_data_types:
778+
return self._find_response_status_code_key(data_types, status_code)
779+
case _:
780+
continue
781+
return None
782+
783+
def _get_primary_response_data_type(
784+
self,
785+
status_code: ResponseStatusCode | None,
786+
data_types: ParsedResponseDataTypes,
787+
) -> DataType:
788+
if (
789+
status_code is None
790+
or (parsed_status_code := self._parse_success_status_code(status_code))
791+
is None
792+
):
793+
return DataType(type='None')
794+
if not (
795+
response_data_types := self._get_response_data_types(
796+
data_types, parsed_status_code
797+
)
798+
):
799+
return DataType(type='None')
800+
data_type = next(iter(response_data_types.values()))
801+
data_type = self._collapse_root_model(data_type)
802+
self.data_types.append(data_type)
803+
return data_type
804+
805+
def _select_route_status_code(
806+
self,
807+
primary_status_code: ResponseStatusCode | None,
808+
data_types: ParsedResponseDataTypes,
809+
success_status_codes: list[int],
810+
) -> int | None:
811+
primary_status_code_value = (
812+
self._parse_success_status_code(primary_status_code)
813+
if primary_status_code is not None
814+
else None
815+
)
816+
if primary_status_code_value and primary_status_code_value != 200:
817+
return primary_status_code_value
818+
match success_status_codes:
819+
case [
820+
status_code
821+
] if status_code != 200 and not self._get_response_data_types(
822+
data_types, status_code
823+
):
824+
return status_code
825+
case _:
826+
return None
827+
828+
def parse_responses(
726829
self,
727830
name: str,
728-
responses: Dict[str, Union[ResponseObject, ReferenceObject]],
831+
responses: Dict[ResponseStatusCode, Union[ResponseObject, ReferenceObject]],
729832
path: List[str],
730-
) -> Dict[Union[str, int], Dict[str, DataType]]:
731-
data_types = super().parse_responses(name, responses, path) # type: ignore[arg-type]
732-
status_code_200 = data_types.get('200')
733-
if status_code_200:
734-
data_type = list(status_code_200.values())[0]
735-
data_type = self._collapse_root_model(data_type)
736-
self.data_types.append(data_type)
737-
else:
738-
data_type = DataType(type='None')
833+
) -> Dict[ResponseStatusCode, Dict[str, DataType]]:
834+
data_types = super().parse_responses(name, responses, path)
835+
success_status_codes = self._get_success_status_codes(responses)
836+
primary_status_code = self._select_primary_response_status_code(
837+
responses, data_types, success_status_codes
838+
)
839+
data_type = self._get_primary_response_data_type(
840+
primary_status_code, data_types
841+
)
739842
type_hint = data_type.type_hint # TODO: change to lazy loading
740843
self._temporary_operation['response'] = type_hint
741-
success_status_codes = [
742-
int(status_code)
743-
for status_code in responses
744-
if str(status_code).isdigit() and 200 <= int(status_code) < 300
745-
]
746-
if '200' not in responses and success_status_codes:
747-
selected_status_code = min(success_status_codes)
748-
if selected_status_code == 204 and not data_types.get(
749-
str(selected_status_code)
750-
):
751-
self._temporary_operation['status_code'] = selected_status_code
844+
if status_code := self._select_route_status_code(
845+
primary_status_code, data_types, success_status_codes
846+
):
847+
self._temporary_operation['status_code'] = status_code
752848
return_types = {type_hint: data_type}
753-
for status_code, additional_responses in data_types.items():
754-
if status_code != '200' and additional_responses: # 200 is processed above
755-
data_type = list(additional_responses.values())[0]
756-
self.data_types.append(data_type)
757-
type_hint = data_type.type_hint # TODO: change to lazy loading
758-
self._temporary_operation.setdefault('additional_responses', {})[
759-
status_code
760-
] = {'model': type_hint}
761-
return_types[type_hint] = data_type
849+
for additional_status_code, additional_responses in data_types.items():
850+
is_primary_response = primary_status_code is not None and str(
851+
additional_status_code
852+
) == str(primary_status_code)
853+
if is_primary_response or not additional_responses:
854+
continue
855+
data_type = next(iter(additional_responses.values()))
856+
self.data_types.append(data_type)
857+
type_hint = data_type.type_hint # TODO: change to lazy loading
858+
self._temporary_operation.setdefault('additional_responses', {})[
859+
additional_status_code
860+
] = {'model': type_hint}
861+
return_types[type_hint] = data_type
762862
if len(return_types) == 1:
763863
return_type = next(iter(return_types.values()))
764864
else:

tests/data/expected/openapi/coverage/callbacks/main.py

Lines changed: 2 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -4,8 +4,6 @@
44

55
from __future__ import annotations
66

7-
from typing import Optional
8-
97
from fastapi import FastAPI
108

119
from .models import Ack, EventPayload, Subscription, SubscriptionRequest
@@ -16,8 +14,6 @@
1614
)
1715

1816

19-
@app.post(
20-
'/subscriptions', response_model=None, responses={'201': {'model': Subscription}}
21-
)
22-
def create_subscription(body: SubscriptionRequest) -> Optional[Subscription]:
17+
@app.post('/subscriptions', response_model=Subscription, status_code=201)
18+
def create_subscription(body: SubscriptionRequest) -> Subscription:
2319
pass

tests/data/expected/openapi/coverage/callbacks_with_operation_id/main.py

Lines changed: 2 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -4,8 +4,6 @@
44

55
from __future__ import annotations
66

7-
from typing import Optional
8-
97
from fastapi import FastAPI
108

119
from .models import Ack, EventPayload, Subscription, SubscriptionRequest
@@ -16,8 +14,6 @@
1614
)
1715

1816

19-
@app.post(
20-
'/subscriptions', response_model=None, responses={'201': {'model': Subscription}}
21-
)
22-
def create_subscription(body: SubscriptionRequest) -> Optional[Subscription]:
17+
@app.post('/subscriptions', response_model=Subscription, status_code=201)
18+
def create_subscription(body: SubscriptionRequest) -> Subscription:
2319
pass

tests/data/expected/openapi/coverage/model_options/main.py

Lines changed: 13 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -51,7 +51,11 @@ def get_foo(foo: Optional[str] = None) -> str:
5151

5252

5353
@app.post(
54-
'/food', response_model=None, responses={'default': {'model': str}}, tags=['pets']
54+
'/food',
55+
response_model=None,
56+
status_code=201,
57+
responses={'default': {'model': str}},
58+
tags=['pets'],
5559
)
5660
def post_food(body: str) -> Optional[str]:
5761
"""
@@ -88,7 +92,11 @@ def list_pets(
8892

8993

9094
@app.post(
91-
'/pets', response_model=None, responses={'default': {'model': Error}}, tags=['pets']
95+
'/pets',
96+
response_model=None,
97+
status_code=201,
98+
responses={'default': {'model': Error}},
99+
tags=['pets'],
92100
)
93101
def post_pets(body: PetForm) -> Optional[Error]:
94102
"""
@@ -113,6 +121,7 @@ def show_pet_by_id(pet_id: str = Path(..., alias='petId')) -> Union[Pet, Error]:
113121
@app.put(
114122
'/pets/{petId}',
115123
response_model=None,
124+
status_code=201,
116125
responses={'default': {'model': Error}},
117126
tags=['pets'],
118127
)
@@ -130,7 +139,7 @@ def get_user() -> UserGetResponse:
130139
pass
131140

132141

133-
@app.post('/user', response_model=None, tags=['user'])
142+
@app.post('/user', response_model=None, status_code=201, tags=['user'])
134143
def post_user(body: UserPostRequest) -> None:
135144
pass
136145

@@ -140,7 +149,7 @@ def get_users() -> List[UsersGetResponseItem]:
140149
pass
141150

142151

143-
@app.post('/users', response_model=None, tags=['user'])
152+
@app.post('/users', response_model=None, status_code=201, tags=['user'])
144153
def post_users(body: List[UsersPostRequestItem]) -> None:
145154
pass
146155

tests/data/expected/openapi/coverage/non_200_responses/main.py

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,7 @@
44

55
from __future__ import annotations
66

7-
from typing import Optional, Union
7+
from typing import Union
88

99
from fastapi import FastAPI
1010

@@ -18,8 +18,9 @@
1818

1919
@app.post(
2020
'/jobs',
21-
response_model=None,
22-
responses={'201': {'model': JobCreated}, '404': {'model': Error}},
21+
response_model=JobCreated,
22+
status_code=201,
23+
responses={'404': {'model': Error}},
2324
)
24-
def create_job() -> Optional[Union[JobCreated, Error]]:
25+
def create_job() -> Union[JobCreated, Error]:
2526
pass

tests/data/expected/openapi/coverage/non_200_status_code/main.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -15,3 +15,8 @@
1515
@app.delete('/items/{item_id}', response_model=None, status_code=204)
1616
def delete_item(item_id: int) -> None:
1717
pass
18+
19+
20+
@app.post('/jobs/{job_id}/start', response_model=None, status_code=202)
21+
def start_job(job_id: int) -> None:
22+
pass

tests/data/expected/openapi/default_template/body_and_parameters/main.py

Lines changed: 13 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -51,7 +51,11 @@ def get_foo(foo: Optional[str] = None) -> str:
5151

5252

5353
@app.post(
54-
'/food', response_model=None, responses={'default': {'model': str}}, tags=['pets']
54+
'/food',
55+
response_model=None,
56+
status_code=201,
57+
responses={'default': {'model': str}},
58+
tags=['pets'],
5559
)
5660
def post_food(body: str) -> Optional[str]:
5761
"""
@@ -88,7 +92,11 @@ def list_pets(
8892

8993

9094
@app.post(
91-
'/pets', response_model=None, responses={'default': {'model': Error}}, tags=['pets']
95+
'/pets',
96+
response_model=None,
97+
status_code=201,
98+
responses={'default': {'model': Error}},
99+
tags=['pets'],
92100
)
93101
def post_pets(body: PetForm) -> Optional[Error]:
94102
"""
@@ -113,6 +121,7 @@ def show_pet_by_id(pet_id: str = Path(..., alias='petId')) -> Union[Pet, Error]:
113121
@app.put(
114122
'/pets/{petId}',
115123
response_model=None,
124+
status_code=201,
116125
responses={'default': {'model': Error}},
117126
tags=['pets'],
118127
)
@@ -130,7 +139,7 @@ def get_user() -> UserGetResponse:
130139
pass
131140

132141

133-
@app.post('/user', response_model=None, tags=['user'])
142+
@app.post('/user', response_model=None, status_code=201, tags=['user'])
134143
def post_user(body: UserPostRequest) -> None:
135144
pass
136145

@@ -140,7 +149,7 @@ def get_users() -> List[UsersGetResponseItem]:
140149
pass
141150

142151

143-
@app.post('/users', response_model=None, tags=['user'])
152+
@app.post('/users', response_model=None, status_code=201, tags=['user'])
144153
def post_users(body: List[UsersPostRequestItem]) -> None:
145154
pass
146155

tests/data/expected/openapi/default_template/duplicate_request_param/main.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -15,7 +15,9 @@
1515
)
1616

1717

18-
@app.post('/pets/{id}/image/octet-stream', response_model=None, tags=['pets'])
18+
@app.post(
19+
'/pets/{id}/image/octet-stream', response_model=None, status_code=201, tags=['pets']
20+
)
1921
def upload_pet_image_with_duplicate_request(
2022
id: str, request: Optional[str] = None
2123
) -> None:

tests/data/expected/openapi/default_template/same_response_model_for_different_status_codes/main.py

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -33,7 +33,11 @@ def list_pets(limit: Optional[int] = None) -> Union[List[Pet], Error]:
3333

3434

3535
@app.post(
36-
'/pets', response_model=None, responses={'default': {'model': Error}}, tags=['pets']
36+
'/pets',
37+
response_model=None,
38+
status_code=201,
39+
responses={'default': {'model': Error}},
40+
tags=['pets'],
3741
)
3842
def create_pets() -> Optional[Error]:
3943
"""

0 commit comments

Comments
 (0)