167 lines
6.8 KiB
Python
167 lines
6.8 KiB
Python
from rest_framework import serializers, status
|
|
from rest_framework.response import Response
|
|
from rest_framework.views import APIView
|
|
|
|
from .algorithms import ProcessingError
|
|
from .registry import operation_metadata
|
|
from .selectors import image_session_get, image_state_get, image_states_list, processing_job_get
|
|
from .services import (
|
|
batch_job_create,
|
|
combine_states,
|
|
image_session_create,
|
|
image_session_process,
|
|
image_state_apply_operation,
|
|
image_state_payload,
|
|
image_states_payload,
|
|
processing_job_payload,
|
|
)
|
|
|
|
|
|
def error_response(message, code=status.HTTP_400_BAD_REQUEST):
|
|
return Response({"detail": message}, status=code)
|
|
|
|
|
|
class ImageUploadView(APIView):
|
|
class InputSerializer(serializers.Serializer):
|
|
image = serializers.ImageField(required=False)
|
|
image_base64 = serializers.CharField(required=False, allow_blank=False)
|
|
|
|
def validate(self, attrs):
|
|
if not attrs.get("image") and not attrs.get("image_base64"):
|
|
raise serializers.ValidationError("Provide image or image_base64.")
|
|
return attrs
|
|
|
|
def post(self, request):
|
|
serializer = self.InputSerializer(data=request.data)
|
|
serializer.is_valid(raise_exception=True)
|
|
try:
|
|
payload = image_session_create(
|
|
uploaded_file=serializer.validated_data.get("image"),
|
|
image_base64=serializer.validated_data.get("image_base64"),
|
|
)
|
|
return Response(payload, status=status.HTTP_201_CREATED)
|
|
except ProcessingError as exc:
|
|
return error_response(str(exc))
|
|
|
|
|
|
class ProcessView(APIView):
|
|
class InputSerializer(serializers.Serializer):
|
|
session_id = serializers.UUIDField()
|
|
operation = serializers.CharField()
|
|
params = serializers.DictField(required=False, default=dict)
|
|
|
|
def post(self, request):
|
|
serializer = self.InputSerializer(data=request.data)
|
|
serializer.is_valid(raise_exception=True)
|
|
session = image_session_get(session_id=serializer.validated_data["session_id"])
|
|
if session is None:
|
|
return error_response("Image session does not exist.", status.HTTP_404_NOT_FOUND)
|
|
|
|
try:
|
|
payload = image_session_process(
|
|
session=session,
|
|
operation=serializer.validated_data["operation"],
|
|
params=serializer.validated_data.get("params", {}),
|
|
)
|
|
return Response(payload)
|
|
except ProcessingError as exc:
|
|
code = status.HTTP_410_GONE if str(exc) == "Image session has expired." else status.HTTP_400_BAD_REQUEST
|
|
return error_response(str(exc), code)
|
|
|
|
|
|
class OperationsView(APIView):
|
|
def get(self, request):
|
|
return Response({"operations": operation_metadata()})
|
|
|
|
|
|
class SessionStatesView(APIView):
|
|
def get(self, request, session_id):
|
|
session = image_session_get(session_id=session_id)
|
|
if session is None:
|
|
return error_response("Image session does not exist.", status.HTTP_404_NOT_FOUND)
|
|
return Response({"session_id": str(session.id), "states": image_states_payload(states=image_states_list(session_id=session.id))})
|
|
|
|
|
|
class StateOperationView(APIView):
|
|
class InputSerializer(serializers.Serializer):
|
|
operation = serializers.CharField()
|
|
params = serializers.DictField(required=False, default=dict)
|
|
|
|
def post(self, request, state_id):
|
|
serializer = self.InputSerializer(data=request.data)
|
|
serializer.is_valid(raise_exception=True)
|
|
state = image_state_get(state_id=state_id)
|
|
if state is None:
|
|
return error_response("Image state does not exist.", status.HTTP_404_NOT_FOUND)
|
|
try:
|
|
payload = image_state_apply_operation(
|
|
state=state,
|
|
operation=serializer.validated_data["operation"],
|
|
params=serializer.validated_data.get("params", {}),
|
|
)
|
|
return Response(payload, status=status.HTTP_201_CREATED)
|
|
except ProcessingError as exc:
|
|
code = status.HTTP_410_GONE if str(exc) == "Image session has expired." else status.HTTP_400_BAD_REQUEST
|
|
return error_response(str(exc), code)
|
|
|
|
|
|
class StateCombineView(APIView):
|
|
class InputSerializer(serializers.Serializer):
|
|
operation = serializers.ChoiceField(choices=["add", "subtract", "dot_product", "average", "and", "or"])
|
|
state_ids = serializers.ListField(child=serializers.UUIDField(), min_length=2)
|
|
params = serializers.DictField(required=False, default=dict)
|
|
|
|
def post(self, request):
|
|
serializer = self.InputSerializer(data=request.data)
|
|
serializer.is_valid(raise_exception=True)
|
|
states = [image_state_get(state_id=state_id) for state_id in serializer.validated_data["state_ids"]]
|
|
if any(state is None for state in states):
|
|
return error_response("One or more image states do not exist.", status.HTTP_404_NOT_FOUND)
|
|
try:
|
|
payload = combine_states(
|
|
states=states,
|
|
operation=serializer.validated_data["operation"],
|
|
params=serializer.validated_data.get("params", {}),
|
|
)
|
|
return Response(payload, status=status.HTTP_201_CREATED)
|
|
except ProcessingError as exc:
|
|
return error_response(str(exc))
|
|
|
|
|
|
class StateHistogramView(APIView):
|
|
def get(self, request, state_id):
|
|
state = image_state_get(state_id=state_id)
|
|
if state is None:
|
|
return error_response("Image state does not exist.", status.HTTP_404_NOT_FOUND)
|
|
return Response({"state_id": str(state.id), "histogram": state.histogram, "state": image_state_payload(state=state, include_image=False)})
|
|
|
|
|
|
class BatchView(APIView):
|
|
class InputSerializer(serializers.Serializer):
|
|
operation = serializers.ChoiceField(choices=["average", "subtract"])
|
|
session_ids = serializers.ListField(child=serializers.UUIDField(), min_length=2)
|
|
params = serializers.DictField(required=False, default=dict)
|
|
|
|
def post(self, request):
|
|
serializer = self.InputSerializer(data=request.data)
|
|
serializer.is_valid(raise_exception=True)
|
|
job = batch_job_create(
|
|
operation=serializer.validated_data["operation"],
|
|
session_ids=serializer.validated_data["session_ids"],
|
|
params=serializer.validated_data.get("params", {}),
|
|
)
|
|
return Response({"job_id": str(job.id), "status": job.status}, status=status.HTTP_202_ACCEPTED)
|
|
|
|
|
|
class JobDetailView(APIView):
|
|
def get(self, request, job_id):
|
|
job = processing_job_get(job_id=job_id)
|
|
if job is None:
|
|
return error_response("Job does not exist.", status.HTTP_404_NOT_FOUND)
|
|
return Response(processing_job_payload(job=job))
|
|
|
|
|
|
class HealthView(APIView):
|
|
def get(self, request):
|
|
return Response({"status": "ok"})
|