diff --git a/api/changelog.d/provider-group-action-scope.security.md b/api/changelog.d/provider-group-action-scope.security.md new file mode 100644 index 0000000000..bf1d003891 --- /dev/null +++ b/api/changelog.d/provider-group-action-scope.security.md @@ -0,0 +1 @@ +Provider deletion and connection checks, scan creation, provider secrets, provider groups, and daily schedules now respect role provider-group visibility diff --git a/api/src/backend/api/tests/test_rbac.py b/api/src/backend/api/tests/test_rbac.py index 92d19bd3c0..9c01e9d448 100644 --- a/api/src/backend/api/tests/test_rbac.py +++ b/api/src/backend/api/tests/test_rbac.py @@ -8,8 +8,10 @@ from api.models import ( Membership, ProviderGroup, ProviderGroupMembership, + ProviderSecret, Role, RoleProviderGroupRelationship, + Scan, User, UserRoleRelationship, ) @@ -666,6 +668,612 @@ class TestLimitedVisibility: limited_admin_user, tenants_fixture[0] ) + @pytest.fixture + def hidden_provider_secret(self, aws_provider_pair): + hidden_provider = aws_provider_pair[1] + return ProviderSecret.objects.create( + tenant_id=hidden_provider.tenant_id, + provider=hidden_provider, + secret_type=ProviderSecret.TypeChoices.STATIC, + secret={ + "aws_access_key_id": "hidden-key", + "aws_secret_access_key": "hidden-secret", + }, + name="Hidden provider secret", + ) + + @pytest.fixture + def limited_provider_group(self, limited_admin_user): + return ProviderGroup.objects.get(name="limited_visibility_group") + + @patch("api.v1.views.enqueue_scan_execution_on_commit") + def test_scan_create_out_of_scope_provider_is_rejected( + self, + mock_enqueue_scan, + authenticated_client_rbac_limited, + aws_provider_pair, + ): + hidden_provider = aws_provider_pair[1] + + response = authenticated_client_rbac_limited.post( + reverse("scan-list"), + data=json.dumps( + { + "data": { + "type": "scans", + "attributes": {"name": "Out of scope scan"}, + "relationships": { + "provider": { + "data": { + "type": "providers", + "id": str(hidden_provider.id), + } + } + }, + } + } + ), + content_type="application/vnd.api+json", + ) + + assert response.status_code == status.HTTP_400_BAD_REQUEST + assert not Scan.objects.filter( + provider=hidden_provider, name="Out of scope scan" + ).exists() + mock_enqueue_scan.assert_not_called() + + @patch("api.v1.views.enqueue_scan_execution_on_commit") + def test_scan_create_in_scope_provider_is_accepted( + self, + mock_enqueue_scan, + authenticated_client_rbac_limited, + aws_provider, + ): + response = authenticated_client_rbac_limited.post( + reverse("scan-list"), + data=json.dumps( + { + "data": { + "type": "scans", + "attributes": {"name": "In scope scan"}, + "relationships": { + "provider": { + "data": { + "type": "providers", + "id": str(aws_provider.id), + } + } + }, + } + } + ), + content_type="application/vnd.api+json", + ) + + assert response.status_code == status.HTTP_202_ACCEPTED + assert Scan.objects.filter(provider=aws_provider, name="In scope scan").exists() + mock_enqueue_scan.assert_called_once() + + def test_provider_secret_retrieve_out_of_scope_returns_404( + self, + authenticated_client_rbac_limited, + hidden_provider_secret, + ): + response = authenticated_client_rbac_limited.get( + reverse( + "providersecret-detail", + kwargs={"pk": hidden_provider_secret.id}, + ) + ) + + assert response.status_code == status.HTTP_404_NOT_FOUND + + def test_provider_secret_list_excludes_out_of_scope_provider( + self, + authenticated_client_rbac_limited, + hidden_provider_secret, + ): + response = authenticated_client_rbac_limited.get(reverse("providersecret-list")) + + assert response.status_code == status.HTTP_200_OK + assert str(hidden_provider_secret.id) not in { + item["id"] for item in response.json()["data"] + } + + def test_provider_secret_create_out_of_scope_provider_is_rejected( + self, + authenticated_client_rbac_limited, + aws_provider_pair, + ): + hidden_provider = aws_provider_pair[1] + response = authenticated_client_rbac_limited.post( + reverse("providersecret-list"), + data=json.dumps( + { + "data": { + "type": "provider-secrets", + "attributes": { + "name": "Out of scope secret", + "secret_type": ProviderSecret.TypeChoices.STATIC, + "secret": { + "aws_access_key_id": "hidden-key", + "aws_secret_access_key": "hidden-secret", + }, + }, + "relationships": { + "provider": { + "data": { + "type": "providers", + "id": str(hidden_provider.id), + } + } + }, + } + } + ), + content_type="application/vnd.api+json", + ) + + assert response.status_code == status.HTTP_400_BAD_REQUEST + assert not ProviderSecret.objects.filter(provider=hidden_provider).exists() + + def test_provider_secret_create_in_scope_provider_is_accepted( + self, + authenticated_client_rbac_limited, + aws_provider, + ): + response = authenticated_client_rbac_limited.post( + reverse("providersecret-list"), + data=json.dumps( + { + "data": { + "type": "provider-secrets", + "attributes": { + "name": "In scope secret", + "secret_type": ProviderSecret.TypeChoices.STATIC, + "secret": { + "aws_access_key_id": "visible-key", + "aws_secret_access_key": "visible-secret", + }, + }, + "relationships": { + "provider": { + "data": { + "type": "providers", + "id": str(aws_provider.id), + } + } + }, + } + } + ), + content_type="application/vnd.api+json", + ) + + assert response.status_code == status.HTTP_201_CREATED + assert ProviderSecret.objects.filter(provider=aws_provider).exists() + + def test_provider_secret_update_out_of_scope_returns_404( + self, + authenticated_client_rbac_limited, + hidden_provider_secret, + ): + response = authenticated_client_rbac_limited.patch( + reverse( + "providersecret-detail", + kwargs={"pk": hidden_provider_secret.id}, + ), + data=json.dumps( + { + "data": { + "type": "provider-secrets", + "id": str(hidden_provider_secret.id), + "attributes": {"name": "Updated hidden secret"}, + } + } + ), + content_type="application/vnd.api+json", + ) + + assert response.status_code == status.HTTP_404_NOT_FOUND + hidden_provider_secret.refresh_from_db() + assert hidden_provider_secret.name == "Hidden provider secret" + + def test_provider_secret_delete_out_of_scope_returns_404( + self, + authenticated_client_rbac_limited, + hidden_provider_secret, + ): + response = authenticated_client_rbac_limited.delete( + reverse( + "providersecret-detail", + kwargs={"pk": hidden_provider_secret.id}, + ) + ) + + assert response.status_code == status.HTTP_404_NOT_FOUND + assert ProviderSecret.objects.filter(id=hidden_provider_secret.id).exists() + + def test_provider_group_create_out_of_scope_provider_is_rejected( + self, + authenticated_client_rbac_limited, + aws_provider_pair, + ): + hidden_provider = aws_provider_pair[1] + response = authenticated_client_rbac_limited.post( + reverse("providergroup-list"), + data=json.dumps( + { + "data": { + "type": "provider-groups", + "attributes": {"name": "Out of scope group"}, + "relationships": { + "providers": { + "data": [ + { + "type": "providers", + "id": str(hidden_provider.id), + } + ] + } + }, + } + } + ), + content_type="application/vnd.api+json", + ) + + assert response.status_code == status.HTTP_400_BAD_REQUEST + assert not ProviderGroup.objects.filter(name="Out of scope group").exists() + + def test_provider_group_create_in_scope_provider_is_accepted( + self, + authenticated_client_rbac_limited, + aws_provider, + ): + response = authenticated_client_rbac_limited.post( + reverse("providergroup-list"), + data=json.dumps( + { + "data": { + "type": "provider-groups", + "attributes": {"name": "In scope group"}, + "relationships": { + "providers": { + "data": [ + { + "type": "providers", + "id": str(aws_provider.id), + } + ] + } + }, + } + } + ), + content_type="application/vnd.api+json", + ) + + assert response.status_code == status.HTTP_201_CREATED + provider_group = ProviderGroup.objects.get(name="In scope group") + assert set(provider_group.providers.all()) == {aws_provider} + + def test_provider_group_update_out_of_scope_provider_is_rejected( + self, + authenticated_client_rbac_limited, + limited_provider_group, + aws_provider_pair, + ): + visible_provider, hidden_provider = aws_provider_pair + response = authenticated_client_rbac_limited.patch( + reverse( + "providergroup-detail", + kwargs={"pk": limited_provider_group.id}, + ), + data=json.dumps( + { + "data": { + "type": "provider-groups", + "id": str(limited_provider_group.id), + "relationships": { + "providers": { + "data": [ + { + "type": "providers", + "id": str(hidden_provider.id), + } + ] + } + }, + } + } + ), + content_type="application/vnd.api+json", + ) + + assert response.status_code == status.HTTP_400_BAD_REQUEST + assert set(limited_provider_group.providers.all()) == {visible_provider} + + def test_provider_group_relationship_create_out_of_scope_provider_is_rejected( + self, + authenticated_client_rbac_limited, + limited_provider_group, + aws_provider_pair, + ): + hidden_provider = aws_provider_pair[1] + response = authenticated_client_rbac_limited.post( + reverse( + "provider_group-providers-relationship", + kwargs={"pk": limited_provider_group.id}, + ), + data={ + "data": [ + {"type": "providers", "id": str(hidden_provider.id)}, + ] + }, + content_type="application/vnd.api+json", + ) + + assert response.status_code == status.HTTP_400_BAD_REQUEST + assert not ProviderGroupMembership.objects.filter( + provider_group=limited_provider_group, + provider=hidden_provider, + ).exists() + + def test_provider_group_relationship_update_out_of_scope_provider_is_rejected( + self, + authenticated_client_rbac_limited, + limited_provider_group, + aws_provider_pair, + ): + visible_provider, hidden_provider = aws_provider_pair + response = authenticated_client_rbac_limited.patch( + reverse( + "provider_group-providers-relationship", + kwargs={"pk": limited_provider_group.id}, + ), + data={ + "data": [ + {"type": "providers", "id": str(hidden_provider.id)}, + ] + }, + content_type="application/vnd.api+json", + ) + + assert response.status_code == status.HTTP_400_BAD_REQUEST + assert set(limited_provider_group.providers.all()) == {visible_provider} + + def test_provider_group_relationship_create_in_scope_provider_is_accepted( + self, + authenticated_client_rbac_limited, + limited_provider_group, + aws_provider_pair, + ): + additional_provider = aws_provider_pair[1] + additional_group = ProviderGroup.objects.create( + tenant_id=additional_provider.tenant_id, + name="Additional visible group", + ) + ProviderGroupMembership.objects.create( + tenant_id=additional_provider.tenant_id, + provider_group=additional_group, + provider=additional_provider, + ) + RoleProviderGroupRelationship.objects.create( + tenant_id=additional_provider.tenant_id, + role=limited_provider_group.roles.get(), + provider_group=additional_group, + ) + + response = authenticated_client_rbac_limited.post( + reverse( + "provider_group-providers-relationship", + kwargs={"pk": limited_provider_group.id}, + ), + data={ + "data": [ + {"type": "providers", "id": str(additional_provider.id)}, + ] + }, + content_type="application/vnd.api+json", + ) + + assert response.status_code == status.HTTP_204_NO_CONTENT + assert ProviderGroupMembership.objects.filter( + provider_group=limited_provider_group, + provider=additional_provider, + ).exists() + + def test_provider_group_relationship_delete_out_of_scope_group_returns_404( + self, + authenticated_client_rbac_limited, + aws_provider_pair, + ): + hidden_provider = aws_provider_pair[1] + hidden_group = ProviderGroup.objects.create( + tenant_id=hidden_provider.tenant_id, + name="Unassigned provider group", + ) + ProviderGroupMembership.objects.create( + tenant_id=hidden_provider.tenant_id, + provider_group=hidden_group, + provider=hidden_provider, + ) + + response = authenticated_client_rbac_limited.delete( + reverse( + "provider_group-providers-relationship", + kwargs={"pk": hidden_group.id}, + ) + ) + + assert response.status_code == status.HTTP_404_NOT_FOUND + assert ProviderGroupMembership.objects.filter( + provider_group=hidden_group, + provider=hidden_provider, + ).exists() + + @patch("api.v1.views.Task.objects.get") + @patch("api.v1.views.delete_provider_task.delay") + def test_provider_delete_out_of_scope_returns_404( + self, + mock_delete_task, + mock_task_get, + authenticated_client_rbac_limited, + aws_provider_pair, + tasks_fixture, + ): + hidden_provider = aws_provider_pair[1] + prowler_task = tasks_fixture[0] + mock_delete_task.return_value.id = prowler_task.id + mock_task_get.return_value = prowler_task + + response = authenticated_client_rbac_limited.delete( + reverse("provider-detail", kwargs={"pk": hidden_provider.id}) + ) + + assert response.status_code == status.HTTP_404_NOT_FOUND + hidden_provider.refresh_from_db() + assert hidden_provider.is_deleted is False + mock_delete_task.assert_not_called() + mock_task_get.assert_not_called() + + @patch("api.v1.views.Task.objects.get") + @patch("api.v1.views.delete_provider_task.delay") + def test_provider_delete_in_scope_returns_202( + self, + mock_delete_task, + mock_task_get, + authenticated_client_rbac_limited, + aws_provider, + tasks_fixture, + ): + prowler_task = tasks_fixture[0] + mock_delete_task.return_value.id = prowler_task.id + mock_task_get.return_value = prowler_task + + response = authenticated_client_rbac_limited.delete( + reverse("provider-detail", kwargs={"pk": aws_provider.id}) + ) + + assert response.status_code == status.HTTP_202_ACCEPTED + mock_delete_task.assert_called_once_with( + provider_id=str(aws_provider.id), tenant_id=ANY + ) + mock_task_get.assert_called_once_with(id=prowler_task.id) + + @patch("api.v1.views.Task.objects.get") + @patch("api.v1.views.check_provider_connection_task.delay") + def test_provider_connection_out_of_scope_returns_404( + self, + mock_provider_connection, + mock_task_get, + authenticated_client_rbac_limited, + aws_provider_pair, + tasks_fixture, + ): + hidden_provider = aws_provider_pair[1] + prowler_task = tasks_fixture[0] + mock_provider_connection.return_value.id = prowler_task.id + mock_task_get.return_value = prowler_task + + response = authenticated_client_rbac_limited.post( + reverse("provider-connection", kwargs={"pk": hidden_provider.id}) + ) + + assert response.status_code == status.HTTP_404_NOT_FOUND + mock_provider_connection.assert_not_called() + mock_task_get.assert_not_called() + + @patch("api.v1.views.Task.objects.get") + @patch("api.v1.views.check_provider_connection_task.delay") + def test_provider_connection_in_scope_returns_202( + self, + mock_provider_connection, + mock_task_get, + authenticated_client_rbac_limited, + aws_provider, + tasks_fixture, + ): + prowler_task = tasks_fixture[0] + mock_provider_connection.return_value.id = prowler_task.id + mock_task_get.return_value = prowler_task + + response = authenticated_client_rbac_limited.post( + reverse("provider-connection", kwargs={"pk": aws_provider.id}) + ) + + assert response.status_code == status.HTTP_202_ACCEPTED + mock_provider_connection.assert_called_once_with( + provider_id=str(aws_provider.id), tenant_id=ANY + ) + mock_task_get.assert_called_once_with(id=prowler_task.id) + + @patch("api.v1.views.Task.objects.get") + @patch("api.v1.views.schedule_provider_scan") + def test_schedule_daily_out_of_scope_returns_404( + self, + mock_schedule_scan, + mock_task_get, + authenticated_client_rbac_limited, + aws_provider_pair, + tasks_fixture, + ): + hidden_provider = aws_provider_pair[1] + prowler_task = tasks_fixture[0] + mock_schedule_scan.return_value.id = prowler_task.id + mock_task_get.return_value = prowler_task + + response = authenticated_client_rbac_limited.post( + reverse("schedule-daily"), + data=json.dumps( + { + "data": { + "type": "daily-schedules", + "attributes": {"provider_id": str(hidden_provider.id)}, + } + } + ), + content_type="application/vnd.api+json", + ) + + assert response.wsgi_request.content_type == "application/vnd.api+json" + assert response.status_code == status.HTTP_404_NOT_FOUND + mock_schedule_scan.assert_not_called() + mock_task_get.assert_not_called() + + @patch("api.v1.views.Task.objects.get") + @patch("api.v1.views.schedule_provider_scan") + def test_schedule_daily_in_scope_returns_202( + self, + mock_schedule_scan, + mock_task_get, + authenticated_client_rbac_limited, + aws_provider, + tasks_fixture, + ): + prowler_task = tasks_fixture[0] + mock_schedule_scan.return_value.id = prowler_task.id + mock_task_get.return_value = prowler_task + + response = authenticated_client_rbac_limited.post( + reverse("schedule-daily"), + data=json.dumps( + { + "data": { + "type": "daily-schedules", + "attributes": {"provider_id": str(aws_provider.id)}, + } + } + ), + content_type="application/vnd.api+json", + ) + + assert response.wsgi_request.content_type == "application/vnd.api+json" + assert response.status_code == status.HTTP_202_ACCEPTED + mock_schedule_scan.assert_called_once_with(aws_provider) + mock_task_get.assert_called_once_with(id=prowler_task.id) + def test_integrations( self, authenticated_client_rbac_limited, integrations_fixture ): diff --git a/api/src/backend/api/v1/mixins.py b/api/src/backend/api/v1/mixins.py index 7645c92f4c..46328fc0fe 100644 --- a/api/src/backend/api/v1/mixins.py +++ b/api/src/backend/api/v1/mixins.py @@ -6,9 +6,11 @@ from api.exceptions import ( TaskNotFoundException, ) from api.models import Provider, StateChoices, Task +from api.rbac.permissions import get_providers from api.v1.serializers import TaskSerializer from django.http import QueryDict from django.urls import reverse +from django.utils.functional import cached_property from django_celery_results.models import TaskResult from rest_framework import status from rest_framework.exceptions import ValidationError @@ -33,6 +35,22 @@ class DisablePaginationMixin: return super().paginate_queryset(queryset) +class ProviderVisibilityMixin: + @cached_property + def provider_queryset(self): + if self.user_role.unlimited_visibility: + return Provider.objects.filter(tenant_id=self.request.tenant_id) + return get_providers(self.user_role) + + def get_provider_queryset(self): + return self.provider_queryset + + def get_serializer_context(self): + context = super().get_serializer_context() + context["provider_queryset"] = self.get_provider_queryset() + return context + + class PaginateByPkMixin: """ Mixin to paginate on a list of PKs (cheaper than heavy JOINs), diff --git a/api/src/backend/api/v1/serializers.py b/api/src/backend/api/v1/serializers.py index 9292a0c633..6a238a1aa5 100644 --- a/api/src/backend/api/v1/serializers.py +++ b/api/src/backend/api/v1/serializers.py @@ -124,6 +124,20 @@ class RLSSerializer(BaseModelSerializerV1): return super().create(validated_data) +class ScopedProviderFieldMixin: + provider_field_name = "provider" + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + provider_queryset = self.context.get("provider_queryset") + provider_field = self.fields.get(self.provider_field_name) + if provider_queryset is None or provider_field is None: + return + + related_field = getattr(provider_field, "child_relation", provider_field) + related_field.queryset = provider_queryset + + class StateEnumSerializerField(serializers.ChoiceField): def __init__(self, **kwargs): kwargs["choices"] = StateChoices.choices @@ -693,7 +707,10 @@ class MembershipIncludeSerializer(serializers.ModelSerializer): # Provider Groups -class ProviderGroupSerializer(RLSSerializer, BaseWriteSerializer): +class ProviderGroupSerializer( + ScopedProviderFieldMixin, RLSSerializer, BaseWriteSerializer +): + provider_field_name = "providers" providers = serializers.ResourceRelatedField( queryset=Provider.objects.all(), many=True, required=False ) @@ -851,9 +868,27 @@ class ProviderGroupMembershipSerializer(RLSSerializer, BaseWriteSerializer): help_text="List of resource identifier objects representing providers.", ) + def get_providers(self, validated_data): + provider_ids = {item["id"] for item in validated_data["providers"]} + provider_queryset = self.context.get("provider_queryset") + if provider_queryset is None: + provider_queryset = Provider.objects.filter( + tenant_id=self.context.get("tenant_id") + ) + + providers = list(provider_queryset.filter(id__in=provider_ids)) + if {provider.id for provider in providers} != provider_ids: + raise serializers.ValidationError( + { + "providers": ( + "One or more providers do not exist or are not accessible." + ) + } + ) + return providers + def create(self, validated_data): - provider_ids = [item["id"] for item in validated_data["providers"]] - providers = Provider.objects.filter(id__in=provider_ids) + providers = self.get_providers(validated_data) tenant_id = self.context.get("tenant_id") new_relationships = [ @@ -869,8 +904,7 @@ class ProviderGroupMembershipSerializer(RLSSerializer, BaseWriteSerializer): return self.context.get("provider_group") def update(self, instance, validated_data): - provider_ids = [item["id"] for item in validated_data["providers"]] - providers = Provider.objects.filter(id__in=provider_ids) + providers = self.get_providers(validated_data) tenant_id = self.context.get("tenant_id") instance.providers.clear() @@ -1109,7 +1143,9 @@ class ScanIncludeSerializer(RLSSerializer): } -class ScanCreateSerializer(RLSSerializer, BaseWriteSerializer): +class ScanCreateSerializer( + ScopedProviderFieldMixin, RLSSerializer, BaseWriteSerializer +): class Meta: model = Scan # TODO: add mutelist when implemented @@ -1974,7 +2010,9 @@ class ProviderSecretSerializer(RLSSerializer): ] -class ProviderSecretCreateSerializer(RLSSerializer, BaseWriteProviderSecretSerializer): +class ProviderSecretCreateSerializer( + ScopedProviderFieldMixin, RLSSerializer, BaseWriteProviderSecretSerializer +): secret = ProviderSecretField(write_only=True) class Meta: diff --git a/api/src/backend/api/v1/views.py b/api/src/backend/api/v1/views.py index db7b76f01a..075dd36829 100644 --- a/api/src/backend/api/v1/views.py +++ b/api/src/backend/api/v1/views.py @@ -146,6 +146,7 @@ from api.v1.mixins import ( JsonApiFilterMixin, PaginateByPkMixin, ProviderFilterParamsMixin, + ProviderVisibilityMixin, TaskManagementMixin, ) from api.v1.serializers import ( @@ -1632,7 +1633,7 @@ class TenantMembersViewSet(BaseTenantViewset): ), update=extend_schema(exclude=True), ) -class ProviderGroupViewSet(BaseRLSViewSet): +class ProviderGroupViewSet(ProviderVisibilityMixin, BaseRLSViewSet): queryset = ProviderGroup.objects.all() serializer_class = ProviderGroupSerializer filterset_class = ProviderGroupFilter @@ -1653,14 +1654,13 @@ class ProviderGroupViewSet(BaseRLSViewSet): self.required_permissions = [Permissions.MANAGE_PROVIDERS] def get_queryset(self): - user_roles = get_role(self.request.user, self.request.tenant_id) - # Check if any of the user's roles have UNLIMITED_VISIBILITY - if user_roles.unlimited_visibility: - # User has unlimited visibility, return all provider groups - return ProviderGroup.objects.prefetch_related("providers", "roles") - - # Collect provider groups associated with the user's roles - return user_roles.provider_groups.all().prefetch_related("providers", "roles") + if self.user_role.unlimited_visibility: + queryset = ProviderGroup.objects.filter(tenant_id=self.request.tenant_id) + else: + queryset = self.user_role.provider_groups.filter( + tenant_id=self.request.tenant_id + ) + return queryset.prefetch_related("providers", "roles") def get_serializer_class(self): if self.action == "create": @@ -1701,7 +1701,9 @@ class ProviderGroupViewSet(BaseRLSViewSet): }, ), ) -class ProviderGroupProvidersRelationshipView(RelationshipView, BaseRLSViewSet): +class ProviderGroupProvidersRelationshipView( + ProviderVisibilityMixin, RelationshipView, BaseRLSViewSet +): queryset = ProviderGroup.objects.all() serializer_class = ProviderGroupMembershipSerializer resource_name = "providers" @@ -1711,7 +1713,9 @@ class ProviderGroupProvidersRelationshipView(RelationshipView, BaseRLSViewSet): required_permissions = [Permissions.MANAGE_PROVIDERS] def get_queryset(self): - return ProviderGroup.objects.filter(tenant_id=self.request.tenant_id) + if self.user_role.unlimited_visibility: + return ProviderGroup.objects.filter(tenant_id=self.request.tenant_id) + return self.user_role.provider_groups.filter(tenant_id=self.request.tenant_id) def create(self, request, *args, **kwargs): provider_group = self.get_object() @@ -1733,6 +1737,7 @@ class ProviderGroupProvidersRelationshipView(RelationshipView, BaseRLSViewSet): data={"providers": request.data}, context={ "provider_group": provider_group, + "provider_queryset": self.get_provider_queryset(), "tenant_id": self.request.tenant_id, "request": request, }, @@ -1747,7 +1752,11 @@ class ProviderGroupProvidersRelationshipView(RelationshipView, BaseRLSViewSet): serializer = self.get_serializer( instance=provider_group, data={"providers": request.data}, - context={"tenant_id": self.request.tenant_id, "request": request}, + context={ + "provider_queryset": self.get_provider_queryset(), + "tenant_id": self.request.tenant_id, + "request": request, + }, ) serializer.is_valid(raise_exception=True) serializer.save() @@ -1864,7 +1873,7 @@ class ProviderViewSet(DisablePaginationMixin, BaseRLSViewSet): ) @action(detail=True, methods=["post"], url_name="connection") def connection(self, request, pk=None): - get_object_or_404(Provider, pk=pk) + self.get_object() with transaction.atomic(): task = check_provider_connection_task.delay( provider_id=pk, tenant_id=self.request.tenant_id @@ -1882,7 +1891,7 @@ class ProviderViewSet(DisablePaginationMixin, BaseRLSViewSet): ) def destroy(self, request, *args, pk=None, **kwargs): - provider = get_object_or_404(Provider, pk=pk) + provider = self.get_object() provider.is_deleted = True provider.save() task_name = f"scan-perform-scheduled-{pk}" @@ -2104,7 +2113,7 @@ class ProviderViewSet(DisablePaginationMixin, BaseRLSViewSet): ) @method_decorator(CACHE_DECORATOR, name="list") @method_decorator(CACHE_DECORATOR, name="retrieve") -class ScanViewSet(BaseRLSViewSet): +class ScanViewSet(ProviderVisibilityMixin, BaseRLSViewSet): queryset = Scan.objects.all() serializer_class = ScanSerializer http_method_names = ["get", "post", "patch"] @@ -2133,13 +2142,7 @@ class ScanViewSet(BaseRLSViewSet): self.required_permissions = [Permissions.MANAGE_SCANS] def get_queryset(self): - user_roles = get_role(self.request.user, self.request.tenant_id) - if user_roles.unlimited_visibility: - # User has unlimited visibility, return all scans - queryset = Scan.objects.filter(tenant_id=self.request.tenant_id) - else: - # User lacks permission, filter providers based on provider groups associated with the role - queryset = Scan.objects.filter(provider__in=get_providers(user_roles)) + queryset = Scan.objects.filter(provider__in=self.get_provider_queryset()) return queryset.select_related("provider", "task") def get_serializer_class(self): @@ -2737,6 +2740,7 @@ class ScanViewSet(BaseRLSViewSet): provider = Provider.objects.select_for_update().get( id=provider.id, tenant_id=self.request.tenant_id, + id__in=self.get_provider_queryset().values("id"), ) active_scan = get_active_provider_scan( self.request.tenant_id, provider.id @@ -4311,7 +4315,7 @@ class FindingViewSet(PaginateByPkMixin, BaseRLSViewSet): ) @method_decorator(CACHE_DECORATOR, name="list") @method_decorator(CACHE_DECORATOR, name="retrieve") -class ProviderSecretViewSet(BaseRLSViewSet): +class ProviderSecretViewSet(ProviderVisibilityMixin, BaseRLSViewSet): queryset = ProviderSecret.objects.all() serializer_class = ProviderSecretSerializer filterset_class = ProviderSecretFilter @@ -4327,7 +4331,7 @@ class ProviderSecretViewSet(BaseRLSViewSet): required_permissions = [Permissions.MANAGE_PROVIDERS] def get_queryset(self): - return ProviderSecret.objects.filter(tenant_id=self.request.tenant_id) + return ProviderSecret.objects.filter(provider__in=self.get_provider_queryset()) def get_serializer_class(self): if self.action == "create": @@ -6608,7 +6612,7 @@ class OverviewViewSet(ProviderFilterParamsMixin, BaseRLSViewSet): responses={202: OpenApiResponse(response=TaskSerializer)}, ) ) -class ScheduleViewSet(BaseRLSViewSet): +class ScheduleViewSet(ProviderVisibilityMixin, BaseRLSViewSet): # TODO: change to Schedule when implemented queryset = Task.objects.none() http_method_names = ["post"] @@ -6635,7 +6639,9 @@ class ScheduleViewSet(BaseRLSViewSet): serializer.is_valid(raise_exception=True) provider_id = serializer.validated_data["provider_id"] - provider_instance = get_object_or_404(Provider, pk=provider_id) + provider_instance = get_object_or_404( + self.get_provider_queryset(), pk=provider_id + ) with transaction.atomic(): task = schedule_provider_scan(provider_instance)