fastapi-typed-client 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- fastapi_typed_client/__init__.py +28 -0
- fastapi_typed_client/__main__.py +4 -0
- fastapi_typed_client/__version__.py +4 -0
- fastapi_typed_client/_core.py +106 -0
- fastapi_typed_client/_generator.py +465 -0
- fastapi_typed_client/_parser.py +229 -0
- fastapi_typed_client/_utils/__init__.py +27 -0
- fastapi_typed_client/_utils/import_.py +271 -0
- fastapi_typed_client/_utils/string.py +29 -0
- fastapi_typed_client/cli.py +125 -0
- fastapi_typed_client/client.py +309 -0
- fastapi_typed_client/py.typed +1 -0
- fastapi_typed_client-0.1.0.dist-info/METADATA +408 -0
- fastapi_typed_client-0.1.0.dist-info/RECORD +17 -0
- fastapi_typed_client-0.1.0.dist-info/WHEEL +4 -0
- fastapi_typed_client-0.1.0.dist-info/entry_points.txt +3 -0
- fastapi_typed_client-0.1.0.dist-info/licenses/LICENSE +201 -0
|
@@ -0,0 +1,28 @@
|
|
|
1
|
+
from . import cli, client
|
|
2
|
+
from .__version__ import __version__
|
|
3
|
+
from ._core import generate_fastapi_typed_client
|
|
4
|
+
from .client import (
|
|
5
|
+
FASTAPI_CLIENT_NOT_REQUIRED,
|
|
6
|
+
FastAPIClientAsyncBase,
|
|
7
|
+
FastAPIClientBase,
|
|
8
|
+
FastAPIClientExtensions,
|
|
9
|
+
FastAPIClientHTTPValidationError,
|
|
10
|
+
FastAPIClientNotDefaultStatusError,
|
|
11
|
+
FastAPIClientResult,
|
|
12
|
+
FastAPIClientValidationError,
|
|
13
|
+
)
|
|
14
|
+
|
|
15
|
+
__all__ = [
|
|
16
|
+
"FASTAPI_CLIENT_NOT_REQUIRED",
|
|
17
|
+
"FastAPIClientAsyncBase",
|
|
18
|
+
"FastAPIClientBase",
|
|
19
|
+
"FastAPIClientExtensions",
|
|
20
|
+
"FastAPIClientHTTPValidationError",
|
|
21
|
+
"FastAPIClientNotDefaultStatusError",
|
|
22
|
+
"FastAPIClientResult",
|
|
23
|
+
"FastAPIClientValidationError",
|
|
24
|
+
"__version__",
|
|
25
|
+
"cli",
|
|
26
|
+
"client",
|
|
27
|
+
"generate_fastapi_typed_client",
|
|
28
|
+
]
|
|
@@ -0,0 +1,106 @@
|
|
|
1
|
+
from collections.abc import Iterable
|
|
2
|
+
from os import PathLike
|
|
3
|
+
from pathlib import Path
|
|
4
|
+
|
|
5
|
+
from fastapi import APIRouter, FastAPI
|
|
6
|
+
|
|
7
|
+
from ._generator import ClientCodeGenerator
|
|
8
|
+
from ._parser import parse_routes
|
|
9
|
+
from ._utils import load_import, to_snake_case, to_upper_camel_case
|
|
10
|
+
from .client import (
|
|
11
|
+
FastAPIClientAsyncBase,
|
|
12
|
+
FastAPIClientBase,
|
|
13
|
+
FastAPIClientExtensions,
|
|
14
|
+
FastAPIClientHTTPValidationError,
|
|
15
|
+
FastAPIClientNotDefaultStatusError,
|
|
16
|
+
FastAPIClientResult,
|
|
17
|
+
FastAPIClientValidationError,
|
|
18
|
+
)
|
|
19
|
+
|
|
20
|
+
# Reserve these names to avoid confusion.
|
|
21
|
+
_RESERVED_TITLES = (
|
|
22
|
+
FastAPIClientExtensions.__name__,
|
|
23
|
+
FastAPIClientResult.__name__,
|
|
24
|
+
FastAPIClientValidationError.__name__,
|
|
25
|
+
FastAPIClientHTTPValidationError.__name__,
|
|
26
|
+
FastAPIClientNotDefaultStatusError.__name__,
|
|
27
|
+
FastAPIClientBase.__name__,
|
|
28
|
+
FastAPIClientAsyncBase.__name__,
|
|
29
|
+
"FASTAPI_CLIENT_NOT_REQUIRED",
|
|
30
|
+
)
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
def generate_fastapi_typed_client(
|
|
34
|
+
app_or_import_str: FastAPI | APIRouter | str,
|
|
35
|
+
*,
|
|
36
|
+
output_path: PathLike[str] | str | None = None,
|
|
37
|
+
title: str | None = None,
|
|
38
|
+
async_: bool = False,
|
|
39
|
+
import_barrier: str | Iterable[str] | None = None,
|
|
40
|
+
import_client_base: bool = False,
|
|
41
|
+
raise_if_not_default_status: bool = False,
|
|
42
|
+
_add_test_markers: bool = False,
|
|
43
|
+
) -> None:
|
|
44
|
+
app = (
|
|
45
|
+
_import_app(app_or_import_str)
|
|
46
|
+
if isinstance(app_or_import_str, str)
|
|
47
|
+
else app_or_import_str
|
|
48
|
+
)
|
|
49
|
+
|
|
50
|
+
if not title:
|
|
51
|
+
title = (
|
|
52
|
+
to_upper_camel_case(app.title) + "Client"
|
|
53
|
+
if isinstance(app, FastAPI)
|
|
54
|
+
else "FastAPIClient"
|
|
55
|
+
)
|
|
56
|
+
if not title.isidentifier():
|
|
57
|
+
raise RuntimeError(f"Title `{title}` is not a valid Python identifier.")
|
|
58
|
+
if title in _RESERVED_TITLES:
|
|
59
|
+
raise RuntimeError(f"Title `{title}` is reserved.")
|
|
60
|
+
|
|
61
|
+
output_path = (
|
|
62
|
+
Path(output_path)
|
|
63
|
+
if output_path
|
|
64
|
+
else Path(to_snake_case(title).replace("fast_api", "fastapi") + ".py")
|
|
65
|
+
)
|
|
66
|
+
|
|
67
|
+
if not import_barrier:
|
|
68
|
+
import_barrier = []
|
|
69
|
+
elif isinstance(import_barrier, str):
|
|
70
|
+
import_barrier = [import_barrier]
|
|
71
|
+
|
|
72
|
+
routes = parse_routes(app.routes)
|
|
73
|
+
code = ClientCodeGenerator(
|
|
74
|
+
title,
|
|
75
|
+
async_,
|
|
76
|
+
import_barrier,
|
|
77
|
+
import_client_base,
|
|
78
|
+
raise_if_not_default_status,
|
|
79
|
+
_add_test_markers,
|
|
80
|
+
).generate(routes)
|
|
81
|
+
|
|
82
|
+
output_path.write_text(code, encoding="utf-8")
|
|
83
|
+
|
|
84
|
+
|
|
85
|
+
def _import_app(app_import_str: str) -> FastAPI | APIRouter:
|
|
86
|
+
module, _, name = app_import_str.partition(":")
|
|
87
|
+
if not module or not name:
|
|
88
|
+
raise RuntimeError(
|
|
89
|
+
"App import string must be in the format `module.submodule:app_name`."
|
|
90
|
+
)
|
|
91
|
+
|
|
92
|
+
try:
|
|
93
|
+
obj = load_import(module, name)
|
|
94
|
+
except ModuleNotFoundError as e:
|
|
95
|
+
if e.name != module:
|
|
96
|
+
raise e from None
|
|
97
|
+
raise RuntimeError(f"Could not import module `{module}`.") from e
|
|
98
|
+
except AttributeError as e:
|
|
99
|
+
raise RuntimeError(f"Attribute `{name}` not found in module `{module}`.") from e
|
|
100
|
+
|
|
101
|
+
if not isinstance(obj, FastAPI) and not isinstance(obj, APIRouter):
|
|
102
|
+
raise RuntimeError(
|
|
103
|
+
f"App import string is not a FastAPI app, but a `{type(obj)}`."
|
|
104
|
+
)
|
|
105
|
+
|
|
106
|
+
return obj
|
|
@@ -0,0 +1,465 @@
|
|
|
1
|
+
from collections import defaultdict
|
|
2
|
+
from collections.abc import AsyncIterator, Collection, Iterable, Iterator, Sequence
|
|
3
|
+
from collections.abc import Set as AbstractSet
|
|
4
|
+
from enum import Enum, auto
|
|
5
|
+
from functools import cache
|
|
6
|
+
from http import HTTPMethod, HTTPStatus
|
|
7
|
+
from importlib.util import find_spec
|
|
8
|
+
from inspect import getsource
|
|
9
|
+
from sys import stdlib_module_names
|
|
10
|
+
from typing import Any, Literal, NamedTuple, overload
|
|
11
|
+
from warnings import warn
|
|
12
|
+
|
|
13
|
+
from ._parser import Route, RouteParam, RouteParamKind, RouteResponse
|
|
14
|
+
from ._utils import Import, ImportRegistry, dq_str_repr, indent, to_constant_case
|
|
15
|
+
from .client import (
|
|
16
|
+
_IMPORTS,
|
|
17
|
+
_IMPORTS_ASYNC_CLIENT,
|
|
18
|
+
_IMPORTS_SYNC_CLIENT,
|
|
19
|
+
_IMPORTS_TYPE_CHECKING,
|
|
20
|
+
_IMPORTS_VALIDATION_ERROR,
|
|
21
|
+
FastAPIClientAsyncBase,
|
|
22
|
+
FastAPIClientBase,
|
|
23
|
+
FastAPIClientExtensions,
|
|
24
|
+
FastAPIClientHTTPValidationError,
|
|
25
|
+
FastAPIClientNotDefaultStatusError,
|
|
26
|
+
FastAPIClientResult,
|
|
27
|
+
FastAPIClientValidationError,
|
|
28
|
+
)
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
class _Identifiers(NamedTuple):
|
|
32
|
+
client_extensions: str
|
|
33
|
+
result: str
|
|
34
|
+
validation_error: str
|
|
35
|
+
http_validation_error: str
|
|
36
|
+
not_default_status_error: str
|
|
37
|
+
not_required: str
|
|
38
|
+
base_class: str
|
|
39
|
+
client_class: str
|
|
40
|
+
|
|
41
|
+
def replace_in_code(self, code: str) -> str:
|
|
42
|
+
replacements = {
|
|
43
|
+
FastAPIClientExtensions.__name__: self.client_extensions,
|
|
44
|
+
FastAPIClientResult.__name__: self.result,
|
|
45
|
+
FastAPIClientValidationError.__name__: self.validation_error,
|
|
46
|
+
FastAPIClientHTTPValidationError.__name__: self.http_validation_error,
|
|
47
|
+
FastAPIClientNotDefaultStatusError.__name__: self.not_default_status_error,
|
|
48
|
+
"FASTAPI_CLIENT_NOT_REQUIRED": self.not_required,
|
|
49
|
+
FastAPIClientBase.__name__: self.base_class,
|
|
50
|
+
FastAPIClientAsyncBase.__name__: self.base_class,
|
|
51
|
+
}
|
|
52
|
+
for old, new in replacements.items():
|
|
53
|
+
code = code.replace(old, new)
|
|
54
|
+
return code
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
class _ImportGroup(Enum):
|
|
58
|
+
STDLIB = auto()
|
|
59
|
+
SITE_PACKAGE = auto()
|
|
60
|
+
LOCAL = auto()
|
|
61
|
+
|
|
62
|
+
|
|
63
|
+
class ClientCodeGenerator:
|
|
64
|
+
def __init__(
|
|
65
|
+
self,
|
|
66
|
+
title: str,
|
|
67
|
+
async_: bool,
|
|
68
|
+
import_barriers: Iterable[str],
|
|
69
|
+
import_client_base: bool,
|
|
70
|
+
raise_if_not_default_status: bool,
|
|
71
|
+
add_test_markers: bool,
|
|
72
|
+
) -> None:
|
|
73
|
+
self._title = title
|
|
74
|
+
self._async = async_
|
|
75
|
+
self._base_class = FastAPIClientBase if not async_ else FastAPIClientAsyncBase
|
|
76
|
+
self._import_client_base = import_client_base
|
|
77
|
+
self._raise_if_not_default_status = raise_if_not_default_status
|
|
78
|
+
self._add_test_markers = add_test_markers
|
|
79
|
+
self._impr = ImportRegistry()
|
|
80
|
+
for import_barrier in import_barriers:
|
|
81
|
+
self._impr.add_barrier(import_barrier)
|
|
82
|
+
self._idents = self._init_identifiers()
|
|
83
|
+
|
|
84
|
+
def _init_identifiers(self) -> _Identifiers:
|
|
85
|
+
if self._import_client_base:
|
|
86
|
+
return _Identifiers(
|
|
87
|
+
client_extensions=FastAPIClientExtensions.__name__,
|
|
88
|
+
result=FastAPIClientResult.__name__,
|
|
89
|
+
validation_error=FastAPIClientValidationError.__name__,
|
|
90
|
+
http_validation_error=FastAPIClientHTTPValidationError.__name__,
|
|
91
|
+
not_default_status_error=FastAPIClientNotDefaultStatusError.__name__,
|
|
92
|
+
not_required="FASTAPI_CLIENT_NOT_REQUIRED",
|
|
93
|
+
base_class=self._base_class.__name__,
|
|
94
|
+
client_class=self._title,
|
|
95
|
+
)
|
|
96
|
+
return _Identifiers(
|
|
97
|
+
client_extensions=f"{self._title}Extensions",
|
|
98
|
+
result=f"{self._title}Result",
|
|
99
|
+
validation_error=f"{self._title}ValidationError",
|
|
100
|
+
http_validation_error=f"{self._title}HTTPValidationError",
|
|
101
|
+
not_default_status_error=f"{self._title}NotDefaultStatusError",
|
|
102
|
+
not_required=(
|
|
103
|
+
to_constant_case(self._title).replace("FAST_API", "FASTAPI")
|
|
104
|
+
+ "_NOT_REQUIRED"
|
|
105
|
+
),
|
|
106
|
+
base_class=self._title,
|
|
107
|
+
client_class=self._title,
|
|
108
|
+
)
|
|
109
|
+
|
|
110
|
+
def generate(self, routes: Sequence[Route]) -> str:
|
|
111
|
+
for route in routes:
|
|
112
|
+
for param in route.params:
|
|
113
|
+
self._impr.add_reserved_ident(param.name)
|
|
114
|
+
|
|
115
|
+
codes = ["", "", self._get_boilerplate_code(routes)]
|
|
116
|
+
codes.extend(indent(self._get_route_code(route)) for route in routes)
|
|
117
|
+
# This relies on the side effects to self._impr of the previous code generating
|
|
118
|
+
# functions, so we can only call it at the end.
|
|
119
|
+
codes[0] = _ImportCodeGenerator(self._impr).generate()
|
|
120
|
+
|
|
121
|
+
return "\n".join(codes)
|
|
122
|
+
|
|
123
|
+
def _get_boilerplate_code(self, routes: Sequence[Route]) -> str:
|
|
124
|
+
return _BoilerplateCodeGenerator(
|
|
125
|
+
self._impr, self._base_class, self._idents, self._add_test_markers
|
|
126
|
+
).generate(routes, self._import_client_base)
|
|
127
|
+
|
|
128
|
+
def _get_route_code(self, route: Route) -> str:
|
|
129
|
+
return self._get_route_signature_code(route) + indent(
|
|
130
|
+
f"return {'await ' if self._async else ''}self._route_handler( # type: ignore\n"
|
|
131
|
+
f" path={dq_str_repr(route.path)},\n"
|
|
132
|
+
f" method={self._impr(HTTPMethod)}.{route.method.name},\n"
|
|
133
|
+
f" default_status={self._impr(HTTPStatus)}.{route.default_status.name},\n"
|
|
134
|
+
+ indent(self._get_models_dict_code(route.responses.values()))
|
|
135
|
+
+ indent(self._get_params_dicts_code(route.params))
|
|
136
|
+
+ indent(self._get_optional_params_code(route))
|
|
137
|
+
+ " raise_if_not_default_status=raise_if_not_default_status,\n"
|
|
138
|
+
" client_exts=client_exts,\n"
|
|
139
|
+
")\n"
|
|
140
|
+
)
|
|
141
|
+
|
|
142
|
+
def _get_route_signature_code(self, route: Route) -> str:
|
|
143
|
+
if len(route.responses) == 1:
|
|
144
|
+
return f"{self._get_route_overload_signature_code(route, route.responses.values(), None)}:\n"
|
|
145
|
+
|
|
146
|
+
return (
|
|
147
|
+
f"@{self._impr(overload)}\n"
|
|
148
|
+
f"{self._get_route_overload_signature_code(route, route.responses[route.default_status], True)}: ...\n"
|
|
149
|
+
f"@{self._impr(overload)}\n"
|
|
150
|
+
f"{self._get_route_overload_signature_code(route, route.responses.values(), False)}: ...\n"
|
|
151
|
+
f"{self._get_route_overload_signature_code(route, None, None)}:\n"
|
|
152
|
+
)
|
|
153
|
+
|
|
154
|
+
def _get_route_overload_signature_code(
|
|
155
|
+
self,
|
|
156
|
+
route: Route,
|
|
157
|
+
responses: RouteResponse | Collection[RouteResponse] | None,
|
|
158
|
+
raise_if_not_default_status: bool | None,
|
|
159
|
+
) -> str:
|
|
160
|
+
return (
|
|
161
|
+
f"{'async ' if self._async else ''}def {route.name}(\n"
|
|
162
|
+
+ " self,\n"
|
|
163
|
+
+ indent(self._get_route_specific_params_code(route.params))
|
|
164
|
+
+ indent(self._get_route_generic_params_code(raise_if_not_default_status))
|
|
165
|
+
+ ") -> "
|
|
166
|
+
+ self._get_route_responses_code(responses, route.is_streaming_json)
|
|
167
|
+
)
|
|
168
|
+
|
|
169
|
+
def _get_route_specific_params_code(self, params: Sequence[RouteParam]) -> str:
|
|
170
|
+
code = ""
|
|
171
|
+
for param in params:
|
|
172
|
+
code += f"{param.name}: {self._impr(param.type_)}"
|
|
173
|
+
if not param.required:
|
|
174
|
+
code += f" = {self._idents.not_required}"
|
|
175
|
+
code += ",\n"
|
|
176
|
+
return code
|
|
177
|
+
|
|
178
|
+
def _get_route_generic_params_code(
|
|
179
|
+
self, raise_if_not_default_status: bool | None
|
|
180
|
+
) -> str:
|
|
181
|
+
raise_if_not_default_status_str = {
|
|
182
|
+
True: self._impr(Literal[True]),
|
|
183
|
+
False: self._impr(Literal[False]),
|
|
184
|
+
None: "bool",
|
|
185
|
+
}[raise_if_not_default_status]
|
|
186
|
+
code = "*,\n"
|
|
187
|
+
code += f"raise_if_not_default_status: {raise_if_not_default_status_str}"
|
|
188
|
+
if (
|
|
189
|
+
raise_if_not_default_status is None
|
|
190
|
+
or raise_if_not_default_status == self._raise_if_not_default_status
|
|
191
|
+
):
|
|
192
|
+
code += f" = {self._raise_if_not_default_status!r}"
|
|
193
|
+
code += ",\n"
|
|
194
|
+
code += f"client_exts: {self._idents.client_extensions} | None = None,\n"
|
|
195
|
+
return code
|
|
196
|
+
|
|
197
|
+
def _get_route_responses_code(
|
|
198
|
+
self,
|
|
199
|
+
responses: RouteResponse | Collection[RouteResponse] | None,
|
|
200
|
+
is_streaming_json: bool,
|
|
201
|
+
) -> str:
|
|
202
|
+
if not responses:
|
|
203
|
+
return f"{self._idents.result}[{self._impr(HTTPStatus)}, {self._impr(Any)}]"
|
|
204
|
+
|
|
205
|
+
if isinstance(responses, RouteResponse):
|
|
206
|
+
responses = (responses,)
|
|
207
|
+
|
|
208
|
+
code = ""
|
|
209
|
+
if len(responses) > 1:
|
|
210
|
+
code += "(\n "
|
|
211
|
+
for i, response in enumerate(responses):
|
|
212
|
+
if i != 0:
|
|
213
|
+
code += "\n | "
|
|
214
|
+
response_type = response.type_
|
|
215
|
+
if i == 0 and is_streaming_json:
|
|
216
|
+
response_type = (Iterator if not self._async else AsyncIterator)[
|
|
217
|
+
response_type
|
|
218
|
+
]
|
|
219
|
+
code += (
|
|
220
|
+
f"{self._idents.result}["
|
|
221
|
+
f"{self._impr(Literal)}[{self._impr(HTTPStatus)}.{response.status.name}], "
|
|
222
|
+
f"{self._get_response_type_code(response_type)}"
|
|
223
|
+
"]"
|
|
224
|
+
)
|
|
225
|
+
if len(responses) > 1:
|
|
226
|
+
code += "\n)"
|
|
227
|
+
return code
|
|
228
|
+
|
|
229
|
+
def _get_response_type_code(self, type_: Any) -> str: # noqa: ANN401
|
|
230
|
+
if type_ is FastAPIClientHTTPValidationError:
|
|
231
|
+
return self._idents.http_validation_error
|
|
232
|
+
return self._impr(type_)
|
|
233
|
+
|
|
234
|
+
def _get_models_dict_code(self, responses: Collection[RouteResponse]) -> str:
|
|
235
|
+
lines = []
|
|
236
|
+
for response in responses:
|
|
237
|
+
status_str = f"{self._impr(HTTPStatus)}.{response.status.name}"
|
|
238
|
+
type_str = self._get_response_type_code(response.type_)
|
|
239
|
+
lines.append(f"{status_str}: {type_str},\n")
|
|
240
|
+
return f"models={{\n{indent(''.join(lines))}}},"
|
|
241
|
+
|
|
242
|
+
@staticmethod
|
|
243
|
+
def _get_params_dicts_code(params: Sequence[RouteParam]) -> str:
|
|
244
|
+
code = ""
|
|
245
|
+
for param_kind in RouteParamKind:
|
|
246
|
+
kind_params = [param for param in params if param.kind is param_kind]
|
|
247
|
+
if not kind_params:
|
|
248
|
+
continue
|
|
249
|
+
code += f"{param_kind.name.lower()}_params={{\n"
|
|
250
|
+
for param in kind_params:
|
|
251
|
+
code += f" {dq_str_repr(param.alias or param.name)}: {param.name},\n"
|
|
252
|
+
code += "},\n"
|
|
253
|
+
return code
|
|
254
|
+
|
|
255
|
+
@staticmethod
|
|
256
|
+
def _get_optional_params_code(route: Route) -> str:
|
|
257
|
+
code = ""
|
|
258
|
+
if route.is_body_embedded:
|
|
259
|
+
code += f"is_body_embedded={route.is_body_embedded},\n"
|
|
260
|
+
if route.is_streaming_json:
|
|
261
|
+
code += f"is_streaming_json={route.is_streaming_json},\n"
|
|
262
|
+
return code
|
|
263
|
+
|
|
264
|
+
|
|
265
|
+
class _BoilerplateCodeGenerator:
|
|
266
|
+
def __init__(
|
|
267
|
+
self,
|
|
268
|
+
impr: ImportRegistry,
|
|
269
|
+
base_class: type[FastAPIClientBase | FastAPIClientAsyncBase],
|
|
270
|
+
idents: _Identifiers,
|
|
271
|
+
add_test_markers: bool,
|
|
272
|
+
) -> None:
|
|
273
|
+
self._impr = impr
|
|
274
|
+
self._base_class = base_class
|
|
275
|
+
self._idents = idents
|
|
276
|
+
self._add_test_markers = add_test_markers
|
|
277
|
+
|
|
278
|
+
def generate(self, routes: Sequence[Route], import_client_base: bool) -> str:
|
|
279
|
+
has_not_required_params = any(
|
|
280
|
+
not param.required for route in routes for param in route.params
|
|
281
|
+
)
|
|
282
|
+
has_validation_errors = any(
|
|
283
|
+
response.type_ is FastAPIClientHTTPValidationError
|
|
284
|
+
for route in routes
|
|
285
|
+
for response in route.responses.values()
|
|
286
|
+
)
|
|
287
|
+
if import_client_base:
|
|
288
|
+
return self._generate_with_import_client_base(
|
|
289
|
+
has_not_required_params, has_validation_errors
|
|
290
|
+
)
|
|
291
|
+
return self._generate_without_import_client_base(has_validation_errors)
|
|
292
|
+
|
|
293
|
+
def _generate_with_import_client_base(
|
|
294
|
+
self,
|
|
295
|
+
has_not_required_params: bool,
|
|
296
|
+
has_validation_errors: bool,
|
|
297
|
+
) -> str:
|
|
298
|
+
# Manually write imports here so that modules are imported from specific
|
|
299
|
+
# submodule instead of top-level module.
|
|
300
|
+
for import_name in [
|
|
301
|
+
self._base_class.__name__,
|
|
302
|
+
self._idents.client_extensions,
|
|
303
|
+
self._idents.result,
|
|
304
|
+
self._idents.http_validation_error if has_validation_errors else None,
|
|
305
|
+
self._idents.not_required if has_not_required_params else None,
|
|
306
|
+
]:
|
|
307
|
+
if import_name:
|
|
308
|
+
self._impr.add_import(
|
|
309
|
+
Import(module=self._base_class.__module__, name=import_name)
|
|
310
|
+
)
|
|
311
|
+
return f"class {self._idents.client_class}({self._impr(self._base_class)}):"
|
|
312
|
+
|
|
313
|
+
def _generate_without_import_client_base(self, has_validation_errors: bool) -> str:
|
|
314
|
+
# Adding Self to one of the *_IMPORTS constant makes type checking fail.
|
|
315
|
+
self._impr.add_import(Import(module="typing", name="Self"))
|
|
316
|
+
|
|
317
|
+
# Can't programmatically look up import location of constants, so have to
|
|
318
|
+
# hard-code those here.
|
|
319
|
+
self._impr.add_import(Import(module="httpx", name="USE_CLIENT_DEFAULT"))
|
|
320
|
+
|
|
321
|
+
# Manually specify where warn is imported from, because otherwise it resolves to
|
|
322
|
+
# `from _warnings import warn`.
|
|
323
|
+
self._impr.add_import_for_type(Import(module="warnings", name="warn"), warn)
|
|
324
|
+
|
|
325
|
+
for type_ in (
|
|
326
|
+
_IMPORTS
|
|
327
|
+
+ (_IMPORTS_VALIDATION_ERROR if has_validation_errors else [])
|
|
328
|
+
+ (
|
|
329
|
+
_IMPORTS_SYNC_CLIENT
|
|
330
|
+
if self._base_class is FastAPIClientBase
|
|
331
|
+
else _IMPORTS_ASYNC_CLIENT
|
|
332
|
+
)
|
|
333
|
+
):
|
|
334
|
+
self._impr(type_)
|
|
335
|
+
|
|
336
|
+
for type_ in _IMPORTS_TYPE_CHECKING:
|
|
337
|
+
self._impr(type_, is_only_for_type_checking=True)
|
|
338
|
+
|
|
339
|
+
def base_class_source_with_test_markers() -> str:
|
|
340
|
+
source = getsource(self._base_class)
|
|
341
|
+
if not self._add_test_markers:
|
|
342
|
+
return source
|
|
343
|
+
source_lines = source.splitlines()
|
|
344
|
+
return (
|
|
345
|
+
f"{source_lines[0]}\n"
|
|
346
|
+
" # TEST_MARKER_BEFORE_BOILERPLATE\n\n"
|
|
347
|
+
f"{'\n'.join(source_lines[1:])}\n\n"
|
|
348
|
+
" # TEST_MARKER_AFTER_BOILERPLATE\n"
|
|
349
|
+
)
|
|
350
|
+
|
|
351
|
+
sources = [
|
|
352
|
+
"# TEST_MARKER_BEFORE_BOILERPLATE\n" if self._add_test_markers else None,
|
|
353
|
+
getsource(FastAPIClientExtensions),
|
|
354
|
+
getsource(FastAPIClientResult),
|
|
355
|
+
getsource(FastAPIClientValidationError) if has_validation_errors else None,
|
|
356
|
+
(
|
|
357
|
+
getsource(FastAPIClientHTTPValidationError)
|
|
358
|
+
if has_validation_errors
|
|
359
|
+
else None
|
|
360
|
+
),
|
|
361
|
+
getsource(FastAPIClientNotDefaultStatusError),
|
|
362
|
+
"FASTAPI_CLIENT_NOT_REQUIRED: Any = ...\n",
|
|
363
|
+
"# TEST_MARKER_AFTER_BOILERPLATE\n" if self._add_test_markers else None,
|
|
364
|
+
base_class_source_with_test_markers(),
|
|
365
|
+
]
|
|
366
|
+
return "\n\n".join(self._idents.replace_in_code(s) for s in sources if s)
|
|
367
|
+
|
|
368
|
+
|
|
369
|
+
class _ImportCodeGenerator:
|
|
370
|
+
def __init__(self, impr: ImportRegistry) -> None:
|
|
371
|
+
self._impr = impr
|
|
372
|
+
|
|
373
|
+
def generate(self) -> str:
|
|
374
|
+
imports = set(self._impr.imports())
|
|
375
|
+
|
|
376
|
+
imports_only_for_type_checking = set(
|
|
377
|
+
self._impr.imports(only_for_type_checking=True)
|
|
378
|
+
)
|
|
379
|
+
if imports_only_for_type_checking:
|
|
380
|
+
for import_ in imports_only_for_type_checking:
|
|
381
|
+
imports.remove(import_)
|
|
382
|
+
|
|
383
|
+
type_checking_import = Import(module="typing", name="TYPE_CHECKING")
|
|
384
|
+
imports.add(type_checking_import)
|
|
385
|
+
imports_only_for_type_checking.discard(type_checking_import)
|
|
386
|
+
|
|
387
|
+
code = self._get_imports_code_for_import_block(imports)
|
|
388
|
+
if imports_only_for_type_checking:
|
|
389
|
+
code += "\nif TYPE_CHECKING:\n"
|
|
390
|
+
for line in self._get_imports_code_for_import_block(
|
|
391
|
+
imports_only_for_type_checking
|
|
392
|
+
).splitlines():
|
|
393
|
+
code += f" {line}\n"
|
|
394
|
+
|
|
395
|
+
return code
|
|
396
|
+
|
|
397
|
+
@classmethod
|
|
398
|
+
def _get_imports_code_for_import_block(cls, imports: AbstractSet[Import]) -> str:
|
|
399
|
+
imports_by_group = defaultdict[_ImportGroup, set[Import]](set)
|
|
400
|
+
for import_ in imports:
|
|
401
|
+
imports_by_group[cls._get_import_group(import_)].add(import_)
|
|
402
|
+
|
|
403
|
+
lines = (
|
|
404
|
+
cls._get_imports_code_for_import_group(imports_by_group[group])
|
|
405
|
+
for group in _ImportGroup
|
|
406
|
+
)
|
|
407
|
+
return "\n".join(filter(None, lines))
|
|
408
|
+
|
|
409
|
+
@classmethod
|
|
410
|
+
@cache
|
|
411
|
+
def _get_import_group(cls, import_: Import) -> _ImportGroup:
|
|
412
|
+
top_level_module = import_.module.split(".", maxsplit=1)[0]
|
|
413
|
+
if top_level_module in stdlib_module_names:
|
|
414
|
+
return _ImportGroup.STDLIB
|
|
415
|
+
spec = find_spec(top_level_module)
|
|
416
|
+
if spec and spec.origin and "site-packages" in spec.origin:
|
|
417
|
+
return _ImportGroup.SITE_PACKAGE
|
|
418
|
+
return _ImportGroup.LOCAL
|
|
419
|
+
|
|
420
|
+
@classmethod
|
|
421
|
+
def _get_imports_code_for_import_group(cls, imports: AbstractSet[Import]) -> str:
|
|
422
|
+
def alias_str(import_: Import) -> str:
|
|
423
|
+
return f" as {import_.alias}" if import_.alias else ""
|
|
424
|
+
|
|
425
|
+
imports_without_name = list[Import]()
|
|
426
|
+
imports_with_name_by_module = defaultdict[str, list[Import]](list)
|
|
427
|
+
for import_ in imports:
|
|
428
|
+
if not import_.name:
|
|
429
|
+
imports_without_name.append(import_)
|
|
430
|
+
else:
|
|
431
|
+
imports_with_name_by_module[import_.module].append(import_)
|
|
432
|
+
|
|
433
|
+
imports_without_name.sort()
|
|
434
|
+
|
|
435
|
+
code = ""
|
|
436
|
+
for import_ in imports_without_name:
|
|
437
|
+
if import_.module == "builtins" and import_.alias is None:
|
|
438
|
+
continue
|
|
439
|
+
code += f"import {import_.module}{alias_str(import_)}\n"
|
|
440
|
+
for module in sorted(imports_with_name_by_module.keys()):
|
|
441
|
+
imports_for_module = imports_with_name_by_module[module]
|
|
442
|
+
if module == "builtins":
|
|
443
|
+
imports_for_module = [
|
|
444
|
+
import_
|
|
445
|
+
for import_ in imports_for_module
|
|
446
|
+
if import_.alias is not None
|
|
447
|
+
]
|
|
448
|
+
if not imports_for_module:
|
|
449
|
+
continue
|
|
450
|
+
|
|
451
|
+
imports_for_module.sort(
|
|
452
|
+
key=lambda import_: (not import_.name.isupper(), import_.name)
|
|
453
|
+
)
|
|
454
|
+
|
|
455
|
+
if len(imports_for_module) == 1:
|
|
456
|
+
import_ = imports_for_module[0]
|
|
457
|
+
code += f"from {module} import {import_.name}{alias_str(import_)}\n"
|
|
458
|
+
else:
|
|
459
|
+
code += f"from {module} import (\n"
|
|
460
|
+
code += "".join(
|
|
461
|
+
f" {import_.name}{alias_str(import_)},\n"
|
|
462
|
+
for import_ in imports_for_module
|
|
463
|
+
)
|
|
464
|
+
code += ")\n"
|
|
465
|
+
return code
|