diff options
Diffstat (limited to 'tests/models')
40 files changed, 412 insertions, 670 deletions
diff --git a/tests/models/innovate/test_question.py b/tests/models/innovate/test_question.py index b0c2964..b206177 100644 --- a/tests/models/innovate/test_question.py +++ b/tests/models/innovate/test_question.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from generalresearch.models import Source from generalresearch.models.innovate.question import ( InnovateQuestion, diff --git a/tests/models/legacy/test_offerwall_parse_response.py b/tests/models/legacy/test_offerwall_parse_response.py index b1c96ad..56ba077 100644 --- a/tests/models/legacy/test_offerwall_parse_response.py +++ b/tests/models/legacy/test_offerwall_parse_response.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import json from generalresearch.models import Source diff --git a/tests/models/legacy/test_profiling_questions.py b/tests/models/legacy/test_profiling_questions.py index 1afaa6b..6f781ae 100644 --- a/tests/models/legacy/test_profiling_questions.py +++ b/tests/models/legacy/test_profiling_questions.py @@ -1,7 +1,11 @@ +from __future__ import annotations + +from generalresearch.models.legacy.questions import UpkQuestionResponse + + class TestUpkQuestionResponse: def test_init(self): - from generalresearch.models.legacy.questions import UpkQuestionResponse s = ( '{"status": "success", "count": 7, "questions": [{"selector": "SL", "validation": {"patterns": [{' diff --git a/tests/models/legacy/test_user_question_answer_in.py b/tests/models/legacy/test_user_question_answer_in.py index 313862c..3fdaa05 100644 --- a/tests/models/legacy/test_user_question_answer_in.py +++ b/tests/models/legacy/test_user_question_answer_in.py @@ -1,9 +1,22 @@ +from __future__ import annotations + import json +from collections.abc import Callable +from datetime import datetime from decimal import Decimal from uuid import uuid4 import pytest +from generalresearch.managers.thl.user_manager.user_manager import UserManager +from generalresearch.models import Source +from generalresearch.models.legacy.questions import ( + UserQuestionAnswers, +) +from generalresearch.models.thl.product import Product +from generalresearch.models.thl.session import Session, Wall +from generalresearch.models.thl.user import User + class TestUserQuestionAnswers: """This is for the GRS POST submission that may contain multiple @@ -15,21 +28,11 @@ class TestUserQuestionAnswers: def test_json_init( self, - product_manager: ProductManager, - user_manager, - session_manager, - wall_manager, user_factory: Callable[..., User], product: Product, - session_factory, - utc_hour_ago, + session_factory: Callable[..., Session], + utc_hour_ago: datetime, ): - from generalresearch.models import Source - from generalresearch.models.legacy.questions import ( - UserQuestionAnswers, - ) - from generalresearch.models.thl.session import Session, Wall - from generalresearch.models.thl.user import User u: User = user_factory(product=product) @@ -61,14 +64,7 @@ class TestUserQuestionAnswers: def test_simple_validation_errors( self, - product_manager: ProductManager, - user_manager, - session_manager, - wall_manager, ): - from generalresearch.models.legacy.questions import ( - UserQuestionAnswers, - ) with pytest.raises(ValueError): UserQuestionAnswers.model_validate( @@ -118,7 +114,7 @@ class TestUserQuestionAnswers: with pytest.raises(ValueError): answers = [ - {"question_id": uuid4().hex, "answer": ["a"]} for i in range(101) + {"question_id": uuid4().hex, "answer": ["a"]} for _ in range(101) ] UserQuestionAnswers.model_validate( { @@ -143,9 +139,6 @@ class TestUserQuestionAnswers: # TODO: depending on if or how many of these types of errors actually # occur, we could get fancy and just drop one of them. I don't # think this is worth exploring yet unless we see if it's a problem. - from generalresearch.models.legacy.questions import ( - UserQuestionAnswers, - ) consistent_qid = uuid4().hex with pytest.raises(ValueError) as cm: @@ -165,11 +158,11 @@ class TestUserQuestionAnswers: def test_allow_answer_failures_silent( self, - user_manager, + user_manager: UserManager, product: Product, user_factory: Callable[..., User], - utc_hour_ago, - session_factory, + utc_hour_ago: datetime, + session_factory: Callable[..., Session], ): """ There are many instances where suppliers may be submitting answers @@ -177,11 +170,6 @@ class TestUserQuestionAnswers: that one QuestionAnswerIn without "loosing" any of the other QuestionAnswerIn items that they provided. """ - from generalresearch.models.legacy.questions import ( - UserQuestionAnswers, - ) - from generalresearch.models.thl.session import Session, Wall - from generalresearch.models.thl.user import User u: User = user_factory(product=product) @@ -286,7 +274,7 @@ class TestUserQuestionAnswerIn: UserQuestionAnswerIn, ) - answer = [uuid4().hex[:6] for i in range(11)] + answer = [uuid4().hex[:6] for _ in range(11)] with pytest.raises(ValueError) as cm: UserQuestionAnswerIn.model_validate( {"question_id": uuid4().hex, "answer": answer} @@ -298,7 +286,7 @@ class TestUserQuestionAnswerIn: UserQuestionAnswerIn, ) - answer = ["aaa" for i in range(5)] + answer = ["aaa" for _ in range(5)] with pytest.raises(ValueError): UserQuestionAnswerIn.model_validate( {"question_id": uuid4().hex, "answer": answer} diff --git a/tests/models/morning/test.py b/tests/models/morning/test.py index 7474766..c1141fb 100644 --- a/tests/models/morning/test.py +++ b/tests/models/morning/test.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from datetime import UTC, datetime from generalresearch.models.morning.question import MorningQuestion diff --git a/tests/models/network/test_mtr.py b/tests/models/network/test_mtr.py index 840a773..7f8a736 100644 --- a/tests/models/network/test_mtr.py +++ b/tests/models/network/test_mtr.py @@ -1,12 +1,15 @@ +from __future__ import annotations + import faker +from generalresearch.managers.network.tool_run import ToolRunManager from generalresearch.models.network.mtr.execute import execute_mtr from generalresearch.models.network.tool_run import ToolClass, ToolName fake = faker.Faker() -def test_execute_mtr(toolrun_manager): +def test_execute_mtr(toolrun_manager: ToolRunManager): ip = "65.19.129.53" run = execute_mtr(ip=ip, report_cycles=3) diff --git a/tests/models/network/test_nmap.py b/tests/models/network/test_nmap.py index a135a13..5e9f4d0 100644 --- a/tests/models/network/test_nmap.py +++ b/tests/models/network/test_nmap.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import subprocess import faker @@ -5,8 +7,8 @@ import faker from generalresearch.managers.network.tool_run import ToolRunManager from generalresearch.models.network.definitions import IPProtocol from generalresearch.models.network.nmap.execute import execute_nmap -from generalresearch.models.network.nmap.result import PortState -from generalresearch.models.network.tool_run import ToolClass, ToolName +from generalresearch.models.network.nmap.result import NmapResult, PortState +from generalresearch.models.network.tool_run import NmapRun, Status, ToolClass, ToolName fake = faker.Faker() @@ -18,10 +20,13 @@ def resolve(host: str): def test_execute_nmap_scanme(toolrun_manager: ToolRunManager): ip = resolve("scanme.nmap.org") - run = execute_nmap(ip=ip, top_ports=None, ports="20-30", enable_advanced=False) + run: NmapRun = execute_nmap( + ip=ip, top_ports=None, ports="20-30", enable_advanced=False + ) assert run.tool_name == ToolName.NMAP assert run.tool_class == ToolClass.PORT_SCAN assert run.ip == ip + assert isinstance(run.parsed, NmapResult) result = run.parsed port22 = result._port_index[(IPProtocol.TCP, 22)] diff --git a/tests/models/network/test_nmap_parser.py b/tests/models/network/test_nmap_parser.py index 7822380..473a63f 100644 --- a/tests/models/network/test_nmap_parser.py +++ b/tests/models/network/test_nmap_parser.py @@ -1,8 +1,14 @@ +from __future__ import annotations + import os import pytest from generalresearch.models.network.nmap.parser import parse_nmap_xml +from generalresearch.models.network.nmap.result import ( + NmapResult, + NmapTrace, +) @pytest.fixture @@ -13,9 +19,11 @@ def nmap_raw_output_2(request) -> str: return data -def test_nmap_xml_parser(nmap_raw_output, nmap_raw_output_2): - n = parse_nmap_xml(nmap_raw_output) +def test_nmap_xml_parser(nmap_raw_output: str, nmap_raw_output_2: str): + n: NmapResult = parse_nmap_xml(nmap_raw_output) assert n.tcp_open_ports == [61232] + + assert isinstance(n.trace, NmapTrace) assert len(n.trace.hops) == 18 n = parse_nmap_xml(nmap_raw_output_2) diff --git a/tests/models/network/test_rdns.py b/tests/models/network/test_rdns.py index 5c3b024..1a15a28 100644 --- a/tests/models/network/test_rdns.py +++ b/tests/models/network/test_rdns.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import faker from generalresearch.managers.network.tool_run import ToolRunManager diff --git a/tests/models/precision/__init__.py b/tests/models/precision/__init__.py index 8006fa3..e69de29 100644 --- a/tests/models/precision/__init__.py +++ b/tests/models/precision/__init__.py @@ -1,115 +0,0 @@ -survey_json = { - "cpi": "1.44", - "country_isos": "ca", - "language_isos": "eng", - "country_iso": "ca", - "language_iso": "eng", - "buyer_id": "7047", - "bid_loi": 1200, - "bid_ir": 0.45, - "source": "e", - "used_question_ids": ["age", "country_iso", "gender", "gender_1"], - "survey_id": "0000", - "group_id": "633473", - "status": "open", - "name": "beauty survey", - "survey_guid": "c7f375c5077d4c6c8209ff0b539d7183", - "category_id": "-1", - "global_conversion": None, - "desired_count": 96, - "achieved_count": 0, - "allowed_devices": "1,2,3", - "entry_link": "https://www.opinionetwork.com/survey/entry.aspx?mid=[%MID%]&project=633473&key=%%key%%", - "excluded_surveys": "470358,633286", - "quotas": [ - { - "name": "25-34,Male,Quebec", - "id": "2324110", - "guid": "23b5760d24994bc08de451b3e62e77c7", - "status": "open", - "desired_count": 48, - "achieved_count": 0, - "termination_count": 0, - "overquota_count": 0, - "condition_hashes": ["b41e1a3", "bc89ee8", "4124366", "9f32c61"], - }, - { - "name": "25-34,Female,Quebec", - "id": "2324111", - "guid": "0706f1a88d7e4f11ad847c03012e68d2", - "status": "open", - "desired_count": 48, - "achieved_count": 0, - "termination_count": 4, - "overquota_count": 0, - "condition_hashes": ["b41e1a3", "0cdc304", "500af2c", "9f32c61"], - }, - ], - "conditions": { - "b41e1a3": { - "logical_operator": "OR", - "value_type": 1, - "negate": False, - "question_id": "country_iso", - "values": ["ca"], - "criterion_hash": "b41e1a3", - "value_len": 1, - "sizeof": 2, - }, - "bc89ee8": { - "logical_operator": "OR", - "value_type": 1, - "negate": False, - "question_id": "gender", - "values": ["male"], - "criterion_hash": "bc89ee8", - "value_len": 1, - "sizeof": 4, - }, - "4124366": { - "logical_operator": "OR", - "value_type": 1, - "negate": False, - "question_id": "gender_1", - "values": ["male"], - "criterion_hash": "4124366", - "value_len": 1, - "sizeof": 4, - }, - "9f32c61": { - "logical_operator": "OR", - "value_type": 1, - "negate": False, - "question_id": "age", - "values": ["25", "26", "27", "28", "29", "30", "31", "32", "33", "34"], - "criterion_hash": "9f32c61", - "value_len": 10, - "sizeof": 20, - }, - "0cdc304": { - "logical_operator": "OR", - "value_type": 1, - "negate": False, - "question_id": "gender", - "values": ["female"], - "criterion_hash": "0cdc304", - "value_len": 1, - "sizeof": 6, - }, - "500af2c": { - "logical_operator": "OR", - "value_type": 1, - "negate": False, - "question_id": "gender_1", - "values": ["female"], - "criterion_hash": "500af2c", - "value_len": 1, - "sizeof": 6, - }, - }, - "expected_end_date": "2024-06-28T10:40:33.000000Z", - "created": None, - "updated": None, - "is_live": True, - "all_hashes": ["0cdc304", "b41e1a3", "9f32c61", "bc89ee8", "4124366", "500af2c"], -} diff --git a/tests/models/precision/test_survey.py b/tests/models/precision/test_survey.py index ff2d6d1..4d671f2 100644 --- a/tests/models/precision/test_survey.py +++ b/tests/models/precision/test_survey.py @@ -1,10 +1,15 @@ -class TestPrecisionQuota: +from __future__ import annotations + +from typing import Any + +from generalresearch.models.precision import PrecisionStatus +from generalresearch.models.precision.survey import PrecisionSurvey - def test_quota_passes(self): - from generalresearch.models.precision.survey import PrecisionSurvey - from tests.models.precision import survey_json - s = PrecisionSurvey.model_validate(survey_json) +class TestPrecisionQuota: + + def test_quota_passes(self, precision_survey_json: dict[str, Any]): + s = PrecisionSurvey.model_validate(precision_survey_json) q = s.quotas[0] ce = {k: True for k in ["b41e1a3", "bc89ee8", "4124366", "9f32c61"]} assert q.matches(ce) @@ -16,12 +21,9 @@ class TestPrecisionQuota: assert not q.matches(ce) assert not q.matches({}) - def test_quota_passes_closed(self): - from generalresearch.models.precision import PrecisionStatus - from generalresearch.models.precision.survey import PrecisionSurvey - from tests.models.precision import survey_json + def test_quota_passes_closed(self, precision_survey_json: dict[str, Any]): - s = PrecisionSurvey.model_validate(survey_json) + s = PrecisionSurvey.model_validate(precision_survey_json) q = s.quotas[0] q.status = PrecisionStatus.CLOSED ce = {k: True for k in ["b41e1a3", "bc89ee8", "4124366", "9f32c61"]} @@ -32,20 +34,15 @@ class TestPrecisionQuota: class TestPrecisionSurvey: - def test_passes(self): - from generalresearch.models.precision.survey import PrecisionSurvey - from tests.models.precision import survey_json + def test_passes(self, precision_survey_json: dict[str, Any]): - s = PrecisionSurvey.model_validate(survey_json) + s = PrecisionSurvey.model_validate(precision_survey_json) ce = {k: True for k in ["b41e1a3", "bc89ee8", "4124366", "9f32c61"]} assert s.determine_eligibility(ce) - def test_elig_closed_quota(self): - from generalresearch.models.precision import PrecisionStatus - from generalresearch.models.precision.survey import PrecisionSurvey - from tests.models.precision import survey_json + def test_elig_closed_quota(self, precision_survey_json: dict[str, Any]): - s = PrecisionSurvey.model_validate(survey_json) + s = PrecisionSurvey.model_validate(precision_survey_json) ce = {k: True for k in ["b41e1a3", "bc89ee8", "4124366", "9f32c61"]} q = s.quotas[0] q.status = PrecisionStatus.CLOSED @@ -57,12 +54,9 @@ class TestPrecisionSurvey: # Now me match an open quota and dont match the closed quota, so we should be eligible assert s.determine_eligibility(ce) - def test_passes_sp(self): - from generalresearch.models.precision import PrecisionStatus - from generalresearch.models.precision.survey import PrecisionSurvey - from tests.models.precision import survey_json + def test_passes_sp(self, precision_survey_json: dict[str, Any]): - s = PrecisionSurvey.model_validate(survey_json) + s = PrecisionSurvey.model_validate(precision_survey_json) ce = {k: True for k in ["b41e1a3", "bc89ee8", "4124366", "9f32c61"]} passes, hashes = s.determine_eligibility_soft(ce) diff --git a/tests/models/prodege/test_survey_participation.py b/tests/models/prodege/test_survey_participation.py index e1ba9ab..10ce884 100644 --- a/tests/models/prodege/test_survey_participation.py +++ b/tests/models/prodege/test_survey_participation.py @@ -1,14 +1,17 @@ +from __future__ import annotations + from datetime import UTC, datetime, timedelta +from generalresearch.models.prodege import ProdegePastParticipationType +from generalresearch.models.prodege.survey import ( + ProdegePastParticipation, + ProdegeUserPastParticipation, +) + class TestProdegeParticipation: def test_exclude(self): - from generalresearch.models.prodege import ProdegePastParticipationType - from generalresearch.models.prodege.survey import ( - ProdegePastParticipation, - ProdegeUserPastParticipation, - ) now = datetime.now(tz=UTC) pp = ProdegePastParticipation.from_api( @@ -84,10 +87,6 @@ class TestProdegeParticipation: assert not pp.is_eligible(upps) def test_include(self): - from generalresearch.models.prodege.survey import ( - ProdegePastParticipation, - ProdegeUserPastParticipation, - ) now = datetime.now(tz=UTC) pp = ProdegePastParticipation.from_api( diff --git a/tests/models/spectrum/test_question.py b/tests/models/spectrum/test_question.py index 57d260d..a44286d 100644 --- a/tests/models/spectrum/test_question.py +++ b/tests/models/spectrum/test_question.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from datetime import UTC, datetime from generalresearch.models import Source @@ -32,6 +34,7 @@ class TestSpectrumQuestion: "mod_on": 1706557247467, } q = SpectrumQuestion.from_api(example_1, "us", "eng") + assert isinstance(q, SpectrumQuestion) expected_q = SpectrumQuestion( question_id="213", @@ -72,6 +75,8 @@ class TestSpectrumQuestion: "mod_on": 1706557249817, } q = SpectrumQuestion.from_api(example_2, "us", "eng") + assert isinstance(q, SpectrumQuestion) + expected_q = SpectrumQuestion( question_id="211", country_iso="us", diff --git a/tests/models/spectrum/test_survey.py b/tests/models/spectrum/test_survey.py index 7365c7e..7ddd407 100644 --- a/tests/models/spectrum/test_survey.py +++ b/tests/models/spectrum/test_survey.py @@ -1,15 +1,25 @@ +from __future__ import annotations + from datetime import UTC, datetime from decimal import Decimal +from generalresearch.models import ( + LogicalOperator, + Source, + TaskCalculationType, +) +from generalresearch.models.spectrum import SpectrumStatus +from generalresearch.models.spectrum.survey import ( + SpectrumCondition, + SpectrumQuota, + SpectrumSurvey, +) +from generalresearch.models.thl.survey.condition import ConditionValueType + class TestSpectrumCondition: def test_condition_create(self): - from generalresearch.models import LogicalOperator - from generalresearch.models.spectrum.survey import ( - SpectrumCondition, - ) - from generalresearch.models.thl.survey.condition import ConditionValueType c = SpectrumCondition.from_api( { @@ -64,10 +74,6 @@ class TestSpectrumCondition: class TestSpectrumQuota: def test_quota_create(self): - from generalresearch.models.spectrum.survey import ( - SpectrumCondition, - SpectrumQuota, - ) d = { "quota_id": "a846b545-4449-4d76-93a2-f8ebdf6e711e", @@ -84,9 +90,6 @@ class TestSpectrumQuota: assert q.is_open def test_quota_passes(self): - from generalresearch.models.spectrum.survey import ( - SpectrumQuota, - ) q = SpectrumQuota(remaining_count=57, condition_hashes=["a"]) assert q.passes({"a": True}) @@ -103,9 +106,6 @@ class TestSpectrumQuota: assert not q.passes({"a": True}) def test_quota_passes_soft(self): - from generalresearch.models.spectrum.survey import ( - SpectrumQuota, - ) q = SpectrumQuota(remaining_count=57, condition_hashes=["a", "b", "c"]) # Pass if we match all @@ -122,18 +122,6 @@ class TestSpectrumQuota: class TestSpectrumSurvey: def test_survey_create(self): - from generalresearch.models import ( - LogicalOperator, - Source, - TaskCalculationType, - ) - from generalresearch.models.spectrum import SpectrumStatus - from generalresearch.models.spectrum.survey import ( - SpectrumCondition, - SpectrumQuota, - SpectrumSurvey, - ) - from generalresearch.models.thl.survey.condition import ConditionValueType # Note: d is the raw response after calling SpectrumAPI.preprocess_survey() on it! d = { @@ -202,6 +190,8 @@ class TestSpectrumSurvey: "exclusion_period": 0, } s = SpectrumSurvey.from_api(d) + assert isinstance(s, SpectrumSurvey) + expected_survey = SpectrumSurvey( cpi=Decimal("1.20000"), country_isos=["fr"], @@ -303,6 +293,8 @@ class TestSpectrumSurvey: "exclusion_period": 0, } s = SpectrumSurvey.from_api(d) + assert isinstance(s, SpectrumSurvey) + assert {"212", "1202", "214"} == s.used_question_ids assert s.is_live assert s.is_open @@ -345,6 +337,8 @@ class TestSpectrumSurvey: "exclusion_period": 0, } s = SpectrumSurvey.from_api(d) + assert isinstance(s, SpectrumSurvey) + s.qualifications = ["a", "b", "c"] s.quotas = [ SpectrumQuota(remaining_count=10, condition_hashes=["a", "b"]), diff --git a/tests/models/spectrum/test_survey_manager.py b/tests/models/spectrum/test_survey_manager.py index ce26c44..11dc01f 100644 --- a/tests/models/spectrum/test_survey_manager.py +++ b/tests/models/spectrum/test_survey_manager.py @@ -1,69 +1,32 @@ -import copy +from __future__ import annotations + import logging from datetime import UTC, datetime from decimal import Decimal +from typing import Any from pymysql import IntegrityError -logger = logging.getLogger() +from generalresearch.config import is_debug +from generalresearch.managers.spectrum.survey import ( + SpectrumSurveyManager, +) +from generalresearch.sql_helper import SqlHelper -example_survey_api_response = { - "survey_id": 29333264, - "survey_name": "#29333264", - "survey_status": 22, - "field_end_date": datetime(2024, 5, 23, 18, 18, 31, tzinfo=UTC), - "category": "Exciting New", - "category_code": 232, - "crtd_on": datetime(2024, 5, 20, 17, 48, 13, tzinfo=UTC), - "mod_on": datetime(2024, 5, 20, 18, 18, 31, tzinfo=UTC), - "soft_launch": False, - "click_balancing": 0, - "price_type": 1, - "pii": False, - "buyer_message": "", - "buyer_id": 4726, - "incl_excl": 0, - "cpi": Decimal("1.20"), - "last_complete_date": None, - "project_last_complete_date": None, - "quotas": [ - { - "quota_id": "c2bc961e-4f26-4223-b409-ebe9165cfdf5", - "quantities": {"currently_open": 491, "remaining": 495, "achieved": 0}, - "criteria": [ - { - "qualification_code": 214, - "range_sets": [{"units": 311, "to": 64, "from": 18}], - } - ], - } - ], - "qualifications": [ - { - "range_sets": [{"units": 311, "to": 64, "from": 18}], - "qualification_code": 212, - }, - {"condition_codes": ["111", "117", "112"], "qualification_code": 1202}, - ], - "country_iso": "fr", - "language_iso": "fre", - "bid_ir": 0.4, - "bid_loi": 600, - "overall_ir": None, - "overall_loi": None, - "last_block_ir": None, - "last_block_loi": None, - "survey_exclusions": set(), - "exclusion_period": 0, -} +logger = logging.getLogger() class TestSpectrumSurvey: - def test_survey_create(self, settings, spectrum_manager, spectrum_rw): + def test_survey_create( + self, + spectrum_survey_manager: SpectrumSurveyManager, + spectrum_rw: SqlHelper, + spectrum_api_survey_json: dict[str, Any], + ): from generalresearch.models.spectrum.survey import SpectrumSurvey - assert settings.debug, "CRITICAL: Do not run this on production." + assert is_debug(), "CRITICAL: Do not run this on production." now = datetime.now(tz=UTC) spectrum_rw.execute_sql_query( @@ -73,24 +36,29 @@ class TestSpectrumSurvey: commit=True, ) - d = example_survey_api_response.copy() - s = SpectrumSurvey.from_api(d) - spectrum_manager.create(s) + s = SpectrumSurvey.from_api(spectrum_api_survey_json) + assert isinstance(s, SpectrumSurvey) + spectrum_survey_manager.create(s) - surveys = spectrum_manager.get_survey_library(updated_since=now) + surveys = spectrum_survey_manager.get_survey_library(updated_since=now) assert len(surveys) == 1 assert "29333264" == surveys[0].survey_id assert s.is_unchanged(surveys[0]) try: - spectrum_manager.create(s) + spectrum_survey_manager.create(s) except IntegrityError as e: print(e.args) - def test_survey_update(self, settings, spectrum_manager, spectrum_rw): + def test_survey_update( + self, + spectrum_survey_manager: SpectrumSurveyManager, + spectrum_rw: SqlHelper, + spectrum_api_survey_json: dict[str, Any], + ): from generalresearch.models.spectrum.survey import SpectrumSurvey - assert settings.debug, "CRITICAL: Do not run this on production." + assert is_debug(), "CRITICAL: Do not run this on production." now = datetime.now(tz=UTC) spectrum_rw.execute_sql_query( @@ -100,14 +68,13 @@ class TestSpectrumSurvey: """, commit=True, ) - d = copy.deepcopy(example_survey_api_response) - s = SpectrumSurvey.from_api(d) - print(s) + s = SpectrumSurvey.from_api(spectrum_api_survey_json) + assert isinstance(s, SpectrumSurvey) - spectrum_manager.create(s) + spectrum_survey_manager.create(s) s.cpi = Decimal("0.50") - spectrum_manager.update([s]) - surveys = spectrum_manager.get_survey_library(updated_since=now) + spectrum_survey_manager.update([s]) + surveys = spectrum_survey_manager.get_survey_library(updated_since=now) assert len(surveys) == 1 assert "29333264" == surveys[0].survey_id assert Decimal("0.50") == surveys[0].cpi @@ -122,8 +89,8 @@ class TestSpectrumSurvey: s.bid_loi = None s.overall_loi = 1000 s.last_block_loi = 1000 - spectrum_manager.update([s]) - surveys = spectrum_manager.get_survey_library(updated_since=now) + spectrum_survey_manager.update([s]) + surveys = spectrum_survey_manager.get_survey_library(updated_since=now) assert 600 == surveys[0].bid_loi assert 1000 == surveys[0].overall_loi assert 1000 == surveys[0].last_block_loi diff --git a/tests/models/test_currency.py b/tests/models/test_currency.py index 9bc2216..e946126 100644 --- a/tests/models/test_currency.py +++ b/tests/models/test_currency.py @@ -3,6 +3,8 @@ functionality is the same, but pasting here so the tests are in the correct spot... """ +from __future__ import annotations + from decimal import Decimal from random import randint diff --git a/tests/models/test_device.py b/tests/models/test_device.py index bf72c81..8e1251a 100644 --- a/tests/models/test_device.py +++ b/tests/models/test_device.py @@ -1,3 +1,5 @@ +from __future__ import annotations + iphone_ua_string = ( "Mozilla/5.0 (iPhone; CPU iPhone OS 5_1 like Mac OS X) AppleWebKit/534.46 (KHTML, like Gecko) " "Version/5.1 Mobile/9B179 Safari/7534.48.3" @@ -13,10 +15,12 @@ chromebook_ua_string = ( ) +from generalresearch.models import DeviceType +from generalresearch.models.device import parse_device_from_useragent + + class TestDeviceUA: def test_device_ua(self): - from generalresearch.models import DeviceType - from generalresearch.models.device import parse_device_from_useragent assert parse_device_from_useragent(iphone_ua_string) == DeviceType.MOBILE assert parse_device_from_useragent(ipad_ua_string) == DeviceType.TABLET diff --git a/tests/models/test_finance.py b/tests/models/test_finance.py index f84d0b6..72f4f4d 100644 --- a/tests/models/test_finance.py +++ b/tests/models/test_finance.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from collections.abc import Callable from datetime import UTC, datetime, timedelta from itertools import product as iter_product @@ -25,14 +27,13 @@ from generalresearch.models.thl.finance import ( POPFinancial, ProductBalances, ) +from generalresearch.models.thl.ledger import LedgerAccount from generalresearch.models.thl.product import Product from generalresearch.models.thl.session import Session from generalresearch.models.thl.user import User +from generalresearch.pg_helper import PostgresConfig from test_utils.incite.collections.conftest import ledger_collection from test_utils.incite.mergers.conftest import pop_ledger_merge -from test_utils.managers.ledger.conftest import ( - session_with_tx_factory: Callable[..., None], -) fake = Faker() @@ -210,6 +211,8 @@ class TestProductBalanceInitialize: # Confirm the @property computed fields show up in openapi. I don't # know how to do that yet... so this is check to confirm they're # known computed fields for now + + assert isinstance(instance, ProductBalances) computed_fields = list(instance.model_computed_fields.keys()) assert "payout" in computed_fields assert "adjustment" in computed_fields @@ -665,17 +668,18 @@ class TestProductFinanceData: def test_base( self, - product: product: Product, + product: Product, user_factory: Callable[..., User], start: datetime, duration: timedelta, - thl_lm: ThlLedgerManager, + thl_ledger_manager: ThlLedgerManager, + session_with_tx_factory: Callable[..., None], ): # -- Build & Setup # assert ledger_collection.start is None # assert ledger_collection.offset is None - u: User = user_factory(product=product: Product, created=ledger_collection.start) + u: User = user_factory(product=product, created=ledger_collection.start) for item in ledger_collection.items: @@ -699,7 +703,7 @@ class TestProductFinanceData: item_finishes.sort(reverse=True) # -- - account = thl_lm.get_account_or_create_bp_wallet(product=u.product) + account = thl_ledger_manager.get_account_or_create_bp_wallet(product=u.product) ddf = pop_ledger_merge.ddf( force_rr_latest=False, @@ -748,7 +752,7 @@ class TestPOPFinancialData: ledger_collection: LedgerDFCollection, pop_ledger_merge: PopLedgerMerge, user_factory: Callable[..., User], - product: product: Product, + product: Product, start: datetime, duration: timedelta, create_main_accounts: Callable[..., None], @@ -791,7 +795,7 @@ class TestPOPFinancialData: last_item_finish = item_finishes[0] accounts = [] - for user in users: + for _ in users: account = thl_lm.get_account_or_create_bp_wallet(product=u.product) accounts.append(account) account_ids = [a.uuid for a in accounts] @@ -808,6 +812,7 @@ class TestPOPFinancialData: ("time_idx", "<", last_item_finish), ], ) + df: pd.DataFrame = client_no_amm.compute(collections=ddf, sync=True) df = df.groupby([pd.Grouper(key="time_idx", freq="D"), "account_id"]).sum() @@ -846,16 +851,15 @@ class TestBusinessBalanceData: ledger_collection: LedgerDFCollection, pop_ledger_merge: PopLedgerMerge, user_factory: Callable[..., User], - product: product: Product, + product: Product, create_main_accounts: Callable[..., None], thl_lm: ThlLedgerManager, thl_web_rr: PostgresConfig, delete_df_collection: Callable[..., None], delete_ledger_db: Callable[..., None], session_with_tx_factory: Callable[..., Session], - rm_ledger_collection, + rm_ledger_collection: Callable[..., None], ): - from generalresearch.models.thl.ledger import LedgerAccount delete_ledger_db() create_main_accounts() @@ -863,7 +867,7 @@ class TestBusinessBalanceData: rm_ledger_collection() for _ in range(5): - u: User = user_factory(product=product: Product, created=ledger_collection.start) + u: User = user_factory(product=product, created=ledger_collection.start) for item in ledger_collection.items: item_time = fake.date_time_between( diff --git a/tests/models/thl/question/test_question_info.py b/tests/models/thl/question/test_question_info.py index b619fc3..af8d2b9 100644 --- a/tests/models/thl/question/test_question_info.py +++ b/tests/models/thl/question/test_question_info.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from generalresearch.models.thl.profiling.upk_property import ( ProfilingInfo, UpkProperty, @@ -6,140 +8,9 @@ from generalresearch.models.thl.profiling.upk_property import ( class TestQuestionInfo: - def test_init(self): + def test_init(self, profiling_info_json: str): - s = ( - '[{"property_label": "hispanic", "cardinality": "*", "prop_type": "i", "country_iso": "us", ' - '"property_id": "05170ae296ab49178a075cab2a2073a6", "item_id": "7911ec1468b146ee870951f8ae9cbac1", ' - '"item_label": "panamanian", "gold_standard": 1, "options": [{"id": "c358c11e72c74fa2880358f1d4be85ab", ' - '"label": "not_hispanic"}, {"id": "b1d6c475770849bc8e0200054975dc9c", "label": "yes_hispanic"}, ' - '{"id": "bd1eb44495d84b029e107c188003c2bd", "label": "other_hispanic"}, ' - '{"id": "f290ad5e75bf4f4ea94dc847f57c1bd3", "label": "mexican"}, ' - '{"id": "49f50f2801bd415ea353063bfc02d252", "label": "puerto_rican"}, ' - '{"id": "dcbe005e522f4b10928773926601f8bf", "label": "cuban"}, ' - '{"id": "467ef8ddb7ac4edb88ba9ef817cbb7e9", "label": "salvadoran"}, ' - '{"id": "3c98e7250707403cba2f4dc7b877c963", "label": "dominican"}, ' - '{"id": "981ee77f6d6742609825ef54fea824a8", "label": "guatemalan"}, ' - '{"id": "81c8057b809245a7ae1b8a867ea6c91e", "label": "colombian"}, ' - '{"id": "513656d5f9e249fa955c3b527d483b93", "label": "honduran"}, ' - '{"id": "afc8cddd0c7b4581bea24ccd64db3446", "label": "ecuadorian"}, ' - '{"id": "61f34b36e80747a89d85e1eb17536f84", "label": "argentinian"}, ' - '{"id": "5330cfa681d44aa8ade3a6d0ea198e44", "label": "peruvian"}, ' - '{"id": "e7bceaffd76e486596205d8545019448", "label": "nicaraguan"}, ' - '{"id": "b7bbb2ebf8424714962e6c4f43275985", "label": "spanish"}, ' - '{"id": "8bf539785e7a487892a2f97e52b1932d", "label": "venezuelan"}, ' - '{"id": "7911ec1468b146ee870951f8ae9cbac1", "label": "panamanian"}], "category": [{"id": ' - '"4fd8381d5a1c4409ab007ca254ced084", "label": "Demographic", "path": "/Demographic", ' - '"adwords_vertical_id": null}]}, {"property_label": "ethnic_group", "cardinality": "*", "prop_type": ' - '"i", "country_iso": "us", "property_id": "15070958225d4132b7f6674fcfc979f6", "item_id": ' - '"64b7114cf08143949e3bcc3d00a5d8a0", "item_label": "other_ethnicity", "gold_standard": 1, "options": [{' - '"id": "a72e97f4055e4014a22bee4632cbf573", "label": "caucasians"}, ' - '{"id": "4760353bc0654e46a928ba697b102735", "label": "black_or_african_american"}, ' - '{"id": "20ff0a2969fa4656bbda5c3e0874e63b", "label": "asian"}, ' - '{"id": "107e0a79e6b94b74926c44e70faf3793", "label": "native_hawaiian_or_other_pacific_islander"}, ' - '{"id": "900fa12691d5458c8665bf468f1c98c1", "label": "native_americans"}, ' - '{"id": "64b7114cf08143949e3bcc3d00a5d8a0", "label": "other_ethnicity"}], "category": [{"id": ' - '"4fd8381d5a1c4409ab007ca254ced084", "label": "Demographic", "path": "/Demographic", ' - '"adwords_vertical_id": null}]}, {"property_label": "educational_attainment", "cardinality": "?", ' - '"prop_type": "i", "country_iso": "us", "property_id": "2637783d4b2b4075b93e2a156e16e1d8", "item_id": ' - '"934e7b81d6744a1baa31bbc51f0965d5", "item_label": "other_education", "gold_standard": 1, "options": [{' - '"id": "df35ef9e474b4bf9af520aa86630202d", "label": "3rd_grade_completion"}, ' - '{"id": "83763370a1064bd5ba76d1b68c4b8a23", "label": "8th_grade_completion"}, ' - '{"id": "f0c25a0670c340bc9250099dcce50957", "label": "not_high_school_graduate"}, ' - '{"id": "02ff74c872bd458983a83847e1a9f8fd", "label": "high_school_completion"}, ' - '{"id": "ba8beb807d56441f8fea9b490ed7561c", "label": "vocational_program_completion"}, ' - '{"id": "65373a5f348a410c923e079ddbb58e9b", "label": "some_college_completion"}, ' - '{"id": "2d15d96df85d4cc7b6f58911fdc8d5e2", "label": "associate_academic_degree_completion"}, ' - '{"id": "497b1fedec464151b063cd5367643ffa", "label": "bachelors_degree_completion"}, ' - '{"id": "295133068ac84424ae75e973dc9f2a78", "label": "some_graduate_completion"}, ' - '{"id": "e64f874faeff4062a5aa72ac483b4b9f", "label": "masters_degree_completion"}, ' - '{"id": "cbaec19a636d476385fb8e7842b044f5", "label": "doctorate_degree_completion"}, ' - '{"id": "934e7b81d6744a1baa31bbc51f0965d5", "label": "other_education"}], "category": [{"id": ' - '"4fd8381d5a1c4409ab007ca254ced084", "label": "Demographic", "path": "/Demographic", ' - '"adwords_vertical_id": null}]}, {"property_label": "household_spoken_language", "cardinality": "*", ' - '"prop_type": "i", "country_iso": "us", "property_id": "5a844571073d482a96853a0594859a51", "item_id": ' - '"62b39c1de141422896ad4ab3c4318209", "item_label": "dut", "gold_standard": 1, "options": [{"id": ' - '"f65cd57b79d14f0f8460761ce41ec173", "label": "ara"}, {"id": "6d49de1f8f394216821310abd29392d9", ' - '"label": "zho"}, {"id": "be6dc23c2bf34c3f81e96ddace22800d", "label": "eng"}, ' - '{"id": "ddc81f28752d47a3b1c1f3b8b01a9b07", "label": "fre"}, {"id": "2dbb67b29bd34e0eb630b1b8385542ca", ' - '"label": "ger"}, {"id": "a747f96952fc4b9d97edeeee5120091b", "label": "hat"}, ' - '{"id": "7144b04a3219433baac86273677551fa", "label": "hin"}, {"id": "e07ff3e82c7149eaab7ea2b39ee6a6dc", ' - '"label": "ita"}, {"id": "b681eff81975432ebfb9f5cc22dedaa3", "label": "jpn"}, ' - '{"id": "5cb20440a8f64c9ca62fb49c1e80cdef", "label": "kor"}, {"id": "171c4b77d4204bc6ac0c2b81e38a10ff", ' - '"label": "pan"}, {"id": "8c3ec18e6b6c4a55a00dd6052e8e84fb", "label": "pol"}, ' - '{"id": "3ce074d81d384dd5b96f1fb48f87bf01", "label": "por"}, {"id": "6138dc951990458fa88a666f6ddd907b", ' - '"label": "rus"}, {"id": "e66e5ecc07df4ebaa546e0b436f034bd", "label": "spa"}, ' - '{"id": "5a981b3d2f0d402a96dd2d0392ec2fcb", "label": "tgl"}, {"id": "b446251bd211403487806c4d0a904981", ' - '"label": "vie"}, {"id": "92fb3ee337374e2db875fb23f52eed46", "label": "xxx"}, ' - '{"id": "8b1f590f12f24cc1924d7bdcbe82081e", "label": "ind"}, {"id": "bf3f4be556a34ff4b836420149fd2037", ' - '"label": "tur"}, {"id": "87ca815c43ba4e7f98cbca98821aa508", "label": "zul"}, ' - '{"id": "0adbf915a7a64d67a87bb3ce5d39ca54", "label": "may"}, {"id": "62b39c1de141422896ad4ab3c4318209", ' - '"label": "dut"}], "category": [{"id": "4fd8381d5a1c4409ab007ca254ced084", "label": "Demographic", ' - '"path": "/Demographic", "adwords_vertical_id": null}]}, {"property_label": "gender", "cardinality": ' - '"?", "prop_type": "i", "country_iso": "us", "property_id": "73175402104741549f21de2071556cd7", ' - '"item_id": "093593e316344cd3a0ac73669fca8048", "item_label": "other_gender", "gold_standard": 1, ' - '"options": [{"id": "b9fc5ea07f3a4252a792fd4a49e7b52b", "label": "male"}, ' - '{"id": "9fdb8e5e18474a0b84a0262c21e17b56", "label": "female"}, ' - '{"id": "093593e316344cd3a0ac73669fca8048", "label": "other_gender"}], "category": [{"id": ' - '"4fd8381d5a1c4409ab007ca254ced084", "label": "Demographic", "path": "/Demographic", ' - '"adwords_vertical_id": null}]}, {"property_label": "age_in_years", "cardinality": "?", "prop_type": ' - '"n", "country_iso": "us", "property_id": "94f7379437874076b345d76642d4ce6d", "item_id": null, ' - '"item_label": null, "gold_standard": 1, "category": [{"id": "4fd8381d5a1c4409ab007ca254ced084", ' - '"label": "Demographic", "path": "/Demographic", "adwords_vertical_id": null}]}, {"property_label": ' - '"children_age_gender", "cardinality": "*", "prop_type": "i", "country_iso": "us", "property_id": ' - '"e926142fcea94b9cbbe13dc7891e1e7f", "item_id": "b7b8074e95334b008e8958ccb0a204f1", "item_label": ' - '"female_18", "gold_standard": 1, "options": [{"id": "16a6448ec24c48d4993d78ebee33f9b4", ' - '"label": "male_under_1"}, {"id": "809c04cb2e3b4a3bbd8077ab62cdc220", "label": "female_under_1"}, ' - '{"id": "295e05bb6a0843bc998890b24c99841e", "label": "no_children"}, ' - '{"id": "142cb948d98c4ae8b0ef2ef10978e023", "label": "male_0"}, ' - '{"id": "5a5c1b0e9abc48a98b3bc5f817d6e9d0", "label": "male_1"}, ' - '{"id": "286b1a9afb884bdfb676dbb855479d1e", "label": "male_2"}, ' - '{"id": "942ca3cda699453093df8cbabb890607", "label": "male_3"}, ' - '{"id": "995818d432f643ec8dd17e0809b24b56", "label": "male_4"}, ' - '{"id": "f38f8b57f25f4cdea0f270297a1e7a5c", "label": "male_5"}, ' - '{"id": "975df709e6d140d1a470db35023c432d", "label": "male_6"}, ' - '{"id": "f60bd89bbe0f4e92b90bccbc500467c2", "label": "male_7"}, ' - '{"id": "6714ceb3ed5042c0b605f00b06814207", "label": "male_8"}, ' - '{"id": "c03c2f8271d443cf9df380e84b4dea4c", "label": "male_9"}, ' - '{"id": "11690ee0f5a54cb794f7ddd010d74fa2", "label": "male_10"}, ' - '{"id": "17bef9a9d14b4197b2c5609fa94b0642", "label": "male_11"}, ' - '{"id": "e79c8338fe28454f89ccc78daf6f409a", "label": "male_12"}, ' - '{"id": "3a4f87acb3fa41f4ae08dfe2858238c1", "label": "male_13"}, ' - '{"id": "36ffb79d8b7840a7a8cb8d63bbc8df59", "label": "male_14"}, ' - '{"id": "1401a508f9664347aee927f6ec5b0a40", "label": "male_15"}, ' - '{"id": "6e0943c5ec4a4f75869eb195e3eafa50", "label": "male_16"}, ' - '{"id": "47d4b27b7b5242758a9fff13d3d324cf", "label": "male_17"}, ' - '{"id": "9ce886459dd44c9395eb77e1386ab181", "label": "female_0"}, ' - '{"id": "6499ccbf990d4be5b686aec1c7353fd8", "label": "female_1"}, ' - '{"id": "d85ceaa39f6d492abfc8da49acfd14f2", "label": "female_2"}, ' - '{"id": "18edb45c138e451d8cb428aefbb80f9c", "label": "female_3"}, ' - '{"id": "bac6f006ed9f4ccf85f48e91e99fdfd1", "label": "female_4"}, ' - '{"id": "5a6a1a8ad00c4ce8be52dcb267b034ff", "label": "female_5"}, ' - '{"id": "6bff0acbf6364c94ad89507bcd5f4f45", "label": "female_6"}, ' - '{"id": "d0d56a0a6b6f4516a366a2ce139b4411", "label": "female_7"}, ' - '{"id": "bda6028468044b659843e2bef4db2175", "label": "female_8"}, ' - '{"id": "dbb6d50325464032b456357b1a6e5e9c", "label": "female_9"}, ' - '{"id": "b87a93d7dc1348edac5e771684d63fb8", "label": "female_10"}, ' - '{"id": "11449d0d98f14e27ba47de40b18921d7", "label": "female_11"}, ' - '{"id": "16156501e97b4263962cbbb743840292", "label": "female_12"}, ' - '{"id": "04ee971c89a345cc8141a45bce96050c", "label": "female_13"}, ' - '{"id": "e818d310bfbc4faba4355e5d2ed49d4f", "label": "female_14"}, ' - '{"id": "440d25e078924ba0973163153c417ed6", "label": "female_15"}, ' - '{"id": "78ff804cc9b441c5a524bd91e3d1f8bf", "label": "female_16"}, ' - '{"id": "4b04d804d7d84786b2b1c22e4ed440f5", "label": "female_17"}, ' - '{"id": "28bc848cd3ff44c3893c76bfc9bc0c4e", "label": "male_18"}, ' - '{"id": "b7b8074e95334b008e8958ccb0a204f1", "label": "female_18"}], "category": [{"id": ' - '"e18ba6e9d51e482cbb19acf2e6f505ce", "label": "Parenting", "path": "/People & Society/Family & ' - 'Relationships/Family/Parenting", "adwords_vertical_id": "58"}]}, {"property_label": "home_postal_code", ' - '"cardinality": "?", "prop_type": "x", "country_iso": "us", "property_id": ' - '"f3b32ebe78014fbeb1ed6ff77d6338bf", "item_id": null, "item_label": null, "gold_standard": 1, ' - '"category": [{"id": "4fd8381d5a1c4409ab007ca254ced084", "label": "Demographic", "path": "/Demographic", ' - '"adwords_vertical_id": null}]}, {"property_label": "household_income", "cardinality": "?", "prop_type": ' - '"n", "country_iso": "us", "property_id": "ff5b1d4501d5478f98de8c90ef996ac1", "item_id": null, ' - '"item_label": null, "gold_standard": 1, "category": [{"id": "4fd8381d5a1c4409ab007ca254ced084", ' - '"label": "Demographic", "path": "/Demographic", "adwords_vertical_id": null}]}]' - ) - instance_list = ProfilingInfo.validate_json(s) + instance_list = ProfilingInfo.validate_json(profiling_info_json) assert isinstance(instance_list, list) for i in instance_list: diff --git a/tests/models/thl/question/test_user_info.py b/tests/models/thl/question/test_user_info.py index 0bbbc78..5410d35 100644 --- a/tests/models/thl/question/test_user_info.py +++ b/tests/models/thl/question/test_user_info.py @@ -1,32 +1,11 @@ +from __future__ import annotations + from generalresearch.models.thl.profiling.user_info import UserInfo class TestUserInfo: - def test_init(self): + def test_init(self, profiling_user_info_json: str): - s = ( - '{"user_profile_knowledge": [], "marketplace_profile_knowledge": [{"source": "d", "question_id": ' - '"1", "answer": ["1"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": "pr", ' - '"question_id": "3", "answer": ["1"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": ' - '"h", "question_id": "60", "answer": ["58"], "created": "2023-11-07T16:41:05.234096Z"}, ' - '{"source": "c", "question_id": "43", "answer": ["1"], "created": "2023-11-07T16:41:05.234096Z"}, ' - '{"source": "s", "question_id": "211", "answer": ["111"], "created": ' - '"2023-11-07T16:41:05.234096Z"}, {"source": "s", "question_id": "1843", "answer": ["111"], ' - '"created": "2023-11-07T16:41:05.234096Z"}, {"source": "h", "question_id": "13959", "answer": [' - '"244155"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": "c", "question_id": "33092", ' - '"answer": ["1"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": "c", "question_id": "gender", ' - '"answer": ["10682"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": "e", "question_id": ' - '"gender", "answer": ["male"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": "f", ' - '"question_id": "gender", "answer": ["male"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": ' - '"i", "question_id": "gender", "answer": ["1"], "created": "2023-11-07T16:41:05.234096Z"}, ' - '{"source": "c", "question_id": "137510", "answer": ["1"], "created": "2023-11-07T16:41:05.234096Z"}, ' - '{"source": "m", "question_id": "gender", "answer": ["1"], "created": ' - '"2023-11-07T16:41:05.234096Z"}, {"source": "o", "question_id": "gender", "answer": ["male"], ' - '"created": "2023-11-07T16:41:05.234096Z"}, {"source": "c", "question_id": "gender_plus", "answer": [' - '"7657644"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": "i", "question_id": ' - '"gender_plus", "answer": ["1"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": "c", ' - '"question_id": "income_level", "answer": ["9071"], "created": "2023-11-07T16:41:05.234096Z"}]}' - ) - instance = UserInfo.model_validate_json(s) + instance = UserInfo.model_validate_json(profiling_user_info_json) assert isinstance(instance, UserInfo) diff --git a/tests/models/thl/test_adjustments.py b/tests/models/thl/test_adjustments.py index 91e5316..c5b3f6b 100644 --- a/tests/models/thl/test_adjustments.py +++ b/tests/models/thl/test_adjustments.py @@ -1,9 +1,13 @@ +from __future__ import annotations + from collections.abc import Callable from datetime import UTC, datetime, timedelta from decimal import Decimal import pytest +from generalresearch.managers.thl.session import SessionManager +from generalresearch.managers.thl.wall import WallManager from generalresearch.models import Source from generalresearch.models.thl.product import Product from generalresearch.models.thl.session import ( @@ -30,7 +34,7 @@ class TestProductAdjustments: @pytest.mark.parametrize("payout", [".6", "1", "1.8", "2", "500.0000"]) def test_determine_bp_payment_no_rounding( - self, product_factory: Callable[..., Product], payout + self, product_factory: Callable[..., Product], payout: str ): p1 = product_factory(commission_pct=Decimal("0.05")) res = p1.determine_bp_payment(thl_net=Decimal(payout)) @@ -39,7 +43,7 @@ class TestProductAdjustments: @pytest.mark.parametrize("payout", [".01", ".05", ".5"]) def test_determine_bp_payment_rounding( - self, product_factory: Callable[..., Product], payout + self, product_factory: Callable[..., Product], payout: str ): p1 = product_factory(commission_pct=Decimal("0.05")) res = p1.determine_bp_payment(thl_net=Decimal(payout)) @@ -73,7 +77,10 @@ class TestSessionAdjustments: class TestAdjustments: def test_finish_with_status( - self, session_factory: Callable[..., Session], user: User, session_manager + self, + session_factory: Callable[..., Session], + user: User, + session_manager: SessionManager, ): # Completed Session with 2 wall events s1 = session_factory( @@ -85,6 +92,7 @@ class TestAdjustments: ) status, status_code_1 = s1.determine_session_status() + assert isinstance(user.product, Product) payout = user.product.determine_bp_payment(Decimal(1)) session_manager.finish_with_status( session=s1, @@ -97,7 +105,10 @@ class TestAdjustments: assert Decimal("0.95") == payout def test_never_adjusted( - self, session_factory: Callable[..., Session], user: User, session_manager + self, + session_factory: Callable[..., Session], + user: User, + session_manager: SessionManager, ): s1 = session_factory( user=user, @@ -130,8 +141,8 @@ class TestAdjustments: self, session_factory: Callable[..., Session], user: User, - session_manager, - wall_manager, + session_manager: SessionManager, + wall_manager: WallManager, ): # Completed Session with 2 wall events s1 = session_factory( @@ -174,13 +185,14 @@ class TestAdjustments: # Because the Product doesn't have the Wallet mode enabled, the # user_payout fields should always be None + assert isinstance(user.product, Product) assert not user.product.user_wallet_config.enabled assert s1.adjusted_user_payout is None def test_adjustment_session_values( self, - wall_manager, - session_manager, + wall_manager: WallManager, + session_manager: SessionManager, session_factory: Callable[..., Session], user: User, ): @@ -218,13 +230,14 @@ class TestAdjustments: # Because the Product doesn't have the Wallet mode enabled, the # user_payout fields should always be None + assert isinstance(user.product, Product) assert not user.product.user_wallet_config.enabled assert s1.adjusted_user_payout is None def test_double_adjustment_session_values( self, - wall_manager, - session_manager, + wall_manager: WallManager, + session_manager: SessionManager, session_factory: Callable[..., Session], user: User, ): @@ -276,8 +289,8 @@ class TestAdjustments: def test_double_adjustment_sm_vs_db_values( self, - wall_manager, - session_manager, + wall_manager: WallManager, + session_manager: SessionManager, session_factory: Callable[..., Session], user: User, ): @@ -343,8 +356,8 @@ class TestAdjustments: def test_double_adjustment_double_completes( self, - wall_manager, - session_manager, + wall_manager: WallManager, + session_manager: SessionManager, session_factory: Callable[..., Session], user: User, ): @@ -419,8 +432,8 @@ class TestAdjustments: self, session_factory: Callable[..., Session], user: User, - session_manager, - wall_manager, + session_manager: SessionManager, + wall_manager: WallManager, utc_hour_ago: datetime, ): s1 = session_factory( @@ -435,6 +448,7 @@ class TestAdjustments: assert status == Status.COMPLETE thl_net = Decimal(sum(w.cpi for w in s1.wall_events if w.is_visible_complete())) + assert isinstance(user.product, Product) payout = user.product.determine_bp_payment(thl_net=thl_net) session_manager.finish_with_status( @@ -459,7 +473,7 @@ class TestAdjustments: assert Status.FAIL == new_status assert Decimal(0) == new_payout - assert isinstance(user.product: Product, Product) + assert isinstance(user.product, Product) assert not user.product.user_wallet_config.enabled assert new_user_payout is None @@ -560,6 +574,7 @@ class TestAdjustments: new_status, new_payout, new_user_payout = s1.determine_new_status_and_payouts() assert Status.COMPLETE == new_status assert Decimal("0.95") == new_payout + assert isinstance(user.product, Product) assert not user.product.user_wallet_config.enabled # assert Decimal("0.48") == new_user_payout assert new_user_payout is None @@ -588,6 +603,7 @@ class TestAdjustments: status, status_code_1 = s1.determine_session_status() thl_net = Decimal(sum(w.cpi for w in s1.wall_events if w.is_visible_complete())) + assert isinstance(user.product, Product) payout = user.product.determine_bp_payment(thl_net=thl_net) s1.update( status=status, @@ -624,7 +640,10 @@ class TestAdjustments: assert s1.adjusted_user_payout is None def test_complete_to_fail_to_complete_adj1( - self, user, session_factory, utc_hour_ago + self, + user: User, + session_factory: Callable[..., Session], + utc_hour_ago: datetime, ): # Same as test_complete_to_fail_to_complete_adj but in opposite order s1 = session_factory( @@ -640,6 +659,7 @@ class TestAdjustments: status, status_code_1 = s1.determine_session_status() thl_net = Decimal(sum(w.cpi for w in s1.wall_events if w.is_visible_complete())) + assert isinstance(user.product, Product) payout = user.product.determine_bp_payment(thl_net) s1.update( status=status, @@ -658,6 +678,7 @@ class TestAdjustments: s1.adjust_status() assert SessionAdjustedStatus.ADJUSTED_TO_FAIL == s1.adjusted_status assert Decimal(0) == s1.adjusted_payout + assert isinstance(user.product, Product) assert not user.product.user_wallet_config.enabled # assert Decimal(0) == s.adjusted_user_payout assert s1.adjusted_user_payout is None @@ -702,6 +723,7 @@ class TestAdjustments: s1.adjust_status() assert SessionAdjustedStatus.ADJUSTED_TO_COMPLETE == s1.adjusted_status assert Decimal("1.90") == s1.adjusted_payout + assert isinstance(user.product, Product) assert not user.product.user_wallet_config.enabled # assert Decimal("0.95") == s1.adjusted_user_payout assert s1.adjusted_user_payout is None diff --git a/tests/models/thl/test_bucket.py b/tests/models/thl/test_bucket.py index 0aa5843..8d2f728 100644 --- a/tests/models/thl/test_bucket.py +++ b/tests/models/thl/test_bucket.py @@ -1,14 +1,17 @@ +from __future__ import annotations + from datetime import timedelta from decimal import Decimal import pytest from pydantic import ValidationError +from generalresearch.models.legacy.bucket import Bucket + class TestBucket: def test_raises_payout(self): - from generalresearch.models.legacy.bucket import Bucket with pytest.raises(expected_exception=ValidationError) as e: Bucket(user_payout_min=123) @@ -27,7 +30,6 @@ class TestBucket: assert "user_payout_min should be <= user_payout_max" in str(e.value) def test_raises_loi(self): - from generalresearch.models.legacy.bucket import Bucket with pytest.raises(expected_exception=ValidationError) as e: Bucket(loi_min=123) @@ -63,7 +65,6 @@ class TestBucket: assert "loi_q1 should be <= loi_q2" in str(e.value) def test_parse_1(self): - from generalresearch.models.legacy.bucket import Bucket b1 = Bucket.parse_from_offerwall({"payout": {"min": 123}}) b_exp = Bucket( @@ -180,7 +181,6 @@ class TestBucket: assert b_exp == b4 def test_parse_3(self): - from generalresearch.models.legacy.bucket import Bucket b1 = Bucket.parse_from_offerwall({"payout": 123}) b_exp = Bucket( diff --git a/tests/models/thl/test_buyer.py b/tests/models/thl/test_buyer.py index eebb828..02093e2 100644 --- a/tests/models/thl/test_buyer.py +++ b/tests/models/thl/test_buyer.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from generalresearch.models import Source from generalresearch.models.thl.survey.buyer import BuyerCountryStat diff --git a/tests/models/thl/test_contest/test_contest.py b/tests/models/thl/test_contest/test_contest.py index acb501c..e1053f4 100644 --- a/tests/models/thl/test_contest/test_contest.py +++ b/tests/models/thl/test_contest/test_contest.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from collections.abc import Callable import pytest diff --git a/tests/models/thl/test_contest/test_leaderboard_contest.py b/tests/models/thl/test_contest/test_leaderboard_contest.py index 5bab060..99cfb37 100644 --- a/tests/models/thl/test_contest/test_leaderboard_contest.py +++ b/tests/models/thl/test_contest/test_leaderboard_contest.py @@ -1,10 +1,14 @@ +from __future__ import annotations + from datetime import UTC from uuid import uuid4 import pytest +from redis import Redis from generalresearch.currency import USDCent from generalresearch.managers.leaderboard.manager import LeaderboardManager +from generalresearch.managers.thl.user_manager.user_manager import UserManager from generalresearch.models.thl.contest import ContestPrize from generalresearch.models.thl.contest.definitions import ( ContestPrizeKind, @@ -18,6 +22,7 @@ from generalresearch.models.thl.contest.utils import ( ) from generalresearch.models.thl.leaderboard import LeaderboardRow from generalresearch.models.thl.product import Product +from generalresearch.models.thl.user import User from tests.models.thl.test_contest.test_contest import TestContest @@ -25,7 +30,7 @@ class TestLeaderboardContest(TestContest): @pytest.fixture def leaderboard_contest( - self, product: product: Product, thl_redis, user_manager + self, product: Product, thl_redis: Redis, user_manager: UserManager ) -> LeaderboardContest: board_key = f"leaderboard:{product.uuid}:us:weekly:2025-05-26:complete_count" @@ -63,7 +68,13 @@ class TestLeaderboardContest(TestContest): c._user_manager = user_manager return c - def test_init(self, leaderboard_contest, thl_redis, user_1, user_2): + def test_init( + self, + leaderboard_contest: LeaderboardContest, + thl_redis: Redis, + user_1: User, + user_2: User, + ): model = leaderboard_contest.leaderboard_model assert leaderboard_contest.end_condition.ends_at is not None @@ -83,7 +94,14 @@ class TestLeaderboardContest(TestContest): lb = leaderboard_contest.get_leaderboard() print(lb) - def test_win(self, leaderboard_contest, thl_redis, user_1, user_2, user_3): + def test_win( + self, + leaderboard_contest: LeaderboardContest, + thl_redis: Redis, + user_1: User, + user_2: User, + user_3: User, + ): model = leaderboard_contest.leaderboard_model lbm = LeaderboardManager( redis_client=thl_redis, @@ -102,10 +120,13 @@ class TestLeaderboardContest(TestContest): lbm.hit_complete_count(product_user_id=user_3.product_user_id) leaderboard_contest.end_contest() + assert isinstance(leaderboard_contest.all_winners, list) assert len(leaderboard_contest.all_winners) == 3 # Prizes are $15, $10, $5. user 2 and 3 ties for 2nd place, so they split (10 + 5) assert leaderboard_contest.all_winners[0].awarded_cash_amount == USDCent(15_00) + + assert isinstance(leaderboard_contest.all_winners[0].user, User) assert ( leaderboard_contest.all_winners[0].user.product_user_id == user_1.product_user_id diff --git a/tests/models/thl/test_contest/test_raffle_contest.py b/tests/models/thl/test_contest/test_raffle_contest.py index f85ba75..8812cb3 100644 --- a/tests/models/thl/test_contest/test_raffle_contest.py +++ b/tests/models/thl/test_contest/test_raffle_contest.py @@ -1,4 +1,7 @@ +from __future__ import annotations + from collections import Counter +from datetime import datetime from uuid import uuid4 import pytest @@ -19,6 +22,7 @@ from generalresearch.models.thl.contest.definitions import ( ) from generalresearch.models.thl.contest.raffle import RaffleContest from generalresearch.models.thl.product import Product +from generalresearch.models.thl.user import User from tests.models.thl.test_contest.test_contest import TestContest @@ -42,7 +46,9 @@ class TestRaffleContest(TestContest): ) @pytest.fixture(scope="function") - def ended_raffle_contest(self, raffle_contest, utc_now) -> RaffleContest: + def ended_raffle_contest( + self, raffle_contest: RaffleContest, utc_now: datetime + ) -> RaffleContest: # Fake ending the contest raffle_contest = raffle_contest.model_copy() raffle_contest.update( @@ -55,7 +61,7 @@ class TestRaffleContest(TestContest): class TestRaffleContestUserView(TestRaffleContest): - def test_user_view(self, raffle_contest, user): + def test_user_view(self, raffle_contest: RaffleContest, user: User): from generalresearch.models.thl.contest.raffle import RaffleUserView data = { @@ -78,7 +84,7 @@ class TestRaffleContestUserView(TestRaffleContest): assert res["current_win_probability"] == approx(0.0099, rel=0.001) assert res["projected_win_probability"] == approx(0.0099, rel=0.001) - def test_win_pct(self, raffle_contest, user): + def test_win_pct(self, raffle_contest: RaffleContest, user: User): from generalresearch.models.thl.contest.raffle import RaffleUserView data = { @@ -124,7 +130,9 @@ class TestRaffleContestUserView(TestRaffleContest): class TestRaffleContestWinners(TestRaffleContest): - def test_winners_1_prize(self, ended_raffle_contest, user_1, user_2, user_3): + def test_winners_1_prize( + self, ended_raffle_contest, user_1: User, user_2: User, user_3: User + ): ended_raffle_contest.entries = [ ContestEntry( user=user_1, @@ -160,7 +168,13 @@ class TestRaffleContestWinners(TestRaffleContest): assert c[user_2.user_id] == approx(10000 * 2 / 6, rel=0.1) assert c[user_3.user_id] == approx(10000 * 3 / 6, rel=0.1) - def test_winners_2_prizes(self, ended_raffle_contest, user_1, user_2, user_3): + def test_winners_2_prizes( + self, + ended_raffle_contest: RaffleContest, + user_1: User, + user_2: User, + user_3: User, + ): ended_raffle_contest.prizes.append( ContestPrize( name="iPod 64GB Black", @@ -193,7 +207,9 @@ class TestRaffleContestWinners(TestRaffleContest): # Same user assert all(w.user.user_id == user_1.user_id for w in winners) - def test_winners_2_prizes_1_entry(self, ended_raffle_contest, user_3): + def test_winners_2_prizes_1_entry( + self, ended_raffle_contest: RaffleContest, user_3: User + ): ended_raffle_contest.prizes = [ ContestPrize( name="iPod 64GB White", @@ -218,7 +234,9 @@ class TestRaffleContestWinners(TestRaffleContest): winners = ended_raffle_contest.select_winners() assert len(winners) == 1 - def test_winners_2_prizes_1_entry_2_pennies(self, ended_raffle_contest, user_3): + def test_winners_2_prizes_1_entry_2_pennies( + self, ended_raffle_contest: RaffleContest, user_3: User + ): ended_raffle_contest.prizes = [ ContestPrize( name="iPod 64GB White", @@ -243,7 +261,12 @@ class TestRaffleContestWinners(TestRaffleContest): assert len(winners) == 2 def test_winners_3_prizes_3_entries( - self, ended_raffle_contest, product: Product, user_1, user_2, user_3 + self, + ended_raffle_contest: RaffleContest, + product: Product, + user_1: User, + user_2: User, + user_3: User, ): ended_raffle_contest.prizes = [ ContestPrize( diff --git a/tests/models/thl/test_ledger.py b/tests/models/thl/test_ledger.py index 7066180..7c48dbd 100644 --- a/tests/models/thl/test_ledger.py +++ b/tests/models/thl/test_ledger.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from datetime import UTC, datetime from uuid import uuid4 diff --git a/tests/models/thl/test_marketplace_condition.py b/tests/models/thl/test_marketplace_condition.py index 8a4b25c..1dd25e8 100644 --- a/tests/models/thl/test_marketplace_condition.py +++ b/tests/models/thl/test_marketplace_condition.py @@ -1,15 +1,18 @@ +from __future__ import annotations + import pytest from pydantic import ValidationError +from generalresearch.models import LogicalOperator +from generalresearch.models.thl.survey.condition import ( + ConditionValueType, + MarketplaceCondition, +) + class TestMarketplaceCondition: def test_list_or(self): - from generalresearch.models import LogicalOperator - from generalresearch.models.thl.survey.condition import ( - ConditionValueType, - MarketplaceCondition, - ) user_qas = {"q1": {"a2"}} c = MarketplaceCondition( @@ -46,11 +49,6 @@ class TestMarketplaceCondition: assert c.evaluate_criterion(user_qas) is None def test_list_or_negate(self): - from generalresearch.models import LogicalOperator - from generalresearch.models.thl.survey.condition import ( - ConditionValueType, - MarketplaceCondition, - ) user_qas = {"q1": {"a2"}} c = MarketplaceCondition( @@ -87,11 +85,6 @@ class TestMarketplaceCondition: assert c.evaluate_criterion(user_qas) is None def test_list_and(self): - from generalresearch.models import LogicalOperator - from generalresearch.models.thl.survey.condition import ( - ConditionValueType, - MarketplaceCondition, - ) user_qas = {"q1": {"a1", "a2"}} c = MarketplaceCondition( @@ -178,11 +171,6 @@ class TestMarketplaceCondition: assert c.evaluate_criterion(user_qas) is None def test_ranges(self): - from generalresearch.models import LogicalOperator - from generalresearch.models.thl.survey.condition import ( - ConditionValueType, - MarketplaceCondition, - ) user_qas = {"q1": {"2", "50"}} c = MarketplaceCondition( @@ -245,12 +233,6 @@ class TestMarketplaceCondition: ) def test_ranges_to_list(self): - from generalresearch.models import LogicalOperator - from generalresearch.models.thl.survey.condition import ( - ConditionValueType, - MarketplaceCondition, - ) - user_qas = {"q1": {"2", "50"}} MarketplaceCondition._CONVERT_LIST_TO_RANGE = ["q1"] c = MarketplaceCondition( @@ -309,10 +291,6 @@ class TestMarketplaceCondition: assert not c.evaluate_criterion({"q1": {"50"}}) def test_answered(self): - from generalresearch.models.thl.survey.condition import ( - ConditionValueType, - MarketplaceCondition, - ) user_qas = {"q1": {"a2"}} c = MarketplaceCondition( diff --git a/tests/models/thl/test_payout.py b/tests/models/thl/test_payout.py index dd0065c..daf1bd7 100644 --- a/tests/models/thl/test_payout.py +++ b/tests/models/thl/test_payout.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from uuid import uuid4 import pytest @@ -5,7 +7,7 @@ from pydantic import ValidationError from generalresearch.currency import USDCent from generalresearch.models.gr import Team -from generalresearch.models.gr.business import business: Business, BusinessAddress, BusinessType +from generalresearch.models.gr.business import Business, BusinessAddress, BusinessType from generalresearch.models.thl.payout import ( BrokerageProductPayoutEvent, BusinessPayoutEvent, diff --git a/tests/models/thl/test_payout_format.py b/tests/models/thl/test_payout_format.py index 83fde25..fe7aea5 100644 --- a/tests/models/thl/test_payout_format.py +++ b/tests/models/thl/test_payout_format.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import pytest from pydantic import BaseModel diff --git a/tests/models/thl/test_product.py b/tests/models/thl/test_product.py index b7ee654..adf276d 100644 --- a/tests/models/thl/test_product.py +++ b/tests/models/thl/test_product.py @@ -13,22 +13,26 @@ from pydantic import ValidationError from generalresearch.currency import USDCent from generalresearch.incite.base import GRLDatasets +from generalresearch.incite.collections.thl_web import LedgerDFCollection from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge from generalresearch.managers.thl.ledger_manager.thl_ledger import ( ThlLedgerManager, ) +from generalresearch.managers.thl.payout import PayoutEventManager from generalresearch.managers.thl.product import ProductManager from generalresearch.models import Source from generalresearch.models.gr.business import Business from generalresearch.models.thl.finance import ProductBalances -from generalresearch.models.thl.product import ( +from generalresearch.models.thl.payout import ( BrokerageProductPayoutEvent, +) +from generalresearch.models.thl.product import ( BrokerageProductPayoutEventManager, IntegrationMode, PayoutConfig, PayoutTransformation, PayoutTransformationPercentArgs, - product: Product, + Product, ProfilingConfig, SourceConfig, SourcesConfig, @@ -37,6 +41,7 @@ from generalresearch.models.thl.product import ( ) from generalresearch.models.thl.session import Session from generalresearch.models.thl.user import User +from generalresearch.redis_helper import RedisConfig class TestProduct: @@ -156,6 +161,11 @@ class TestProduct: assert ( "payout_transformation_percent" == p.payout_config.payout_transformation.f ) + + assert isinstance( + p.payout_config.payout_transformation.kwargs, + PayoutTransformationPercentArgs, + ) assert 0.5 == p.payout_config.payout_transformation.kwargs.pct assert ( Decimal("0.10") == p.payout_config.payout_transformation.kwargs.min_payout @@ -287,10 +297,10 @@ class TestProduct: p.profiling_config = ProfilingConfig(max_questions=1) assert p.profiling_config.max_questions == 1 - def test_bp_account(self, product: Product, thl_lm): + def test_bp_account(self, product: Product, thl_ledger_manager: ThlLedgerManager): assert product.bp_account is None - product.prefetch_bp_account(thl_lm=thl_lm) + product.prefetch_bp_account(thl_lm=thl_ledger_manager) from generalresearch.models.thl.ledger import LedgerAccount @@ -391,7 +401,7 @@ class TestGlobalProduct: random_product = uuid4().hex random_team = uuid4().hex res = instance.sources_config.get_policies_for( - product_id=random_product: Product, team_id=random_team + product_id=random_product, team_id=random_team ) assert res == s.global_scoped_policies_dict @@ -598,7 +608,7 @@ class TestProductFinancials: def test_balance( self, - business: business: Business, + business: Business, product_factory: Callable[..., Product], user_factory: Callable[..., User], mnt_filepath: GRLDatasets, @@ -610,7 +620,7 @@ class TestProductFinancials: delete_ledger_db: Callable[..., None], create_main_accounts: Callable[..., None], client_no_amm: DaskClient, - ledger_collection, + ledger_collection: LedgerDFCollection, pop_ledger_merge: PopLedgerMerge, delete_df_collection: Callable[..., None], ): @@ -781,20 +791,20 @@ class TestProductBalance: def test_inconsistent( self, - product: product: Product, + product: Product, mnt_filepath: GRLDatasets, thl_lm: ThlLedgerManager, client_no_amm: DaskClient, delete_ledger_db: Callable[..., None], create_main_accounts: Callable[..., None], delete_df_collection: Callable[..., None], - ledger_collection, + ledger_collection: LedgerDFCollection, user_factory: Callable[..., User], session_with_tx_factory: Callable[..., Session], - pop_ledger_merge, + pop_ledger_merge: PopLedgerMerge, start: datetime, - bp_payout_factory, - payout_event_manager, + bp_payout_factory: Callable[..., BrokerageProductPayoutEvent], + payout_event_manager: PayoutEventManager, ): # Now let's load it up and actually test some things delete_ledger_db() @@ -815,7 +825,7 @@ class TestProductBalance: # 2. Payout and build Parquets 2nd time payout_event_manager.set_account_lookup_table(thl_lm=thl_lm) bp_payout_factory( - product=product: Product, + product=product, amount=USDCent(71), ext_ref_id=uuid4().hex, created=start + timedelta(days=1, minutes=1), @@ -833,20 +843,20 @@ class TestProductBalance: def test_not_inconsistent( self, - product: product: Product, + product: Product, mnt_filepath: GRLDatasets, thl_lm: ThlLedgerManager, client_no_amm: DaskClient, delete_ledger_db: Callable[..., None], create_main_accounts: Callable[..., None], delete_df_collection: Callable[..., None], - ledger_collection, + ledger_collection: LedgerDFCollection, user_factory: Callable[..., User], session_with_tx_factory: Callable[..., None], pop_ledger_merge: PopLedgerMerge, start: datetime, - bp_payout_factory, - payout_event_manager, + bp_payout_factory: Callable[..., BrokerageProductPayoutEvent], + payout_event_manager: PayoutEventManager, ): # This is very similar to the test_complete_payout_pq_inconsistent # test, however this time we're only going to assign the payout @@ -874,7 +884,7 @@ class TestProductBalance: # so it hasn't already been archived payout_event_manager.set_account_lookup_table(thl_lm=thl_lm) bp_payout_factory( - product=product: Product, + product=product, amount=USDCent(71), ext_ref_id=uuid4().hex, created=datetime.now(tz=UTC), @@ -904,14 +914,14 @@ class TestProductPOPFinancial: def test_base( self, - product: product: Product, + product: Product, mnt_filepath: GRLDatasets, thl_lm: ThlLedgerManager, client_no_amm: DaskClient, delete_ledger_db: Callable[..., None], create_main_accounts: Callable[..., None], delete_df_collection: Callable[..., None], - ledger_collection, + ledger_collection: LedgerDFCollection, user_factory: Callable[..., User], session_with_tx_factory: Callable[..., None], pop_ledger_merge: PopLedgerMerge, @@ -977,16 +987,16 @@ class TestProductCache: def test_basic( self, - product: product: Product, - mnt_filepath, - thl_lm, + product: Product, + mnt_filepath: GRLDatasets, + thl_ledger_manager: ThlLedgerManager, client_no_amm: DaskClient, thl_redis_config: RedisConfig, - brokerage_product_payout_event_manager, + brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager, delete_ledger_db: Callable[..., None], create_main_accounts: Callable[..., None], delete_df_collection: Callable[..., None], - ledger_collection, + ledger_collection: LedgerDFCollection, user_factory: Callable[..., User], session_with_tx_factory: Callable[..., None], pop_ledger_merge: PopLedgerMerge, @@ -1007,7 +1017,7 @@ class TestProductCache: ds=mnt_filepath, client=client_no_amm, bp_pem=brokerage_product_payout_event_manager, - redis_config=thl_redis_config: RedisConfig, + redis_config=thl_redis_config, ) from generalresearch.models.thl.product import Product @@ -1029,7 +1039,7 @@ class TestProductCache: ds=mnt_filepath, client=client_no_amm, bp_pem=brokerage_product_payout_event_manager, - redis_config=thl_redis_config: RedisConfig, + redis_config=thl_redis_config, ) # Fetch from cache and assert the instance loaded from redis @@ -1048,23 +1058,23 @@ class TestProductCache: def test_neg_balance_cache( self, - product: product: Product, + product: Product, mnt_filepath: GRLDatasets, - thl_lm, + thl_ledger_manager: ThlLedgerManager, client_no_amm: DaskClient, thl_redis_config: RedisConfig, - brokerage_product_payout_event_manager, + brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager, delete_ledger_db: Callable[..., None], create_main_accounts: Callable[..., None], delete_df_collection: Callable[..., None], - ledger_collection, + ledger_collection: LedgerDFCollection, user_factory: Callable[..., User], session_with_tx_factory: Callable[..., None], pop_ledger_merge: PopLedgerMerge, start: datetime, - bp_payout_factory, - payout_event_manager, - adj_to_fail_with_tx_factory, + bp_payout_factory: Callable[..., BrokerageProductPayoutEvent], + payout_event_manager: PayoutEventManager, + adj_to_fail_with_tx_factory: Callable[..., None], ): # Now let's load it up and actually test some things delete_ledger_db() @@ -1083,9 +1093,9 @@ class TestProductCache: ) # 2. Payout - payout_event_manager.set_account_lookup_table(thl_lm=thl_lm) + payout_event_manager.set_account_lookup_table(thl_lm=thl_ledger_manager) bp_payout_factory( - product=product: Product, + product=product, amount=USDCent(71), ext_ref_id=uuid4().hex, created=start + timedelta(days=1, minutes=1), @@ -1104,11 +1114,11 @@ class TestProductCache: pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) product.set_cache( - thl_lm=thl_lm, + thl_lm=thl_ledger_manager, ds=mnt_filepath, client=client_no_amm, bp_pem=brokerage_product_payout_event_manager, - redis_config=thl_redis_config: RedisConfig, + redis_config=thl_redis_config, ) # Fetch from cache and assert the instance loaded from redis diff --git a/tests/models/thl/test_product_userwalletconfig.py b/tests/models/thl/test_product_userwalletconfig.py index 4f6a6cc..b348981 100644 --- a/tests/models/thl/test_product_userwalletconfig.py +++ b/tests/models/thl/test_product_userwalletconfig.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from itertools import groupby from random import shuffle as rshuffle @@ -7,7 +9,7 @@ from generalresearch.models.thl.product import ( from generalresearch.models.thl.wallet import PayoutType -def all_equal(iterable): +def all_equal(iterable: list[str]) -> bool: g = groupby(iterable) return next(g, True) and not next(g, False) @@ -35,13 +37,13 @@ class TestProductUserWalletConfig: # in the same order because they're the same assert isinstance(instance.model_dump_json(), str) res = [] - for idx in range(100): + for _ in range(100): res.append(instance.model_dump_json()) assert all_equal(res) def test_model_dump_payout_types(self): res = [] - for idx in range(100): + for _ in range(100): # Generate a random order of PayoutTypes each time payout_types = [e for e in PayoutType] diff --git a/tests/models/thl/test_soft_pair.py b/tests/models/thl/test_soft_pair.py index 588847e..3cf835e 100644 --- a/tests/models/thl/test_soft_pair.py +++ b/tests/models/thl/test_soft_pair.py @@ -1,12 +1,14 @@ +from __future__ import annotations + from generalresearch.models import Source +from generalresearch.models.dynata.survey import ( + ConditionValueType, + DynataCondition, +) from generalresearch.models.thl.soft_pair import SoftPairResult, SoftPairResultType def test_model(): - from generalresearch.models.dynata.survey import ( - ConditionValueType, - DynataCondition, - ) c1 = DynataCondition( question_id="1", value_type=ConditionValueType.LIST, values=["a", "b"] diff --git a/tests/models/thl/test_upkquestion.py b/tests/models/thl/test_upkquestion.py index 99d7871..719fcff 100644 --- a/tests/models/thl/test_upkquestion.py +++ b/tests/models/thl/test_upkquestion.py @@ -1,13 +1,30 @@ +from __future__ import annotations + import pytest from pydantic import ValidationError +from generalresearch.models.morning.question import ( + MorningQuestion, + MorningQuestionType, +) +from generalresearch.models.thl.profiling.upk_question import ( + PatternValidation, + UPKImportance, + UpkQuestion, + UpkQuestionChoice, + UpkQuestionConfigurationMC, + UpkQuestionConfigurationTE, + UpkQuestionSelectorMC, + UpkQuestionSelectorTE, + UpkQuestionType, + UpkQuestionValidation, + order_exclusive_options, +) + class TestUpkQuestion: def test_importance(self): - from generalresearch.models.thl.profiling.upk_question import ( - UPKImportance, - ) res = UPKImportance(task_score=1, task_count=None) assert isinstance(res, UPKImportance) @@ -20,9 +37,6 @@ class TestUpkQuestion: assert "Input should be greater than or equal to 0" in str(e.value) def test_pattern(self): - from generalresearch.models.thl.profiling.upk_question import ( - PatternValidation, - ) s = PatternValidation(message="hi", pattern="x") with pytest.raises(ValidationError) as e: @@ -30,13 +44,6 @@ class TestUpkQuestion: assert "Instance is frozen" in str(e.value) def test_mc(self): - from generalresearch.models.thl.profiling.upk_question import ( - UpkQuestion, - UpkQuestionChoice, - UpkQuestionConfigurationMC, - UpkQuestionSelectorMC, - UpkQuestionType, - ) q = UpkQuestion( id="601377a0d4c74529afc6293a8e5c3b5e", @@ -126,14 +133,6 @@ class TestUpkQuestion: assert "Extra inputs are not permitted" in str(e.value) def test_te(self): - from generalresearch.models.thl.profiling.upk_question import ( - PatternValidation, - UpkQuestion, - UpkQuestionConfigurationTE, - UpkQuestionSelectorTE, - UpkQuestionType, - UpkQuestionValidation, - ) q = UpkQuestion( id="601377a0d4c74529afc6293a8e5c3b5e", @@ -152,9 +151,6 @@ class TestUpkQuestion: assert q.choices is None def test_deserialization(self): - from generalresearch.models.thl.profiling.upk_question import ( - UpkQuestion, - ) q = UpkQuestion.model_validate( { @@ -195,16 +191,18 @@ class TestUpkQuestion: assert q == UpkQuestion.model_validate(q.model_dump(mode="json")) def test_from_morning(self): - from generalresearch.models.morning.question import ( - MorningQuestion, - MorningQuestionType, - ) q = MorningQuestion( - id="gender", country_iso="us", language_iso="eng", name="Gender", text="What is your gender?", type="s", options=[ - {"id": "1", "text": "yes", "order": 1}, - {"id": "2", "text": "no", "order": 2}, - ] + id="gender", + country_iso="us", + language_iso="eng", + name="Gender", + text="What is your gender?", + type="s", + options=[ + {"id": "1", "text": "yes", "order": 1}, + {"id": "2", "text": "no", "order": 2}, + ], ) q.to_upk_question() q = MorningQuestion( @@ -218,13 +216,6 @@ class TestUpkQuestion: q.to_upk_question() def test_order(self): - from generalresearch.models.thl.profiling.upk_question import ( - UpkQuestion, - UpkQuestionChoice, - UpkQuestionSelectorMC, - UpkQuestionType, - order_exclusive_options, - ) q = UpkQuestion( country_iso="us", @@ -258,9 +249,6 @@ class TestUpkQuestion: class TestUpkQuestionValidateAnswer: def test_validate_answer_SA(self): - from generalresearch.models.thl.profiling.upk_question import ( - UpkQuestion, - ) question = UpkQuestion.model_validate( { @@ -296,9 +284,6 @@ class TestUpkQuestionValidateAnswer: ) def test_validate_answer_MA(self): - from generalresearch.models.thl.profiling.upk_question import ( - UpkQuestion, - ) question = UpkQuestion.model_validate( { @@ -368,9 +353,6 @@ class TestUpkQuestionValidateAnswer: ) def test_validate_answer_TE(self): - from generalresearch.models.thl.profiling.upk_question import ( - UpkQuestion, - ) question = UpkQuestion.model_validate( { diff --git a/tests/models/thl/test_user.py b/tests/models/thl/test_user.py index 0b8634a..9c4b548 100644 --- a/tests/models/thl/test_user.py +++ b/tests/models/thl/test_user.py @@ -1,4 +1,7 @@ +from __future__ import annotations + import json +from collections.abc import Callable from datetime import UTC, datetime, timedelta, timezone from decimal import Decimal from random import choice as rand_choice @@ -8,18 +11,21 @@ from uuid import uuid4 import pytest from pydantic import ValidationError +from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager +from generalresearch.managers.thl.userhealth import AuditLogManager +from generalresearch.models.thl.product import Product +from generalresearch.models.thl.user import User + class TestUserUserID: def test_valid(self): - from generalresearch.models.thl.user import User val = randint(1, 2**30) user = User(user_id=val) assert user.user_id == val def test_type(self): - from generalresearch.models.thl.user import User # It will cast str to int assert User(user_id="1").user_id == 1 @@ -44,7 +50,6 @@ class TestUserUserID: assert "Input should be a valid integer," in str(cm.value) def test_zero(self): - from generalresearch.models.thl.user import User with pytest.raises(expected_exception=ValidationError) as cm: User(user_id=0) @@ -52,7 +57,6 @@ class TestUserUserID: assert "Input should be greater than 0" in str(cm.value) def test_negative(self): - from generalresearch.models.thl.user import User with pytest.raises(expected_exception=ValidationError) as cm: User(user_id=-1) @@ -60,7 +64,6 @@ class TestUserUserID: assert "Input should be greater than 0" in str(cm.value) def test_too_big(self): - from generalresearch.models.thl.user import User val = 2**31 with pytest.raises(expected_exception=ValidationError) as cm: @@ -69,7 +72,6 @@ class TestUserUserID: assert "Input should be less than 2147483648" in str(cm.value) def test_identifiable(self): - from generalresearch.models.thl.user import User val = randint(1, 2**30) user = User(user_id=val) @@ -80,7 +82,6 @@ class TestUserProductID: user_id = randint(1, 2**30) def test_valid(self): - from generalresearch.models.thl.user import User product_id = uuid4().hex @@ -89,7 +90,6 @@ class TestUserProductID: assert user.product_id == product_id def test_type(self): - from generalresearch.models.thl.user import User with pytest.raises(expected_exception=ValueError) as cm: User(user_id=self.user_id, product_id=0) @@ -102,7 +102,6 @@ class TestUserProductID: assert "Input should be a valid string" in str(cm.value) def test_empty(self): - from generalresearch.models.thl.user import User with pytest.raises(expected_exception=ValueError) as cm: User(user_id=self.user_id, product_id="") @@ -110,7 +109,6 @@ class TestUserProductID: assert "String should have at least 32 characters" in str(cm.value) def test_invalid_len(self): - from generalresearch.models.thl.user import User # Valid uuid4s are 32 char long product_id = uuid4().hex[:31] @@ -133,7 +131,6 @@ class TestUserProductID: assert "String should have at most 32 characters" in str(cm.value) def test_invalid_uuid(self): - from generalresearch.models.thl.user import User # Modify the UUID to break it product_id = uuid4().hex[:31] + "x" @@ -144,7 +141,6 @@ class TestUserProductID: assert "Invalid UUID" in str(cm.value) def test_invalid_hex_form(self): - from generalresearch.models.thl.user import User # Sure not in hex form, but it'll get caught for being the # wrong length before anything else @@ -157,7 +153,6 @@ class TestUserProductID: def test_identifiable(self): """Can't create a User with only a product_id because it also needs to the product_user_id""" - from generalresearch.models.thl.user import User product_id = uuid4().hex with pytest.raises(expected_exception=ValueError) as cm: @@ -172,10 +167,9 @@ class TestUserProductUserID: def randomword(self, length: int = 50): # Raw so nothing is escaped to add additional backslashes _bpuid_allowed = r"0123456789abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ!#$%&()*+,-.:;<=>?@[]^_{|}~" - return "".join(rand_choice(_bpuid_allowed) for i in range(length)) + return "".join(rand_choice(_bpuid_allowed) for _ in range(length)) def test_valid(self): - from generalresearch.models.thl.user import User product_user_id = uuid4().hex[:12] user = User(user_id=self.user_id, product_user_id=product_user_id) @@ -184,7 +178,6 @@ class TestUserProductUserID: assert user.product_user_id == product_user_id def test_type(self): - from generalresearch.models.thl.user import User with pytest.raises(expected_exception=ValueError) as cm: User(user_id=self.user_id, product_user_id=0) @@ -202,7 +195,6 @@ class TestUserProductUserID: assert "Input should be a valid string" in str(cm.value) def test_empty(self): - from generalresearch.models.thl.user import User with pytest.raises(expected_exception=ValueError) as cm: User(user_id=self.user_id, product_user_id="") @@ -210,7 +202,6 @@ class TestUserProductUserID: assert "String should have at least 3 characters" in str(cm.value) def test_invalid_len(self): - from generalresearch.models.thl.user import User product_user_id = self.randomword(251) with pytest.raises(expected_exception=ValueError) as cm: @@ -225,7 +216,6 @@ class TestUserProductUserID: assert "String should have at least 3 characters" in str(cm.value) def test_invalid_chars_space(self): - from generalresearch.models.thl.user import User product_user_id = f"{self.randomword(50)} {self.randomword(50)}" with pytest.raises(expected_exception=ValueError) as cm: @@ -234,7 +224,6 @@ class TestUserProductUserID: assert "String cannot contain spaces" in str(cm.value) def test_invalid_chars_slash(self): - from generalresearch.models.thl.user import User product_user_id = rf"{self.randomword(50)}\{self.randomword(50)}" with pytest.raises(expected_exception=ValueError) as cm: @@ -253,7 +242,6 @@ class TestUserProductUserID: I wanted a test that made sure the regex was hit. I do not know how we want to provide with the level of specific String checks we do in here for specific error messages.""" - from generalresearch.models.thl.user import User product_user_id = f"{self.randomword(50)}`{self.randomword(50)}" with pytest.raises(expected_exception=ValueError) as cm: @@ -275,7 +263,6 @@ class TestUserProductUserID: def test_identifiable(self): """Can't create a User with only a product_user_id because it also needs to the product_id""" - from generalresearch.models.thl.user import User product_user_id = uuid4().hex with pytest.raises(ValueError) as cm: @@ -288,7 +275,6 @@ class TestUserUUID: user_id = randint(1, 2**30) def test_valid(self): - from generalresearch.models.thl.user import User uuid_pk = uuid4().hex @@ -297,7 +283,6 @@ class TestUserUUID: assert user.uuid == uuid_pk def test_type(self): - from generalresearch.models.thl.user import User with pytest.raises(ValueError) as cm: User(user_id=self.user_id, uuid=0) @@ -315,7 +300,6 @@ class TestUserUUID: assert "Input should be a valid string" in str(cm.value) def test_empty(self): - from generalresearch.models.thl.user import User with pytest.raises(ValueError) as cm: User(user_id=self.user_id, uuid="") @@ -323,7 +307,6 @@ class TestUserUUID: assert "String should have at least 32 characters" in str(cm.value) def test_invalid_len(self): - from generalresearch.models.thl.user import User # Valid uuid4s are 32 char long uuid_pk = uuid4().hex[:31] @@ -341,7 +324,6 @@ class TestUserUUID: assert "String should have at most 32 characters" in str(cm.value) def test_invalid_uuid(self): - from generalresearch.models.thl.user import User # Modify the UUID to break it uuid_pk = uuid4().hex[:31] + "x" @@ -352,7 +334,6 @@ class TestUserUUID: assert "Invalid UUID" in str(cm.value) def test_invalid_hex_form(self): - from generalresearch.models.thl.user import User # Sure not in hex form, but it'll get caught for being the # wrong length before anything else @@ -369,7 +350,6 @@ class TestUserUUID: assert "Invalid UUID" in str(cm.value) def test_identifiable(self): - from generalresearch.models.thl.user import User user_uuid = uuid4().hex user = User(uuid=user_uuid) @@ -380,7 +360,6 @@ class TestUserCreated: user_id = randint(1, 2**30) def test_valid(self): - from generalresearch.models.thl.user import User user = User(user_id=self.user_id) dt = datetime.now(tz=UTC) @@ -389,7 +368,6 @@ class TestUserCreated: assert user.created == dt def test_tz_naive_throws_init(self): - from generalresearch.models.thl.user import User with pytest.raises(ValueError) as cm: User(user_id=self.user_id, created=datetime.now(tz=None)) # noqa @@ -397,7 +375,6 @@ class TestUserCreated: assert "Input should have timezone info" in str(cm.value) def test_tz_naive_throws_setter(self): - from generalresearch.models.thl.user import User user = User(user_id=self.user_id) with pytest.raises(ValueError) as cm: @@ -406,7 +383,6 @@ class TestUserCreated: assert "Input should have timezone info" in str(cm.value) def test_tz_utc(self): - from generalresearch.models.thl.user import User with pytest.raises(ValueError) as cm: User( @@ -417,7 +393,6 @@ class TestUserCreated: assert "Timezone is not UTC" in str(cm.value) def test_not_in_future(self): - from generalresearch.models.thl.user import User the_future = datetime.now(tz=UTC) + timedelta(minutes=1) with pytest.raises(ValueError) as cm: @@ -426,7 +401,6 @@ class TestUserCreated: assert "Input is in the future" in str(cm.value) def test_after_anno_domini(self): - from generalresearch.models.thl.user import User before_ad = datetime(year=2015, month=1, day=1, tzinfo=UTC) + timedelta( minutes=1 @@ -441,7 +415,6 @@ class TestUserLastSeen: user_id = randint(1, 2**30) def test_valid(self): - from generalresearch.models.thl.user import User user = User(user_id=self.user_id) dt = datetime.now(tz=UTC) @@ -450,7 +423,6 @@ class TestUserLastSeen: assert user.last_seen == dt def test_tz_naive_throws_init(self): - from generalresearch.models.thl.user import User with pytest.raises(ValueError) as cm: User(user_id=self.user_id, last_seen=datetime.now(tz=None)) # noqa @@ -458,7 +430,6 @@ class TestUserLastSeen: assert "Input should have timezone info" in str(cm.value) def test_tz_naive_throws_setter(self): - from generalresearch.models.thl.user import User user = User(user_id=self.user_id) with pytest.raises(ValueError) as cm: @@ -467,7 +438,6 @@ class TestUserLastSeen: assert "Input should have timezone info" in str(cm.value) def test_tz_utc(self): - from generalresearch.models.thl.user import User with pytest.raises(ValueError) as cm: User( @@ -478,7 +448,6 @@ class TestUserLastSeen: assert "Timezone is not UTC" in str(cm.value) def test_not_in_future(self): - from generalresearch.models.thl.user import User the_future = datetime.now(tz=UTC) + timedelta(minutes=1) with pytest.raises(ValueError) as cm: @@ -487,7 +456,6 @@ class TestUserLastSeen: assert "Input is in the future" in str(cm.value) def test_after_anno_domini(self): - from generalresearch.models.thl.user import User before_ad = datetime(year=2015, month=1, day=1, tzinfo=UTC) + timedelta( minutes=1 @@ -502,7 +470,6 @@ class TestUserBlocked: user_id = randint(1, 2**30) def test_valid(self): - from generalresearch.models.thl.user import User user = User(user_id=self.user_id, blocked=True) assert user.blocked @@ -510,7 +477,6 @@ class TestUserBlocked: def test_str_casting(self): """We don't want any of these to work, and that's why we set strict=True on the column""" - from generalresearch.models.thl.user import User with pytest.raises(ValueError) as cm: User(user_id=self.user_id, blocked="true") @@ -547,7 +513,6 @@ class TestUserTiming: user_id = randint(1, 2**30) def test_valid(self): - from generalresearch.models.thl.user import User created = datetime.now(tz=UTC) - timedelta(minutes=60) last_seen = datetime.now(tz=UTC) - timedelta(minutes=59) @@ -557,7 +522,6 @@ class TestUserTiming: assert user.last_seen == last_seen def test_created_first(self): - from generalresearch.models.thl.user import User created = datetime.now(tz=UTC) - timedelta(minutes=60) last_seen = datetime.now(tz=UTC) - timedelta(minutes=59) @@ -572,7 +536,6 @@ class TestUserModelVerification: """Tests that may be dependent on more than 1 attribute""" def test_identifiable(self): - from generalresearch.models.thl.user import User product_id = uuid4().hex product_user_id = uuid4().hex @@ -580,7 +543,6 @@ class TestUserModelVerification: assert user.is_identifiable def test_valid_helper(self): - from generalresearch.models.thl.user import User user_bool = User.is_valid_ubp( product_id=uuid4().hex, product_user_id=uuid4().hex @@ -594,7 +556,6 @@ class TestUserModelVerification: class TestUserSerialization: def test_basic_json(self): - from generalresearch.models.thl.user import User product_id = uuid4().hex product_user_id = uuid4().hex @@ -615,7 +576,6 @@ class TestUserSerialization: assert d.get("created").endswith("Z") def test_basic_dict(self): - from generalresearch.models.thl.user import User product_id = uuid4().hex product_user_id = uuid4().hex @@ -633,10 +593,11 @@ class TestUserSerialization: assert not d.get("blocked") assert d.get("product") is None - assert d.get("created").tzinfo == UTC + created = d.get("created") + assert isinstance(created, datetime) + assert created.tzinfo == UTC def test_from_json(self): - from generalresearch.models.thl.user import User product_id = uuid4().hex product_user_id = uuid4().hex @@ -651,12 +612,13 @@ class TestUserSerialization: u = User.model_validate_json(user.to_json()) assert u.product_id == product_id assert u.product is None + assert isinstance(u.created, datetime) assert u.created.tzinfo == UTC class TestUserMethods: - def test_audit_log(self, user, audit_log_manager): + def test_audit_log(self, user: User, audit_log_manager: AuditLogManager): assert user.audit_log is None user.prefetch_audit_log(audit_log_manager=audit_log_manager) assert user.audit_log == [] @@ -668,21 +630,21 @@ class TestUserMethods: def test_transactions( self, user_factory: Callable[..., User], - thl_lm, + thl_ledger_manager: ThlLedgerManager, session_with_tx_factory: Callable[..., None], - product_user_wallet_yes, + product_user_wallet_yes: Product, ): u1 = user_factory(product=product_user_wallet_yes) assert u1.transactions is None - u1.prefetch_transactions(thl_lm=thl_lm) + u1.prefetch_transactions(thl_lm=thl_ledger_manager) assert u1.transactions == [] session_with_tx_factory(user=u1) - u1.prefetch_transactions(thl_lm=thl_lm) + u1.prefetch_transactions(thl_lm=thl_ledger_manager) assert len(u1.transactions) == 1 @pytest.mark.skip(reason="TODO") - def test_location_history(self, user): + def test_location_history(self, user: User): assert user.location_history is None diff --git a/tests/models/thl/test_user_iphistory.py b/tests/models/thl/test_user_iphistory.py index d6ade9d..b8a0be3 100644 --- a/tests/models/thl/test_user_iphistory.py +++ b/tests/models/thl/test_user_iphistory.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from datetime import UTC, datetime, timedelta from generalresearch.models.thl.user_iphistory import ( diff --git a/tests/models/thl/test_user_metadata.py b/tests/models/thl/test_user_metadata.py index 3d851dc..a7b479d 100644 --- a/tests/models/thl/test_user_metadata.py +++ b/tests/models/thl/test_user_metadata.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import pytest from generalresearch.models import MAX_INT32 diff --git a/tests/models/thl/test_user_streak.py b/tests/models/thl/test_user_streak.py index 26c5e25..8300474 100644 --- a/tests/models/thl/test_user_streak.py +++ b/tests/models/thl/test_user_streak.py @@ -71,6 +71,7 @@ def test_user_streak_remaining(): ) print(f"{now.isoformat()=}, {end_of_today.isoformat()=}") expected = (end_of_today - now).total_seconds() + assert isinstance(us.time_remaining_in_period, timedelta) assert us.time_remaining_in_period.total_seconds() == pytest.approx(expected, abs=1) @@ -92,5 +93,6 @@ def test_user_streak_remaining_month(): ).replace(day=1) print(f"{now.isoformat()=}, {end_of_month.isoformat()=}") expected = (end_of_month - now).total_seconds() + assert isinstance(us.time_remaining_in_period, timedelta) assert us.time_remaining_in_period.total_seconds() == pytest.approx(expected, abs=1) print(us.time_remaining_in_period) diff --git a/tests/models/thl/test_wall.py b/tests/models/thl/test_wall.py index 88914ac..58e9825 100644 --- a/tests/models/thl/test_wall.py +++ b/tests/models/thl/test_wall.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from datetime import UTC, datetime, timedelta from decimal import Decimal from uuid import uuid4 diff --git a/tests/models/thl/test_wall_session.py b/tests/models/thl/test_wall_session.py index b39ad31..48b89ea 100644 --- a/tests/models/thl/test_wall_session.py +++ b/tests/models/thl/test_wall_session.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from datetime import UTC, datetime, timedelta from decimal import Decimal |
