Skip to content

Commit 8cb6ba7

Browse files
committed
refactor: moved imports under TYPE_CHECKING
1 parent 538d1f9 commit 8cb6ba7

2 files changed

Lines changed: 92 additions & 40 deletions

File tree

tortoise/fields/data.py

Lines changed: 39 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,6 @@
44
import datetime
55
import functools
66
import json
7-
import sys
87
import warnings
98
from collections.abc import Callable
109
from decimal import Decimal
@@ -45,13 +44,14 @@
4544
_PydanticModelMetaclass = None # type: ignore[assignment,misc]
4645

4746
if TYPE_CHECKING: # pragma: nocoverage
48-
from tortoise.models import Model
47+
import sys
4948

49+
from tortoise.models import Model
5050

51-
if sys.version_info >= (3, 11):
52-
from typing import Unpack
53-
else: # pragma: no cover
54-
from typing_extensions import Unpack
51+
if sys.version_info >= (3, 11):
52+
from typing import Unpack
53+
else: # pragma: no cover
54+
from typing_extensions import Unpack
5555

5656
__all__ = (
5757
"BigIntField",
@@ -333,7 +333,9 @@ def __init__(
333333
stacklevel=2,
334334
)
335335
if index or db_index:
336-
raise ConfigurationError("TextField can't be indexed, consider CharField")
336+
raise ConfigurationError(
337+
"TextField can't be indexed, consider CharField"
338+
)
337339
elif db_index:
338340
raise ConfigurationError("TextField can't be indexed, consider CharField")
339341

@@ -429,7 +431,9 @@ def __init__(self, max_digits: int, decimal_places: int, **kwargs: Any) -> None:
429431
super().__init__(**kwargs)
430432
self.max_digits = max_digits
431433
self.decimal_places = decimal_places
432-
self.quant = Decimal("1" if decimal_places == 0 else f"1.{('0' * decimal_places)}")
434+
self.quant = Decimal(
435+
"1" if decimal_places == 0 else f"1.{('0' * decimal_places)}"
436+
)
433437

434438
def to_python_value(self, value: Any) -> Decimal | None:
435439
if value is not None:
@@ -454,7 +458,9 @@ def function_cast(self, term: Term) -> Term:
454458
DatetimeFieldQueryValueType = TypeVar(
455459
"DatetimeFieldQueryValueType", datetime.datetime, int, float, str
456460
)
457-
DateFieldQueryValueType = TypeVar("DateFieldQueryValueType", datetime.date, int, float, str)
461+
DateFieldQueryValueType = TypeVar(
462+
"DateFieldQueryValueType", datetime.date, int, float, str
463+
)
458464

459465

460466
class DatetimeField(Field[T_DATETIME], datetime.datetime):
@@ -504,7 +510,9 @@ def __init__(
504510
**kwargs: Unpack[FieldKwargs],
505511
) -> None: ...
506512

507-
def __init__(self, auto_now: bool = False, auto_now_add: bool = False, **kwargs: Any) -> None:
513+
def __init__(
514+
self, auto_now: bool = False, auto_now_add: bool = False, **kwargs: Any
515+
) -> None:
508516
if auto_now_add and auto_now:
509517
raise ConfigurationError("You can choose only 'auto_now' or 'auto_now_add'")
510518
super().__init__(**kwargs)
@@ -644,7 +652,9 @@ def __init__(
644652
**kwargs: Unpack[FieldKwargs],
645653
) -> None: ...
646654

647-
def __init__(self, auto_now: bool = False, auto_now_add: bool = False, **kwargs: Any) -> None:
655+
def __init__(
656+
self, auto_now: bool = False, auto_now_add: bool = False, **kwargs: Any
657+
) -> None:
648658
if auto_now_add and auto_now:
649659
raise ConfigurationError("You can choose only 'auto_now' or 'auto_now_add'")
650660
super().__init__(**kwargs)
@@ -743,7 +753,9 @@ def to_db_value(
743753

744754
if value is None:
745755
return None
746-
return (value.days * 86400000000) + (value.seconds * 1000000) + value.microseconds
756+
return (
757+
(value.days * 86400000000) + (value.seconds * 1000000) + value.microseconds
758+
)
747759

748760

749761
class FloatField(Field[T_FLOAT], float):
@@ -907,7 +919,9 @@ def __init__(
907919
) -> None: ...
908920

909921
def __init__(self, **kwargs: Any) -> None:
910-
if (kwargs.get("primary_key") or kwargs.get("pk", False)) and "default" not in kwargs:
922+
if (
923+
kwargs.get("primary_key") or kwargs.get("pk", False)
924+
) and "default" not in kwargs:
911925
kwargs["default"] = uuid4
912926
super().__init__(**kwargs)
913927

@@ -982,7 +996,9 @@ def __init__(
982996

983997
# Automatic description for the field if not specified by the user
984998
if description is None:
985-
description = "\n".join([f"{e.name}: {int(e.value)}" for e in enum_type])[:2048]
999+
description = "\n".join([f"{e.name}: {int(e.value)}" for e in enum_type])[
1000+
:2048
1001+
]
9861002

9871003
super().__init__(description=description, **kwargs)
9881004
self.enum_type = enum_type
@@ -991,7 +1007,9 @@ def to_python_value(self, value: int | None) -> IntEnum | None:
9911007
value = self.enum_type(value) if value is not None else None
9921008
return value
9931009

994-
def to_db_value(self, value: IntEnum | None | int, instance: type[Model] | Model) -> int | None:
1010+
def to_db_value(
1011+
self, value: IntEnum | None | int, instance: type[Model] | Model
1012+
) -> int | None:
9951013
if isinstance(value, IntEnum):
9961014
value = int(value.value)
9971015
if isinstance(value, int):
@@ -1038,7 +1056,9 @@ def __init__(
10381056
) -> None:
10391057
# Automatic description for the field if not specified by the user
10401058
if description is None:
1041-
description = "\n".join([f"{e.name}: {str(e.value)}" for e in enum_type])[:2048]
1059+
description = "\n".join([f"{e.name}: {str(e.value)}" for e in enum_type])[
1060+
:2048
1061+
]
10421062

10431063
# Automatic CharField max_length
10441064
if max_length == 0:
@@ -1053,7 +1073,9 @@ def __init__(
10531073
def to_python_value(self, value: str | None) -> Enum | None:
10541074
return self.enum_type(value) if value is not None else None
10551075

1056-
def to_db_value(self, value: Enum | None | str, instance: type[Model] | Model) -> str | None:
1076+
def to_db_value(
1077+
self, value: Enum | None | str, instance: type[Model] | Model
1078+
) -> str | None:
10571079
self.validate(value)
10581080
if isinstance(value, Enum):
10591081
return str(value.value)

tortoise/fields/relational.py

Lines changed: 53 additions & 23 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,5 @@
11
from __future__ import annotations
22

3-
import sys
43
import warnings
54
from collections.abc import AsyncGenerator, Generator, Iterator
65
from typing import TYPE_CHECKING, Any, Generic, Literal, TypeVar, overload
@@ -17,16 +16,18 @@
1716
RelationalFieldKwargs,
1817
)
1918

20-
if sys.version_info >= (3, 11):
21-
from typing import Unpack
22-
else: # pragma: no cover
23-
from typing_extensions import Unpack
24-
2519
if TYPE_CHECKING: # pragma: nocoverage
20+
import sys
21+
2622
from tortoise.backends.base.client import BaseDBAsyncClient
2723
from tortoise.models import Model
2824
from tortoise.queryset import Q, QuerySet
2925

26+
if sys.version_info >= (3, 11):
27+
from typing import Unpack
28+
else: # pragma: no cover
29+
from typing_extensions import Unpack
30+
3031
MODEL = TypeVar("MODEL", bound="Model")
3132

3233

@@ -132,7 +133,9 @@ def offset(self, offset: int) -> QuerySet[MODEL]:
132133
"""
133134
return self._query.offset(offset)
134135

135-
async def create(self, using_db: BaseDBAsyncClient | None = None, **kwargs: Any) -> MODEL:
136+
async def create(
137+
self, using_db: BaseDBAsyncClient | None = None, **kwargs: Any
138+
) -> MODEL:
136139
"""
137140
Create a related record in the DB and returns the object, automatically setting the
138141
foreign key relationship to the parent instance.
@@ -164,7 +167,9 @@ async def create(self, using_db: BaseDBAsyncClient | None = None, **kwargs: Any)
164167
# Call remote model's create method
165168
return await self.remote_model.create(using_db=using_db, **kwargs)
166169

167-
def _set_result_for_query(self, sequence: list[MODEL], attr: str | None = None) -> None:
170+
def _set_result_for_query(
171+
self, sequence: list[MODEL], attr: str | None = None
172+
) -> None:
168173
self._fetched = True
169174
self.related_objects = sequence
170175
if attr:
@@ -182,12 +187,18 @@ class ManyToManyRelation(ReverseRelation[MODEL]):
182187
Many-to-many relation container for :func:`.ManyToManyField`.
183188
"""
184189

185-
def __init__(self, instance: Model, m2m_field: ManyToManyFieldInstance[MODEL]) -> None:
186-
super().__init__(m2m_field.related_model, m2m_field.related_name, instance, "pk")
190+
def __init__(
191+
self, instance: Model, m2m_field: ManyToManyFieldInstance[MODEL]
192+
) -> None:
193+
super().__init__(
194+
m2m_field.related_model, m2m_field.related_name, instance, "pk"
195+
)
187196
self.field = m2m_field
188197
self.instance = instance
189198

190-
async def add(self, *instances: MODEL, using_db: BaseDBAsyncClient | None = None) -> None:
199+
async def add(
200+
self, *instances: MODEL, using_db: BaseDBAsyncClient | None = None
201+
) -> None:
191202
"""
192203
Adds one or more of ``instances`` to the relation.
193204
@@ -206,16 +217,25 @@ async def add(self, *instances: MODEL, using_db: BaseDBAsyncClient | None = None
206217
pks_f: list = []
207218
for instance_to_add in instances:
208219
if not instance_to_add._saved_in_db:
209-
raise OperationalError(f"You should first call .save() on {instance_to_add}")
220+
raise OperationalError(
221+
f"You should first call .save() on {instance_to_add}"
222+
)
210223
pk_f = related_pk_formatting_func(instance_to_add.pk, instance_to_add)
211224
pks_f.append(pk_f)
212225
through_table = Table(self.field.through, schema=self.field.through_schema)
213226
backward_key, forward_key = self.field.backward_key, self.field.forward_key
214-
backward_field, forward_field = through_table[backward_key], through_table[forward_key]
227+
backward_field, forward_field = (
228+
through_table[backward_key],
229+
through_table[forward_key],
230+
)
215231
select_query = (
216-
db.query_class.from_(through_table).where(backward_field == pk_b).select(forward_key)
232+
db.query_class.from_(through_table)
233+
.where(backward_field == pk_b)
234+
.select(forward_key)
235+
)
236+
criterion = (
237+
forward_field == pks_f[0] if len(pks_f) == 1 else forward_field.isin(pks_f)
217238
)
218-
criterion = forward_field == pks_f[0] if len(pks_f) == 1 else forward_field.isin(pks_f)
219239
select_query = select_query.where(criterion)
220240

221241
_, already_existing_relations_raw = await db.execute_query(
@@ -227,7 +247,9 @@ async def add(self, *instances: MODEL, using_db: BaseDBAsyncClient | None = None
227247
}
228248

229249
if pks_f_to_insert := set(pks_f) - already_existing_forward_pks:
230-
query = db.query_class.into(through_table).columns(forward_field, backward_field)
250+
query = db.query_class.into(through_table).columns(
251+
forward_field, backward_field
252+
)
231253
for pk_f in pks_f_to_insert:
232254
query = query.insert(pk_f, pk_b)
233255
await db.execute_query(*query.get_parameterized_sql())
@@ -238,7 +260,9 @@ async def clear(self, using_db: BaseDBAsyncClient | None = None) -> None:
238260
"""
239261
await self._remove_or_clear(using_db=using_db)
240262

241-
async def remove(self, *instances: MODEL, using_db: BaseDBAsyncClient | None = None) -> None:
263+
async def remove(
264+
self, *instances: MODEL, using_db: BaseDBAsyncClient | None = None
265+
) -> None:
242266
"""
243267
Removes one or more of ``instances`` from the relation.
244268
@@ -263,9 +287,9 @@ async def _remove_or_clear(
263287
if instances:
264288
related_pk_formatting_func = type(instances[0])._meta.pk.to_db_value
265289
if len(instances) == 1:
266-
condition &= through_table[self.field.forward_key] == related_pk_formatting_func(
267-
instances[0].pk, instances[0]
268-
)
290+
condition &= through_table[
291+
self.field.forward_key
292+
] == related_pk_formatting_func(instances[0].pk, instances[0])
269293
else:
270294
condition &= through_table[self.field.forward_key].isin(
271295
[related_pk_formatting_func(i.pk, i) for i in instances]
@@ -293,7 +317,9 @@ def __init__(
293317
if TYPE_CHECKING:
294318

295319
@overload
296-
def __get__(self, instance: None, owner: type[Model]) -> RelationalField[MODEL]: ...
320+
def __get__(
321+
self, instance: None, owner: type[Model]
322+
) -> RelationalField[MODEL]: ...
297323

298324
@overload
299325
def __get__(self, instance: Model, owner: type[Model]) -> MODEL: ...
@@ -322,7 +348,9 @@ def validate_model_name(cls, model_name: str | type[Model]) -> None:
322348
) from None
323349
elif len(model_name.split(".")) != 2:
324350
field_type = cls.__name__.replace("Instance", "")
325-
raise ConfigurationError(f'{field_type} accepts model name in format "app.Model"')
351+
raise ConfigurationError(
352+
f'{field_type} accepts model name in format "app.Model"'
353+
)
326354

327355

328356
class ForeignKeyFieldInstance(RelationalField[MODEL]):
@@ -342,7 +370,9 @@ def __init__(
342370
"on_delete can only be CASCADE, RESTRICT, SET_NULL, SET_DEFAULT or NO_ACTION"
343371
)
344372
if on_delete == SET_NULL and not bool(kwargs.get("null")):
345-
raise ConfigurationError("If on_delete is SET_NULL, then field must have null=True set")
373+
raise ConfigurationError(
374+
"If on_delete is SET_NULL, then field must have null=True set"
375+
)
346376
self.on_delete = on_delete
347377

348378
def describe(self, serializable: bool) -> dict:

0 commit comments

Comments
 (0)