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.
@@ -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,4 @@
1
+ from .cli import app
2
+
3
+ if __name__ == "__main__":
4
+ app()
@@ -0,0 +1,4 @@
1
+ from importlib.metadata import version
2
+ from typing import cast
3
+
4
+ __version__ = version(cast(str, __package__))
@@ -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