mirror of
https://github.com/inventree/InvenTree.git
synced 2026-08-10 15:36:17 +00:00
Row locking for data import / update (#12515)
- Concurrency protection for data imports
This commit is contained in:
@@ -9,7 +9,7 @@ from django.contrib.auth.models import User
|
||||
from django.core.exceptions import FieldDoesNotExist
|
||||
from django.core.exceptions import ValidationError as DjangoValidationError
|
||||
from django.core.validators import FileExtensionValidator
|
||||
from django.db import models
|
||||
from django.db import models, transaction
|
||||
from django.urls import reverse
|
||||
from django.utils.translation import gettext_lazy as _
|
||||
|
||||
@@ -951,6 +951,64 @@ class DataImportRow(models.Model):
|
||||
context={'request': request},
|
||||
)
|
||||
|
||||
def _get_update_instance(self, instance_id, queryset):
|
||||
"""Fetch the target instance for an update, recording any error against this row.
|
||||
|
||||
Arguments:
|
||||
instance_id: The primary key of the instance to fetch
|
||||
queryset: The queryset to fetch the instance from (may be locked)
|
||||
|
||||
Returns:
|
||||
The fetched instance, or None if it could not be found/resolved
|
||||
"""
|
||||
try:
|
||||
return queryset.get(pk=instance_id)
|
||||
except self.session.model_class.DoesNotExist:
|
||||
self.errors = {
|
||||
'non_field_errors': _('No record found with the provided ID')
|
||||
+ f': {instance_id}'
|
||||
}
|
||||
except ValueError:
|
||||
self.errors = {
|
||||
'non_field_errors': _('Invalid ID format provided') + f': {instance_id}'
|
||||
}
|
||||
except Exception as e:
|
||||
self.errors = {'non_field_errors': str(e)}
|
||||
|
||||
return None
|
||||
|
||||
def _process_serializer(self, serializer, commit) -> bool:
|
||||
"""Run is_valid() against the provided serializer, and save it if requested."""
|
||||
if not serializer:
|
||||
self.errors = {
|
||||
'non_field_errors': 'No serializer class linked to this import session'
|
||||
}
|
||||
return False
|
||||
|
||||
result = False
|
||||
|
||||
try:
|
||||
result = serializer.is_valid(raise_exception=True)
|
||||
except (DjangoValidationError, DRFValidationError) as e:
|
||||
self.errors = e.detail
|
||||
|
||||
if result:
|
||||
self.errors = None
|
||||
|
||||
if commit:
|
||||
try:
|
||||
serializer.save()
|
||||
self.complete = True
|
||||
|
||||
except ValueError as e: # Exception as e:
|
||||
self.errors = {'non_field_errors': str(e)}
|
||||
result = False
|
||||
|
||||
self.save()
|
||||
self.session.check_complete()
|
||||
|
||||
return result
|
||||
|
||||
def validate(self, commit=False, request=None) -> bool:
|
||||
"""Validate the data in this row against the linked serializer.
|
||||
|
||||
@@ -981,22 +1039,33 @@ class DataImportRow(models.Model):
|
||||
_('ID is required for updating existing records.')
|
||||
)
|
||||
|
||||
try:
|
||||
instance = self.session.model_class.objects.get(pk=instance_id)
|
||||
except self.session.model_class.DoesNotExist:
|
||||
self.errors = {
|
||||
'non_field_errors': _('No record found with the provided ID')
|
||||
+ f': {instance_id}'
|
||||
}
|
||||
return False
|
||||
except ValueError:
|
||||
self.errors = {
|
||||
'non_field_errors': _('Invalid ID format provided')
|
||||
+ f': {instance_id}'
|
||||
}
|
||||
return False
|
||||
except Exception as e:
|
||||
self.errors = {'non_field_errors': str(e)}
|
||||
queryset = self.session.model_class.objects
|
||||
|
||||
if commit:
|
||||
# Lock the target row for the duration of the transaction. Without
|
||||
# this, another process could modify the record between this read
|
||||
# and the serializer.save() call in _process_serializer() - and as
|
||||
# save() writes the *entire* instance (not just the fields present
|
||||
# in the imported data), that concurrent change would be silently
|
||||
# lost (a "lost update" race). The lock is held until the save
|
||||
# completes below, so a concurrent writer blocks until this commit
|
||||
# finishes rather than racing against it.
|
||||
with transaction.atomic():
|
||||
instance = self._get_update_instance(
|
||||
instance_id, queryset.select_for_update()
|
||||
)
|
||||
|
||||
if instance is None:
|
||||
return False
|
||||
|
||||
serializer = self.construct_serializer(
|
||||
instance=instance, request=request
|
||||
)
|
||||
return self._process_serializer(serializer, commit)
|
||||
|
||||
instance = self._get_update_instance(instance_id, queryset)
|
||||
|
||||
if instance is None:
|
||||
return False
|
||||
|
||||
serializer = self.construct_serializer(instance=instance, request=request)
|
||||
@@ -1004,32 +1073,4 @@ class DataImportRow(models.Model):
|
||||
else:
|
||||
serializer = self.construct_serializer(request=request)
|
||||
|
||||
if not serializer:
|
||||
self.errors = {
|
||||
'non_field_errors': 'No serializer class linked to this import session'
|
||||
}
|
||||
return False
|
||||
|
||||
result = False
|
||||
|
||||
try:
|
||||
result = serializer.is_valid(raise_exception=True)
|
||||
except (DjangoValidationError, DRFValidationError) as e:
|
||||
self.errors = e.detail
|
||||
|
||||
if result:
|
||||
self.errors = None
|
||||
|
||||
if commit:
|
||||
try:
|
||||
serializer.save()
|
||||
self.complete = True
|
||||
|
||||
except ValueError as e: # Exception as e:
|
||||
self.errors = {'non_field_errors': str(e)}
|
||||
result = False
|
||||
|
||||
self.save()
|
||||
self.session.check_complete()
|
||||
|
||||
return result
|
||||
return self._process_serializer(serializer, commit)
|
||||
|
||||
@@ -1,9 +1,13 @@
|
||||
"""Unit tests for the 'importer' app."""
|
||||
|
||||
import os
|
||||
import threading
|
||||
from unittest import mock
|
||||
|
||||
from django.contrib.auth.models import User
|
||||
from django.core.files.base import ContentFile
|
||||
from django.db import connection
|
||||
from django.test import TransactionTestCase, skipUnlessDBFeature
|
||||
from django.urls import reverse
|
||||
|
||||
from importer.models import DataImportColumnMap, DataImportRow, DataImportSession
|
||||
@@ -323,6 +327,103 @@ class ImporterTest(ImporterMixin, InvenTreeTestCase):
|
||||
quantity_mapping.save()
|
||||
|
||||
|
||||
@skipUnlessDBFeature('has_select_for_update')
|
||||
class DataImportRowConcurrencyTest(ImporterMixin, TransactionTestCase):
|
||||
"""Genuine cross-transaction regression test for row locking during a data-import update.
|
||||
|
||||
Uses two real threads (each with its own database connection) to reproduce the
|
||||
lost-update race: two concurrent import rows updating *different* fields on the
|
||||
same StockItem could both read the pre-update instance before either had
|
||||
committed its save() - and since save() writes the *entire* instance, whichever
|
||||
thread committed second would silently discard the first thread's change.
|
||||
|
||||
DataImportRow.validate() now locks the target row (select_for_update, inside
|
||||
transaction.atomic()) for the duration of the read-modify-write cycle, so the
|
||||
second thread's read is forced to wait until the first thread's transaction
|
||||
commits - and therefore sees (and preserves) the first thread's change.
|
||||
"""
|
||||
|
||||
def setUp(self):
|
||||
"""Create a single StockItem to update concurrently."""
|
||||
super().setUp()
|
||||
|
||||
from part.models import Part, PartCategory
|
||||
from stock.models import StockItem
|
||||
|
||||
category = PartCategory.objects.create(
|
||||
name='Concurrency Category', description='Test category'
|
||||
)
|
||||
self.part = Part.objects.create(
|
||||
category=category, name='Concurrency Part', description='Test part'
|
||||
)
|
||||
self.item = StockItem.objects.create(part=self.part, quantity=10)
|
||||
|
||||
def helper_session(self) -> DataImportSession:
|
||||
"""Construct a minimal DataImportSession configured for stock item updates."""
|
||||
return DataImportSession.objects.create(
|
||||
data_file=self.helper_file('companies.csv'),
|
||||
model_type='stockitem',
|
||||
update_records=True,
|
||||
)
|
||||
|
||||
def test_concurrent_update_does_not_lose_writes(self):
|
||||
"""Two concurrent updates to different fields must not clobber one another."""
|
||||
start_barrier = threading.Barrier(2, timeout=5)
|
||||
errors = []
|
||||
|
||||
# Wrap DataImportRow._get_update_instance() so both threads reach the
|
||||
# (real, database-level) row lock at the same time - one wins the lock
|
||||
# and proceeds through to commit, the other blocks until the winner's
|
||||
# transaction completes.
|
||||
original_get_update_instance = DataImportRow._get_update_instance
|
||||
|
||||
def synced_get_update_instance(row, instance_id, queryset):
|
||||
start_barrier.wait(timeout=5)
|
||||
return original_get_update_instance(row, instance_id, queryset)
|
||||
|
||||
def update(field, value):
|
||||
try:
|
||||
row = DataImportRow(
|
||||
session=self.helper_session(),
|
||||
# 'part' is a required field on StockItemSerializer, so it must
|
||||
# be included even though this row is only actually changing
|
||||
# 'field' - the fix under test is *locking*, not partial-update
|
||||
# support, so we work within the serializer's existing rules
|
||||
data={'id': self.item.pk, 'part': self.part.pk, field: value},
|
||||
)
|
||||
if not row.validate(commit=True):
|
||||
errors.append(row.errors)
|
||||
except Exception as exc: # pragma: no cover - surfaced via errors list
|
||||
errors.append(exc)
|
||||
finally:
|
||||
connection.close()
|
||||
|
||||
with mock.patch.object(
|
||||
DataImportRow, '_get_update_instance', synced_get_update_instance
|
||||
):
|
||||
# Note: 'batch' is deliberately avoided here - it has a model-level
|
||||
# default (generate_batch_code), and DRF applies a field's default to
|
||||
# *any* non-partial update where the field is omitted, regardless of
|
||||
# the existing instance value. That's a separate, already-tracked
|
||||
# issue (GH #12499) and would confound this test, which is only
|
||||
# about proving the row lock closes the read/write race.
|
||||
thread_a = threading.Thread(target=update, args=('notes', 'notes-a'))
|
||||
thread_b = threading.Thread(
|
||||
target=update, args=('packaging', 'packaging-b')
|
||||
)
|
||||
|
||||
thread_a.start()
|
||||
thread_b.start()
|
||||
thread_a.join(timeout=10)
|
||||
thread_b.join(timeout=10)
|
||||
|
||||
self.assertEqual(errors, [])
|
||||
|
||||
self.item.refresh_from_db()
|
||||
self.assertEqual(self.item.notes, 'notes-a')
|
||||
self.assertEqual(self.item.packaging, 'packaging-b')
|
||||
|
||||
|
||||
class ImportAPITest(ImporterMixin, InvenTreeAPITestCase):
|
||||
"""End-to-end tests for the importer API."""
|
||||
|
||||
|
||||
Reference in New Issue
Block a user