Skip to content

Commit 26ea1e5

Browse files
committed
Add more filters
1 parent f3988c5 commit 26ea1e5

5 files changed

Lines changed: 75 additions & 2 deletions

File tree

cats/api/filters.py

Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,11 +2,23 @@
22

33
from cats.models import SimulationRun
44

5+
56
class SimulationFilter(filters.FilterSet):
67
status = filters.CharFilter(field_name="status")
78
user = filters.CharFilter(field_name="user")
89
created_after = filters.IsoDateTimeFilter(field_name="created_at", lookup_expr="gte")
910
created_before = filters.IsoDateTimeFilter(field_name="created_at", lookup_expr="lte")
11+
iterations = filters.NumberFilter(field_name="params__iterations", method="filter_int_params")
12+
cat_amount = filters.NumberFilter(field_name="params__cat_amount", method="filter_int_params")
13+
node_amount = filters.NumberFilter(field_name="params__node_amount", method="filter_int_params")
14+
15+
def filter_int_params(self, queryset, name, value):
16+
value = int(value)
17+
matching_filter = next((filter for filter in self.filters.values() if filter.field_name == name), None)
18+
if not matching_filter:
19+
return queryset
20+
lookup = getattr(matching_filter, "lookup_expr", "exact")
21+
return queryset.filter(**{f"{name}__{lookup}": value})
1022

1123
class Meta:
1224
model = SimulationRun
@@ -15,4 +27,7 @@ class Meta:
1527
'user',
1628
'created_after',
1729
'created_before',
30+
'iterations',
31+
'cat_amount',
32+
'node_amount',
1833
]

cats/api/serializers.py

Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,4 @@
1+
from decimal import Decimal
12
from rest_framework import serializers
23

34
from cats.models import SimulationResults, SimulationRun
@@ -88,9 +89,23 @@ def validate(self, data):
8889

8990

9091
class SimulationStatusSerializer(serializers.ModelSerializer):
92+
params = serializers.SerializerMethodField()
9193
class Meta:
9294
model = SimulationRun
9395
fields = ["id", "status", "created_at", "started_at", "finished_at", "params", "user"]
96+
97+
def get_params(self, obj):
98+
# Recursively convert any Decimal in JSON to float
99+
def convert(value):
100+
if isinstance(value, dict):
101+
return {k: convert(v) for k, v in value.items()}
102+
elif isinstance(value, list):
103+
return [convert(v) for v in value]
104+
elif isinstance(value, Decimal):
105+
return float(value)
106+
return value
107+
108+
return convert(obj.params)
94109

95110

96111
class SimulationErrorSerializer(serializers.ModelSerializer):

cats/api/views.py

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -24,7 +24,6 @@
2424
NOT_FAILED_RESPONSE = {"detail": "Simulation has not failed"}
2525
NOT_COMPLETED_RESPONSE = {"detail": "Simulation has not completed"}
2626

27-
2827
class SimulationStartView(APIView):
2928
permission_classes = [IsAuthenticated, IsOwnerOrAdmin]
3029

@@ -100,3 +99,8 @@ def get_queryset(self):
10099
if user.is_staff:
101100
return SimulationRun.objects.all()
102101
return SimulationRun.objects.filter(user=user)
102+
103+
def list(self, request, *args, **kwargs):
104+
queryset = self.filter_queryset(self.get_queryset())
105+
serializer = self.get_serializer(queryset, many=True)
106+
return Response(serializer.data)

cats/tests/conftest.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -33,7 +33,7 @@ def create_simulation_list(db,create_user, create_simulation):
3333
def _create_simulation_list(user=None, number_of_lists=5):
3434
if not user:
3535
user = create_user()
36-
return [create_simulation(user=user, params={"iterations":10*i}) for i in range(1,number_of_lists +1)]
36+
return [create_simulation(user=user, params={"iterations":10*i, "cat_amount":3*i, "node_amount": 5*i}) for i in range(1,number_of_lists +1)]
3737
return _create_simulation_list
3838

3939
@pytest.fixture

cats/tests/test_filters.py

Lines changed: 39 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -112,3 +112,42 @@ def test_filter_by_user(auth_client_with_refresh, create_user, create_superuser,
112112
sims = response4.data
113113
assert len(sims) == 5, "Length of response list is not 5"
114114
assert sims[0] == SimulationStatusSerializer(sim_list2[0]).data, "First Simulation of List response is different from first simulation of list in memory"
115+
116+
def test_filter_by_iterations(auth_client_with_refresh, create_user, create_simulation_list):
117+
user = create_user(email="test1@email.com",password="test1password")
118+
sim_list = create_simulation_list(user = user)
119+
120+
auth_client, _ = auth_client_with_refresh(user=user, password="test1password")
121+
url = reverse("simulation-list")
122+
response1 = auth_client.get(url, {"iterations": 10})
123+
124+
assert response1.status_code == 200, "response status code not 200"
125+
sims = response1.data
126+
assert len(sims) == 1, "Length of response list is not 1"
127+
assert sims[0] == SimulationStatusSerializer(sim_list[0]).data, "First Simulation of List response is different from first simulation of list in memory"
128+
129+
def test_filter_by_cat_amounts(auth_client_with_refresh, create_user, create_simulation_list):
130+
user = create_user(email="test1@email.com",password="test1password")
131+
sim_list = create_simulation_list(user = user)
132+
133+
auth_client, _ = auth_client_with_refresh(user=user, password="test1password")
134+
url = reverse("simulation-list")
135+
response1 = auth_client.get(url, {"cat_amount": 6})
136+
137+
assert response1.status_code == 200, "response status code not 200"
138+
sims = response1.data
139+
assert len(sims) == 1, "Length of response list is not 1"
140+
assert sims[0] == SimulationStatusSerializer(sim_list[1]).data, "First Simulation of List response is different from first simulation of list in memory"
141+
142+
def test_filter_by_node_amounts(auth_client_with_refresh, create_user, create_simulation_list):
143+
user = create_user(email="test1@email.com",password="test1password")
144+
sim_list = create_simulation_list(user = user)
145+
146+
auth_client, _ = auth_client_with_refresh(user=user, password="test1password")
147+
url = reverse("simulation-list")
148+
response1 = auth_client.get(url, {"node_amount": 10})
149+
150+
assert response1.status_code == 200, "response status code not 200"
151+
sims = response1.data
152+
assert len(sims) == 1, "Length of response list is not 1"
153+
assert sims[0] == SimulationStatusSerializer(sim_list[1]).data, "First Simulation of List response is different from first simulation of list in memory"

0 commit comments

Comments
 (0)