-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathcrop_run_views.py
More file actions
91 lines (78 loc) · 3.11 KB
/
Copy pathcrop_run_views.py
File metadata and controls
91 lines (78 loc) · 3.11 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
"""Views for CropRun — unified segmentation-inference persistence and human corrections."""
from django.shortcuts import get_object_or_404
from rest_framework import generics, mixins, status, viewsets
from rest_framework.exceptions import PermissionDenied
from rest_framework.permissions import IsAuthenticated
from rest_framework.response import Response
from .models import CropRun, PieceStateImage
from .serializers import CropRunCreateSerializer, CropRunSerializer
class CropRunViewSet(
mixins.CreateModelMixin,
mixins.ListModelMixin,
viewsets.GenericViewSet,
):
permission_classes = [IsAuthenticated]
def get_serializer_class(self):
if self.action == "create":
return CropRunCreateSerializer
return CropRunSerializer
def create(self, request, *args, **kwargs):
serializer = self.get_serializer(data=request.data)
serializer.is_valid(raise_exception=True)
crop_run = self.perform_create(serializer)
output_serializer = CropRunSerializer(
crop_run, context=self.get_serializer_context()
)
headers = self.get_success_headers(output_serializer.data)
return Response(
output_serializer.data,
status=status.HTTP_201_CREATED,
headers=headers,
)
def get_queryset(self):
user = self.request.user
assert user.is_authenticated
qs = CropRun.objects.select_related(
"piece_state_image", "piece_state_image__image", "submitter"
)
if user.is_staff:
return qs
return qs.filter(piece_state_image__piece_state__piece__user=user)
def perform_create(self, serializer):
user = self.request.user
assert user.is_authenticated
piece_state_image_id = serializer.validated_data["piece_state_image_id"]
piece_state_image = get_object_or_404(PieceStateImage, id=piece_state_image_id)
if not user.is_staff:
if piece_state_image.piece_state.piece.user_id != user.id:
raise PermissionDenied(
"You may only submit crop runs for your own pieces."
)
source = {
"type": "human",
"backend": None,
"deployment": "web-ui",
"version": None,
}
return serializer.save(
piece_state_image=piece_state_image,
submitter=user,
source=source,
status=CropRun.Status.SUCCESS,
)
class ImageCropRunsView(generics.ListAPIView):
serializer_class = CropRunSerializer
permission_classes = [IsAuthenticated]
def get_queryset(self):
image_id = self.kwargs["image_id"]
user = self.request.user
assert user.is_authenticated
qs = CropRun.objects.filter(piece_state_image__image_id=image_id)
if not user.is_staff:
if not PieceStateImage.objects.filter(
image_id=image_id, piece_state__piece__user=user
).exists():
return CropRun.objects.none()
if self.request.query_params.get("latest"):
qs = qs[:1]
return qs