diff options
Diffstat (limited to 'tests/managers')
| -rw-r--r-- | tests/managers/gr/test_business.py | 38 | ||||
| -rw-r--r-- | tests/managers/thl/test_contest/test_milestone.py | 2 | ||||
| -rw-r--r-- | tests/managers/thl/test_contest/test_raffle.py | 10 | ||||
| -rw-r--r-- | tests/managers/thl/test_ipinfo.py | 20 | ||||
| -rw-r--r-- | tests/managers/thl/test_ledger/test_lm_accounts.py | 4 | ||||
| -rw-r--r-- | tests/managers/thl/test_ledger/test_lm_tx_locks.py | 2 | ||||
| -rw-r--r-- | tests/managers/thl/test_ledger/test_thl_lm_tx.py | 3 | ||||
| -rw-r--r-- | tests/managers/thl/test_ledger/test_wallet.py | 6 | ||||
| -rw-r--r-- | tests/managers/thl/test_product.py | 76 | ||||
| -rw-r--r-- | tests/managers/thl/test_task_status.py | 18 | ||||
| -rw-r--r-- | tests/managers/thl/test_user_manager/test_base.py | 20 | ||||
| -rw-r--r-- | tests/managers/thl/test_wall_manager.py | 2 |
12 files changed, 124 insertions, 77 deletions
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, |
