aboutsummaryrefslogtreecommitdiff
path: root/tests/managers/thl/test_survey.py
diff options
context:
space:
mode:
Diffstat (limited to 'tests/managers/thl/test_survey.py')
-rw-r--r--tests/managers/thl/test_survey.py96
1 files changed, 60 insertions, 36 deletions
diff --git a/tests/managers/thl/test_survey.py b/tests/managers/thl/test_survey.py
index 58c4577..e114b70 100644
--- a/tests/managers/thl/test_survey.py
+++ b/tests/managers/thl/test_survey.py
@@ -1,29 +1,41 @@
+from __future__ import annotations
+
import uuid
-from datetime import datetime, timezone
+from collections.abc import Callable
+from datetime import UTC, datetime
from decimal import Decimal
+from typing import TYPE_CHECKING
import pytest
-from generalresearch.models import Source
+from generalresearch.models.definitions import Source
from generalresearch.models.legacy.bucket import (
- SurveyEligibilityCriterion,
- TopNPlusBucket,
DurationSummary,
PayoutSummary,
+ SurveyEligibilityCriterion,
+ TopNPlusBucket,
)
from generalresearch.models.thl.profiling.user_question_answer import (
UserQuestionAnswer,
)
from generalresearch.models.thl.survey.model import (
Survey,
- SurveyStat,
SurveyCategoryModel,
SurveyEligibilityDefinition,
+ SurveyStat,
)
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.buyer import BuyerManager
+ from generalresearch.managers.thl.profiling.question import (
+ QuestionManager,
+ )
+ from generalresearch.managers.thl.profiling.uqa import UQAManager
+ from generalresearch.managers.thl.survey import SurveyManager, SurveyStatManager
+
@pytest.fixture(scope="session")
-def surveys_fixture():
+def surveys_fixture() -> list[Survey]:
return [
Survey(source=Source.TESTING, survey_id="a", buyer_code="buyer1"),
Survey(source=Source.TESTING, survey_id="b", buyer_code="buyer2"),
@@ -73,11 +85,13 @@ class TestSurvey:
def test(
self,
- delete_buyers_surveys,
- buyer_manager,
- survey_manager,
- surveys_fixture,
+ delete_buyers_surveys: Callable[..., None],
+ buyer_manager: BuyerManager,
+ survey_manager: SurveyManager,
+ surveys_fixture: list[Survey],
):
+ delete_buyers_surveys()
+
survey_manager.create_or_update(surveys_fixture)
survey_ids = {s.survey_id for s in surveys_fixture}
res = survey_manager.filter_by_natural_key(
@@ -98,7 +112,7 @@ class TestSurvey:
assert res2[0] == res[0]
assert len(res2) == len(surveys2)
- def test_category(self, survey_manager):
+ def test_category(self, survey_manager: SurveyManager):
survey1 = Survey(id=562289, survey_id="a", source=Source.TESTING)
survey2 = Survey(id=562290, survey_id="a", source=Source.TESTING)
categories = list(survey_manager.category_manager.categories.values())
@@ -110,8 +124,14 @@ class TestSurvey:
survey_manager.update_surveys_categories(surveys)
def test_survey_eligibility(
- self, survey_manager, upk_data, question_manager, uqa_manager
+ self,
+ survey_manager: SurveyManager,
+ upk_data: Callable[..., None],
+ question_manager: QuestionManager,
+ uqa_manager: UQAManager,
):
+ upk_data()
+
bucket = TopNPlusBucket(
id="c82cf98c578a43218334544ab376b00e",
contents=[],
@@ -161,9 +181,10 @@ class TestSurvey:
calc_answers={"i:adhoc_13126": ("3", "4")},
),
]
- uqad = dict()
+ uqad = {}
for uqa in uqas:
- for k, v in uqa.calc_answers.items():
+ assert uqa.calc_answers
+ for k in uqa.calc_answers:
if k in qualifying_questions:
uqad[k] = uqa
uqad[uqa.property_code] = uqa
@@ -174,9 +195,9 @@ class TestSurvey:
qs = sorted(qs, key=lambda x: x.importance.task_count if x.importance else 0)
# qd = {q.id: q for q in qs}
- q = [x for x in qs if x.ext_question_id == "i:adhoc_13126"][0]
+ q = next(x for x in qs if x.ext_question_id == "i:adhoc_13126")
q.explanation_template = "You have been diagnosed with: {answer}."
- q = [x for x in qs if x.ext_question_id == "gr:gender"][0]
+ q = next(x for x in qs if x.ext_question_id == "gr:gender")
q.explanation_template = "Your gender is {answer}."
ecs = []
@@ -205,10 +226,9 @@ class TestSurvey:
class TestSurveyStat:
def test(
self,
- delete_buyers_surveys,
surveystat_manager,
- survey_manager,
- surveys_fixture,
+ survey_manager: SurveyManager,
+ surveys_fixture: list[Survey],
):
survey_manager.create_or_update(surveys_fixture)
ss = [ssa, ssb]
@@ -234,7 +254,7 @@ class TestSurveyStat:
):
survey = surveys_fixture[0].model_copy()
surveys = []
- for idx in range(20_000):
+ for _ in range(20_000):
s = survey.model_copy()
s.survey_id = uuid.uuid4().hex
surveys.append(s)
@@ -251,14 +271,14 @@ class TestSurveyStat:
survey_stats.append(ss)
print(len(survey_stats))
print(survey_stats[12].natural_key, survey_stats[2000].natural_key)
- print(f"----a-----: {datetime.now().isoformat()}")
+ print(f"----a-----: {datetime.now(tz=UTC).isoformat()}")
res = surveystat_manager.update_or_create(survey_stats)
- print(f"----b-----: {datetime.now().isoformat()}")
+ print(f"----b-----: {datetime.now(tz=UTC).isoformat()}")
assert len(res) == 20_000
return
# 1,000 of the 20,000 are "new"
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
for s in ss[:1000]:
s.survey__survey_id = "b"
s.updated_at = now
@@ -269,22 +289,24 @@ class TestSurveyStat:
s.conv_beta = 20
s.updated_at = now
# and 1,000 don't change
- print(f"----c-----: {datetime.now().isoformat()}")
+ print(f"----c-----: {datetime.now(tz=UTC).isoformat()}")
res2 = surveystat_manager.update_or_create(ss)
- print(f"----d-----: {datetime.now().isoformat()}")
+ print(f"----d-----: {datetime.now(tz=UTC).isoformat()}")
assert len(res2) == 20_000
def test_ymsp(
self,
- delete_buyers_surveys,
- surveys_fixture,
- survey_manager,
- surveystat_manager,
+ delete_buyers_surveys: Callable[..., None],
+ surveys_fixture: list[Survey],
+ survey_manager: SurveyManager,
+ surveystat_manager: SurveyStatManager,
):
+ delete_buyers_surveys()
+
source = Source.TESTING
survey = surveys_fixture[0].model_copy()
surveys = []
- for idx in range(100):
+ for _ in range(100):
s = survey.model_copy()
s.survey_id = uuid.uuid4().hex
surveys.append(s)
@@ -298,14 +320,14 @@ class TestSurveyStat:
source=source, surveys=surveys, survey_stats=survey_stats
)
# UPDATE -------
- since = datetime.now(tz=timezone.utc)
+ since = datetime.now(tz=UTC)
print(f"{since=}")
# 10 survey disappear
surveys = surveys[10:]
# and 2 new ones are created
- for idx in range(2):
+ for _ in range(2):
s = survey.model_copy()
s.survey_id = uuid.uuid4().hex
surveys.append(s)
@@ -329,11 +351,13 @@ class TestSurveyStat:
def test_filter(
self,
- delete_buyers_surveys,
- surveys_fixture,
- survey_manager,
- surveystat_manager,
+ delete_buyers_surveys: Callable[..., None],
+ surveys_fixture: list[Survey],
+ survey_manager: SurveyManager,
+ surveystat_manager: SurveyStatManager,
):
+ delete_buyers_surveys()
+
surveys = []
survey = surveys_fixture[0].model_copy()
survey.source = Source.TESTING