# SPDX-FileCopyrightText: Copyright (C) 2026 Omid Jafari <omidjafari.com>
# SPDX-License-Identifier: AGPL-3.0-or-later
#
# This program is free software: you can redistribute it and/or modify
# it under the terms of the GNU Affero General Public License as published by
# the Free Software Foundation, either version 3 of the License, or
# (at your option) any later version.
#
# This program is distributed in the hope that it will be useful,
# but WITHOUT ANY WARRANTY; without even the implied warranty of
# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
# GNU Affero General Public License for more details.
#
# You should have received a copy of the GNU Affero General Public License
# along with this program. If not, see <http://www.gnu.org/licenses/>.
from django.contrib.auth.models import Group
from django.db import transaction
from django.db.models import Exists
from django.db.models import OuterRef
from django.db.models import Q
from rest_framework import status
from rest_framework.decorators import action
from rest_framework.exceptions import ValidationError
from rest_framework.permissions import IsAuthenticated
from rest_framework.response import Response
from nlpmed_portal.annotations.api.serializers import BulkCreateTaskSerializer
from nlpmed_portal.annotations.api.serializers import ProjectMembershipDetailSerializer
from nlpmed_portal.annotations.api.serializers import ProjectMembershipListSerializer
from nlpmed_portal.annotations.api.serializers import TaskDetailSerializer
from nlpmed_portal.annotations.api.serializers import TaskListSerializer
from nlpmed_portal.annotations.api.viewsets.base import BaseViewSet
from nlpmed_portal.annotations.models import Patient
from nlpmed_portal.annotations.models import Project
from nlpmed_portal.annotations.models import ProjectMembership
from nlpmed_portal.annotations.models import Task
from nlpmed_portal.annotations.permissions import MethodBasePermission
from nlpmed_portal.annotations.permissions import has_permission
from nlpmed_portal.users.models import User
[docs]
class ProjectMembershipViewSet(BaseViewSet):
queryset = ProjectMembership.objects.select_related(
"project",
"group",
"assignee",
"assigner",
)
permission_classes = [IsAuthenticated, MethodBasePermission]
required_perms = {
"GET": "view_projectmembership",
"HEAD": "view_projectmembership",
"OPTIONS": "view_projectmembership",
"POST": "add_projectmembership",
"PUT": "change_projectmembership",
"PATCH": "change_projectmembership",
"DELETE": "delete_projectmembership",
}
[docs]
def get_serializer_class(self):
if self.action == "list":
return ProjectMembershipListSerializer
return ProjectMembershipDetailSerializer
[docs]
def get_object(self, *args, **kwargs):
if not hasattr(self, "_cached_obj"):
self._cached_obj = super().get_object(*args, **kwargs)
return self._cached_obj
[docs]
def get_queryset(self):
required_perm = self.required_perms.get("GET")
queryset = super().get_queryset()
project_id = self.request.query_params.get("project_id")
if project_id:
if not project_id.isdigit():
raise ValidationError(
{"project_id": "Invalid project_id. Must be an integer."},
)
queryset = queryset.filter(project_id=project_id)
return self.filter_queryset_by_scope(
queryset,
required_perm,
project_field="project",
)
[docs]
def get_serializer_context(self):
ctx = dict(super().get_serializer_context())
if self.action != "list":
ctx.setdefault("all_groups", list(Group.objects.values("id", "name")))
ctx.setdefault(
"all_users",
list(
User.objects.filter(is_staff=False).values("id", "username"),
),
)
return ctx
[docs]
class TaskViewSet(BaseViewSet):
queryset = Task.objects.select_related(
"assignee",
"assigner",
"patient",
).all()
permission_classes = [IsAuthenticated, MethodBasePermission]
required_perms = {
"GET": "view_task",
"HEAD": "view_task",
"OPTIONS": "view_task",
"POST": "add_task",
"PUT": "change_task",
"PATCH": "change_task",
"DELETE": "delete_task",
}
[docs]
def get_serializer_class(self):
if self.action == "list":
return TaskListSerializer
return TaskDetailSerializer
[docs]
def get_queryset(self):
user = self.request.user
queryset = super().get_queryset()
has_view_all_perm = has_permission(user, "view_all_tasks")
if not has_view_all_perm:
queryset = queryset.filter(assignee=user)
assignee_id = self.request.query_params.get("assignee_id")
if assignee_id:
if not assignee_id.isdigit():
raise ValidationError(
{"assignee_id": "Invalid assignee_id. Must be an integer."},
)
queryset = queryset.filter(assignee_id=assignee_id)
return queryset
[docs]
@action(detail=False, methods=["POST"])
def bulk_create(self, request):
serializer = BulkCreateTaskSerializer(data=request.data)
serializer.is_valid(raise_exception=True)
validated_data = serializer.validated_data
assignee = validated_data["assignee"]
patients = validated_data["patients"]
task_type = validated_data["task_type"]
assigner = request.user
# Can't use bulk_create because of simple-history and auditlog
with transaction.atomic():
created_tasks = []
for patient in patients:
t = Task(
assignee=assignee,
assigner=assigner,
patient=patient,
task_type=task_type,
)
t.save()
created_tasks.append(t)
return Response(
{"created_count": len(created_tasks)},
status=status.HTTP_201_CREATED,
)
[docs]
@action(detail=False, methods=["DELETE"])
def delete_all(self, request, *args, **kwargs):
patient_id = request.query_params.get("patient_id")
project_id = request.query_params.get("project_id")
if not patient_id and not project_id:
return Response(
{
"error": "Either patient_id or project_id must be \
provided for delete all operation.",
},
status=status.HTTP_400_BAD_REQUEST,
)
queryset = super().get_queryset()
project = None
if patient_id:
if not patient_id.isdigit():
raise ValidationError(
{"patient_id": "Invalid patient_id. Must be an integer."},
)
queryset = queryset.filter(patient_id=patient_id)
project = Patient.objects.get(id=patient_id).project
elif project_id:
if not project_id.isdigit():
raise ValidationError(
{"project_id": "Invalid project_id. Must be an integer."},
)
queryset = queryset.filter(patient__project_id=project_id)
project = Project.objects.get(id=project_id)
if not has_permission(
user=request.user,
codename="delete_all_tasks",
project=project,
):
return Response(
{"error": "No permission to delete all tasks."},
status=status.HTTP_403_FORBIDDEN,
)
_, deleted_details = queryset.delete()
model_label = Task._meta.label # ruff: ignore[private-member-access]
model_deleted_count = deleted_details.get(model_label, 0)
return Response(
{"deleteCount": model_deleted_count},
status=status.HTTP_200_OK,
)