fix(api): respect provider group scope in provider actions (#12216)

This commit is contained in:
Adrián Peña
2026-07-30 10:06:57 +02:00
committed by GitHub
parent a57a507cee
commit ecf7ec8e85
5 changed files with 704 additions and 33 deletions
@@ -0,0 +1 @@
Provider deletion and connection checks, scan creation, provider secrets, provider groups, and daily schedules now respect role provider-group visibility
+608
View File
@@ -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
):
+18
View File
@@ -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),
+45 -7
View File
@@ -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:
+32 -26
View File
@@ -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)