diff --git a/docs/domains/shifts.md b/docs/domains/shifts.md new file mode 100644 index 00000000..9547ef82 --- /dev/null +++ b/docs/domains/shifts.md @@ -0,0 +1,74 @@ +# Shifts + +## Domain + +### Models + +- `BusStop`: defines arrival + departure time from a `Place` +- `BusShift`: driver and bus assigned to a list of `BusStop` + +Possible improvements (depending on business logic): +- Make `BusStop` a true JOIN table connecting `Place` and `BusShift` +if a `BusStop` can only belong to one `BusShift`. +- Remove `date` logic and keep only `time` if a `BusShift` is expected +to run the same on different day (could also add logic differentiating +weekdays/weekends/holidays/etc.) +- Add check that `arrival_time`/`departure_time` on `BusStop` objects +is coherent with average bus speed and location of stops. + +### Signals + +- `create_shift_managers_group`: creates Shift managers group +and assigns `can_manage_shifts` permission. Runs when migrations +for `shifts` app have run. Caveat: will run after each new migration. +Should be updated to run only once after relevant migration has gone +through. Perhaps a separate data migration would be a better choice. +- `validate_bus_shift_stops`: uses `validate_stops` validator +when updating/deleting/clearing a `BusShift`. Currently subject to +a possible TOCTOU race condition. Could be fixed with a `select_for_update()` +and `transaction.atomic()` in the form save path, but won't be effective +on SQLite, so did not implement. + +### Admin + +* Separate admin panel allowing access to `Shift managers` group (RBAC logic) only to +`BusStop` and `BusShift` objects with permission to edit/add, but not +delete (sensitive action). Can be updated depending on business logic. +* Main admin panel could let `staff` users access the objects if +users were granted the auto-created permission on the model (`view_busshift`) +outside of the group. Not particularly problematic for this test. +* In a real-life case, an actual front-end interface would be better +suited to allowing users to perform this task, hence the choice to +separate the admin sites. The form aspect to modify `BusStop` objects +from a `BusShift` could be improved. + +### Validators + +* `validate_stops` is used both by the signal and the admin panel +form to have a single source of truth. Using `Min/Max` annotations +to avoid running N+1 queries. `ValidationError` in the validator applies to +admin panel whilst signal catches it and raises `IntegrityError` for +DB-level. + +### Commands + +* `create_shifts`: creates bus shifts with associated stops + +Improvements: +* Add command to create user in `ShiftManager` group. + +## Run + +Access the shift admin panel in local at `http://localhost:8000/shift-admin`. +You need a `superuser` or `shift manager` account to connect. + +Run tests: +```commandline +python3 manage.py tests padam_django.apps.shifts +``` + +Create `N` shifts (with `N^2` stops, +`N` of those randomly assigned to each shift): +```commandline +python3 manage.py create_shifts -n N +``` diff --git a/padam_django/apps/common/management/commands/create_data.py b/padam_django/apps/common/management/commands/create_data.py index a149a937..045f5667 100644 --- a/padam_django/apps/common/management/commands/create_data.py +++ b/padam_django/apps/common/management/commands/create_data.py @@ -12,3 +12,4 @@ def handle(self, *args, **options): management.call_command('create_drivers', number=5) management.call_command('create_buses', number=10) management.call_command('create_places', number=30) + management.call_command('create_shifts', number=5) diff --git a/padam_django/apps/shifts/__init__.py b/padam_django/apps/shifts/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/padam_django/apps/shifts/admin.py b/padam_django/apps/shifts/admin.py new file mode 100644 index 00000000..58f47656 --- /dev/null +++ b/padam_django/apps/shifts/admin.py @@ -0,0 +1,120 @@ +from django.contrib import admin +from django.contrib.admin.forms import AuthenticationForm +from django.db.models import Max, Min +from django.db.models.query import QuerySet +from django.http.request import HttpRequest + +from . import models +from .forms import BusShiftForm + + +class ShiftAdminSite(admin.AdminSite): + site_header = "Shift Admin Portal" + site_title = "Shift Admin" + index_title = "Shift management" + login_form = AuthenticationForm + + def has_permission(self, request: HttpRequest) -> bool: + return request.user.is_authenticated and ( + request.user.is_superuser + or request.user.groups.filter(name="Shift managers").exists() + ) + + +shift_admin = ShiftAdminSite(name="shift_admin") + + +@admin.register(models.BusStop) +class BusStopAdmin(admin.ModelAdmin): + list_display = ["id", "get_place_name", "arrival_time", "departure_time"] + + def get_place_name(self, obj: models.BusStop) -> str: + return obj.place.name + + get_place_name.short_description = "Place" + + def has_change_permission( + self, request: HttpRequest, obj: models.BusShift | None = None + ) -> bool: + if request.user.has_perm("shifts.can_manage_shifts"): + return True + return False + + def has_delete_permission( + self, request: HttpRequest, obj: models.BusShift | None = None + ) -> bool: + return False + + def has_add_permission( + self, request: HttpRequest, obj: models.BusShift | None = None + ) -> bool: + if request.user.has_perm("shifts.can_manage_shifts"): + return True + return False + + +@admin.register(models.BusShift) +class BusShiftAdmin(admin.ModelAdmin): + list_display = [ + "id", + "get_bus_licence_plate", + "get_driver_username", + "get_start_time", + "get_end_time", + ] + list_filter = ["bus__licence_plate"] + search_fields = ["bus__licence_plate", "driver__user__username"] + form = BusShiftForm + list_select_related = ["bus", "driver__user"] + + def get_queryset(self, request: HttpRequest) -> QuerySet[models.BusShift]: + return ( + super() + .get_queryset(request) + .annotate( + _start_time=Min("stops__departure_time"), + _end_time=Max("stops__arrival_time"), + ) + ) + + def get_bus_licence_plate(self, obj: models.BusShift) -> str: + return obj.bus.licence_plate + + def get_driver_username(self, obj: models.BusShift) -> str: + return obj.driver.user.username + + def get_start_time(self, obj: models.BusShift): + return obj._start_time + + def get_end_time(self, obj: models.BusShift): + return obj._end_time + + get_bus_licence_plate.short_description = "Bus" + get_driver_username.short_description = "Driver" + get_start_time.short_description = "Start time" + get_start_time.admin_order_field = "_start_time" + get_end_time.short_description = "End time" + get_end_time.admin_order_field = "_end_time" + + def has_change_permission( + self, request: HttpRequest, obj: models.BusShift | None = None + ) -> bool: + if request.user.has_perm("shifts.can_manage_shifts"): + return True + return False + + def has_delete_permission( + self, request: HttpRequest, obj: models.BusShift | None = None + ) -> bool: + return False + + def has_add_permission( + self, request: HttpRequest, obj: models.BusShift | None = None + ) -> bool: + if request.user.has_perm("shifts.can_manage_shifts"): + return True + return False + + +shift_admin.register(models.BusStop, BusStopAdmin) +shift_admin.register(models.BusShift, BusShiftAdmin) diff --git a/padam_django/apps/shifts/apps.py b/padam_django/apps/shifts/apps.py new file mode 100644 index 00000000..f170196c --- /dev/null +++ b/padam_django/apps/shifts/apps.py @@ -0,0 +1,10 @@ +from django.apps import AppConfig + + +class ShiftsConfig(AppConfig): + name = "padam_django.apps.shifts" + models_module = "padam_django.apps.shifts" + label = "shifts" + + def ready(self): + import padam_django.apps.shifts.signals # noqa: F401 diff --git a/padam_django/apps/shifts/factories.py b/padam_django/apps/shifts/factories.py new file mode 100644 index 00000000..cba79701 --- /dev/null +++ b/padam_django/apps/shifts/factories.py @@ -0,0 +1,37 @@ +import datetime + +import factory +from django.utils.timezone import now +from faker import Faker + +from . import models + +fake = Faker(["fr"]) + + +class BusStopFactory(factory.django.DjangoModelFactory): + place = factory.SubFactory("padam_django.apps.geography.factories.PlaceFactory") + arrival_time = factory.Faker( + "date_time_between", start_date=now().date(), tzinfo=datetime.timezone.utc + ) + departure_time = factory.LazyAttribute( + lambda o: o.arrival_time + + datetime.timedelta(minutes=fake.random_int(min=1, max=3)) + ) + + class Meta: + model = models.BusStop + + +class BusShiftFactory(factory.django.DjangoModelFactory): + driver = factory.SubFactory("padam_django.apps.fleet.factories.DriverFactory") + bus = factory.SubFactory("padam_django.apps.fleet.factories.BusFactory") + + @factory.post_generation + def stops(self, create: bool, extracted: list, **kwargs: dict): + if not create or not extracted: + return + self.stops.add(*extracted) + + class Meta: + model = models.BusShift diff --git a/padam_django/apps/shifts/forms.py b/padam_django/apps/shifts/forms.py new file mode 100644 index 00000000..4b395d02 --- /dev/null +++ b/padam_django/apps/shifts/forms.py @@ -0,0 +1,28 @@ +from typing import Any + +from django import forms +from django.core.exceptions import ValidationError + +from padam_django.apps.shifts.models import BusShift +from padam_django.apps.shifts.validators import validate_stops + + +class BusShiftForm(forms.ModelForm): + class Meta: + model = BusShift + fields = "__all__" + + def clean(self) -> dict[str, Any] | None: + cleaned_data = super().clean() + if not cleaned_data: + return None + stops = cleaned_data.get("stops") + driver = cleaned_data.get("driver") + bus = cleaned_data.get("bus") + + if stops and driver and bus: + try: + validate_stops(driver, bus, list(stops), exclude_pk=self.instance.pk) + except ValidationError as e: + raise forms.ValidationError(e.message) + return cleaned_data diff --git a/padam_django/apps/shifts/management/__init__.py b/padam_django/apps/shifts/management/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/padam_django/apps/shifts/management/commands/__init__.py b/padam_django/apps/shifts/management/commands/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/padam_django/apps/shifts/management/commands/create_shifts.py b/padam_django/apps/shifts/management/commands/create_shifts.py new file mode 100644 index 00000000..c5928532 --- /dev/null +++ b/padam_django/apps/shifts/management/commands/create_shifts.py @@ -0,0 +1,19 @@ +import random + +from padam_django.apps.common.management.base import CreateDataBaseCommand +from padam_django.apps.shifts.factories import BusShiftFactory, BusStopFactory + + +class Command(CreateDataBaseCommand): + + help = "Create a few shifts" + + def handle(self, *args, **options): + super().handle(*args, **options) + self.stdout.write( + f"Creating {self.number} shifts and {self.number ** 2} stops ..." + ) + stops = BusStopFactory.create_batch(size=self.number**2) + # Doing this to have different stops, not ideal + for ind in range(self.number): + BusShiftFactory.create(stops=random.sample(stops, self.number)) diff --git a/padam_django/apps/shifts/migrations/0001_initial.py b/padam_django/apps/shifts/migrations/0001_initial.py new file mode 100644 index 00000000..c9eabbe6 --- /dev/null +++ b/padam_django/apps/shifts/migrations/0001_initial.py @@ -0,0 +1,90 @@ +# Generated by Django 4.2.16 on 2026-08-04 17:25 + +import django.db.models.deletion +from django.db import migrations, models + + +class Migration(migrations.Migration): + + initial = True + + dependencies = [ + ("geography", "0001_initial"), + ("fleet", "0002_auto_20211109_1456"), + ] + + operations = [ + migrations.CreateModel( + name="BusStop", + fields=[ + ( + "id", + models.BigAutoField( + auto_created=True, + primary_key=True, + serialize=False, + verbose_name="ID", + ), + ), + ("arrival_time", models.DateTimeField()), + ("departure_time", models.DateTimeField()), + ( + "place", + models.ForeignKey( + on_delete=django.db.models.deletion.CASCADE, + related_name="bus_stop", + to="geography.place", + ), + ), + ], + options={ + "db_table": "bus_stop", + "ordering": ["arrival_time"], + }, + ), + migrations.CreateModel( + name="BusShift", + fields=[ + ( + "id", + models.BigAutoField( + auto_created=True, + primary_key=True, + serialize=False, + verbose_name="ID", + ), + ), + ( + "bus", + models.ForeignKey( + on_delete=django.db.models.deletion.CASCADE, + related_name="bus_shift", + to="fleet.bus", + ), + ), + ( + "driver", + models.ForeignKey( + on_delete=django.db.models.deletion.CASCADE, + related_name="bus_shift", + to="fleet.driver", + ), + ), + ("stops", models.ManyToManyField(to="shifts.busstop")), + ], + options={ + "db_table": "bus_shift", + }, + ), + migrations.AddConstraint( + model_name="busstop", + constraint=models.CheckConstraint( + check=models.Q(("arrival_time__lt", models.F("departure_time"))), + name="bus_stop_arrival_time__lt_departure_time", + ), + ), + migrations.AlterModelOptions( + name="busshift", + options={"permissions": (("can_manage_shifts", "Can manage shifts"),)}, + ), + ] diff --git a/padam_django/apps/shifts/migrations/__init__.py b/padam_django/apps/shifts/migrations/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/padam_django/apps/shifts/models.py b/padam_django/apps/shifts/models.py new file mode 100644 index 00000000..f46af540 --- /dev/null +++ b/padam_django/apps/shifts/models.py @@ -0,0 +1,49 @@ +from datetime import datetime, timedelta + +from django.db import models + +from padam_django.apps.fleet.models import Bus, Driver +from padam_django.apps.geography.models import Place + + +class BusStop(models.Model): + place = models.ForeignKey(Place, on_delete=models.CASCADE, related_name="bus_stop") + arrival_time = models.DateTimeField() + departure_time = models.DateTimeField() + + class Meta: + db_table = "bus_stop" + ordering = ["arrival_time"] + constraints = [ + models.CheckConstraint( + check=models.Q(arrival_time__lt=models.F("departure_time")), + name="bus_stop_arrival_time__lt_departure_time", + ) + ] + + def __str__(self) -> str: + return f"{self.place.name} - {self.arrival_time} - {self.departure_time}" + + +class BusShift(models.Model): + driver = models.ForeignKey( + Driver, on_delete=models.CASCADE, related_name="bus_shift" + ) + bus = models.ForeignKey(Bus, on_delete=models.CASCADE, related_name="bus_shift") + stops = models.ManyToManyField(BusStop) + + class Meta: + db_table = "bus_shift" + permissions = (("can_manage_shifts", "Can manage shifts"),) + + @property + def start_time(self) -> datetime: + return self.stops.order_by("departure_time").first().departure_time + + @property + def end_time(self) -> datetime: + return self.stops.order_by("arrival_time").last().arrival_time + + @property + def duration(self) -> timedelta: + return self.end_time - self.start_time diff --git a/padam_django/apps/shifts/signals.py b/padam_django/apps/shifts/signals.py new file mode 100644 index 00000000..c2fbf9e1 --- /dev/null +++ b/padam_django/apps/shifts/signals.py @@ -0,0 +1,36 @@ +from django.apps import apps +from django.contrib.auth.models import Group, Permission +from django.core.exceptions import ValidationError +from django.db import models +from django.db.models.signals import post_migrate +from django.db.utils import IntegrityError +from django.dispatch import receiver + +from .models import BusShift +from .validators import validate_stops + + +@receiver(post_migrate, sender=apps.get_app_config("shifts")) +def create_shift_managers_group(sender, **kwargs) -> None: + permission = Permission.objects.get( + codename="can_manage_shifts", content_type__app_label="shifts" + ) + shift_managers, _ = Group.objects.get_or_create(name="Shift managers") + shift_managers.permissions.add(permission) + + +@receiver(models.signals.m2m_changed, sender=BusShift.stops.through) +def validate_bus_shift_stops( + sender: object | None, instance: BusShift, action: str, **kwargs +) -> None: + if action not in ["post_add", "post_remove", "post_clear"]: + return + try: + validate_stops( + instance.driver, + instance.bus, + list(instance.stops.all()), + exclude_pk=instance.pk, + ) + except ValidationError as e: + raise IntegrityError(str(e)) diff --git a/padam_django/apps/shifts/tests/__init__.py b/padam_django/apps/shifts/tests/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/padam_django/apps/shifts/tests/test_admin.py b/padam_django/apps/shifts/tests/test_admin.py new file mode 100644 index 00000000..4d4e8088 --- /dev/null +++ b/padam_django/apps/shifts/tests/test_admin.py @@ -0,0 +1,58 @@ +import django.test +from django.contrib.auth.models import Group + +from padam_django.apps.fleet.factories import BusFactory, DriverFactory +from padam_django.apps.shifts.factories import BusStopFactory +from padam_django.apps.shifts.models import BusShift +from padam_django.apps.users.factories import UserFactory + +from .utils import _at + + +class TestBusShiftAdminView(django.test.TestCase): + def setUp(self) -> None: + self.manager = UserFactory() + self.manager.groups.add(Group.objects.get(name="Shift managers")) + self.client.force_login(self.manager) + self.driver = DriverFactory.create() + self.bus = BusFactory.create() + self.stops = [ + BusStopFactory.create(arrival_time=_at(10, 0), departure_time=_at(10, 5)), + BusStopFactory.create(arrival_time=_at(11, 0), departure_time=_at(11, 5)), + ] + + def test_add_view_rejects_overlap_without_creating_a_row(self) -> None: + BusShift.objects.create(driver=self.driver, bus=self.bus).stops.set(self.stops) + clashing_stops = [ + BusStopFactory.create(arrival_time=_at(10, 15), departure_time=_at(10, 30)), + BusStopFactory.create(arrival_time=_at(20, 0), departure_time=_at(20, 5)), + ] + + response = self.client.post( + "/shift-admin/shifts/busshift/add/", + { + "driver": self.driver.pk, + "bus": self.bus.pk, + "stops": [stop.pk for stop in clashing_stops], + }, + ) + + self.assertEqual(response.status_code, 200) # re-rendered form, not a redirect + self.assertContains(response, "overlapping shift") + self.assertEqual(BusShift.objects.count(), 1) + + def test_add_view_accepts_valid_shift(self) -> None: + driver, bus = DriverFactory.create(), BusFactory.create() + stops = BusStopFactory.create_batch(2) + + response = self.client.post( + "/shift-admin/shifts/busshift/add/", + { + "driver": driver.pk, + "bus": bus.pk, + "stops": [stop.pk for stop in stops], + }, + ) + + self.assertEqual(response.status_code, 302) + self.assertEqual(BusShift.objects.count(), 1) diff --git a/padam_django/apps/shifts/tests/test_bus_shifts.py b/padam_django/apps/shifts/tests/test_bus_shifts.py new file mode 100644 index 00000000..844bcab6 --- /dev/null +++ b/padam_django/apps/shifts/tests/test_bus_shifts.py @@ -0,0 +1,86 @@ +import django.test +from django.db.utils import IntegrityError + +from padam_django.apps.fleet.factories import BusFactory, DriverFactory +from padam_django.apps.shifts.factories import BusStopFactory +from padam_django.apps.shifts.models import BusShift + + +class TestBusShift(django.test.TestCase): + def setUp(self) -> None: + self.stops = BusStopFactory.create_batch(10) + self.driver = DriverFactory.create() + self.bus = BusFactory.create() + self.shift = BusShift.objects.create( + driver=self.driver, + bus=self.bus, + ) + self.stops = sorted(self.stops, key=lambda stop: stop.arrival_time) + self.shift.stops.set(self.stops) + + def test_bus_shift_has_correct_stops(self) -> None: + self.assertEqual(len(self.shift.stops.all()), len(self.stops)) + self.assertEqual(list(self.shift.stops.all()), self.stops) + + def test_bus_shift_start_time_is_first_stop_departure(self) -> None: + self.assertEqual(self.shift.start_time, self.stops[0].departure_time) + + def test_bus_shift_end_time_is_last_stop_arrival(self) -> None: + self.assertEqual(self.shift.end_time, self.stops[-1].arrival_time) + + def test_bus_shift_duration_is_difference_between_start_and_end(self) -> None: + self.assertEqual( + self.shift.duration, self.shift.end_time - self.shift.start_time + ) + + def test_bus_shift_cannot_be_assigned_to_already_driving_driver(self) -> None: + other_bus = BusFactory.create() + other_stops = [ + BusStopFactory.create(arrival_time=self.stops[0].arrival_time), + BusStopFactory.create( + departure_time=self.stops[2].departure_time, + arrival_time=self.stops[2].arrival_time, + ), + ] + other_shift = BusShift.objects.create( + driver=self.driver, + bus=other_bus, + ) + with self.assertRaises(IntegrityError): + other_shift.stops.set(other_stops) + + def test_bus_shift_cannot_be_assigned_to_already_running_bus(self) -> None: + other_driver = DriverFactory.create() + other_stops = [ + BusStopFactory.create(arrival_time=self.stops[0].arrival_time), + BusStopFactory.create( + departure_time=self.stops[2].departure_time, + arrival_time=self.stops[2].arrival_time, + ), + ] + other_shift = BusShift.objects.create(driver=other_driver, bus=self.bus) + with self.assertRaises(IntegrityError): + other_shift.stops.set(other_stops) + + def test_bus_shift_must_have_two_or_more_stops(self) -> None: + stop = BusStopFactory.create() + shift = BusShift.objects.create( + driver=self.driver, + bus=self.bus, + ) + with self.assertRaises(IntegrityError): + shift.stops.set([stop]) + + def test_bus_shift_update_cannot_remove_stop_if_less_than_two(self) -> None: + with self.assertRaises(IntegrityError): + self.shift.stops.set([self.stops[0]]) + + def test_bus_shift_can_be_assigned_if_no_shift_running(self) -> None: + other_driver = DriverFactory.create() + other_bus = BusFactory.create() + shift = BusShift.objects.create( + driver=other_driver, + bus=other_bus, + ) + shift.stops.set(self.stops) + self.assertEqual(len(shift.stops.all()), len(self.stops)) diff --git a/padam_django/apps/shifts/tests/test_bus_stops.py b/padam_django/apps/shifts/tests/test_bus_stops.py new file mode 100644 index 00000000..dff9d4aa --- /dev/null +++ b/padam_django/apps/shifts/tests/test_bus_stops.py @@ -0,0 +1,18 @@ +import django.test +from django.db.utils import IntegrityError + +from padam_django.apps.geography.factories import PlaceFactory +from padam_django.apps.shifts.models import BusStop + +from .utils import _at + + +class BusStopTest(django.test.TestCase): + def setUp(self) -> None: + self.place = PlaceFactory.create() + + def test_bus_stop_departure_time_must_be_smaller_than_arrival(self) -> None: + with self.assertRaises(IntegrityError): + BusStop.objects.create( + place=self.place, arrival_time=_at(15, 0), departure_time=_at(14, 0) + ) diff --git a/padam_django/apps/shifts/tests/test_forms.py b/padam_django/apps/shifts/tests/test_forms.py new file mode 100644 index 00000000..fd791f22 --- /dev/null +++ b/padam_django/apps/shifts/tests/test_forms.py @@ -0,0 +1,71 @@ +import django.test + +from padam_django.apps.fleet.factories import BusFactory, DriverFactory +from padam_django.apps.shifts.factories import BusStopFactory +from padam_django.apps.shifts.forms import BusShiftForm +from padam_django.apps.shifts.models import BusShift + +from .utils import _at + + +class TestBusShiftForm(django.test.TestCase): + def setUp(self) -> None: + self.driver = DriverFactory.create() + self.bus = BusFactory.create() + self.stops = [ + BusStopFactory.create(arrival_time=_at(10, 0), departure_time=_at(10, 5)), + BusStopFactory.create(arrival_time=_at(11, 0), departure_time=_at(11, 5)), + ] + self.clashing_stops = [ + BusStopFactory.create(arrival_time=_at(10, 15), departure_time=_at(10, 30)), + BusStopFactory.create(arrival_time=_at(20, 0), departure_time=_at(20, 5)), + ] + + def _data(self, **overrides) -> dict: + data = { + "driver": self.driver.pk, + "bus": self.bus.pk, + "stops": [stop.pk for stop in self.stops], + } + data.update(overrides) + return data + + def test_valid_submission_is_valid(self) -> None: + form = BusShiftForm(data=self._data()) + self.assertTrue(form.is_valid()) + + def test_single_stop_is_rejected_with_field_error(self) -> None: + one_stop = BusStopFactory.create() + form = BusShiftForm(data=self._data(stops=[one_stop.pk])) + self.assertFalse(form.is_valid()) + self.assertIn("at least two stops", str(form.errors)) + + def test_overlapping_driver_is_rejected_and_creates_nothing(self) -> None: + BusShift.objects.create(driver=self.driver, bus=BusFactory.create()).stops.set( + self.clashing_stops + ) + + form = BusShiftForm(data=self._data()) + + self.assertFalse(form.is_valid()) + self.assertIn("overlapping shift", str(form.errors)) + self.assertEqual(BusShift.objects.count(), 1) + + def test_overlapping_bus_is_rejected_and_creates_nothing(self) -> None: + BusShift.objects.create(driver=DriverFactory.create(), bus=self.bus).stops.set( + self.clashing_stops + ) + + form = BusShiftForm(data=self._data()) + + self.assertFalse(form.is_valid()) + self.assertIn("overlapping shift", str(form.errors)) + self.assertEqual(BusShift.objects.count(), 1) + + def test_editing_existing_shift_excludes_itself_from_clash_check(self) -> None: + shift = BusShift.objects.create(driver=self.driver, bus=self.bus) + shift.stops.set(self.stops) + + form = BusShiftForm(instance=shift, data=self._data()) + + self.assertTrue(form.is_valid()) diff --git a/padam_django/apps/shifts/tests/test_permission.py b/padam_django/apps/shifts/tests/test_permission.py new file mode 100644 index 00000000..2c1c2e57 --- /dev/null +++ b/padam_django/apps/shifts/tests/test_permission.py @@ -0,0 +1,82 @@ +import django.test +from django.apps import apps +from django.contrib.auth.models import Group, Permission +from django.db.models.signals import post_migrate + +from padam_django.apps.shifts.admin import ( + BusShiftAdmin, + BusStopAdmin, + ShiftAdminSite, + shift_admin, +) +from padam_django.apps.shifts.apps import ShiftsConfig +from padam_django.apps.shifts.models import BusShift, BusStop +from padam_django.apps.users.factories import UserFactory + + +class MockRequest: + pass + + +class BusShiftPermissionSignalTest(django.test.TestCase): + def test_creation_signal_sent(self) -> None: + Group.objects.filter(name="Shift managers").delete() + + post_migrate.send(sender=apps.get_app_config("shifts"), app_config=ShiftsConfig) + + group = Group.objects.get(name="Shift managers") + self.assertTrue(group.permissions.filter(codename="can_manage_shifts").exists()) + + +class BusShiftPermissionTest(django.test.TestCase): + def setUp(self): + self.permission = Permission.objects.get(codename="can_manage_shifts") + self.non_manager_user = UserFactory() + self.manager_user = UserFactory() + self.group = Group.objects.get(name="Shift managers") + self.manager_user.groups.add(self.group) + self.bus_shift_admin = BusShiftAdmin( + model=BusShift, admin_site=ShiftAdminSite() + ) + self.bus_stop_admin = BusStopAdmin(model=BusStop, admin_site=ShiftAdminSite()) + + def test_user_in_shift_managers_group_has_permission(self) -> None: + self.assertTrue(self.manager_user.has_perm("shifts.can_manage_shifts")) + + def test_user_not_in_shift_managers_group_does_not_have_access(self) -> None: + self.assertFalse(self.non_manager_user.has_perm("shifts.can_manage_shifts")) + + def test_manager_can_add_and_change_bus_shifts(self) -> None: + mock_request = MockRequest() + mock_request.user = self.manager_user + self.assertTrue(self.bus_shift_admin.has_add_permission(mock_request)) + self.assertTrue(self.bus_shift_admin.has_change_permission(mock_request)) + self.assertFalse(self.bus_shift_admin.has_delete_permission(mock_request)) + + def test_non_manager_cannot_add_or_change_bus_shifts(self) -> None: + mock_request = MockRequest() + mock_request.user = self.non_manager_user + self.assertFalse(self.bus_shift_admin.has_add_permission(mock_request)) + self.assertFalse(self.bus_shift_admin.has_change_permission(mock_request)) + self.assertFalse(self.bus_shift_admin.has_delete_permission(mock_request)) + + def test_superuser_can_add_and_change_bus_shifts(self) -> None: + mock_request = MockRequest() + mock_request.user = UserFactory(is_superuser=True) + self.assertTrue(self.bus_shift_admin.has_add_permission(mock_request)) + self.assertTrue(self.bus_shift_admin.has_change_permission(mock_request)) + + def test_manager_can_access_admin_site(self) -> None: + mock_request = MockRequest() + mock_request.user = self.manager_user + self.assertTrue(shift_admin.has_permission(mock_request)) + + def test_non_manager_cannot_access_admin_site(self) -> None: + mock_request = MockRequest() + mock_request.user = self.non_manager_user + self.assertFalse(shift_admin.has_permission(mock_request)) + + def test_superuser_can_access_admin_site(self) -> None: + mock_request = MockRequest() + mock_request.user = UserFactory(is_superuser=True) + self.assertTrue(shift_admin.has_permission(mock_request)) diff --git a/padam_django/apps/shifts/tests/test_validators.py b/padam_django/apps/shifts/tests/test_validators.py new file mode 100644 index 00000000..3c98155c --- /dev/null +++ b/padam_django/apps/shifts/tests/test_validators.py @@ -0,0 +1,57 @@ +import django.test +from django.core.exceptions import ValidationError + +from padam_django.apps.fleet.factories import BusFactory, DriverFactory +from padam_django.apps.shifts.factories import BusStopFactory +from padam_django.apps.shifts.models import BusShift +from padam_django.apps.shifts.validators import validate_stops + +from .utils import _at + + +class TestValidateStops(django.test.TestCase): + def setUp(self) -> None: + self.driver = DriverFactory.create() + self.bus = BusFactory.create() + + def test_interleaved_shifts_with_no_stop_fully_nested_are_still_detected( + self, + ) -> None: + # Windows overlapping but not fully nested + BusShift.objects.create(driver=self.driver, bus=self.bus).stops.set( + [ + BusStopFactory.create( + arrival_time=_at(8, 0), departure_time=_at(8, 10) + ), + BusStopFactory.create( + arrival_time=_at(8, 20), departure_time=_at(8, 30) + ), + ] + ) + + clashing_stops = [ + BusStopFactory.create(arrival_time=_at(8, 5), departure_time=_at(8, 15)), + BusStopFactory.create(arrival_time=_at(8, 25), departure_time=_at(8, 35)), + ] + + with self.assertRaises(ValidationError): + validate_stops(self.driver, self.bus, clashing_stops) + + def test_non_overlapping_shifts_are_accepted(self) -> None: + BusShift.objects.create(driver=self.driver, bus=self.bus).stops.set( + [ + BusStopFactory.create( + arrival_time=_at(8, 0), departure_time=_at(8, 10) + ), + BusStopFactory.create( + arrival_time=_at(8, 20), departure_time=_at(8, 30) + ), + ] + ) + + later_stops = [ + BusStopFactory.create(arrival_time=_at(9, 0), departure_time=_at(9, 10)), + BusStopFactory.create(arrival_time=_at(9, 20), departure_time=_at(9, 30)), + ] + + validate_stops(self.driver, self.bus, later_stops) # does not raise diff --git a/padam_django/apps/shifts/tests/utils.py b/padam_django/apps/shifts/tests/utils.py new file mode 100644 index 00000000..1981dc29 --- /dev/null +++ b/padam_django/apps/shifts/tests/utils.py @@ -0,0 +1,5 @@ +from datetime import datetime, timezone + + +def _at(hour: int, minute: int = 0) -> datetime: + return datetime(2026, 1, 1, hour, minute, tzinfo=timezone.utc) diff --git a/padam_django/apps/shifts/validators.py b/padam_django/apps/shifts/validators.py new file mode 100644 index 00000000..87c6be0a --- /dev/null +++ b/padam_django/apps/shifts/validators.py @@ -0,0 +1,34 @@ +from typing import List + +from django.core.exceptions import ValidationError +from django.db import models +from django.db.models import Max, Min +from django.db.models.query import QuerySet + +from padam_django.apps.fleet.models import Bus, Driver + +from .models import BusShift, BusStop + + +def validate_stops( + driver: Driver, + bus: Bus, + stops: QuerySet[BusStop] | List[BusStop], + exclude_pk: int | None = None, +) -> None: + if len(stops) < 2: + raise ValidationError("A bus shift needs at least two stops.") + + start_time = min(s.departure_time for s in stops) + end_time = max(s.arrival_time for s in stops) + clashing = ( + BusShift.objects.filter(models.Q(driver=driver) | models.Q(bus=bus)) + .exclude(pk=exclude_pk) + .annotate( + other_start_time=Min("stops__departure_time"), + other_end_time=Max("stops__arrival_time"), + ) + .filter(other_start_time__lt=end_time, other_end_time__gt=start_time) + ) + if clashing.exists(): + raise ValidationError("Driver or bus already assigned to an overlapping shift.") diff --git a/padam_django/settings.py b/padam_django/settings.py index 129e922c..22b9b57d 100644 --- a/padam_django/settings.py +++ b/padam_django/settings.py @@ -45,6 +45,7 @@ 'padam_django.apps.fleet', 'padam_django.apps.geography', 'padam_django.apps.users', + 'padam_django.apps.shifts', ] MIDDLEWARE = [ diff --git a/padam_django/urls.py b/padam_django/urls.py index 7ecf590e..b2c2c7f4 100644 --- a/padam_django/urls.py +++ b/padam_django/urls.py @@ -16,6 +16,9 @@ from django.contrib import admin from django.urls import path +from padam_django.apps.shifts.admin import shift_admin + urlpatterns = [ path('admin/', admin.site.urls), + path("shift-admin/", shift_admin.urls), ]