diff --git a/docs/contributing.md b/docs/contributing.md index cb5faa13..088f2887 100644 --- a/docs/contributing.md +++ b/docs/contributing.md @@ -63,7 +63,7 @@ $ hatch env create docs ### Enable pre-commit -The project comes with a pre-commit hook configuration. To enable it, just run inside the clone: +The project comes with a pre-commit (or prek) hook configuration. To enable it, just run inside the clone: ```shell $ hatch run pre-commit install @@ -92,6 +92,9 @@ $ hatch test tests/test_apiviews.py Pytest native arguments can be passed after passing `--`. +!!! Warning + You need pytest >= 9.0 for subtests. This can be an issue with old hatch versions or old hatch environments. + To run the linting, use: ```shell diff --git a/docs/queries/bulk.md b/docs/queries/bulk.md index 3a28c4c9..b60c2f59 100644 --- a/docs/queries/bulk.md +++ b/docs/queries/bulk.md @@ -13,6 +13,9 @@ The returned array is in the same order as the values/objects provided. And cont Input for all bulk operations are models of the right type or dictionaries. They can be intermixed and must be provided in an `Iterable`. +!!! Warning + When using `SkipOperation` in a signal, `None` values are returned. + ## Operations ### Bulk create @@ -32,6 +35,9 @@ assert not returned_objs[0].can_load # the pks are incomplete assert returned_objs[1].can_load # the pks are complete ``` +Output: +The array can contain `None` when either an `SkipOperation` is raised or if `ignore_conflicts=True` is used. + #### `ignore_conflicts` When the database is compatible and we don't need the returned values, we can use `ignore_conflicts=True` instead `bulk_get_or_create`. @@ -155,6 +161,9 @@ This mode has two effects: - It is ensured that all returned instances `can_load` when `embed_parent` is active. If necessary, it will issue serialized single inserts. - The embedding is resolved if `embed_parent` is set. You get the child with the embedded parent. +!!! Note + When pointing to an attribute which is `None` and you use `resolve_embed` you get this in the output. This can be confusing. + ### loadable An instances is loadable if its `can_load` property signals it is loadable. For dictionaries the on the fly generated instance is used. @@ -164,3 +173,11 @@ You can set the `identifying_db_fields` so a provided instance becomes suddenly Other effects are that `resolve_embed` succeeds for such crafted instances. It is planned to add signals so you will be able to manipulate the immediate instances via signals so you can do this trick also for dict inputs. + +### Signals + +When raising `SkipOperation` in a `pre_bulk` signal, the operation is cancelled and returned is an array with `None` values or with `(None, False)` tuples. + +By default the invariant: length and order of the input = length and order of the output, is kept. You can however manipulate this by signals. E.g. clearing the whole output array (values). This is **not** recommended. + +You can however remove elements from the output by setting them to `None`. diff --git a/docs/release-notes.md b/docs/release-notes.md index 38b69aec..84a1062d 100644 --- a/docs/release-notes.md +++ b/docs/release-notes.md @@ -10,6 +10,8 @@ - Add the `ignore_conflicts` parameter to `bulk_create`. - Add relationship signals (`pre/post_relation_add` and `pre/post_relation_remove`). - Add bulk signals (`pre/post_bulk`). +- Add `SkipOperation` exception for signals. +- Add `injected_filters` parameter for `pre_delete` to dynamically inject protection rules. ### Changed @@ -22,12 +24,15 @@ - The typings changed for QuerySetType: `EdgyEmbedTarget` and `EdgyModel` (the queryset model) are switched in the `Generic` definition. - Bulk operations are now keywords only (except the first `objs` parameter). - Dedupe `bulk_create`, `bulk_get_or_create`, and `bulk_update_or_create` inputs. +- Refactor `QueryExecutor` so it can update itself. ### Fixed - Relations did not use the tenancy/used schema properly when querying. - Bulk operations with reflected fields did not always work properly. - `run_concurrently` now properly cleans up not executed coroutines in case of an error. +- `SuspiciousFileOperation` inherits now correctly from `EdgyException`. +- Return row count for update. ### Removed @@ -44,6 +49,7 @@ - Bulk operations return now a result instead `None`. For `bulk_get_or_create` the returned list format changes to `(instance, created)` tuples. - The typings changed for QuerySetType: `EdgyEmbedTarget` and `EdgyModel` (the queryset model) are switched in the `Generic` definition. - Stop issuing `pre_delete` and `post_delete` signals during relation operations; use `pre_relation_remove` and `post_relation_remove` signals instead. +- `bulk_get_or_create` can return `None` values if signals are used. ## 0.35.11 diff --git a/docs/signals.md b/docs/signals.md index 1e7075ec..4159acaa 100644 --- a/docs/signals.md +++ b/docs/signals.md @@ -24,13 +24,20 @@ from edgy.core.signals import ( post_update, post_migrate, pre_migrate, - pre_relation, - post_relation, + pre_relation_add, + post_relation_add, + pre_relation_remove, + post_relation_remove, pre_bulk, post_bulk, ) ``` +#### Pre Operations special exception + +If you just want to **skip** the operation without causing a bigger error, you can raise `edgy.exceptions.SkipOperation` to stop the operation and +returning an empty value and send the corresponding post signal with `operation_skipped=True`. + #### pre_save Triggered before a model is saved (during `Model.save()` and `Model.query.create()`). @@ -69,12 +76,14 @@ post_update(sender: type["Model"], instance: Union["Model", "QuerySet"], model_i The receiver function receives following parameters: -- instance - The model or QuerySet instance. -- model_instance -The model instance if available. For save signals always available -- values - The passed values. -- column_values - The parsed values which are used for the db. -- is_update - Is it an update? This is also set for `*_update` to match the save parameters. -- is_migration - Called from `apply_default_force_nullable_fields` which is mostly for migrations. Here we have model instances. +- `instance` - The model or QuerySet instance. +- `model_instance` -The model instance if available. For save signals always available +- `values` - The passed values. +- `column_values` - The parsed values which are used for the db. +- `is_update` - Is it an update? This is also set for `*_update` to match the save parameters. +- `is_migration` - Called from `apply_default_force_nullable_fields` which is mostly for migrations. Here we have model instances. +- `row_count` - (post only) The rows updated. +- `operation_skipped` - (post only) If the operation was skipped. #### pre_delete @@ -84,10 +93,26 @@ Triggered before a model is deleted (during `Model.delete()` and `Model.query.de pre_delete(send: type["Model"], instance: Union["Model", "QuerySet"], model_instance: Optional["Model"]) ``` +A more advanced example is: + +```python +{!> ../docs_src/signals/prevent_deletion.py !} +``` + ##### pre_delete parameters -- instance - The model or QuerySet instance. -- model_instance -The model instance if available. +- `instance` - The model or QuerySet instance. +- `model_instance` - The model instance if available otherwise `None`. +- `injected_filters` (query only) - You can insert or remove (when inserted by another signal) extra filter parameters for query deletions. + +**Example for the insertion of new parameters** + +```python +{!> ../docs_src/signals/excempt_from_deletion.py !} +``` + +You can also add `or_`, `and_` or other clauses valid for the `filter` method of `QuerySet`. +By default they are combined like with `and_`. #### post_delete @@ -150,17 +175,20 @@ And for revision: #### pre_bulk -The `pre_bulk` signal is issued before the database modifications and allows before executing bulk operations to manipulate -the instances. +The `pre_bulk` signal is issued before the database modifications and allows before executing bulk operations to manipulate the instances. Sender is the queryset. When using bulk operations on the relation queryset the sender is either the `target` model (one-to-many) or `through` model (many-to-many). + +```python +{!> ../docs_src/signals/manipulate_bulk.py !} +``` #### post_bulk -The `post_bulk` signal is issued after the database modifications and contains information about how many rows were created and/or updated. +The `post_bulk` signal is issued after the database modifications and contains information about how many rows were created and/or updated. When using bulk operations on the relation queryset the sender is either the `target` model (one-to-many) or `through` model (many-to-many). #### Parameters of `*_bulk` signals - `raw_values`: Raw model instances with created flag. No resolving of `embed_parent`. -- `values` (post only): Resolved model instances with created flag. When not using `resolve_embed`, the raw model instances. +- `values` (post only, only when operation_skipped=False): Resolved model instances with created flag. When not using `resolve_embed`, the raw model instances. - `operation`: `bulk_create`, `bulk_update`, `bulk_update_or_create`, `bulk_get_or_create`. - `resolve_embed`: Value of `resolve_embed`. - `create_params`: `(raw_instance, position in raw_values, set of input kwarg names)` tuple. You can prevent an insert operation by removing an tuple. You can move an tuple to update_params if the instance should be updated instead. @@ -181,13 +209,19 @@ Methods where this trick can be applied are: `bulk_create` (with `ignore_conflic The `pre_relation_add` signal is issued before the database modifications and allows before executing changing the relations to manipulate the instances. -For Many-to-Many relations the sender is the through model. +The sender is either the `target` model (one-to-many) or `through` model (many-to-many). Signals are also issued on overwrites of `pre_relation_add` in source.meta.signals, through.meta.signals and target.meta.signals. However if the default signal or a shared signal is used it is only issued per different signal object, so you can expect when listening to one of the signals, you get notified only once. + +!!! Note + This signal is not issued if `add_many` is called with empty arguments or `save_related` is called when staged are empty. #### post_relation_add The `post_relation_add` signal is issued after the database modifications and contains information about how many rows were changed. -For Many-to-Many relations the sender is the through model. +The sender is either the `target` model (one-to-many) or `through` model (many-to-many). Signals are also issued on overwrites of `pre_relation_add` in source.meta.signals, through.meta.signals and target.meta.signals. However if the default signal or a shared signal is used it is only issued per different signal object, so you can expect when listening to one of the signals, you get notified only once. + +!!! Note + This signal is not issued if `add_many` is called with empty arguments or `save_related` is called when staged are empty. #### Parameters of `*_relation_add` signals @@ -195,7 +229,7 @@ For Many-to-Many relations the sender is the through model. - `row_count`: How many rows were updated/created? `None` for db systems not supporting it. - `row_count_create`: How many rows were created? `None` for db systems not supporting it. - `raw_values`: Raw model instances with created flag of either the source model (`one_to_many`) or the `through` model (`many_to_many`). Useful in combination with `create_params` and `update_params` to tweak output. There is **no** resolving via `resolve_embed` -- `values` (post only): The resolved counterpart instances with created flag. +- `values` (post only, only when add, add_many and `operation_skipped=False`): The resolved counterpart instances with created flag. - `operation`: `save_related` and `add` (also issued for `add_many` and `create`). - `field`: RelationField name on `source` triggering this signal. - `source`: Source model which contains the RelationField triggering the signals. @@ -203,6 +237,7 @@ For Many-to-Many relations the sender is the through model. - `relation`: Relation type. `one_to_many`, `many_to_many`. - `create_params`: See bulk signal parameter. - `update_params`: See bulk signal parameter. +- `operation_skipped`: Is the operation skipped? This will lead to missing parameters (`values`) **Replacing raw_values/values** @@ -212,7 +247,10 @@ Here every operation allows `None` values instead of instances. So it is no prob The `pre_relation_remove` signal is issued before the database modifications for the removal of connections and allow customizations including to stop the deletion by issuing an exception. -For Many-to-Many relations the sender is the through model. +The sender is either the `target` model (one-to-many) or `through` model (many-to-many). Signals are also issued on overwrites of `pre_relation_add` in source.meta.signals, through.meta.signals and target.meta.signals. However if the default signal or a shared signal is used it is only issued per different signal object, so you can expect when listening to one of the signals, you get notified only once. + +!!! Note + This signal is not issued if `remove_many` is called with empty arguments. **Remove from removal list** @@ -220,17 +258,28 @@ You have two options to block the removal of an instance 1. raise an exception 2. remove the instance from `raw_values` +```python +{!> ../docs_src/signals/prevent_deletion.py !} +``` + +```python +{!> ../docs_src/signals/excempt_from_deletion.py !} +``` + #### post_relation_remove The `post_relation_remove` signal is issued after the database modifications for the removal of connections and contains information about how many rows were changed/removed. -For Many-to-Many relations the sender is the through model. +The sender is either the `target` model (one-to-many) or `through` model (many-to-many). Signals are also issued on overwrites of `pre_relation_add` in source.meta.signals, through.meta.signals and target.meta.signals. However if the default signal or a shared signal is used it is only issued per different signal object, so you can expect when listening to one of the signals, you get notified only once. + +!!! Note + This signal is not issued if `remove_many` is called with empty arguments. #### Parameters of `*_relation_remove signals` - `instance`: Source instance. - `row_count` (post): How many rows were updated/deleted? `None` for db systems not supporting it. -- `raw_values`: Raw model instances **without created flag** of either the source model (`one_to_many`) or the `through` model (`many_to_many`). There is **no** resolving via `resolve_embed`. You can +- `raw_values`: Raw model instances **without created flag** of either the source model (`one_to_many`) or the `through` model (`many_to_many`). There is **no** resolving via `resolve_embed`. You can clear the list and readd the models you want to delete (or doing position based modifications (harder)) as long you don't await during the modifications. You should recheck if something changes if you fetch something with await. - `field`: RelationField name on `source` triggering this signal. - `source`: Source model which contains the RelationField triggering the signals. - `target`: Target model. @@ -331,14 +380,24 @@ To prevent default lifecycle signals from being called, you can overwrite them p ### How to Use It +**Using a custom signal** Use the custom signal in your logic: ```python hl_lines="17" {!> ../docs_src/signals/logic.py !} ``` - The `on_verify` signal is triggered only when the user is verified. +**Log changes** + +An other useful usecase is logging user actions: + +```python +{!> ../docs_src/signals/log_changes.py !} +``` + +Of course there are better ways for serialization. + ### Disconnect the Signal Disconnecting a custom signal is the same as disconnecting a default signal: diff --git a/docs_src/signals/excempt_from_deletion.py b/docs_src/signals/excempt_from_deletion.py new file mode 100644 index 00000000..e3c30fa1 --- /dev/null +++ b/docs_src/signals/excempt_from_deletion.py @@ -0,0 +1,44 @@ +import edgy +from edgy.exceptions import SkipOperation + + +class BaseModel(edgy.StrictModel): + protected = edgy.BooleanField(default=False) + + class Meta: + registry = ... + abstract = True + + +class Friend(BaseModel): + name = edgy.CharField(max_length=100) + + +class Profile(BaseModel): + name = edgy.CharField(max_length=100) + + +class User(BaseModel): + name = edgy.CharField(max_length=100) + profile = edgy.ForeignKey("Profile", null=True, on_delete=edgy.CASCADE, related_name="users") + friends = edgy.ManyToMany("Friend", related_name="users") + + +@User.meta.signals.pre_relation_remove.connect_via(User) +@Profile.meta.signals.pre_relation_remove.connect_via(Profile) +async def excempt_removal_relation(sender, raw_values, **kwargs): + new_raw_values = list(raw_values) + raw_values.clear() + for value in new_raw_values: + if not value.protected: + raw_values.append(value) + + +@User.meta.signals.pre_delete.connect_via(User) +@Profile.meta.signals.pre_delete.connect_via(Profile) +async def abort_removal_relation(sender, model_instance, injected_filters, **kwargs): + if model_instance is None: + injected_filters.append({"protected": True}) + else: + if model_instance.protected: + raise SkipOperation() diff --git a/docs_src/signals/log_changes.py b/docs_src/signals/log_changes.py new file mode 100644 index 00000000..1f377d9e --- /dev/null +++ b/docs_src/signals/log_changes.py @@ -0,0 +1,59 @@ +from sqlalchemy import ForeignKey +import edgy +from contextvars import ContextVar +from edgy.core import signals + +models = edgy.Registry(...) +current_user = ContextVar("current_user", default=None) + + +class BaseModel(edgy.StrictModel): + class Meta: + registry = models + abstract = True + + +class Friend(BaseModel): + name = edgy.CharField(max_length=100) + + +class Profile(BaseModel): + name = edgy.CharField(max_length=100) + + +class User(BaseModel): + name = edgy.CharField(max_length=100) + profile = edgy.ForeignKey("Profile", null=True, on_delete=edgy.CASCADE, related_name="users") + friends = edgy.ManyToMany("Friend", related_name="users") + + +class Log(BaseModel): + signal = edgy.CharField(max_length=255) + class_name = edgy.CharField(max_length=255) + params = edgy.JSONField() + user = edgy.ForeignKey(User, null=True, on_delete=edgy.CASCADE, related_name="logs") + + def __str__(self) -> str: + return str(self.extract_db_fields()) + + def __repr__(self) -> str: + return f"Log<{self}>" + + +for signal_name in dir(signals): + if not signal_name.startswith("post_") or signal_name == "post_migrate": + continue + signal: signals.Signal = getattr(signals, signal_name) + + async def log(sender, _signal_name=signal_name, **kwargs): + await Log.query.create( + signal=_signal_name, + class_name=sender.__name__, + params={k: str(v) for k, v in kwargs.items()}, + user=current_user.get(), + ) + + for model in models.models.values(): + if model is not Log: + # weak must be False otherwise the receivers vanish + signal.connect(log, model, weak=False) diff --git a/docs_src/signals/manipulate_bulk.py b/docs_src/signals/manipulate_bulk.py new file mode 100644 index 00000000..f112f242 --- /dev/null +++ b/docs_src/signals/manipulate_bulk.py @@ -0,0 +1,38 @@ +import edgy + + +class BaseModel(edgy.StrictModel): + active = edgy.IntegerField(default=0) + duplicate: bool = False + + class Meta: + registry = ... + abstract = True + + +class Friend(BaseModel): + name = edgy.CharField(max_length=100) + + +class Profile(BaseModel): + name = edgy.CharField(max_length=100) + + +class User(BaseModel): + name = edgy.CharField(max_length=100) + profile = edgy.ForeignKey("Profile", null=True, on_delete=edgy.CASCADE, related_name="users") + friends = edgy.ManyToMany("Friend", related_name="users") + + +@User.meta.signals.pre_bulk.connect_via(User) +@Profile.meta.signals.pre_bulk.connect_via(Profile) +async def handle_active_duplicate_bulk(sender, raw_values, create_params, update_params, **kwargs): + for item in create_params: + if item[0].active > 0: + item[0].active = 0 + + for item in update_params: + if item[0].active > 0: + item[0].active -= 1 + if item[0].duplicate: + create_params.append((item[0].model_copy(), item[1], item[2])) diff --git a/docs_src/signals/prevent_deletion.py b/docs_src/signals/prevent_deletion.py new file mode 100644 index 00000000..d460d408 --- /dev/null +++ b/docs_src/signals/prevent_deletion.py @@ -0,0 +1,42 @@ +import edgy + + +class BaseModel(edgy.StrictModel): + protected = edgy.BooleanField(default=False) + __deletion_with_signals__ = True + __require_model_based_deletion__ = True + + class Meta: + registry = ... + abstract = True + + +class Friend(BaseModel): + name = edgy.CharField(max_length=100) + + +class Profile(BaseModel): + name = edgy.CharField(max_length=100) + + +class User(BaseModel): + name = edgy.CharField(max_length=100) + profile = edgy.ForeignKey("Profile", null=True, on_delete=edgy.CASCADE, related_name="users") + friends = edgy.ManyToMany("Friend", related_name="users") + + +@User.meta.signals.pre_relation_remove.connect_via(User) +@Friend.meta.signals.pre_relation_remove.connect_via(Friend.meta.fields["users"].through) +@Profile.meta.signals.pre_relation_remove.connect_via(Profile) +async def abort_removal_relation(sender, raw_values, **kwargs): + for value in raw_values: + if value.protected: + raise Exception() + + +@User.meta.signals.pre_delete.connect_via(User) +@Friend.meta.signals.pre_delete.connect_via(Friend.meta.fields["users"].through) +@Profile.meta.signals.pre_delete.connect_via(Profile) +async def abort_removal_relation(sender, model_instance, **kwargs): + if model_instance.protected: + raise Exception() diff --git a/edgy/contrib/contenttypes/models.py b/edgy/contrib/contenttypes/models.py index b736c826..c90763d4 100644 --- a/edgy/contrib/contenttypes/models.py +++ b/edgy/contrib/contenttypes/models.py @@ -1,8 +1,10 @@ from __future__ import annotations +from contextlib import suppress from typing import TYPE_CHECKING, ClassVar, cast import edgy +from edgy.exceptions import SkipOperation from .metaclasses import ContentTypeMeta @@ -51,7 +53,8 @@ async def raw_delete( referenced_obs = cast("QuerySet", getattr(self, reverse_name)) fk = cast("BaseForeignKeyField", self.meta.fields[reverse_name].foreign_key) if fk.force_cascade_deletion_relation: - await referenced_obs.using(schema=self.schema_name).raw_delete( - use_models=fk.use_model_based_deletion, remove_referenced_call=reverse_name - ) + with suppress(SkipOperation): + await referenced_obs.using(schema=self.schema_name).raw_delete( + use_models=fk.use_model_based_deletion, remove_referenced_call=reverse_name + ) return row_count diff --git a/edgy/core/db/datastructures.py b/edgy/core/db/datastructures.py index e631e0fd..ba93c1d7 100644 --- a/edgy/core/db/datastructures.py +++ b/edgy/core/db/datastructures.py @@ -234,7 +234,7 @@ def create_sub_cache(self, attrs: Sequence[str], prefix: str = "") -> Self: Returns: Self: A new `QueryModelResultCache` instance. """ - return self.__class__(attrs, prefix=prefix, cache=self.cache) + return type(self)(attrs, prefix=prefix, cache=self.cache) def clear( self, model_class: type[BaseModelType] | None = None, prefix: str | None = None @@ -600,3 +600,7 @@ async def _helper(row_or_model: Any) -> Any: prefix=prefix, ) return results + + def __bool__(self) -> bool: + """Copy empty check of internal cache dictionary.""" + return bool(self.cache) diff --git a/edgy/core/db/fields/foreign_keys.py b/edgy/core/db/fields/foreign_keys.py index e96b4402..cbb8ca8b 100644 --- a/edgy/core/db/fields/foreign_keys.py +++ b/edgy/core/db/fields/foreign_keys.py @@ -21,7 +21,7 @@ VirtualCascadeDeletionSingleRelation, ) from edgy.core.terminal import Print -from edgy.exceptions import FieldDefinitionError +from edgy.exceptions import FieldDefinitionError, SkipOperation from edgy.protocols.many_relationship import ManyRelationProtocol if TYPE_CHECKING: @@ -158,6 +158,8 @@ async def _notset_post_delete_callback(self, value: Any) -> None: # triggered by a reverse relation, it doesn't cause a loop. remove_referenced_call=self.reverse_name or True, ) + except SkipOperation: + ... finally: # Reset the current instance context. CURRENT_INSTANCE.reset(token) diff --git a/edgy/core/db/models/mixins/db.py b/edgy/core/db/models/mixins/db.py index 482a152d..00bdf043 100644 --- a/edgy/core/db/models/mixins/db.py +++ b/edgy/core/db/models/mixins/db.py @@ -32,7 +32,12 @@ from edgy.core.db.relationships.related_field import RelatedField from edgy.core.utils.db import check_db_connection, hash_names from edgy.core.utils.models import create_edgy_model -from edgy.exceptions import ForeignKeyBadConfigured, ModelCollisionError, ObjectNotFound +from edgy.exceptions import ( + ForeignKeyBadConfigured, + ModelCollisionError, + ObjectNotFound, + SkipOperation, +) from edgy.types import Undefined if sys.version_info >= (3, 11): # pragma: no cover @@ -779,13 +784,25 @@ async def _update( model_instance=self, evaluate_values=True, ) - await pre_fn( - real_class, - model_instance=self, - instance=instance, - values=kwargs, - column_values=column_values, - ) + try: + await pre_fn( + real_class, + model_instance=self, + instance=instance, + values=kwargs, + column_values=column_values, + ) + except SkipOperation: + await post_fn( + real_class, + model_instance=self, + instance=instance, + values=kwargs, + column_values=column_values, + operation_skipped=True, + row_count=0, + ) + return 0 # empty updates shouldn't cause an error. E.g. only model references are updated clauses = self.identifying_clauses() row_count: int | None = None @@ -819,6 +836,8 @@ async def _update( instance=instance, values=kwargs, column_values=column_values, + operation_skipped=False, + row_count=row_count, ) return row_count @@ -882,9 +901,20 @@ async def raw_delete( instance is not self or remove_referenced_call ) if with_signals: - await self.meta.signals.pre_delete.send_async( - real_class, instance=instance, model_instance=self - ) + try: + await self.meta.signals.pre_delete.send_async( + real_class, instance=instance, model_instance=self, injected_filters=None + ) + except SkipOperation as exc: + await self.meta.signals.post_delete.send_async( + real_class, + instance=CURRENT_INSTANCE.get(), + model_instance=self, + row_count=0, + operation_skipped=True, + ) + # reraises + raise exc ignore_fields: set[str] = set() if remove_referenced_call and isinstance(remove_referenced_call, str): ignore_fields.add(remove_referenced_call) @@ -937,6 +967,7 @@ async def raw_delete( instance=CURRENT_INSTANCE.get(), model_instance=self, row_count=row_count, + operation_skipped=False, ) return row_count @@ -950,19 +981,31 @@ async def delete(self: Model, skip_post_delete_hooks: bool = False) -> int: skip_post_delete_hooks: If True, post-delete hooks will not be executed. """ real_class = self.get_real_class() - await self.meta.signals.pre_delete.send_async( - real_class, instance=self, model_instance=self - ) + try: + await self.meta.signals.pre_delete.send_async( + real_class, instance=self, model_instance=self, injected_filters=None + ) + except SkipOperation: + await self.meta.signals.post_delete.send_async( + real_class, instance=self, model_instance=self, row_count=0, operation_skipped=True + ) + return 0 token = CURRENT_INSTANCE.set(self) try: row_count = await self.raw_delete( skip_post_delete_hooks=skip_post_delete_hooks, remove_referenced_call=False, ) + except SkipOperation: + row_count = 0 finally: CURRENT_INSTANCE.reset(token) await self.meta.signals.post_delete.send_async( - real_class, instance=self, model_instance=self, row_count=row_count + real_class, + instance=self, + model_instance=self, + row_count=row_count, + operation_skipped=False, ) return row_count @@ -1070,13 +1113,24 @@ async def _insert( model_instance=self, evaluate_values=evaluate_values, ) - await pre_fn( - real_class, - model_instance=self, - instance=instance, - column_values=column_values, - values=kwargs, - ) + try: + await pre_fn( + real_class, + model_instance=self, + instance=instance, + column_values=column_values, + values=kwargs, + ) + except SkipOperation: + await post_fn( + real_class, + model_instance=self, + instance=instance, + column_values=column_values, + values=kwargs, + operation_skipped=True, + ) + return check_db_connection(self.database, stacklevel=4) table: sqlalchemy.Table = self.table async with ( @@ -1132,6 +1186,7 @@ async def _insert( instance=instance, column_values=column_values, values=kwargs, + operation_skipped=False, ) async def real_save( @@ -1177,6 +1232,7 @@ async def real_save( if value is None and self.table.columns[pkcolumn].autoincrement: extracted_fields.pop(pkcolumn, None) force_insert = True + break field = self.meta.fields.get(pkcolumn) # this is an IntegerField/DateTime with primary_key set if field is not None: @@ -1185,10 +1241,12 @@ async def real_save( ): # we create a new revision. force_insert = True + break elif getattr(field, "auto_now_add", False): # noqa: SIM102 # force_insert if auto_now_add field is empty if value is None: force_insert = True + break # check if it exists if not force_insert and not await self.check_exist_in_db(only_needed=True): @@ -1203,8 +1261,8 @@ async def real_save( extracted_fields.update(values) # force save must ensure a complete mapping await self._insert( - bool(values), - extracted_fields, + evaluate_values=bool(values), + kwargs=extracted_fields, pre_fn=partial( self.meta.signals.pre_save.send_async, is_update=False, is_migration=False ), @@ -1215,9 +1273,9 @@ async def real_save( ) else: await self._update( - # assume partial when values are None - values is not None, - extracted_fields if values is None else values, + # assume partial when values are not None + is_partial=values is not None, + kwargs=extracted_fields if values is None else values, pre_fn=partial( self.meta.signals.pre_save.send_async, is_update=True, is_migration=False ), diff --git a/edgy/core/db/models/types.py b/edgy/core/db/models/types.py index 00afccff..d4cf1a8a 100644 --- a/edgy/core/db/models/types.py +++ b/edgy/core/db/models/types.py @@ -259,6 +259,8 @@ async def raw_delete( The `remove_referenced_call` as a string is crucial when traversing related fields for deletions, as it helps in trimming stub back references that might otherwise lead to incorrect model deletions. + + The model type raw_delete can raise SkipOperation. """ @abstractmethod diff --git a/edgy/core/db/querysets/base.py b/edgy/core/db/querysets/base.py index 9f388530..e5aa7285 100644 --- a/edgy/core/db/querysets/base.py +++ b/edgy/core/db/querysets/base.py @@ -8,6 +8,7 @@ Iterable, Sequence, ) +from contextvars import ContextVar from functools import cached_property from inspect import isawaitable from itertools import chain @@ -34,7 +35,6 @@ from .compiler import QueryCompiler from .executor import QueryExecutor, get_current_row from .mixins import QuerySetPropsMixin, TenancyMixin -from .parser import ResultParser from .prefetch import Prefetch, PrefetchMixin from .types import ( EdgyEmbedTarget, @@ -50,6 +50,9 @@ from edgy.core.db.querysets.queryset import QuerySet _empty_set = cast(set[Any], frozenset()) +_injected_filters_deletion: ContextVar[Iterable] = ContextVar( + "_injected_filters_deletion", default=() +) class BaseQuerySet( @@ -223,8 +226,8 @@ def _clear_cache( tuple[Any, dict[str, tuple[sqlalchemy.Table, type[BaseModelType]]]] | None ) = None self._cache_count: int | None = None - self._cache_first: tuple[BaseModelType, Any] | None = None - self._cache_last: tuple[BaseModelType, Any] | None = None + self._cache_first: tuple[EdgyModel, EdgyEmbedTarget] | None = None + self._cache_last: tuple[EdgyModel, EdgyEmbedTarget] | None = None self._cache_fetch_all: bool = False def _build_order_by_iterable( @@ -482,9 +485,7 @@ async def _execute_iterate( (Refactored: Now delegates to the Executor) """ # Create the specialists - compiler = QueryCompiler(self) - parser = ResultParser(self) - executor = QueryExecutor(self, compiler, parser) + executor = QueryExecutor(self) # Delegate the work async for model in executor.iterate(fetch_all_at_once=fetch_all_at_once): # type: ignore @@ -665,38 +666,43 @@ async def raw_delete( self, use_models: bool = False, remove_referenced_call: str | bool = False ) -> int: """ - (Refactored: Now delegates to the Executor) + Executes delete without raising an extra signal. + + Delegates to QueryExecutor.delete. """ # We must create new specialists *every time* because the queryset # state might have changed (e.g., in _model_based_delete) - compiler = QueryCompiler(self) - parser = ResultParser(self) # Delete doesn't use parser, but good practice - executor = QueryExecutor(self, compiler, parser) + executor = QueryExecutor(self) return await executor.delete( - use_models=use_models, remove_referenced_call=remove_referenced_call + use_models=use_models, + remove_referenced_call=remove_referenced_call, + injected_filters=_injected_filters_deletion.get(), ) - async def _get_raw(self, **kwargs: Any) -> tuple[BaseModelType, Any]: + async def _get_raw( + self, kwargs: dict | None = None, no_update_result_cache: bool = False + ) -> tuple[EdgyModel, EdgyEmbedTarget]: """ - (Refactored: Builder logic stays, execution logic delegates) + Base method used by get like methods. """ if kwargs: cached = cast( - tuple[BaseModelType, Any] | None, self._cache.get(self.model_class, kwargs) + "tuple[EdgyModel, EdgyEmbedTarget] | None", + self._cache.get(self.model_class, kwargs), ) if cached is not None: return cached filter_query = cast("BaseQuerySet", self.filter(**kwargs)) filter_query._cache = self._cache - return await filter_query._get_raw() + return await filter_query._get_raw(no_update_result_cache=no_update_result_cache) elif self._cache_count == 1: if self._cache_first is not None: return self._cache_first elif self._cache_last is not None: return self._cache_last - - compiler = QueryCompiler(self) - parser = ResultParser(self) - executor = QueryExecutor(self, compiler, parser) - return await executor.get_one() + executor = QueryExecutor(self) + return cast( + "tuple[EdgyModel, EdgyEmbedTarget]", + await executor.get_one(no_update_result_cache=no_update_result_cache), + ) diff --git a/edgy/core/db/querysets/bulk.py b/edgy/core/db/querysets/bulk.py index 36afb1b6..6b850d3e 100644 --- a/edgy/core/db/querysets/bulk.py +++ b/edgy/core/db/querysets/bulk.py @@ -17,7 +17,7 @@ from edgy.core.db.context_vars import CURRENT_INSTANCE from edgy.core.utils.concurrency import run_concurrently from edgy.core.utils.db import check_db_connection -from edgy.exceptions import QuerySetError +from edgy.exceptions import QuerySetError, SkipOperation from .types import ( EdgyEmbedTarget, @@ -354,7 +354,13 @@ async def send_pre_signal(self) -> None: **self.signal_params, ) ) - await asyncio.gather(*ops) + try: + await asyncio.gather(*ops) + except SkipOperation as exc: + self.execution_step = 4 # cache + self.signal_params["operation_skipped"] = True + self.signal_params["values"] = None + raise exc async def apply_db(self) -> None: """ @@ -618,6 +624,7 @@ async def _iterate_update(obj: EdgyModel) -> dict[str, Any]: "values": self.result, "create_params": self.create_params, "update_params": self.update_params, + "operation_skipped": False, **self.provided_signal_params, } if self.create: diff --git a/edgy/core/db/querysets/executor.py b/edgy/core/db/querysets/executor.py index a396fd55..32e39d5b 100644 --- a/edgy/core/db/querysets/executor.py +++ b/edgy/core/db/querysets/executor.py @@ -1,7 +1,7 @@ from __future__ import annotations import warnings -from collections.abc import AsyncGenerator, Sequence +from collections.abc import AsyncGenerator, Iterable, Sequence from contextvars import ContextVar from typing import TYPE_CHECKING, Any, cast @@ -9,17 +9,17 @@ from edgy.core.db.context_vars import CURRENT_INSTANCE from edgy.core.db.datastructures import QueryModelResultCache +from edgy.core.db.querysets.clauses import and_, or_ +from edgy.core.db.querysets.parser import ResultParser from edgy.core.db.querysets.prefetch import Prefetch, check_prefetch_collision from edgy.core.db.relationships.utils import crawl_relationship from edgy.core.utils.db import check_db_connection, hash_tablekey -from edgy.exceptions import MultipleObjectsReturned, ObjectNotFound, QuerySetError +from edgy.exceptions import MultipleObjectsReturned, ObjectNotFound, QuerySetError, SkipOperation from .types import EdgyEmbedTarget, EdgyModel if TYPE_CHECKING: # pragma: no cover from edgy.core.db.querysets.base import BaseQuerySet - from edgy.core.db.querysets.compiler import QueryCompiler - from edgy.core.db.querysets.parser import ResultParser from edgy.core.db.querysets.queryset import QuerySet from .types import tables_and_models_type @@ -47,28 +47,33 @@ class QueryExecutor: def __init__( self, queryset: BaseQuerySet, - compiler: QueryCompiler, - parser: ResultParser, ): """ Initializes the QueryExecutor. Args: queryset: The BaseQuerySet instance holding the query state. - compiler: The QueryCompiler to be used for WHERE clauses (e.g., in deletes). - parser: The ResultParser to be used for turning rows into models. """ + self.set_queryset(queryset) + + def set_queryset(self, queryset: BaseQuerySet) -> None: + from .compiler import QueryCompiler + # we need so many internals, so we just cast to a QuerySet self.queryset = cast("QuerySet", queryset) - self.compiler = compiler - self.parser = parser + self.compiler = QueryCompiler(self.queryset) + self.parser: None | ResultParser = None self.database = queryset.database self.model_class = queryset.model_class + def init_parser(self, tables_and_models: tables_and_models_type) -> None: + from .parser import ResultParser + + self.parser = ResultParser(self.queryset, tables_and_models) + async def _process_and_yield_batch( self, batch: Sequence[sqlalchemy.Row], - tables_and_models: tables_and_models_type, new_cache: QueryModelResultCache, ) -> AsyncGenerator[tuple[tuple[EdgyModel, EdgyEmbedTarget], sqlalchemy.Row], None]: """ @@ -87,9 +92,10 @@ async def _process_and_yield_batch( - (result_tuple): The (raw_model, embed_target) tuple. - (row): The raw SQLAlchemy Row. """ - prefetches = await self._prepare_prefetches_for_batch(batch, tables_and_models) + assert self.parser is not None, "parser not initialized" + prefetches = await self._prepare_prefetches_for_batch(batch, self.parser.tables_and_models) results: Sequence[tuple[EdgyModel, EdgyEmbedTarget]] = await self.parser.batch_to_models( - batch, tables_and_models, prefetches, new_cache + batch, prefetches, new_cache ) for row_num, result_tuple in enumerate(results): @@ -122,6 +128,7 @@ async def iterate(self, fetch_all_at_once: bool = False) -> AsyncGenerator[EdgyM qs = qs.distinct() expression, tables_and_models = await qs.as_select_with_tables() + self.init_parser(tables_and_models) if not fetch_all_at_once and bool(self.database.force_rollback): warnings.warn( @@ -146,18 +153,18 @@ async def iterate(self, fetch_all_at_once: bool = False) -> AsyncGenerator[EdgyM batch = cast(Sequence[sqlalchemy.Row], await database.fetch_all(expression)) # Use the new helper to process the single, large batch - async for result, row in self._process_and_yield_batch( - batch, tables_and_models, new_cache - ): + async for result, row in self._process_and_yield_batch(batch, new_cache): if counter == 0: - qs._cache_first = result + # qs is maybe a copy, so update cache on self.queryset + self.queryset._cache_first = result last_element = result counter += 1 current_row[0] = row yield result[1] - qs._cache_fetch_all = True - qs._cache = new_cache + # qs is maybe a copy, so update cache on self.queryset + self.queryset._cache_fetch_all = True + self.queryset._cache = new_cache else: batch_num: int = 0 new_cache = QueryModelResultCache(qs._cache.attrs) @@ -170,11 +177,9 @@ async def iterate(self, fetch_all_at_once: bool = False) -> AsyncGenerator[EdgyM qs._cache_fetch_all = False # Use the new helper to process each batch - async for result, row in self._process_and_yield_batch( - batch, tables_and_models, new_cache - ): + async for result, row in self._process_and_yield_batch(batch, new_cache): if counter == 0: - qs._cache_first = result + self.queryset._cache_first = result last_element = result counter += 1 current_row[0] = row @@ -182,15 +187,19 @@ async def iterate(self, fetch_all_at_once: bool = False) -> AsyncGenerator[EdgyM batch_num += 1 if batch_num <= 1: - qs._cache = new_cache - qs._cache_fetch_all = True + # qs is maybe a copy, so update cache on self.queryset + self.queryset._cache = new_cache + self.queryset._cache_fetch_all = True finally: _current_row_holder.reset(token) - - qs._cache_count = counter - qs._cache_last = last_element - - async def get_one(self) -> tuple[EdgyModel, EdgyEmbedTarget]: + # qs is maybe a copy, so update cache on self.queryset + self.queryset._cache_count = counter + self.queryset._cache_last = last_element + self.parser = None + + async def get_one( + self, no_update_result_cache: bool = False + ) -> tuple[EdgyModel, EdgyEmbedTarget]: """ Fetches a single unique record from the database. This is the refactored _get_raw (when no kwargs are present). @@ -203,6 +212,7 @@ async def get_one(self) -> tuple[EdgyModel, EdgyEmbedTarget]: MultipleObjectsReturned: If more than one record is found. """ expression, tables_and_models = await self.queryset.as_select_with_tables() + self.init_parser(tables_and_models) check_db_connection(self.database, stacklevel=4) async with self.database as database: @@ -215,14 +225,17 @@ async def get_one(self) -> tuple[EdgyModel, EdgyEmbedTarget]: raise MultipleObjectsReturned() self.queryset._cache_count = 1 + if no_update_result_cache: + resultsingle: EdgyModel = await self.parser.row_to_model_uncached(rows[0]) + return resultsingle, cast(EdgyEmbedTarget, resultsingle) - result: tuple[EdgyModel, EdgyEmbedTarget] = await self.parser.row_to_model( - rows[0], tables_and_models - ) + result: tuple[EdgyModel, EdgyEmbedTarget] = await self.parser.row_to_model(rows[0]) # Update cache attributes + self.queryset._cache_fetch_all = True self.queryset._cache_first = result self.queryset._cache_last = result + self.parser = None return result async def _prepare_prefetches_for_batch( @@ -295,7 +308,10 @@ async def _prepare_prefetches_for_batch( return prepared_prefetches async def delete( - self, use_models: bool = False, remove_referenced_call: str | bool = False + self, + use_models: bool = False, + remove_referenced_call: str | bool = False, + injected_filters: Iterable = (), ) -> int: """ Executes a delete operation. @@ -317,6 +333,8 @@ async def delete( or self.model_class.meta.post_delete_fields ): use_models = True + if injected_filters: + self.set_queryset(self.queryset.filter(*injected_filters)) if use_models: row_count = await self._model_based_delete( @@ -331,7 +349,8 @@ async def delete( row_count = cast(int, await database.execute(expression)) # clear cache after deletion. - self.queryset._clear_cache(keep_cached_selected=True) + if row_count != 0: + self.queryset._clear_cache(keep_cached_selected=True) return row_count async def _model_based_delete(self, remove_referenced_call: str | bool) -> int: @@ -347,38 +366,44 @@ async def _model_based_delete(self, remove_referenced_call: str | bool) -> int: Returns: The total number of models deleted. """ - from edgy.core.db.querysets.compiler import QueryCompiler - from edgy.core.db.querysets.parser import ResultParser queryset = ( self.queryset.limit(self.queryset._batch_size) - if not self.queryset._cache_fetch_all - else self.queryset + if self.queryset._batch_size is not None + else self.queryset.all() ) queryset.embed_parent = None row_count = 0 - compiler = QueryCompiler(queryset) - parser = ResultParser(queryset) - - # Instantiate the QueryExecutor recursively for the new queryset - executor = QueryExecutor(queryset, compiler, parser) - - # Uuse the new executor's iterate method - models = [model async for model in executor.iterate(fetch_all_at_once=True)] # type: ignore - + executor = QueryExecutor(queryset) token = CURRENT_INSTANCE.set(self.queryset) try: + # Use the new executor's iterate method + models = [model async for model in executor.iterate(fetch_all_at_once=True)] # type: ignore while models: + exclusion_filters = [] for model in models: - await model.raw_delete( - skip_post_delete_hooks=False, remove_referenced_call=remove_referenced_call - ) - row_count += 1 - + try: + _row_count = await model.raw_delete( + skip_post_delete_hooks=False, + remove_referenced_call=remove_referenced_call, + ) + except SkipOperation: + # raised from raw_delete + exclusion_filters.append(and_(*model.identifying_clauses())) + continue + if _row_count != 0: + row_count += 1 + + # we fetched all + if self.queryset._batch_size is None: + break + if exclusion_filters: + executor.set_queryset(executor.queryset.exclude(or_(*exclusion_filters))) + # clear again + exclusion_filters.clear() # clear cache and fetch new batch - # reuse cached query - queryset._clear_cache(keep_cached_selected=True) + executor.queryset._clear_cache(keep_cached_selected=True) models = [model async for model in executor.iterate(fetch_all_at_once=True)] # type: ignore finally: CURRENT_INSTANCE.reset(token) diff --git a/edgy/core/db/querysets/parser.py b/edgy/core/db/querysets/parser.py index abd18cfd..f76b59f2 100644 --- a/edgy/core/db/querysets/parser.py +++ b/edgy/core/db/querysets/parser.py @@ -20,36 +20,47 @@ class ResultParser: including caching and relationship embedding. """ - def __init__(self, queryset: BaseQuerySet | Any) -> None: + def __init__(self, queryset: BaseQuerySet, tables_and_models: tables_and_models_type) -> None: self.queryset = queryset self.model_class = queryset.model_class + self.tables_and_models = tables_and_models + self.is_defer_fields = bool(self.queryset._defer) - async def row_to_model( + async def row_to_model_uncached( self, row: sqlalchemy.Row | Any, - tables_and_models: tables_and_models_type, - ) -> tuple[EdgyModel, EdgyEmbedTarget]: + ) -> EdgyModel: """ - Parses a single row into a model instance, using the cache. - (Refactored from _get_or_cache_row) + Parses a single row into a model instance, without using the cache. """ - is_defer_fields = bool(self.queryset._defer) - - result = await self.queryset._cache.aget_or_cache_many( - self.model_class, - [row], - cache_fn=lambda _row: self.model_class.from_sqla_row( - _row, - tables_and_models=tables_and_models, + return cast( + "EdgyModel", + await self.model_class.from_sqla_row( + row, + tables_and_models=self.tables_and_models, select_related=self.queryset._select_related, only_fields=self.queryset._only, - is_defer_fields=is_defer_fields, + is_defer_fields=self.is_defer_fields, prefetch_related=self.queryset._prefetch_related, exclude_secrets=self.queryset._exclude_secrets, using_schema=self.queryset.active_schema, database=self.queryset.database, reference_select=self.queryset._reference_select, ), + ) + + async def row_to_model( + self, + row: sqlalchemy.Row | Any, + ) -> tuple[EdgyModel, EdgyEmbedTarget]: + """ + Parses a single row into a model instance, using the cache. + (Refactored from _get_or_cache_row) + """ + result = await self.queryset._cache.aget_or_cache_many( + self.model_class, + [row], + cache_fn=self.row_to_model_uncached, transform_fn=self.queryset._embed_parent_in_result, ) return cast(tuple[EdgyModel, EdgyEmbedTarget], result[0]) @@ -57,7 +68,6 @@ async def row_to_model( async def batch_to_models( self, batch: Sequence[sqlalchemy.Row], - tables_and_models: tables_and_models_type, prefetch_list: list[Prefetch], new_cache: QueryModelResultCache, ) -> Sequence[tuple[EdgyModel, EdgyEmbedTarget]]: @@ -65,7 +75,6 @@ async def batch_to_models( Parses a batch of rows into model instances. (This is the parsing half of the original _handle_batch method) """ - is_defer_fields = bool(self.queryset._defer) qs = self.queryset return await new_cache.aget_or_cache_many( @@ -73,10 +82,10 @@ async def batch_to_models( batch, cache_fn=lambda row: self.model_class.from_sqla_row( row, - tables_and_models=tables_and_models, + tables_and_models=self.tables_and_models, select_related=qs._select_related, only_fields=qs._only, - is_defer_fields=is_defer_fields, + is_defer_fields=self.is_defer_fields, prefetch_related=prefetch_list, # Use the prepared list exclude_secrets=qs._exclude_secrets, using_schema=qs.active_schema, diff --git a/edgy/core/db/querysets/queryset.py b/edgy/core/db/querysets/queryset.py index 61fedbf5..0fc85dba 100644 --- a/edgy/core/db/querysets/queryset.py +++ b/edgy/core/db/querysets/queryset.py @@ -20,11 +20,11 @@ from edgy.core.db.models.model_reference import ModelRef from edgy.core.db.models.types import BaseModelType from edgy.core.db.models.utils import apply_instance_extras -from edgy.core.db.querysets.base import BaseQuerySet +from edgy.core.db.querysets.base import BaseQuerySet, _injected_filters_deletion from edgy.core.db.querysets.parser import ResultParser from edgy.core.utils.db import CHECK_DB_CONNECTION_SILENCED, check_db_connection from edgy.core.utils.sync import run_sync -from edgy.exceptions import ObjectNotFound, QuerySetError +from edgy.exceptions import ObjectNotFound, QuerySetError, SkipOperation from .bulk import BulkOperation from .types import ( @@ -887,7 +887,7 @@ async def get(self, **kwargs: Any) -> EdgyEmbedTarget: ObjectNotFound: If no object is found. MultipleObjectsReturned: If more than one object is found (implicitly handled by underlying `_get_raw`). """ - return cast(EdgyEmbedTarget, (await self._get_raw(**kwargs))[1]) + return (await self._get_raw(kwargs=kwargs))[1] select = get @@ -900,10 +900,10 @@ async def first(self) -> EdgyEmbedTarget | None: Returns: The first model instance, or `None` if the QuerySet is empty. """ - if self._cache_count is not None and self._cache_count == 0: + if self._cache_count == 0: return None if self._cache_first is not None: - return cast(EdgyEmbedTarget, self._cache_first[1]) + return self._cache_first[1] queryset = self if not queryset._order_by: queryset = queryset.order_by(*self.model_class.pkcolumns) @@ -913,22 +913,22 @@ async def first(self) -> EdgyEmbedTarget | None: async with queryset.database as database: row = await database.fetch_one(expression, pos=0) if row: - parser = ResultParser(self) - result_tuple: tuple[Any, EdgyEmbedTarget] = await parser.row_to_model( - row, tables_and_models - ) + parser = ResultParser(self, tables_and_models) + result_tuple: tuple[Any, EdgyEmbedTarget] = await parser.row_to_model(row) self._cache_first = result_tuple return result_tuple[1] + else: + self._cache_count = 0 return None async def last(self) -> EdgyEmbedTarget | None: """ Returns the last record from the QuerySet... """ - if self._cache_count is not None and self._cache_count == 0: + if self._cache_count == 0: return None if self._cache_last is not None: - return cast(EdgyEmbedTarget, self._cache_last[1]) + return self._cache_last[1] queryset = self if not queryset._order_by: queryset = queryset.order_by(*self.model_class.pkcolumns) @@ -940,12 +940,12 @@ async def last(self) -> EdgyEmbedTarget | None: row = await database.fetch_one(expression, pos=0) if row: # NEW FIXED LINES: - parser = ResultParser(self) - result_tuple: tuple[Any, EdgyEmbedTarget] = await parser.row_to_model( - row, tables_and_models - ) + parser = ResultParser(self, tables_and_models) + result_tuple: tuple[Any, EdgyEmbedTarget] = await parser.row_to_model(row) self._cache_last = result_tuple return result_tuple[1] + else: + self._cache_count = 0 return None async def create(self, *args: Any, **kwargs: Any) -> EdgyEmbedTarget: @@ -1032,16 +1032,39 @@ async def delete(self, use_models: bool = False) -> int: Returns: int: The number of rows deleted. """ - await self.model_class.meta.signals.pre_delete.send_async( - self.model_class, instance=self, model_instance=None - ) - row_count = await self.raw_delete(use_models=use_models, remove_referenced_call=False) + injected_filters: list[Any] = [] + try: + await self.model_class.meta.signals.pre_delete.send_async( + self.model_class, + instance=self, + model_instance=None, + injected_filters=injected_filters, + ) + except SkipOperation: + await self.model_class.meta.signals.post_delete.send_async( + self.model_class, + instance=self, + model_instance=None, + row_count=0, + operation_skipped=True, + ) + + return 0 + token = _injected_filters_deletion.set(injected_filters) + try: + row_count = await self.raw_delete(use_models=use_models, remove_referenced_call=False) + finally: + _injected_filters_deletion.reset(token) await self.model_class.meta.signals.post_delete.send_async( - self.model_class, instance=self, model_instance=None, row_count=row_count + self.model_class, + instance=self, + model_instance=None, + row_count=row_count, + operation_skipped=False, ) return row_count - async def update(self, **kwargs: Any) -> None: + async def update(self, **kwargs: Any) -> None | int: """ Updates records in a specific table with the given keyword arguments, matching the QuerySet's filters. @@ -1053,6 +1076,8 @@ async def update(self, **kwargs: Any) -> None: Args: **kwargs: The field names and new values to apply to the matching records. + Returns: + int | None: Amount of rows changed if known for the database. """ column_values = self.model_class.extract_column_values( @@ -1065,21 +1090,36 @@ async def update(self, **kwargs: Any) -> None: # Broadcast the initial update details # add is_update to match save - await self.model_class.meta.signals.pre_update.send_async( - self.model_class, - instance=self, - model_instance=None, - values=kwargs, - column_values=column_values, - is_update=True, - is_migration=False, - ) + try: + await self.model_class.meta.signals.pre_update.send_async( + self.model_class, + instance=self, + model_instance=None, + values=kwargs, + column_values=column_values, + is_update=True, + is_migration=False, + ) + except SkipOperation: + await self.model_class.meta.signals.post_update.send_async( + self.model_class, + instance=self, + model_instance=None, + values=kwargs, + column_values=column_values, + is_update=True, + is_migration=False, + row_count=0, + operation_skipped=True, + ) + return 0 expression = self.table.update().values(**column_values) expression = expression.where(await self.build_where_clause()) check_db_connection(self.database) + row_count: int | None async with self.database as database: - await database.execute(expression) + row_count = cast(int | None, await database.execute(expression)) # Broadcast the update executed # add is_update to match save @@ -1091,8 +1131,11 @@ async def update(self, **kwargs: Any) -> None: column_values=column_values, is_update=True, is_migration=False, + row_count=row_count, + operation_skipped=False, ) self._clear_cache() + return row_count async def get_or_create( self, defaults: dict[str, Any] | Any | None = None, *args: Any, **kwargs: Any @@ -1117,7 +1160,7 @@ async def get_or_create( defaults = {} try: - raw_instance, get_instance = await self._get_raw(**kwargs) + raw_instance, resolved = await self._get_raw(kwargs=kwargs) except ObjectNotFound: kwargs.update(defaults) instance: EdgyEmbedTarget = await self.create(*args, **kwargs) @@ -1141,7 +1184,7 @@ async def get_or_create( ) relation = getattr(raw_instance, arg.__related_name__) await relation.add(model) - return cast(EdgyEmbedTarget, get_instance), False + return resolved, False select_or_insert = get_or_create @@ -1166,12 +1209,14 @@ async def update_or_create( args = (defaults, *args) defaults = {} try: - raw_instance, get_instance = await self._get_raw(**kwargs) + # bypass cache + raw_instance = (await self._get_raw(kwargs=kwargs, no_update_result_cache=True))[0] except ObjectNotFound: kwargs.update(defaults) instance: EdgyEmbedTarget = await self.create(*args, **kwargs) return instance, True - await get_instance.update(**defaults) + # when updating the resolved, we can end up with a complete different model type + await raw_instance.update(**defaults) for arg in args: if isinstance(arg, ModelRef): relation_field = self.model_class.meta.fields[arg.__related_name__] @@ -1191,8 +1236,22 @@ async def update_or_create( ) relation = getattr(raw_instance, arg.__related_name__) await relation.add(model) - self._clear_cache() - return cast(EdgyEmbedTarget, get_instance), False + # we can keep the result cache because we update it and the results are only used for parsing + self._clear_cache(keep_cached_selected=True, keep_result_cache=True) + # now resolve again + resolved = (await self._embed_parent_in_result(raw_instance))[1] + # update the result cache now + self._cache.update( + self.model_class, + values=[(raw_instance, resolved)], + cache_keys=[self._cache.create_cache_key(self.model_class, raw_instance)], + ) + if not kwargs: + # 1. no extra filters, 2. only one result => we can fill the cache + self._cache_first = (raw_instance, resolved) + self._cache_last = (raw_instance, resolved) + self._cache_fetch_all = True + return resolved, False update_or_insert = update_or_create @@ -1228,7 +1287,7 @@ async def bulk_create( self, objs: Iterable[dict[str, Any] | EdgyModel], *, - ignore_conflicts: Literal[True], + ignore_conflicts: bool = False, resolve_embed: Literal[True], ) -> list[EdgyEmbedTarget | None]: ... @@ -1237,27 +1296,10 @@ async def bulk_create( self, objs: Iterable[dict[str, Any] | EdgyModel], *, - ignore_conflicts: Literal[False] = False, - resolve_embed: Literal[True], - ) -> list[EdgyEmbedTarget]: ... - - @overload - async def bulk_create( - self, - objs: Iterable[dict[str, Any] | EdgyModel], - *, - ignore_conflicts: Literal[True], + ignore_conflicts: bool = False, resolve_embed: Literal[False] = False, ) -> list[EdgyModel | None]: ... - @overload - async def bulk_create( - self, - objs: Iterable[dict[str, Any] | EdgyModel], - *, - ignore_conflicts: Literal[False] = False, - resolve_embed: Literal[False] = False, - ) -> list[EdgyModel]: ... async def bulk_create( self, objs: Iterable[dict[str, Any] | EdgyModel], @@ -1279,7 +1321,7 @@ async def bulk_create( resolve_embed (bool): Triggers mode in which embedding is applied when True. Returns: - list[EdgyModel] | list[EdgyEmbedTarget]: + list[EdgyModel | None] | list[EdgyEmbedTarget | None]: A list of created objects. Warning: for performance reasons no embedding is applied by default and the returned objects are maybe incomplete (check `can_load` property). @@ -1301,7 +1343,14 @@ async def bulk_create( ignore_create_conflicts=ignore_conflicts, ) await operation.prepare(objs) - await operation.send_pre_signal() + try: + await operation.send_pre_signal() + except SkipOperation: + await operation.send_post_signal() + return cast( + "list[EdgyModel | None] | list[EdgyEmbedTarget | None]", + [None for _ in operation.instances_and_created], + ) await operation.apply_db() operation.update_cache() await operation.send_post_signal() @@ -1405,7 +1454,14 @@ async def bulk_update( resolve_embed=resolve_embed, ) await operation.prepare(objs) - await operation.send_pre_signal() + try: + await operation.send_pre_signal() + except SkipOperation: + await operation.send_post_signal() + return cast( + "list[EdgyModel | None] | list[EdgyEmbedTarget | None]", + [None for _ in operation.instances_and_created], + ) await operation.apply_db() operation.update_cache() await operation.send_post_signal() @@ -1508,7 +1564,14 @@ async def bulk_update_or_create( resolve_embed=resolve_embed, ) await operation.prepare(objs) - await operation.send_pre_signal() + try: + await operation.send_pre_signal() + except SkipOperation: + await operation.send_post_signal() + return cast( + "list[tuple[EdgyModel | None, bool]] | list[tuple[EdgyEmbedTarget | None, bool]]", + [(None, False) for _ in operation.instances_and_created], + ) await operation.apply_db() operation.update_cache() await operation.send_post_signal() @@ -1521,7 +1584,7 @@ async def bulk_get_or_create( *, unique_fields: Iterable[str] | None = None, resolve_embed: Literal[True], - ) -> list[tuple[EdgyEmbedTarget, bool]]: ... + ) -> list[tuple[EdgyEmbedTarget | None, bool]]: ... @overload async def bulk_get_or_create( @@ -1530,7 +1593,7 @@ async def bulk_get_or_create( *, unique_fields: Iterable[str] | None = None, resolve_embed: Literal[False] = False, - ) -> list[tuple[EdgyModel, bool]]: ... + ) -> list[tuple[EdgyModel | None, bool]]: ... async def bulk_get_or_create( self, @@ -1538,7 +1601,7 @@ async def bulk_get_or_create( *, unique_fields: Iterable[str] | None = None, resolve_embed: bool = False, - ) -> list[tuple[EdgyModel, bool]] | list[tuple[EdgyEmbedTarget, bool]]: + ) -> list[tuple[EdgyModel | None, bool]] | list[tuple[EdgyEmbedTarget | None, bool]]: """ Bulk gets or creates records in a table. @@ -1553,7 +1616,7 @@ async def bulk_get_or_create( resolve_embed (bool): Triggers mode in which embedding is applied when True. Returns: - list[tuple[EdgyModel, bool]] | list[tuple[EdgyEmbedTarget, bool]]: + list[tuple[EdgyModel | None, bool]] | list[tuple[EdgyEmbedTarget | None, bool]]: A list of tuples with retrieved or newly created objects and created flag. Warning: for performance reasons no embedding is applied by default and the returned objects are maybe incomplete (check `can_load` property). @@ -1588,7 +1651,14 @@ async def bulk_get_or_create( resolve_embed=resolve_embed, ) await operation.prepare(objs) - await operation.send_pre_signal() + try: + await operation.send_pre_signal() + except SkipOperation: + await operation.send_post_signal() + return cast( + "list[tuple[EdgyModel | None, bool]] | list[tuple[EdgyEmbedTarget | None, bool]]", + [(None, False) for _ in operation.instances_and_created], + ) await operation.apply_db() operation.update_cache() await operation.send_post_signal() @@ -1610,6 +1680,9 @@ def transaction(self, *, force_rollback: bool = False, **kwargs: Any) -> Transac """ return self.database.transaction(force_rollback=force_rollback, **kwargs) + def __repr__(self) -> str: + return f"QuerySet at {hex(id(self))}>" + def __await__( self, ) -> Generator[Any, None, list[EdgyEmbedTarget]]: diff --git a/edgy/core/db/querysets/types.py b/edgy/core/db/querysets/types.py index 7e0b9d90..83497a74 100644 --- a/edgy/core/db/querysets/types.py +++ b/edgy/core/db/querysets/types.py @@ -521,25 +521,16 @@ async def bulk_create( self, objs: Iterable[dict[str, Any] | EdgyModel], *, - ignore_conflicts: Literal[True], - resolve_embed: Literal[True], - ) -> list[EdgyEmbedTarget | None]: ... - - @overload - async def bulk_create( - self, - objs: Iterable[dict[str, Any] | EdgyModel], - *, - ignore_conflicts: Literal[False] = False, + ignore_conflicts: bool = False, resolve_embed: Literal[True], - ) -> list[EdgyEmbedTarget]: + ) -> list[EdgyEmbedTarget | None]: """ Args: ... resolve_embed (True): Enables embedding. Returns: - list[EdgyEmbedTarget]: A list of created objects. + list[EdgyEmbedTarget | None]: A list of created objects. Warning: Instances which are not compatible (`can_load` is False) will execute an extra save. """ @@ -549,24 +540,15 @@ async def bulk_create( self, objs: Iterable[dict[str, Any] | EdgyModel], *, - ignore_conflicts: Literal[True], - resolve_embed: Literal[False] = False, - ) -> list[EdgyModel | None]: ... - - @overload - async def bulk_create( - self, - objs: Iterable[dict[str, Any] | EdgyModel], - *, - ignore_conflicts: Literal[False] = False, + ignore_conflicts: bool = False, resolve_embed: Literal[False] = False, - ) -> list[EdgyModel]: + ) -> list[EdgyModel | None]: """ Args: ... resolve_embed (False): Disables embedding. Default. Returns: - list[EdgyModel]: A list of created objects. + list[EdgyModel | None]: A list of created objects. Warning: for performance reasons no embedding is applied and the returned objects are maybe incomplete (check `can_load` property). """ @@ -590,7 +572,7 @@ async def bulk_create( resolve_embed (bool): Triggers mode in which embedding is applied when True. Returns: - list[EdgyModel] | list[EdgyEmbedTarget]: + list[EdgyModel | None] | list[EdgyEmbedTarget | None]: A list of created objects. Warning: for performance reasons no embedding is applied by default and the returned objects are maybe incomplete (check `can_load` property). @@ -612,7 +594,7 @@ async def bulk_update( resolve_embed (True): Enables embedding. Returns: - list[EdgyEmbedTarget]: A list of updated objects. + list[EdgyEmbedTarget | None]: A list of updated objects. Warning: All models must be loadable (`can_load` property is true) otherwise an `QuerySetError` is raised. """ @@ -631,7 +613,7 @@ async def bulk_update( ... resolve_embed (False): Disables embedding. Default. Returns: - list[EdgyModel]: A list of updated objects. + list[EdgyModel | None]: A list of updated objects. Warning: for performance reasons no embedding is applied and the returned objects are maybe incomplete (check `can_load` property). """ @@ -680,7 +662,7 @@ async def bulk_update_or_create( resolve_embed (True): Enables embedding. Returns: - list[tuple[EdgyEmbedTarget, bool]]: A list of `(instance, created)` tuples. + list[tuple[EdgyEmbedTarget | None, bool]]: A list of `(instance, created)` tuples. Warning: All models must be loadable (`can_load` property is true) otherwise an `QuerySetError` is raised. """ @@ -699,7 +681,7 @@ async def bulk_update_or_create( ... resolve_embed (False): Disables embedding. Default. Returns: - list[tuple[EdgyModel, bool]]: A list of tuples with updated or created objects and created flag. + list[tuple[EdgyModel | None, bool]]: A list of tuples with updated or created objects and created flag. Warning: for performance reasons no embedding is applied and the returned objects are maybe incomplete (check `can_load` property). """ @@ -744,14 +726,14 @@ async def bulk_get_or_create( *, unique_fields: Iterable[str] | None = None, resolve_embed: Literal[True], - ) -> list[tuple[EdgyEmbedTarget, bool]]: + ) -> list[tuple[EdgyEmbedTarget | None, bool]]: """ Kwargs: ... resolve_embed (True): Enables embedding. Returns: - list[tuple[EdgyEmbedTarget, bool]]: A list of `(instance, created)` tuples. + list[tuple[EdgyEmbedTarget | None, bool]]: A list of `(instance, created)` tuples. Warning: Instances which are not compatible (`can_load` is False) will execute an extra save. """ @@ -763,13 +745,13 @@ async def bulk_get_or_create( *, unique_fields: Iterable[str] | None = None, resolve_embed: Literal[False] = False, - ) -> list[tuple[EdgyModel, bool]]: + ) -> list[tuple[EdgyModel | None, bool]]: """ Kwargs: ... resolve_embed (False): Disables embedding. Default. Returns: - list[tuple[EdgyModel, bool]]: A list of `(instance, created)` tuples. + list[tuple[EdgyModel | None, bool]]: A list of `(instance, created)` tuples. Warning: for performance reasons no embedding is applied and the returned objects are maybe incomplete (check `can_load` property). """ @@ -781,7 +763,7 @@ async def bulk_get_or_create( *, unique_fields: Iterable[str] | None = None, resolve_embed: bool = False, - ) -> list[tuple[EdgyModel, bool]] | list[tuple[EdgyEmbedTarget, bool]]: + ) -> list[tuple[EdgyModel | None, bool]] | list[tuple[EdgyEmbedTarget | None, bool]]: """ Abstract method to bulk get or create records in a table. @@ -796,7 +778,7 @@ async def bulk_get_or_create( resolve_embed (bool): Triggers mode in which embedding is applied when True. Returns: - list[tuple[EdgyModel, bool]] | list[tuple[EdgyEmbedTarget, bool]]: + list[tuple[EdgyModel | None, bool]] | list[tuple[EdgyEmbedTarget | None, bool]]: A list of tuples with retrieved or newly created objects and created flag. Warning: for performance reasons no embedding is applied by default and the returned objects are maybe incomplete (check `can_load` property). @@ -813,12 +795,14 @@ async def delete(self) -> int: ... @abstractmethod - async def update(self, **kwargs: Any) -> None: + async def update(self, **kwargs: Any) -> int | None: """ Abstract method to update fields for all objects matching the QuerySet criteria. Args: **kwargs: Keyword arguments representing the fields and their new values to update. + Returns: + int | None: Amount of rows changed if known for the database. """ ... diff --git a/edgy/core/db/relationships/related_field.py b/edgy/core/db/relationships/related_field.py index ca74ced0..98caa77f 100644 --- a/edgy/core/db/relationships/related_field.py +++ b/edgy/core/db/relationships/related_field.py @@ -326,4 +326,5 @@ async def _notset_post_delete_callback(self, value: ManyRelationProtocol) -> Non # Await the post_delete_callback if it exists. await value.post_delete_callback() - def reverse_clean(self, name: str, value: Any, for_query: bool = False) -> dict[str, Any]: ... + def reverse_clean(self, name: str, value: Any, for_query: bool = False) -> dict[str, Any]: + raise NotImplementedError() diff --git a/edgy/core/db/relationships/relation.py b/edgy/core/db/relationships/relation.py index 5d657742..2e42d87a 100644 --- a/edgy/core/db/relationships/relation.py +++ b/edgy/core/db/relationships/relation.py @@ -10,7 +10,12 @@ from edgy.core.db.fields.base import RelationshipField from edgy.core.db.querysets.bulk import BulkOperation from edgy.core.db.querysets.clauses import and_, or_ -from edgy.exceptions import ObjectNotFound, RelationshipIncompatible, RelationshipNotFound +from edgy.exceptions import ( + ObjectNotFound, + RelationshipIncompatible, + RelationshipNotFound, + SkipOperation, +) from edgy.protocols.many_relationship import ManyRelationProtocol if TYPE_CHECKING: @@ -174,7 +179,13 @@ async def save_related(self) -> None: resolve_embed=False, ) await operation.prepare(refs) - await operation.send_pre_signal() + try: + await operation.send_pre_signal() + except SkipOperation: + operation.signal_params["row_count"] = 0 + operation.signal_params["row_count_create"] = 0 + await operation.send_post_signal() + return await operation.apply_db() # no cache update because the queryset is temporarysignal # both parameters are the same here @@ -312,7 +323,7 @@ async def add_many(self, *children: BaseModelType) -> list[BaseModelType | None] Returns: list[BaseModelType | None]: A list of saved intermediate model instances, or None for each record that already exists - (IntegrityError). + or for each child None when operation was skipped. """ prepared = [] through = self.through @@ -342,7 +353,13 @@ async def add_many(self, *children: BaseModelType) -> list[BaseModelType | None] resolve_embed=True, ) await operation.prepare(prepared) - await operation.send_pre_signal() + try: + await operation.send_pre_signal() + except SkipOperation: + operation.signal_params["row_count"] = 0 + operation.signal_params["row_count_create"] = 0 + await operation.send_post_signal() + return [None for _ in prepared] await operation.apply_db() # no cache update because the queryset is temporary # we can just rename the signals parameters for the post signal @@ -371,7 +388,9 @@ async def add(self, child: BaseModelType) -> BaseModelType | None: Raises: RelationshipIncompatible: If the child type is not compatible. """ - return (await self.add_many(child))[0] + results = await self.add_many(child) + # bail out if no results are returned, in case the result list is modified in signals + return results[0] if results else None async def remove_many(self, *children: BaseModelType) -> None: """ @@ -423,16 +442,43 @@ def _helper_prepare(child: Any) -> Any: **self._shared_relation_signals_params, ) ) - await asyncio.gather(*ops) + try: + await asyncio.gather(*ops) + except SkipOperation: + ops = [] + # not really necessary but be safe + seen_signals.clear() + for model_class in [ + self._shared_relation_signals_params["source"], + self._shared_relation_signals_params["target"], + through, + ]: + signal = model_class.meta.signals.post_relation_remove + if (signal_id := id(signal)) in seen_signals: + continue + seen_signals.add(signal_id) + ops.append( + signal.send_async( + through, + instance=self.instance, + raw_values=prepared, + row_count=0, + operation_skipped=True, + model_based_deletion=model_based_deletion, + **self._shared_relation_signals_params, + ) + ) + await asyncio.gather(*ops) + return if prepared: queryset = self.get_queryset().update_embed_parent(None) # children can be removed by setting them to None clauses = [ and_(*child.identifying_clauses()) for child in prepared if child is not None ] - query = queryset.filter(or_(*clauses)) + queryset = queryset.filter(or_(*clauses)) async with queryset.transaction(): - row_count = await query.raw_delete(use_models=model_based_deletion) + row_count = await queryset.raw_delete(use_models=model_based_deletion) if row_count is not None and row_count != len(clauses): related_name = fk_source.name raise RelationshipNotFound( @@ -462,6 +508,7 @@ def _helper_prepare(child: Any) -> Any: instance=self.instance, raw_values=prepared, row_count=row_count, + operation_skipped=False, model_based_deletion=model_based_deletion, **self._shared_relation_signals_params, ) @@ -771,7 +818,13 @@ async def save_related(self) -> None: ) await operation.prepare(refs) operation.signal_params["instance"] = self.instance - await operation.send_pre_signal() + try: + await operation.send_pre_signal() + except SkipOperation: + operation.signal_params["row_count"] = 0 + operation.signal_params["row_count_create"] = 0 + await operation.send_post_signal() + return await operation.apply_db() # no cache update because the queryset is temporary # we can just rename the signals parameters for the post signal @@ -811,7 +864,9 @@ async def add(self, child: BaseModelType) -> BaseModelType | None: Raises: RelationshipIncompatible: If the child type is not compatible. """ - return (await self.add_many(child))[0] + results = await self.add_many(child) + # bail out if no results are returned, in case the result list is modified in signals + return results[0] if results else None async def add_many(self, *children: BaseModelType) -> list[BaseModelType | None]: """ @@ -824,7 +879,8 @@ async def add_many(self, *children: BaseModelType) -> list[BaseModelType | None] model or a dictionary. Returns: - list[BaseModelType | None]: A list of saved child model instances. + list[BaseModelType | None]: A list of saved intermediate model instances, + or None for each record when the operation was skipped. Raises: RelationshipIncompatible: If a child type is not compatible. @@ -854,7 +910,11 @@ async def add_many(self, *children: BaseModelType) -> list[BaseModelType | None] resolve_embed=True, ) await operation.prepare(prepared) - await operation.send_pre_signal() + try: + await operation.send_pre_signal() + except SkipOperation: + await operation.send_post_signal() + return [None for _ in prepared] await operation.apply_db() # no cache update because the queryset is temporary # we can just rename the signals parameters for the post signal @@ -921,7 +981,12 @@ def _helper_prepare(child: Any) -> Any: del operation.signal_params["create_params"] del operation.signal_params["update_params"] operation.signal_params["raw_values"] = raw_values - await operation.send_pre_signal() + try: + await operation.send_pre_signal() + except SkipOperation: + operation.signal_params["row_count"] = 0 + await operation.send_post_signal() + return # allow modification via raw_values obj_ids = [id(obj) for obj in raw_values] operation.update_params = [tup for tup in operation.update_params if id(tup[0]) in obj_ids] @@ -1026,5 +1091,5 @@ async def post_delete_callback(self) -> None: # Determine whether to use model-based deletion from the foreign key's configuration. use_models=self.to.meta.fields[self.to_foreign_key].use_model_based_deletion, # Specify the foreign key that references the deleted instance to ensure correct removal. - remove_referenced_call=self.to_foreign_key, + remove_referenced_call=self.to_foreign_key or True, ) diff --git a/edgy/core/files/storage/base.py b/edgy/core/files/storage/base.py index 7507b748..eadf7766 100644 --- a/edgy/core/files/storage/base.py +++ b/edgy/core/files/storage/base.py @@ -224,7 +224,9 @@ def get_available_name( # Check for path traversal attempts in the directory component. if ".." in pathlib.PurePath(dir_name).parts: - raise SuspiciousFileOperation(f"Detected path traversal attempt in '{dir_name}'") + raise SuspiciousFileOperation( + detail=f"Detected path traversal attempt in '{dir_name}'" + ) # Sanitize the file name component. validate_file_name(file_name) @@ -257,9 +259,11 @@ def get_available_name( if not file_root: # If file_root becomes empty after truncation, raise an error. raise SuspiciousFileOperation( - f'Storage can not find an available filename for "{name}". ' - "Please make sure that the corresponding file field " - 'allows sufficient "max_length".' + detail=( + f'Storage can not find an available filename for "{name}". ' + "Please make sure that the corresponding file field " + 'allows sufficient "max_length".' + ) ) # Regenerate the name with the truncated file_root. if not overwrite: diff --git a/edgy/exceptions.py b/edgy/exceptions.py index e4845fc6..15e19d15 100644 --- a/edgy/exceptions.py +++ b/edgy/exceptions.py @@ -216,7 +216,7 @@ class ModelSchemaError(EdgyException): """ -class SuspiciousFileOperation(Exception): +class SuspiciousFileOperation(EdgyException): """ Exception raised for suspicious file operations, typically for security reasons. @@ -224,6 +224,12 @@ class SuspiciousFileOperation(Exception): """ +class SkipOperation(Exception): + """ + Exception raised in signals to skip operation and return empty values. + """ + + class InvalidStorageError(ImproperlyConfigured): """ Exception raised when an invalid storage backend is configured. diff --git a/tests/foreign_keys/test_many_to_many_field.py b/tests/foreign_keys/test_many_to_many_field.py index d027680a..35b8001a 100644 --- a/tests/foreign_keys/test_many_to_many_field.py +++ b/tests/foreign_keys/test_many_to_many_field.py @@ -7,7 +7,9 @@ pytestmark = pytest.mark.anyio -database = DatabaseTestClient(DATABASE_URL, full_isolation=False) +database = DatabaseTestClient( + DATABASE_URL, force_rollback=False, use_existing=False, full_isolation=False +) models = edgy.Registry(database=database) @@ -47,7 +49,7 @@ class Meta: @pytest.fixture(scope="function") async def create_test_database(): - async with database: + async with models: await models.create_all() yield if not database.drop: @@ -388,7 +390,40 @@ async def test_values_list(create_test_database): assert arr == ["The Bird", "Heart don't stand a chance", "The Waters"] -def test_assertation_error_on_embed_through_double_underscore_attr(): +async def test_relation_load(create_test_database, subtests): + for i in range(10): + with subtests.test(msg=f"Iteration: {i}"): + await Album.query.delete() + await Track.query.delete() + album = await Album.query.create( + name="Malibu", + tracks=[ + Track(title="The Bird", position=1), + Track(title="Heart don't stand a chance", position=2), + Track(title="The Waters", position=3), + ], + ) + assert await Track.query.count() == 3 + assert await Album.query.count() == 1 + arr = await album.tracks.order_by("position").values_list("title", flat=True) + assert arr == ["The Bird", "Heart don't stand a chance", "The Waters"] + tracks = await album.tracks.all() + await album.tracks.remove_many(*tracks) + assert await album.tracks.count() == 0 + assert await Track.query.count() == 3 + await album.tracks.add_many( + Track(title="The Bird", position=1), + Track(title="Heart don't stand a chance", position=2), + Track(title="The Waters", position=3), + ) + assert await Track.query.count() == 6 + arr = await album.tracks.order_by("position").values_list("title", flat=True) + assert arr == ["The Bird", "Heart don't stand a chance", "The Waters"] + await Album.query.delete() + await Track.query.delete() + + +async def test_assertation_error_on_embed_through_double_underscore_attr(): with pytest.raises(FieldDefinitionError) as raised: class MyModel(edgy.StrictModel): diff --git a/tests/foreign_keys/test_ref_foreignkey.py b/tests/foreign_keys/test_ref_foreignkey.py index 12702078..92d001c0 100644 --- a/tests/foreign_keys/test_ref_foreignkey.py +++ b/tests/foreign_keys/test_ref_foreignkey.py @@ -9,8 +9,8 @@ pytestmark = pytest.mark.anyio -database = DatabaseTestClient(DATABASE_URL, full_isolation=False) -models = edgy.Registry(database=database) +database = DatabaseTestClient(DATABASE_URL) +models = edgy.Registry(database=edgy.Database(database, force_rollback=True)) pydantic_version = ".".join(__version__.split(".")[:2]) @@ -64,16 +64,19 @@ class Meta: @pytest.fixture(autouse=True, scope="module") async def create_test_database(): - await models.create_all() - yield - await models.drop_all() - - -@pytest.fixture(autouse=True) -async def rollback_connections(): - with database.force_rollback(): - async with database: - yield + # this creates and drops the database + async with database: + await models.create_all() + yield + if not database.drop: + await models.drop_all() + + +@pytest.fixture(autouse=True, scope="function") +async def rollback_transactions(): + # this rolls back + async with models: + yield async def test_conversion(): diff --git a/tests/models/test_model_class.py b/tests/models/test_model_class.py index a46c54e9..07b97e71 100644 --- a/tests/models/test_model_class.py +++ b/tests/models/test_model_class.py @@ -1,3 +1,5 @@ +from typing import cast + import pytest import edgy @@ -129,3 +131,28 @@ async def test_eq(create_test_database): assert not user.__eq__("") assert user != User assert not user.__eq__(User) + + +@pytest.mark.parametrize("model_based", [True, False]) +async def test_queryset_cache(model_based, create_test_database): + queryset = User.query.all() + cache = queryset._cache + assert not queryset._cache + assert not queryset._cache_fetch_all + await queryset.create(name="Test 1") + assert cache is queryset._cache + assert len(queryset._cache.cache[queryset._cache.create_category(User)]) == 1 + assert not queryset._cache_fetch_all + await queryset.create(name="Test 2") + assert len(queryset._cache.cache[queryset._cache.create_category(User)]) == 2 + assert not queryset._cache_fetch_all + await queryset + assert cast(bool, queryset._cache_fetch_all) + assert len(queryset._cache.cache[queryset._cache.create_category(User)]) == 2 + await queryset.create(name="Test 3") + assert not queryset._cache_fetch_all + assert len(queryset._cache.cache[queryset._cache.create_category(User)]) == 3 + await queryset.create(name="Test 4") + assert len(queryset._cache.cache[queryset._cache.create_category(User)]) == 4 + await queryset.delete(model_based) + assert not queryset._cache diff --git a/tests/models/test_model_first.py b/tests/models/test_model_first.py index 8173b012..a8fc173f 100644 --- a/tests/models/test_model_first.py +++ b/tests/models/test_model_first.py @@ -37,12 +37,21 @@ async def rollback_transactions(): async def test_model_first(): - Test = await User.query.create(name="Test") + tester = await User.query.create(name="Test") jane = await User.query.create(name="Jane") query = User.query.all() assert query._cache_first is None - assert await query.first() == Test + assert await query.first() == tester assert query._cache_first is not None assert await User.query.filter(name="Jane").first() == jane assert await User.query.filter(name="Lucy").first() is None + + +async def test_model_first_none(): + query = User.query.all() + assert query._cache_first is None + assert query._cache_count is None + assert await query.first() is None + assert query._cache_first is None + assert query._cache_count == 0 diff --git a/tests/models/test_model_last.py b/tests/models/test_model_last.py index 9f6aebb6..f1c0c004 100644 --- a/tests/models/test_model_last.py +++ b/tests/models/test_model_last.py @@ -47,3 +47,12 @@ async def test_model_last(): assert await User.query.filter(name="Jane").last() == jane assert await User.query.filter(name="Test").last() == Test assert await User.query.filter(name="Lucy").last() is None + + +async def test_model_last_none(): + query = User.query.all() + assert query._cache_last is None + assert query._cache_count is None + assert await query.last() is None + assert query._cache_last is None + assert query._cache_count == 0 diff --git a/tests/models/test_model_queryset_delete.py b/tests/models/test_model_queryset_delete.py index ee98a737..86b4d6a6 100644 --- a/tests/models/test_model_queryset_delete.py +++ b/tests/models/test_model_queryset_delete.py @@ -48,3 +48,24 @@ async def test_queryset_delete(): await Product.query.delete() assert await Product.query.count() == 0 + + +async def test_queryset_delete_cache(): + queryset = Product.query.all() + await queryset.create(name="Belt", rating=5) + await queryset.create(name="Tie", rating=5) + assert queryset._cache + assert await queryset.delete() + assert not queryset._cache + + +async def test_queryset_no_delete_cache(): + await Product.query.create(name="Belt", rating=5) + await Product.query.create(name="Tie", rating=5) + queryset = Product.query.filter(name="test") + await queryset + assert queryset._cache_count == 0 + assert queryset._cache_fetch_all + await queryset.delete() + assert queryset._cache_count == 0 + assert queryset._cache_fetch_all diff --git a/tests/signals/test_deletion_signals.py b/tests/signals/test_deletion_signals.py index 44e1be27..ae5116e0 100644 --- a/tests/signals/test_deletion_signals.py +++ b/tests/signals/test_deletion_signals.py @@ -1,10 +1,10 @@ +from typing import cast + import pytest import edgy -from edgy.core.signals import ( - post_delete, - pre_delete, -) +from edgy.core.signals import Signal, post_delete, pre_delete +from edgy.exceptions import SkipOperation from edgy.testclient import DatabaseTestClient from tests.settings import DATABASE_URL @@ -29,6 +29,7 @@ class User(edgy.StrictModel): class Meta: registry = models + signals = {"pre_delete": Signal()} class Profile(edgy.StrictModel): @@ -37,14 +38,13 @@ class Profile(edgy.StrictModel): class Meta: registry = models + signals = {"pre_delete": Signal()} class Log(edgy.StrictModel): signal = edgy.CharField(max_length=255) - is_queryset: bool = edgy.BooleanField() - model_instance_id = edgy.BigIntegerField(null=True) - row_count = edgy.BigIntegerField(null=True) - class_name: str = edgy.CharField(max_length=255) + class_name = edgy.CharField(max_length=255) + params = edgy.JSONField() class Meta: registry = models @@ -65,41 +65,39 @@ async def create_test_database(): @pytest.fixture(autouse=True, scope="function") async def connect_signals(): - @pre_delete.connect_via(Profile, weak=True) - @pre_delete.connect_via(User, weak=True) - async def pre_deleting(sender, instance, model_instance, **kwargs): + @Profile.meta.signals.pre_delete.connect_via(Profile, weak=True) + @User.meta.signals.pre_delete.connect_via(User, weak=True) + async def pre_deleting(sender, **kwargs): await Log.query.create( signal="pre_delete", - is_queryset=model_instance is None, - model_instance_id=None if model_instance is None else model_instance.id, - class_name=instance.model_class.__name__ - if model_instance is None - else type(model_instance).__name__, + class_name=sender.__name__, + params={k: str(v) for k, v in kwargs.items()}, ) - @post_delete.connect_via(Profile, weak=True) - @post_delete.connect_via(User, weak=True) - async def post_deleting(sender, instance, model_instance, row_count, **kwargs): + @Profile.meta.signals.post_delete.connect_via(Profile, weak=True) + @User.meta.signals.post_delete.connect_via(User, weak=True) + async def post_deleting(sender, **kwargs): await Log.query.create( signal="post_delete", - is_queryset=model_instance is None, - model_instance_id=None if model_instance is None else model_instance.id, - class_name=instance.model_class.__name__ - if model_instance is None - else type(model_instance).__name__, - row_count=row_count, + class_name=sender.__name__, + params={k: str(v) for k, v in kwargs.items()}, ) try: yield finally: - pre_delete.disconnect(pre_deleting) - post_delete.disconnect(post_deleting) + Profile.meta.signals.pre_delete.disconnect(pre_deleting) + User.meta.signals.pre_delete.disconnect(pre_deleting) + Profile.meta.signals.post_delete.disconnect(post_deleting) + User.meta.signals.post_delete.disconnect(post_deleting) @pytest.mark.parametrize("klass", [User, Profile]) async def test_correct_connection(klass): - assert pre_delete.has_receivers_for(klass) + assert klass.meta.signals.pre_delete is not pre_delete + assert klass.meta.signals.post_delete is post_delete + assert not pre_delete.has_receivers_for(klass) + assert klass.meta.signals.pre_delete.has_receivers_for(klass) assert post_delete.has_receivers_for(klass) @@ -113,23 +111,57 @@ async def test_deletion_called_once_model(klass): assert len(logs) == 2 assert logs[0].signal == "pre_delete" assert logs[0].class_name == klass.__name__ + assert logs[0].params["instance"].startswith(f"{klass.__name__}") + assert logs[0].params["model_instance"].startswith(f"{klass.__name__}") + assert "row_count" not in logs[0].params assert logs[1].signal == "post_delete" assert logs[1].class_name == klass.__name__ + assert logs[1].params["instance"].startswith(f"{klass.__name__}") + assert logs[1].params["model_instance"].startswith(f"{klass.__name__}") + assert logs[1].params["row_count"] == "1" -async def test_deletion_called_once_query(): - await User.query.create(name="Edgy") +@pytest.mark.parametrize("klass", [User, Profile]) +async def test_deletion_called_once_query(klass): + await klass.query.create(name="Edgy") logs = await Log.query.all() assert len(logs) == 0 - await User.query.delete() + await klass.query.delete() logs = await Log.query.all() - assert len(logs) == 2 - assert logs[0].signal == "pre_delete" - assert logs[0].class_name == "User" - assert logs[0].is_queryset - assert logs[1].signal == "post_delete" - assert logs[1].class_name == "User" - assert logs[1].is_queryset + if klass.__deletion_with_signals__: + assert len(logs) == 4 + assert logs[0].signal == "pre_delete" + assert logs[0].class_name == klass.__name__ + assert logs[0].params["instance"].startswith(f"QuerySet") + assert logs[0].params["model_instance"] == "None" + assert "row_count" not in logs[0].params + assert logs[1].signal == "pre_delete" + assert logs[2].params["instance"].startswith(f"QuerySet") + assert logs[1].params["model_instance"].startswith(f"{klass.__name__}") + assert "row_count" not in logs[0].params + assert logs[2].signal == "post_delete" + assert logs[2].class_name == klass.__name__ + assert logs[2].params["instance"].startswith(f"QuerySet") + assert logs[2].params["model_instance"].startswith(f"{klass.__name__}") + assert logs[2].params["row_count"] == "1" + assert logs[3].signal == "post_delete" + assert logs[3].class_name == klass.__name__ + assert logs[3].params["instance"].startswith(f"QuerySet") + assert logs[3].params["model_instance"] == "None" + assert logs[3].params["row_count"] == "1" + + else: + assert len(logs) == 2 + assert logs[0].signal == "pre_delete" + assert logs[0].class_name == klass.__name__ + assert logs[0].params["instance"].startswith(f"QuerySet") + assert logs[0].params["model_instance"] == "None" + assert "row_count" not in logs[0].params + assert logs[1].signal == "post_delete" + assert logs[1].class_name == klass.__name__ + assert logs[1].params["instance"].startswith(f"QuerySet") + assert logs[1].params["model_instance"] == "None" + assert logs[1].params["row_count"] == "1" async def test_deletion_called_referenced(): @@ -142,9 +174,13 @@ async def test_deletion_called_referenced(): logs = await Log.query.all() assert len(logs) == 4 assert logs[0].signal == "pre_delete" + assert logs[0].class_name == "User" assert logs[1].signal == "pre_delete" + assert logs[1].class_name == "Profile" assert logs[2].signal == "post_delete" + assert logs[2].class_name == "Profile" assert logs[3].signal == "post_delete" + assert logs[3].class_name == "User" async def test_deletion_called_cascade(): @@ -187,3 +223,34 @@ async def test_deletion_called_cascade_with_signals(): assert logs[4].class_name == "User" assert logs[5].signal == "post_delete" assert logs[5].class_name == "Profile" + + +async def test_model_base_deletion_with_skip(): + queryset = Profile.query.batch_size(2) + + @Profile.meta.signals.pre_delete.connect_via(Profile, weak=True) + async def pre_deleting(sender, instance, model_instance, **kwargs): + assert instance is queryset + if model_instance and model_instance.name.startswith("raise"): + raise SkipOperation() + + try: + await queryset.create(name="raise Edgy") + await queryset.create(name="noraise Edgy") + await queryset.create(name="raise Saffier") + await queryset.create(name="noraise Saffier") + assert bool(queryset._cache) + await queryset + # now everything is fetched + assert queryset._cache_fetch_all + await queryset.delete() + finally: + Profile.meta.signals.pre_delete.disconnect(pre_deleting) + assert not queryset._cache.cache + assert queryset._cache_count is None + assert queryset._cache_first is None + assert queryset._cache_last is None + # prevent confused type checkers + assert not cast(bool, queryset._cache_fetch_all) + assert await queryset.count() == 2 + assert await Profile.query.count() == 2 diff --git a/tests/signals/test_deletion_signals_skip.py b/tests/signals/test_deletion_signals_skip.py new file mode 100644 index 00000000..2d815c19 --- /dev/null +++ b/tests/signals/test_deletion_signals_skip.py @@ -0,0 +1,275 @@ +import pytest + +import edgy +from edgy.core.signals import Signal, post_delete, pre_delete +from edgy.exceptions import SkipOperation +from edgy.testclient import DatabaseTestClient +from tests.settings import DATABASE_URL + +pytestmark = pytest.mark.anyio + +database = DatabaseTestClient( + DATABASE_URL, drop_database=True, force_rollback=False, full_isolation=False +) +models = edgy.Registry(database=database) + + +class BaseModelWithDeletionHandling(edgy.StrictModel): + protection = edgy.BooleanField(default=True) + + class Meta: + registry = models + abstract = True + + +class Unrelated(BaseModelWithDeletionHandling): + name = edgy.CharField(max_length=100) + + +class User(BaseModelWithDeletionHandling): + name = edgy.CharField(max_length=100) + profile = edgy.ForeignKey( + "Profile", + null=True, + on_delete=edgy.CASCADE, + no_constraint=True, + remove_referenced=True, + use_model_based_deletion=True, + ) + + class Meta: + signals = {"pre_delete": Signal()} + + +class Profile(BaseModelWithDeletionHandling): + name = edgy.CharField(max_length=100) + __deletion_with_signals__ = True + + class Meta: + signals = {"pre_delete": Signal()} + + +class Log(edgy.StrictModel): + signal = edgy.CharField(max_length=255) + class_name = edgy.CharField(max_length=255) + params = edgy.JSONField() + + class Meta: + registry = models + + def __str__(self) -> str: + return str(self.extract_db_fields()) + + def __repr__(self) -> str: + return f"Log<{self}>" + + +@pytest.fixture(autouse=True, scope="function") +async def create_test_database(): + async with models: + await models.create_all() + yield + + +@pytest.fixture(autouse=True, scope="function") +async def connect_signals(): + @Unrelated.meta.signals.pre_delete.connect_via(Unrelated, weak=True) + @Profile.meta.signals.pre_delete.connect_via(Profile, weak=True) + @User.meta.signals.pre_delete.connect_via(User, weak=True) + async def pre_deleting(sender, model_instance, injected_filters=None, **kwargs): + if model_instance is not None: + if model_instance.protection: + raise SkipOperation() + elif injected_filters is not None: + # only those without protection are retrieved + injected_filters.append({"protection": False}) + + @Unrelated.meta.signals.post_delete.connect_via(Unrelated, weak=True) + @Profile.meta.signals.post_delete.connect_via(Profile, weak=True) + @User.meta.signals.post_delete.connect_via(User, weak=True) + async def post_deleting(sender, **kwargs): + await Log.query.create( + signal="post_delete", + class_name=sender.__name__, + params={k: str(v) for k, v in kwargs.items()}, + ) + + try: + yield + finally: + Unrelated.meta.signals.pre_delete.disconnect(pre_deleting) + Profile.meta.signals.pre_delete.disconnect(pre_deleting) + User.meta.signals.pre_delete.disconnect(pre_deleting) + Unrelated.meta.signals.post_delete.disconnect(post_deleting) + Profile.meta.signals.post_delete.disconnect(post_deleting) + User.meta.signals.post_delete.disconnect(post_deleting) + + +@pytest.mark.parametrize("klass", [User, Profile]) +async def test_correct_connection(klass): + assert klass.meta.signals.pre_delete is not pre_delete + assert klass.meta.signals.post_delete is post_delete + assert not pre_delete.has_receivers_for(klass) + assert klass.meta.signals.pre_delete.has_receivers_for(klass) + assert post_delete.has_receivers_for(klass) + + +@pytest.mark.parametrize("klass", [User, Profile, Unrelated]) +async def test_deletion_called_once_model(klass): + obj = await klass.query.create(name="Edgy") + assert not obj._db_deleted + logs = await Log.query.all() + assert len(logs) == 0 + await obj.delete() + assert not obj._db_deleted + assert await klass.query.count() == 1 + logs = await Log.query.all() + assert len(logs) == 1 + assert logs[0].signal == "post_delete" + assert logs[0].class_name == klass.__name__ + assert logs[0].params["model_instance"] != "None" + assert logs[0].params["row_count"] == "0" + + +@pytest.mark.parametrize("klass", [User, Profile, Unrelated]) +@pytest.mark.parametrize("model_based", [True, False]) +async def test_deletion_called_once_query(klass, model_based): + await klass.query.create(name="Edgy") + logs = await Log.query.all() + assert len(logs) == 0 + await klass.query.delete(model_based) + assert await klass.query.count() == 1 + logs = await Log.query.all() + assert len(logs) == 1 + assert logs[0].signal == "post_delete" + assert logs[0].class_name == klass.__name__ + assert logs[0].params["instance"].startswith(f"QuerySet") + assert logs[0].params["model_instance"] == "None" + assert logs[0].params["row_count"] == "0" + + +@pytest.mark.parametrize("klass", [User, Profile]) +async def test_deletion_called_once_query_model_based(klass): + await klass.query.create(name="Edgy") + logs = await Log.query.all() + assert len(logs) == 0 + await klass.query.delete() + logs = await Log.query.all() + assert await klass.query.count() == 1 + assert len(logs) == 1 + assert logs[0].signal == "post_delete" + assert logs[0].class_name == klass.__name__ + assert logs[0].params["instance"].startswith(f"QuerySet") + assert logs[0].params["model_instance"] == "None" + assert logs[0].params["row_count"] == "0" + + +async def test_deletion_called_referenced(): + user = await User.query.create(name="Edgy", profile=Profile(name="Edgy")) + logs = await Log.query.all() + assert len(logs) == 0 + await user.delete() + logs = await Log.query.all() + assert len(logs) == 1 + assert logs[0].class_name == "User" + assert logs[0].signal == "post_delete" + assert logs[0].params["row_count"] == "0" + assert logs[0].params["operation_skipped"] == "True" + # now it really deletes + with User.meta.signals.pre_delete.muted(): + await user.delete() + logs = await Log.query.offset(1) + assert len(logs) == 2 + assert logs[0].class_name == "Profile" + assert logs[0].signal == "post_delete" + assert logs[0].params["row_count"] == "0" + assert logs[0].params["operation_skipped"] == "True" + assert logs[0].params["model_instance"] != "None" + + assert logs[1].class_name == "User" + assert logs[1].signal == "post_delete" + assert logs[1].params["row_count"] == "1" + assert logs[1].params["model_instance"] != "None" + assert logs[1].params["operation_skipped"] == "False" + + +async def test_deletion_called_referenced_query(): + await User.query.create(name="Edgy", profile=Profile(name="Edgy")) + logs = await Log.query.all() + assert len(logs) == 0 + await User.query.delete() + logs = await Log.query.all() + assert len(logs) == 1 + assert logs[0].class_name == "User" + assert logs[0].signal == "post_delete" + assert logs[0].params["row_count"] == "0" + assert logs[0].params["operation_skipped"] == "False" + assert await User.query.count() == 1 + # now it really deletes + with User.meta.signals.pre_delete.muted(): + await User.query.delete() + assert await User.query.count() == 0 + logs = await Log.query.offset(1) + assert len(logs) == 2 + + assert logs[0].class_name == "Profile" + assert logs[0].signal == "post_delete" + assert logs[0].params["row_count"] == "0" + assert logs[0].params["operation_skipped"] == "True" + assert logs[0].params["model_instance"] != "None" + + assert logs[1].class_name == "User" + assert logs[1].signal == "post_delete" + assert logs[1].params["row_count"] == "1" + assert logs[1].params["model_instance"] == "None" + assert logs[1].params["operation_skipped"] == "False" + + +async def test_deletion_called_cascade(): + profile = await Profile.query.create(name="Edgy") + await User.query.create(name="Edgy", profile=profile) + await User.query.create(name="Edgy2", profile=profile) + logs = await Log.query.all() + assert len(logs) == 0 + await profile.delete() + logs = await Log.query.all() + assert len(logs) == 1 + assert logs[0].class_name == "Profile" + assert logs[0].signal == "post_delete" + assert logs[0].params["row_count"] == "0" + assert logs[0].params["operation_skipped"] == "True" + # now it really deletes + with Profile.meta.signals.pre_delete.muted(): + await profile.delete() + logs = await Log.query.offset(1) + assert len(logs) == 1 + assert logs[0].class_name == "Profile" + assert logs[0].signal == "post_delete" + assert logs[0].params["row_count"] == "1" + assert logs[0].params["operation_skipped"] == "False" + assert logs[0].params["model_instance"] != "None" + + +async def test_deletion_called_cascade_query(): + profile = await Profile.query.create(name="Edgy") + await User.query.create(name="Edgy", profile=profile) + await User.query.create(name="Edgy2", profile=profile) + logs = await Log.query.all() + assert len(logs) == 0 + await Profile.query.delete() + logs = await Log.query.all() + assert len(logs) == 1 + assert logs[0].class_name == "Profile" + assert logs[0].signal == "post_delete" + assert logs[0].params["row_count"] == "0" + # now it really deletes + with Profile.meta.signals.pre_delete.muted(): + await Profile.query.delete() + logs = await Log.query.offset(1) + assert len(logs) == 2 + assert logs[0].class_name == "Profile" + assert logs[0].signal == "post_delete" + assert logs[0].params["row_count"] == "1" + assert logs[1].class_name == "Profile" + assert logs[1].signal == "post_delete" + assert logs[1].params["row_count"] == "1" diff --git a/tests/signals/test_relation_signals.py b/tests/signals/test_relation_signals.py index 65bec988..a3dcf530 100644 --- a/tests/signals/test_relation_signals.py +++ b/tests/signals/test_relation_signals.py @@ -18,9 +18,7 @@ pytestmark = pytest.mark.anyio -database = DatabaseTestClient( - DATABASE_URL, drop_database=True, force_rollback=False, full_isolation=False -) +database = DatabaseTestClient(DATABASE_URL, force_rollback=False, use_existing=False) models = edgy.Registry(database=database) @@ -67,6 +65,8 @@ async def create_test_database(): async with models: await models.create_all() yield + if not database.drop: + await models.drop_all() async def pre_add(sender, **kwargs): @@ -179,8 +179,7 @@ async def test_basic_m2m(): assert logs[0].signal == "pre_relation_add" # see docs for parameters assert len(logs[0].params) == 9 - assert logs[1].signal == "post_relation_add" - assert len(logs[1].params) == 12 + assert len(logs[1].params) == 13 assert logs[0].params["raw_values"] assert "User" in logs[0].params["source"] assert "Friend" in logs[0].params["target"] @@ -194,8 +193,9 @@ async def test_basic_m2m(): assert "values" not in logs[0].params assert "row_count" not in logs[0].params assert "row_count_create" not in logs[0].params - assert logs[1].params["instance"] == str(user) assert logs[1].signal == "post_relation_add" + assert logs[1].params["operation_skipped"] == "False" + assert logs[1].params["instance"] == str(user) assert logs[1].params["values"] assert logs[1].params["row_count"] == "1" assert logs[1].params["row_count_create"] == "1" @@ -215,7 +215,7 @@ async def test_basic_m2m(): assert "update_params" not in logs[2].params assert logs[3].signal == "post_relation_remove" # see docs for parameters - assert len(logs[3].params) == 8 + assert len(logs[3].params) == 9 assert logs[3].params["instance"] == str(user) assert "operation" not in logs[3].params assert "values" not in logs[3].params @@ -223,6 +223,7 @@ async def test_basic_m2m(): assert "create_params" not in logs[3].params assert "update_params" not in logs[3].params assert logs[3].params["row_count"] == "1" + assert logs[3].params["operation_skipped"] == "False" assert await Log.query.filter(signal="post_delete").count() == 0 assert await Log.query.filter(signal="pre_delete").count() == 0 @@ -264,7 +265,7 @@ async def test_basic_one_to_many(): assert "model_based_deletion" not in logs[4].params assert logs[5].signal == "post_relation_remove" # see docs for parameters - assert len(logs[5].params) == 7 + assert len(logs[5].params) == 8 assert "operation" not in logs[5].params assert "values" not in logs[5].params assert "create_params" not in logs[5].params @@ -332,15 +333,33 @@ async def test_nullify_many_to_many(): assert len(await user.friends.all()) == 2 Friend.meta.signals.pre_relation_remove.connect(nullify_removal, through) - # shared - assert User.meta.signals.pre_relation_remove.has_receivers_for(through) try: + # shared + assert User.meta.signals.pre_relation_remove.has_receivers_for(through) await user.friends.remove_many(*(await user.friends.all())) finally: Friend.meta.signals.pre_relation_remove.disconnect(nullify_removal) assert len(await user.friends.all()) == 2 +async def test_nullify_many_to_many_load(subtests): + through = User.meta.fields["friends"].through + user = await User.query.create(name="Edgy") + assert len(await user.friends.all()) == 0 + + Friend.meta.signals.pre_relation_remove.connect(nullify_removal, through) + try: + for i in range(10): + with subtests.test(msg=f"iteration: {i}"): + await user.friends.add_many(Friend(name=f"saffier_{i}"), {"name": "saffier2"}) + assert len(await user.friends.all()) == 2 + await user.friends.remove_many(*(await user.friends.all())) + assert len(await user.friends.all()) == 2 + await Friend.query.delete() + finally: + Friend.meta.signals.pre_relation_remove.disconnect(nullify_removal) + + async def test_nullify_one_to_many(): profile = await Profile.query.create(name="Edgy") await profile.users.create(name="saffier") diff --git a/tests/signals/test_signals.py b/tests/signals/test_signals.py index 73ac37db..5a6912a4 100644 --- a/tests/signals/test_signals.py +++ b/tests/signals/test_signals.py @@ -1,6 +1,7 @@ import pytest import edgy +from edgy.core import signals from edgy.core.signals import ( Broadcaster, post_delete, @@ -59,7 +60,7 @@ async def test_invalid_signal(): broadcaster.save = 1 -async def test_signals(): +async def test_signals_simple(): try: @pre_save.connect_via(User) @@ -67,38 +68,32 @@ async def pre_saving(sender, instance, model_instance, **kwargs): await Log.query.create( signal="pre_save", instance=model_instance.model_dump(), params=kwargs ) - print(f"pre_save signal broadcasted for {model_instance.get_instance_name()}") @post_save.connect_via(User) async def post_saving(sender, instance, model_instance, **kwargs): await Log.query.create( signal="post_save", instance=model_instance.model_dump(), params=kwargs ) - print(f"post_save signal broadcasted for {model_instance.get_instance_name()}") @pre_update.connect_via(User) async def pre_updating(sender, instance, model_instance, **kwargs): await Log.query.create( signal="pre_update", instance=model_instance.model_dump(), params=kwargs ) - print(f"pre_update signal broadcasted for {model_instance.get_instance_name()}") @post_update.connect_via(User) async def post_updating(sender, instance, model_instance, **kwargs): await Log.query.create( signal="post_update", instance=model_instance.model_dump(), params=kwargs ) - print(f"post_update signal broadcasted for {model_instance.get_instance_name()}") @pre_delete.connect_via(User) async def pre_deleting(sender, instance, model_instance, **kwargs): await Log.query.create(signal="pre_delete", instance=model_instance.model_dump()) - print(f"pre_delete signal broadcasted for {model_instance.get_instance_name()}") @post_delete.connect_via(User) async def post_deleting(sender, instance, model_instance, **kwargs): await Log.query.create(signal="post_delete", instance=model_instance.model_dump()) - print(f"post_delete signal broadcasted for {model_instance.get_instance_name()}") # Signals for the create user = await User.query.create(name="Edgy") @@ -150,6 +145,70 @@ async def post_deleting(sender, instance, model_instance, **kwargs): assert len(users) == 1 +async def test_signals_advanced(): + cleanup_array = [] + logs = await Log.query.all() + assert len(logs) == 0 + try: + for signal_name in dir(signals): + if not signal_name.startswith("post_") or signal_name == "post_migrate": + continue + signal: signals.Signal = getattr(signals, signal_name) + + async def log(sender, model_instance, _signal_name=signal_name, **kwargs): + await Log.query.create( + signal=_signal_name, + instance=model_instance.model_dump(), + params={k: str(v) for k, v in kwargs.items()}, + ) + + for model in models.models.values(): + if model is not Log: + signal.connect(log, model, weak=False) + cleanup_array.append(lambda _signal=signal, _log=log: _signal.disconnect(_log)) + + assert signal.has_receivers_for(User) + # Signals for the create + user = await User.query.create(name="Edgy") + logs = await Log.query.all() + + assert len(logs) == 1 + assert logs[0].signal == "post_save" + assert logs[0].instance["name"] == user.name + + user = await User.query.create(name="Saffier") + logs = await Log.query.offset(1) + + assert len(logs) == 1 + assert logs[0].signal == "post_save" + assert logs[0].instance["name"] == user.name + + # For the updates + user = await user.update(name="Another Saffier") + logs = await Log.query.filter(signal__icontains="update").all() + + assert len(logs) == 1 + assert logs[0].signal == "post_update" + assert logs[0].instance["name"] == "Another Saffier" + + # Delete + await user.delete() + logs = await Log.query.filter(signal__icontains="delete").all() + assert len(logs) == 1 + assert logs[0].signal == "post_delete" + assert logs[0].instance["name"] == "Another Saffier" + finally: + for cleanup_ob in cleanup_array: + cleanup_ob() + + users = await User.query.all() + assert len(users) == 1 + await Log.query.delete() + await User.query.create(name="Another Edgy") + logs = await Log.query.all() + assert len(logs) == 0 + + async def test_staticmethod_signals(): class Static: @staticmethod @@ -182,10 +241,10 @@ async def processing(sender, instance, **kwargs): await instance.save() User.meta.signals.custom.connect(receiver=processing) + try: + user = await User.query.create(name="Edgy") + await User.meta.signals.custom.send_async(User, instance=user) - user = await User.query.create(name="Edgy") - await User.meta.signals.custom.send_async(User, instance=user) - - assert user.name == "Edgy ORM" - - User.meta.signals.custom.disconnect(processing) + assert user.name == "Edgy ORM" + finally: + User.meta.signals.custom.disconnect(processing)