mirror of
https://github.com/prowler-cloud/prowler.git
synced 2026-10-04 02:04:06 +00:00
fix(api): respect provider group scope in provider actions (#12216)
This commit is contained in:
@@ -0,0 +1 @@
|
||||
Provider deletion and connection checks, scan creation, provider secrets, provider groups, and daily schedules now respect role provider-group visibility
|
||||
@@ -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
|
||||
):
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user