From 17ff15c06655717627da820417337c6b0b97de42 Mon Sep 17 00:00:00 2001 From: Max Nanis Date: Wed, 2 Sep 2026 23:31:52 -0700 Subject: Lots more tests/managers/thl - doing all the factory organization from create_dummy --- tests/grliq/managers/test_forensic_data.py | 65 ++++++++++++------ tests/grliq/managers/test_forensic_results.py | 11 ++-- tests/managers/gr/test_business.py | 38 ++++++----- tests/managers/thl/test_contest/test_milestone.py | 2 +- tests/managers/thl/test_contest/test_raffle.py | 10 +-- tests/managers/thl/test_ipinfo.py | 20 +++--- tests/managers/thl/test_ledger/test_lm_accounts.py | 4 +- tests/managers/thl/test_ledger/test_lm_tx_locks.py | 2 +- tests/managers/thl/test_ledger/test_thl_lm_tx.py | 3 +- tests/managers/thl/test_ledger/test_wallet.py | 6 +- tests/managers/thl/test_product.py | 76 ++++++++++++++-------- tests/managers/thl/test_task_status.py | 18 +++-- tests/managers/thl/test_user_manager/test_base.py | 20 ++++-- tests/managers/thl/test_wall_manager.py | 2 +- tests/models/gr/test_business.py | 3 - tests/models/gr/test_team.py | 28 ++++++-- tests/models/thl/test_product.py | 6 +- tests/models/thl/test_user.py | 10 ++- 18 files changed, 211 insertions(+), 113 deletions(-) (limited to 'tests') diff --git a/tests/grliq/managers/test_forensic_data.py b/tests/grliq/managers/test_forensic_data.py index e4854e8..1b83757 100644 --- a/tests/grliq/managers/test_forensic_data.py +++ b/tests/grliq/managers/test_forensic_data.py @@ -1,5 +1,6 @@ from __future__ import annotations +from collections.abc import Callable from datetime import timedelta from typing import TYPE_CHECKING from uuid import uuid4 @@ -16,6 +17,8 @@ from generalresearch.grliq.models.forensic_result import ( if TYPE_CHECKING: from generalresearch.grliq.managers.forensic_data import ( GrlIqDataManager, + ) + from generalresearch.grliq.managers.forensic_events import ( GrlIqEventManager, ) from generalresearch.models.thl.product import Product @@ -28,10 +31,13 @@ except ImportError: class TestGrlIqDataManager: - def test_create_dummy(self, grliq_dm: GrlIqDataManager): + def test_create_dummy( + self, + grliq_data_factory: Callable[..., GrlIqData], + ): from generalresearch.grliq.models.forensic_data import GrlIqData - gd1: GrlIqData = grliq_dm.create_dummy(is_attempt_allowed=True) + gd1: GrlIqData = grliq_data_factory(is_attempt_allowed=True) assert isinstance(gd1, GrlIqData) assert isinstance(gd1.results, GrlIqCheckerResults) @@ -119,7 +125,9 @@ class TestGrlIqDataManager: class TestForensicDataGetAndFilter: - def test_events(self, grliq_dm: GrlIqDataManager): + def test_events( + self, grliq_dm: GrlIqDataManager, grliq_data_factory: Callable[..., GrlIqData] + ): """If load_events=True, the events and mouse_events attributes should be an array no matter what. An empty array means that the events were loaded, but there were no events available. @@ -129,7 +137,7 @@ class TestForensicDataGetAndFilter: """ # Load Events == False forensic_uuid = uuid4().hex - grliq_dm.create_dummy(is_attempt_allowed=True, uuid=forensic_uuid) + grliq_data_factory(is_attempt_allowed=True, uuid=forensic_uuid) instance = grliq_dm.filter_data(uuids=[forensic_uuid])[0] assert isinstance(instance, GrlIqData) @@ -144,41 +152,53 @@ class TestForensicDataGetAndFilter: assert len(instance.events) == 0 assert len(instance.mouse_events) == 0 - def test_timing(self, grliq_dm: GrlIqDataManager, grliq_em: GrlIqEventManager): + def test_timing( + self, + grliq_data_factory: Callable[..., GrlIqData], + grliq_data_manager: GrlIqDataManager, + grliq_event_manager: GrlIqEventManager, + ): forensic_uuid = uuid4().hex - grliq_dm.create_dummy(is_attempt_allowed=True, uuid=forensic_uuid) + grliq_data_factory(is_attempt_allowed=True, uuid=forensic_uuid) - instance = grliq_dm.filter_data(uuids=[forensic_uuid])[0] + instance = grliq_data_manager.filter_data(uuids=[forensic_uuid])[0] - grliq_em.update_or_create_timing( + grliq_event_manager.update_or_create_timing( session_uuid=instance.mid, timing_data=TimingData( client_rtts=[100, 200, 150], server_rtts=[150, 120, 120] ), ) - instance = grliq_dm.get_data(forensic_uuid=forensic_uuid, load_events=True) + instance = grliq_data_manager.get_data( + forensic_uuid=forensic_uuid, load_events=True + ) assert isinstance(instance, GrlIqData) assert isinstance(instance.events, list) assert isinstance(instance.mouse_events, list) assert isinstance(instance.timing_data, TimingData) def test_events_events( - self, grliq_dm: GrlIqDataManager, grliq_em: GrlIqEventManager + self, + grliq_data_factory: Callable[..., GrlIqData], + grliq_data_manager: GrlIqDataManager, + grliq_event_manager: GrlIqEventManager, ): forensic_uuid = uuid4().hex - grliq_dm.create_dummy(is_attempt_allowed=True, uuid=forensic_uuid) + grliq_data_factory(is_attempt_allowed=True, uuid=forensic_uuid) - instance = grliq_dm.filter_data(uuids=[forensic_uuid])[0] + instance = grliq_data_manager.filter_data(uuids=[forensic_uuid])[0] - grliq_em.update_or_create_events( + grliq_event_manager.update_or_create_events( session_uuid=instance.mid, events=[{"a": "b"}], mouse_events=[], event_start=instance.created_at, event_end=instance.created_at + timedelta(minutes=1), ) - instance = grliq_dm.get_data(forensic_uuid=forensic_uuid, load_events=True) + instance = grliq_data_manager.get_data( + forensic_uuid=forensic_uuid, load_events=True + ) assert isinstance(instance, GrlIqData) assert isinstance(instance.events, list) assert isinstance(instance.mouse_events, list) @@ -189,11 +209,16 @@ class TestForensicDataGetAndFilter: assert len(instance.keyboard_events) == 0 def test_events_click( - self, grliq_dm: GrlIqDataManager, grliq_em: GrlIqEventManager + self, + grliq_data_factory: Callable[..., GrlIqData], + grliq_data_manager: GrlIqDataManager, + grliq_event_manager: GrlIqEventManager, ): forensic_uuid = uuid4().hex - grliq_dm.create_dummy(is_attempt_allowed=True, uuid=forensic_uuid) - instance = grliq_dm.get_data(forensic_uuid=forensic_uuid, load_events=True) + grliq_data_factory(is_attempt_allowed=True, uuid=forensic_uuid) + instance = grliq_data_manager.get_data( + forensic_uuid=forensic_uuid, load_events=True + ) click_event = { "type": "click", @@ -203,14 +228,16 @@ class TestForensicDataGetAndFilter: "pointerType": "mouse", } me = MouseEvent.from_dict(click_event) - grliq_em.update_or_create_events( + grliq_event_manager.update_or_create_events( session_uuid=instance.mid, events=[click_event], mouse_events=[], event_start=instance.created_at, event_end=instance.created_at + timedelta(minutes=1), ) - instance = grliq_dm.get_data(forensic_uuid=forensic_uuid, load_events=True) + instance = grliq_data_manager.get_data( + forensic_uuid=forensic_uuid, load_events=True + ) assert isinstance(instance, GrlIqData) assert isinstance(instance.events, list) assert isinstance(instance.mouse_events, list) diff --git a/tests/grliq/managers/test_forensic_results.py b/tests/grliq/managers/test_forensic_results.py index a030451..86834d0 100644 --- a/tests/grliq/managers/test_forensic_results.py +++ b/tests/grliq/managers/test_forensic_results.py @@ -1,18 +1,21 @@ from __future__ import annotations +from collections.abc import Callable from typing import TYPE_CHECKING if TYPE_CHECKING: - from generalresearch.grliq.managers.forensic_data import GrlIqDataManager from generalresearch.grliq.managers.forensic_results import ( GrlIqCategoryResultsReader, ) + from generalresearch.grliq.models.forensic_data import GrlIqData class TestGrlIqCategoryResultsReader: def test_filter_category_results( - self, grliq_dm: GrlIqDataManager, grliq_crr: GrlIqCategoryResultsReader + self, + grliq_data_factory: Callable[..., GrlIqData], + grliq_crr: GrlIqCategoryResultsReader, ): from generalresearch.grliq.models.forensic_result import ( GrlIqForensicCategoryResult, @@ -20,8 +23,8 @@ class TestGrlIqCategoryResultsReader: ) # this is just testing that it doesn't fail - grliq_dm.create_dummy(is_attempt_allowed=True) - grliq_dm.create_dummy(is_attempt_allowed=True) + grliq_data_factory(is_attempt_allowed=True) + grliq_data_factory(is_attempt_allowed=True) res = grliq_crr.filter_category_results(limit=2, phase=Phase.OFFERWALL_ENTER)[0] assert res.get("category_result") diff --git a/tests/managers/gr/test_business.py b/tests/managers/gr/test_business.py index 0d5b0d5..6a930b4 100644 --- a/tests/managers/gr/test_business.py +++ b/tests/managers/gr/test_business.py @@ -76,30 +76,30 @@ class TestBusinessManager: assert isinstance(instance, Business) assert isinstance(instance.id, int) - def test_get_or_create(self, business_manager: BusinessManager): + def test_get_or_create(self, gr_business_manager: BusinessManager): uuid_key = uuid4().hex - assert business_manager.get_by_uuid(business_uuid=uuid_key) is None + assert gr_business_manager.get_by_uuid(business_uuid=uuid_key) is None - instance = business_manager.get_or_create( + instance = gr_business_manager.get_or_create( uuid=uuid_key, name=f"name-{uuid4().hex[:6]}", ) - res = business_manager.get_by_uuid(business_uuid=uuid_key) + res = gr_business_manager.get_by_uuid(business_uuid=uuid_key) assert isinstance(res, Business) assert res.id == instance.id def test_get_all( self, - business_manager: BusinessManager, + gr_business_manager: BusinessManager, gr_business_factory: Callable[..., Business], ): - res1 = business_manager.get_all() + res1 = gr_business_manager.get_all() assert isinstance(res1, list) gr_business_factory() - res2 = business_manager.get_all() + res2 = gr_business_manager.get_all() assert len(res1) == len(res2) - 1 @pytest.mark.skip(reason="TODO") @@ -108,42 +108,42 @@ class TestBusinessManager: def test_get_by_user_id( self, - business_manager: BusinessManager, + gr_business_manager: BusinessManager, gr_user: GRUser, team_manager: TeamManager, membership_manager: MembershipManager, gr_business_factory: Callable[..., Business], gr_team_factory: Callable[..., Team], ): - res = business_manager.get_by_user_id(user_id=gr_user.id) + res = gr_business_manager.get_by_user_id(user_id=gr_user.id) assert len(res) == 0 # Create a business: Business, but don't add it to anything b1 = gr_business_factory() - res = business_manager.get_by_user_id(user_id=gr_user.id) + res = gr_business_manager.get_by_user_id(user_id=gr_user.id) assert len(res) == 0 # Create a Team, but don't create any Memberships t1 = gr_team_factory() - res = business_manager.get_by_user_id(user_id=gr_user.id) + res = gr_business_manager.get_by_user_id(user_id=gr_user.id) assert len(res) == 0 # Create a Membership for the gr_user to the Team... but it doesn't # matter because the Team doesn't have any Business yet _ = membership_manager.create(team=t1, gr_user=gr_user) - res = business_manager.get_by_user_id(user_id=gr_user.id) + res = gr_business_manager.get_by_user_id(user_id=gr_user.id) assert len(res) == 0 # Add the Business to the Team... now the Business should be available # to the gr_user team_manager.add_business(team=t1, business=b1) - res = business_manager.get_by_user_id(user_id=gr_user.id) + res = gr_business_manager.get_by_user_id(user_id=gr_user.id) assert len(res) == 1 # Add another Business to the Team! b2 = gr_business_factory() team_manager.add_business(team=t1, business=b2) - res = business_manager.get_by_user_id(user_id=gr_user.id) + res = gr_business_manager.get_by_user_id(user_id=gr_user.id) assert len(res) == 2 @pytest.mark.skip(reason="TODO") @@ -151,14 +151,16 @@ class TestBusinessManager: pass def test_get_by_uuid( - self, gr_business: Business, business_manager: BusinessManager + self, gr_business: Business, gr_business_manager: BusinessManager ): - instance = business_manager.get_by_uuid(business_uuid=gr_business.uuid) + instance = gr_business_manager.get_by_uuid(business_uuid=gr_business.uuid) assert isinstance(instance, Business) assert gr_business.id == instance.id - def test_get_by_id(self, gr_business: Business, business_manager: BusinessManager): - instance = business_manager.get_by_id(business_id=gr_business.id) + def test_get_by_id( + self, gr_business: Business, gr_business_manager: BusinessManager + ): + instance = gr_business_manager.get_by_id(business_id=gr_business.id) assert isinstance(instance, Business) assert gr_business.uuid == instance.uuid diff --git a/tests/managers/thl/test_contest/test_milestone.py b/tests/managers/thl/test_contest/test_milestone.py index dbb2016..dab02e7 100644 --- a/tests/managers/thl/test_contest/test_milestone.py +++ b/tests/managers/thl/test_contest/test_milestone.py @@ -6,10 +6,10 @@ from typing import TYPE_CHECKING from generalresearch.models.thl.contest.definitions import ( ContestEndReason, + ContestEntryTrigger, ContestStatus, ) from generalresearch.models.thl.contest.milestone import ( - ContestEntryTrigger, MilestoneContest, MilestoneUserView, ) diff --git a/tests/managers/thl/test_contest/test_raffle.py b/tests/managers/thl/test_contest/test_raffle.py index 7803952..0b2b852 100644 --- a/tests/managers/thl/test_contest/test_raffle.py +++ b/tests/managers/thl/test_contest/test_raffle.py @@ -17,17 +17,17 @@ from generalresearch.models.thl.contest import ( ContestEntryRule, ContestPrize, ) +from generalresearch.models.thl.contest.contest_entry import ( + ContestEntry, + ContestEntryType, +) from generalresearch.models.thl.contest.definitions import ( ContestEndReason, ContestPrizeKind, ContestStatus, ) from generalresearch.models.thl.contest.exceptions import ContestError -from generalresearch.models.thl.contest.raffle import ( - ContestEntry, - ContestEntryType, - RaffleContest, -) +from generalresearch.models.thl.contest.raffle import RaffleContest if TYPE_CHECKING: from generalresearch.managers.thl.contest_manager import ContestManager diff --git a/tests/managers/thl/test_ipinfo.py b/tests/managers/thl/test_ipinfo.py index 6954163..47b1712 100644 --- a/tests/managers/thl/test_ipinfo.py +++ b/tests/managers/thl/test_ipinfo.py @@ -31,14 +31,16 @@ class TestIPGeonameManager: assert isinstance(instance, IPGeonameManager) assert isinstance(ip_geoname_manager, IPGeonameManager) - def test_create(self, ip_geoname_manager: IPGeonameManager): - - instance = ip_geoname_manager.create_dummy() + def test_create( + self, + ip_geoname_factory: Callable[..., IPGeoname], + ip_geoname_manager: IPGeonameManager, + ): + instance = ip_geoname_factory() assert isinstance(instance, IPGeoname) res = ip_geoname_manager.fetch_geoname_ids(filter_ids=[instance.geoname_id]) - assert res[0].model_dump_json() == instance.model_dump_json() @@ -51,13 +53,15 @@ class TestIPInformationManager: assert isinstance(instance, IPInformationManager) assert isinstance(ip_information_manager, IPInformationManager) - def test_create(self, ip_information_manager: IPInformationManager): - instance = ip_information_manager.create_dummy() - + def test_create( + self, + ip_geoname_factory: Callable[..., IPGeoname], + ip_information_manager: IPInformationManager, + ): + instance = ip_geoname_factory() assert isinstance(instance, IPInformation) res = ip_information_manager.fetch_ip_information(filter_ips=[instance.ip]) - assert res[0].model_dump_json() == instance.model_dump_json() def test_prefetch_geoname( diff --git a/tests/managers/thl/test_ledger/test_lm_accounts.py b/tests/managers/thl/test_ledger/test_lm_accounts.py index 57a2261..3af10e7 100644 --- a/tests/managers/thl/test_ledger/test_lm_accounts.py +++ b/tests/managers/thl/test_ledger/test_lm_accounts.py @@ -14,8 +14,10 @@ from generalresearch.managers.thl.ledger_manager.exceptions import ( LedgerAccountDoesntExistError, ) from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager -from generalresearch.models.custom_types import AccountType, Direction, UUIDStr +from generalresearch.models.custom_types import UUIDStr from generalresearch.models.thl.ledger import ( + AccountType, + Direction, LedgerAccount, LedgerEntry, ) diff --git a/tests/managers/thl/test_ledger/test_lm_tx_locks.py b/tests/managers/thl/test_ledger/test_lm_tx_locks.py index 9ecc1bc..166598e 100644 --- a/tests/managers/thl/test_ledger/test_lm_tx_locks.py +++ b/tests/managers/thl/test_ledger/test_lm_tx_locks.py @@ -4,10 +4,10 @@ import logging from collections.abc import Callable, Generator from datetime import UTC, datetime, timedelta from decimal import Decimal -from logging import LogCaptureFixture from typing import TYPE_CHECKING import pytest +from pytest import LogCaptureFixture from generalresearch.managers.thl.ledger_manager.conditions import ( generate_condition_mp_payment, diff --git a/tests/managers/thl/test_ledger/test_thl_lm_tx.py b/tests/managers/thl/test_ledger/test_thl_lm_tx.py index b0484ae..cda88da 100644 --- a/tests/managers/thl/test_ledger/test_thl_lm_tx.py +++ b/tests/managers/thl/test_ledger/test_thl_lm_tx.py @@ -128,10 +128,11 @@ class TestThlLedgerTxManager: thl_ledger_manager: ThlLedgerManager, ledger_manager: LedgerManager, session_manager: SessionManager, + product_factory: Callable[..., Product], ): delete_ledger_db() create_main_accounts() - product = product_manager.create_dummy( + product = product_factory( payout_config=PayoutConfig( payout_transformation=PayoutTransformation( f="payout_transformation_amt" diff --git a/tests/managers/thl/test_ledger/test_wallet.py b/tests/managers/thl/test_ledger/test_wallet.py index 1ee9bf9..dc1feec 100644 --- a/tests/managers/thl/test_ledger/test_wallet.py +++ b/tests/managers/thl/test_ledger/test_wallet.py @@ -22,8 +22,10 @@ if TYPE_CHECKING: @pytest.fixture() -def schrute_product(product_manager: ProductManager) -> Product: - return product_manager.create_dummy( +def schrute_product( + product_factory: Callable[..., Product], product_manager: ProductManager +) -> Product: + return product_factory( user_wallet_config=UserWalletConfig(enabled=True, amt=False), payout_config=PayoutConfig( payout_transformation=PayoutTransformation( diff --git a/tests/managers/thl/test_product.py b/tests/managers/thl/test_product.py index 644dc90..81e0122 100644 --- a/tests/managers/thl/test_product.py +++ b/tests/managers/thl/test_product.py @@ -24,8 +24,12 @@ if TYPE_CHECKING: class TestProductManagerGetMethods: - def test_get_by_uuid(self, product_manager: ProductManager): - product: Product = product_manager.create_dummy( + def test_get_by_uuid( + self, + product_manager: ProductManager, + product_factory: Callable[..., Product], + ): + product: Product = product_factory( product_id=uuid4().hex, team_id=uuid4().hex, name=f"Test Product ID #{uuid4().hex[:6]}", @@ -44,12 +48,14 @@ class TestProductManagerGetMethods: product_manager.get_by_uuid(product_uuid=uuid4().hex) assert "product not found" in str(cm.value) - def test_get_by_uuids(self, product_manager: ProductManager): + def test_get_by_uuids( + self, product_factory: Callable[..., Product], product_manager: ProductManager + ): cnt = 5 product_uuids = [uuid4().hex for _ in range(cnt)] for product_id in product_uuids: - product_manager.create_dummy( + product_factory( product_id=product_id, team_id=uuid4().hex, name=f"Test Product ID #{uuid4().hex[:6]}", @@ -69,8 +75,10 @@ class TestProductManagerGetMethods: product_manager.get_by_uuids(product_uuids=product_uuids + ["abc123"]) assert "invalid uuid" in str(cm.value) - def test_get_by_uuid_if_exists(self, product_manager: ProductManager): - product: Product = product_manager.create_dummy( + def test_get_by_uuid_if_exists( + self, product_factory: Callable[..., Product], product_manager: ProductManager + ): + product: Product = product_factory( product_id=uuid4().hex, team_id=uuid4().hex, name=f"Test Product ID #{uuid4().hex[:6]}", @@ -81,10 +89,12 @@ class TestProductManagerGetMethods: instance = product_manager.get_by_uuid_if_exists(product_uuid="abc123") assert instance == None - def test_get_by_uuids_if_exists(self, product_manager: ProductManager): + def test_get_by_uuids_if_exists( + self, product_factory: Callable[..., Product], product_manager: ProductManager + ): product_uuids = [uuid4().hex for _ in range(2)] for product_id in product_uuids: - product_manager.create_dummy( + product_factory( product_id=product_id, team_id=uuid4().hex, name=f"Test Product ID #{uuid4().hex[:6]}", @@ -113,13 +123,15 @@ class TestProductManagerGetMethods: # for instance in res: # assert isinstance(instance, Product) - def test_get_by_business_ids(self, product_manager: ProductManager): + def test_get_by_business_ids( + self, product_factory: Callable[..., Product], product_manager: ProductManager + ): business_ids = [uuid4().hex for _ in range(5)] product_manager.fetch_uuids(business_uuids=business_ids) for business_id in business_ids: - product_manager.create( + product_factory( product_id=uuid4().hex, team_id=None, business_id=business_id, @@ -131,8 +143,10 @@ class TestProductManagerGetMethods: class TestProductManagerCreation: - def test_base(self, product_manager: ProductManager): - instance = product_manager.create_dummy( + def test_base( + self, product_factory: Callable[..., Product], product_manager: ProductManager + ): + instance = product_factory( product_id=uuid4().hex, team_id=uuid4().hex, name=f"New Test Product {uuid4().hex[:6]}", @@ -235,10 +249,12 @@ class TestProductManager: assert instance.user_create_config.max_hourly_create_limit is None assert not instance.user_wallet_config.enabled - def test_sources(self, product_manager: ProductManager): + def test_sources( + self, product_factory: Callable[..., Product], product_manager: ProductManager + ): user_defined = [SourceConfig(name=Source.DYNATA, active=False)] sources_config = SourcesConfig(user_defined=user_defined) - p = product_manager.create_dummy(sources_config=sources_config) + p = product_factory(sources_config=sources_config) p2 = product_manager.get_by_uuid(p.id) @@ -250,7 +266,9 @@ class TestProductManager: assert not dynata.active assert all(x.active is True for x in p2.sources if x.name != Source.DYNATA) - def test_global_sources(self, product_manager: ProductManager): + def test_global_sources( + self, product_factory: Callable[..., Product], product_manager: ProductManager + ): sources_config = SupplyConfig( policies=[ SupplyPolicy( @@ -261,7 +279,7 @@ class TestProductManager: ) ] ) - p1 = product_manager.create_dummy(sources_config=sources_config) + p1 = product_factory(sources_config=sources_config) p2 = product_manager.get_by_uuid(p1.id) assert p1 == p2 @@ -277,8 +295,10 @@ class TestProductManager: p2 = product_manager.get_by_uuid(p1.id) assert p1 == p2 - def test_user_health_config(self, product_manager: ProductManager): - p = product_manager.create_dummy( + def test_user_health_config( + self, product_factory: Callable[..., Product], product_manager: ProductManager + ): + p = product_factory( user_health_config=UserHealthConfig(banned_countries=["ng", "in"]) ) @@ -288,10 +308,10 @@ class TestProductManager: assert p2.user_health_config.banned_countries == ["in", "ng"] assert p2.user_health_config.allow_ban_iphist - def test_profiling_config(self, product_manager: ProductManager): - p = product_manager.create_dummy( - profiling_config=ProfilingConfig(max_questions=1) - ) + def test_profiling_config( + self, product_factory: Callable[..., Product], product_manager: ProductManager + ): + p = product_factory(profiling_config=ProfilingConfig(max_questions=1)) p2 = product_manager.get_by_uuid(p.id) assert p == p2 @@ -335,8 +355,10 @@ class TestProductManager: class TestProductManagerUpdate: - def test_update(self, product_manager: ProductManager): - p = product_manager.create_dummy() + def test_update( + self, product_factory: Callable[..., Product], product_manager: ProductManager + ): + p = product_factory() p.name = "new name" p.enabled = False p.user_create_config = UserCreateConfig(min_hourly_create_limit=200) @@ -356,8 +378,10 @@ class TestProductManagerUpdate: class TestProductManagerCacheClear: - def test_cache_clear(self, product_manager: ProductManager): - p = product_manager.create_dummy() + def test_cache_clear( + self, product_factory: Callable[..., Product], product_manager: ProductManager + ): + p = product_factory() product_manager.get_by_uuid(product_uuid=p.id) product_manager.get_by_uuid(product_uuid=p.id) product_manager.pg_config.execute_write( diff --git a/tests/managers/thl/test_task_status.py b/tests/managers/thl/test_task_status.py index 9846ce0..4a401fa 100644 --- a/tests/managers/thl/test_task_status.py +++ b/tests/managers/thl/test_task_status.py @@ -40,18 +40,22 @@ finish3 = start3 + timedelta(minutes=5) @pytest.fixture(scope="session") -def bp1(product_manager: ProductManager) -> Product: +def bp1( + product_factory: Callable[..., Product], product_manager: ProductManager +) -> Product: # user wallet disabled, payout xform NULL - return product_manager.create_dummy( + return product_factory( user_wallet_config=UserWalletConfig(enabled=False), payout_config=PayoutConfig(), ) @pytest.fixture(scope="session") -def bp2(product_manager: ProductManager) -> Product: +def bp2( + product_factory: Callable[..., Product], product_manager: ProductManager +) -> Product: # user wallet disabled, payout xform 40% - return product_manager.create_dummy( + return product_factory( user_wallet_config=UserWalletConfig(enabled=False), payout_config=PayoutConfig( payout_transformation=PayoutTransformation( @@ -63,9 +67,11 @@ def bp2(product_manager: ProductManager) -> Product: @pytest.fixture(scope="session") -def bp3(product_manager: ProductManager) -> Product: +def bp3( + product_factory: Callable[..., Product], product_manager: ProductManager +) -> Product: # user wallet enabled, payout xform 50% - return product_manager.create_dummy( + return product_factory( user_wallet_config=UserWalletConfig(enabled=True), payout_config=PayoutConfig( payout_transformation=PayoutTransformation( diff --git a/tests/managers/thl/test_user_manager/test_base.py b/tests/managers/thl/test_user_manager/test_base.py index 4a9750e..c69f297 100644 --- a/tests/managers/thl/test_user_manager/test_base.py +++ b/tests/managers/thl/test_user_manager/test_base.py @@ -1,4 +1,5 @@ import logging +from collections.abc import Callable from datetime import UTC, datetime from random import randint from typing import TYPE_CHECKING @@ -7,9 +8,11 @@ from uuid import uuid4 import pytest from generalresearch.managers.thl.user_manager import ( - UserCreateNotAllowedError, get_bp_user_create_limit_hourly, ) +from generalresearch.managers.thl.user_manager.exceptions import ( + UserCreateNotAllowedError, +) from generalresearch.managers.thl.user_manager.mysql_user_manager import ( MysqlUserManager, ) @@ -152,11 +155,11 @@ class TestCreateUserManager: def test_create_user( self, - product_manager: ProductManager, + product_factory: Callable[..., Product], thl_web_rw: PostgresConfig, user_manager: UserManager, ): - product: Product = product_manager.create_dummy( + product: Product = product_factory( user_create_config=UserCreateConfig( min_hourly_create_limit=10, max_hourly_create_limit=69 ), @@ -195,11 +198,11 @@ class TestCreateUserManager: def test_create_user_integrity_error( self, - product_manager: ProductManager, user_manager: UserManager, + product_factory: Callable[..., Product], caplog, ): - product: Product = product_manager.create_dummy( + product: Product = product_factory( product_id=uuid4().hex, team_id=uuid4().hex, name=f"Test Product ID #{uuid4().hex[:6]}", @@ -241,10 +244,13 @@ class TestCreateUserManager: assert user1 == user2 def test_raise_allow_user_create( - self, product_manager: ProductManager, user_manager: UserManager + self, + product_manager: ProductManager, + user_manager: UserManager, + product_factory: Callable[..., Product], ): rand_num = randint(25, 200) - product: Product = product_manager.create_dummy( + product: Product = product_factory( product_id=uuid4().hex, team_id=uuid4().hex, name=f"Test Product ID #{uuid4().hex[:6]}", diff --git a/tests/managers/thl/test_wall_manager.py b/tests/managers/thl/test_wall_manager.py index 3215de8..58de7a2 100644 --- a/tests/managers/thl/test_wall_manager.py +++ b/tests/managers/thl/test_wall_manager.py @@ -10,7 +10,7 @@ import pytest from pydantic import PositiveInt from generalresearch.models.definitions import Source -from generalresearch.models.thl.session import ( +from generalresearch.models.thl.definitions import ( ReportValue, Status, StatusCode1, diff --git a/tests/models/gr/test_business.py b/tests/models/gr/test_business.py index 4e0b4e1..e942be5 100644 --- a/tests/models/gr/test_business.py +++ b/tests/models/gr/test_business.py @@ -551,7 +551,6 @@ class TestBusinessBalance: ledger_manager: LedgerManager, product_manager: ProductManager, start: datetime, - thl_web_rr: PostgresConfig, session_with_tx_factory: Callable[..., Session], delete_ledger_db: Callable[..., None], create_main_accounts: Callable[..., None], @@ -964,7 +963,6 @@ class TestBusinessBalance: ledger_manager: LedgerManager, product_manager: ProductManager, start: datetime, - thl_web_rr: PostgresConfig, payout_event_manager, session_with_tx_factory: Callable[..., None], delete_ledger_db: Callable[..., None], @@ -1194,7 +1192,6 @@ class TestBusinessMethods: def test_set_cache( self, gr_business: Business, - gr_db: PostgresConfig, thl_web_rr: PostgresConfig, client_no_amm: DaskClient, mnt_filepath: GRLDatasets, diff --git a/tests/models/gr/test_team.py b/tests/models/gr/test_team.py index a94d53f..0ca9b11 100644 --- a/tests/models/gr/test_team.py +++ b/tests/models/gr/test_team.py @@ -113,13 +113,13 @@ class TestTeam: assert gr_team.businesses is None - gr_team.prefetch_businesses(business_manager=gr_business_manager) + gr_team.prefetch_businesses(gr_business_manager=gr_business_manager) assert isinstance(gr_team.businesses, list) assert len(gr_team.businesses) == 0 team_manager.add_business(team=gr_team, business=business) assert len(gr_team.businesses) == 0 - gr_team.prefetch_businesses(business_manager=gr_business_manager) + gr_team.prefetch_businesses(gr_business_manager=gr_business_manager) assert len(gr_team.businesses) == 1 assert isinstance(gr_team.businesses[0], Business) assert gr_team.businesses[0].uuid == business.uuid @@ -163,12 +163,19 @@ class TestTeamMethods: mnt_gr_api_dir: Path, enriched_wall_merge: EnrichedWallMerge, enriched_session_merge: EnrichedSessionMerge, + product_manager: ProductManager, + gr_user_manager: GRUserManager, + gr_business_manager: BusinessManager, + gr_membership_manager: MembershipManager, ): client = gr_redis_config.create_redis_client() assert client.get(name=gr_team.cache_key) is None gr_team.set_cache( - pg_config=gr_db, + product_manager=product_manager, + gr_user_manager=gr_user_manager, + gr_business_manager=gr_business_manager, + gr_membership_manager=gr_membership_manager, thl_web_rr=thl_web_rr, redis_config=gr_redis_config, client=client_no_amm, @@ -193,6 +200,10 @@ class TestTeamMethods: mnt_gr_api_dir: Path, enriched_wall_merge: EnrichedWallMerge, enriched_session_merge: EnrichedSessionMerge, + product_manager: ProductManager, + gr_user_manager: GRUserManager, + gr_business_manager: BusinessManager, + gr_membership_manager: MembershipManager, ): from generalresearch.models.gr.team import Team @@ -200,7 +211,10 @@ class TestTeamMethods: membership_factory(team=gr_team, gr_user=gr_user) gr_team.set_cache( - pg_config=gr_db, + product_manager=product_manager, + gr_user_manager=gr_user_manager, + gr_business_manager=gr_business_manager, + gr_membership_manager=gr_membership_manager, thl_web_rr=thl_web_rr, redis_config=gr_redis_config, client=client_no_amm, @@ -239,6 +253,7 @@ class TestTeamMethods: mnt_filepath: GRLDatasets, mnt_gr_api_dir: Path, gr_team: Team, + product_manager: ProductManager, ): delete_df_collection(coll=wall_collection) @@ -267,7 +282,7 @@ class TestTeamMethods: ) gr_team.prebuild_enriched_session_parquet( - thl_pg_config=thl_web_rr, + product_manager=product_manager, ds=mnt_filepath, client=client_no_amm, mnt_gr_api=mnt_gr_api_dir, @@ -295,6 +310,7 @@ class TestTeamMethods: mnt_filepath: GRLDatasets, mnt_gr_api_dir: Path, gr_team: Team, + product_manager: ProductManager, ): delete_df_collection(coll=wall_collection) @@ -323,7 +339,7 @@ class TestTeamMethods: ) gr_team.prebuild_enriched_wall_parquet( - thl_pg_config=thl_web_rr, + product_manager=product_manager, ds=mnt_filepath, client=client_no_amm, mnt_gr_api=mnt_gr_api_dir, diff --git a/tests/models/thl/test_product.py b/tests/models/thl/test_product.py index 446b59f..f1050bb 100644 --- a/tests/models/thl/test_product.py +++ b/tests/models/thl/test_product.py @@ -64,9 +64,11 @@ class TestProduct: # We're not excluding anything here, only in the "*Out" variants assert "id_int" in res - def test_init_db(self, product_manager: ProductManager): + def test_init_db( + self, product_factory: Callable[..., Product], product_manager: ProductManager + ): # By default, just a Pydantic instance doesn't have an id_int - instance = product_manager.create_dummy() + instance = product_factory() assert isinstance(instance.id_int, int) res = instance.model_dump_json() diff --git a/tests/models/thl/test_user.py b/tests/models/thl/test_user.py index bc941d4..68b413c 100644 --- a/tests/models/thl/test_user.py +++ b/tests/models/thl/test_user.py @@ -18,6 +18,7 @@ if TYPE_CHECKING: 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.userhealth import AuditLog class TestUserUserID: @@ -621,12 +622,17 @@ class TestUserSerialization: class TestUserMethods: - def test_audit_log(self, user: User, audit_log_manager: AuditLogManager): + def test_audit_log( + self, + audit_log_factory: Callable[..., AuditLog], + 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 == [] - audit_log_manager.create_dummy(user_id=user.user_id) + audit_log_factory(user_id=user.user_id) user.prefetch_audit_log(audit_log_manager=audit_log_manager) assert len(user.audit_log) == 1 -- cgit v1.2.3