|
48 | 48 | from pydantic import BaseModel, ConfigDict, ValidationInfo |
49 | 49 |
|
50 | 50 | 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]] |
51 | 57 | RE_SNAKECASE_REPLACE_PATTERN: Pattern[str] = re.compile(r"[\-\.\s]") |
52 | 58 | RE_UPPERCASE_PATTERN: Pattern[str] = re.compile(r"[A-Z]") |
53 | 59 | 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: |
722 | 728 | and schema.format == 'binary' |
723 | 729 | ) |
724 | 730 |
|
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( |
726 | 829 | self, |
727 | 830 | name: str, |
728 | | - responses: Dict[str, Union[ResponseObject, ReferenceObject]], |
| 831 | + responses: Dict[ResponseStatusCode, Union[ResponseObject, ReferenceObject]], |
729 | 832 | 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 | + ) |
739 | 842 | type_hint = data_type.type_hint # TODO: change to lazy loading |
740 | 843 | 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 |
752 | 848 | 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 |
762 | 862 | if len(return_types) == 1: |
763 | 863 | return_type = next(iter(return_types.values())) |
764 | 864 | else: |
|
0 commit comments