diff options
| author | Max Nanis | 2026-08-24 17:12:54 -0700 |
|---|---|---|
| committer | Max Nanis | 2026-08-24 17:12:54 -0700 |
| commit | 5682f48a96b0713929a4bb52e72eec5907d5dd32 (patch) | |
| tree | 3e5ba728145cf69f4ea0fde48e630b341c164cd0 | |
| parent | 979efd01b0c61d493d1fef3dff9885f5d38a8975 (diff) | |
| download | generalresearch-5682f48a96b0713929a4bb52e72eec5907d5dd32.tar.gz generalresearch-5682f48a96b0713929a4bb52e72eec5907d5dd32.zip | |
pytest fixture annotations, ruf manual edits
| -rw-r--r-- | tests/managers/thl/test_ledger/test_thl_lm_tx.py | 14 | ||||
| -rw-r--r-- | tests/managers/thl/test_ledger/test_user_txs.py | 23 | ||||
| -rw-r--r-- | tests/managers/thl/test_payout.py | 125 | ||||
| -rw-r--r-- | tests/managers/thl/test_survey.py | 8 | ||||
| -rw-r--r-- | tests/managers/thl/test_user_manager/test_base.py | 30 | ||||
| -rw-r--r-- | tests/managers/thl/test_user_streak.py | 4 | ||||
| -rw-r--r-- | tests/managers/thl/test_wall_manager.py | 6 | ||||
| -rw-r--r-- | tests/models/custom_types/test_aware_datetime.py | 2 | ||||
| -rw-r--r-- | tests/models/gr/test_business.py | 100 | ||||
| -rw-r--r-- | tests/models/test_currency.py | 114 |
10 files changed, 179 insertions, 247 deletions
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 90e5469..6fb0a0f 100644 --- a/tests/managers/thl/test_ledger/test_thl_lm_tx.py +++ b/tests/managers/thl/test_ledger/test_thl_lm_tx.py @@ -1181,7 +1181,7 @@ class TestThlLedgerManagerAdj: thl_lm.create_tx_task_complete(wall=w2, user=user, created=w2.started) status, status_code_1 = s1.determine_session_status() - thl_net, commission_amount, bp_pay, user_pay = s1.determine_payments() + _, _, bp_pay, user_pay = s1.determine_payments() session_manager.finish_with_status( session=s1, status=status, @@ -1251,7 +1251,7 @@ class TestThlLedgerManagerAdj: user=user, created=utc_hour_ago + timedelta(minutes=45), ) - new_status, new_payout, new_user_payout = s1.determine_new_status_and_payouts() + _, _, _ = s1.determine_new_status_and_payouts() s1.adjust_status() thl_lm.create_tx_bp_adjustment(session=s1) assert 380 == thl_lm.get_account_balance(bp_wallet_account) @@ -1261,13 +1261,13 @@ class TestThlLedgerManagerAdj: def test_create_tx_bp_adjustment_small( self, - user_factory, + user_factory: Callable[..., User], product_user_wallet_no, create_main_accounts, delete_ledger_db, - thl_lm, - lm, - utc_hour_ago, + thl_ledger_manager, + ledger_manager, + utc_hour_ago: datetime, currency, ): delete_ledger_db() @@ -1296,7 +1296,7 @@ class TestThlLedgerManagerAdj: session = Session(started=wall1.started, user=user, wall_events=[wall1]) status, status_code_1 = session.determine_session_status() - thl_net, commission_amount, bp_pay, user_pay = session.determine_payments() + _, _, bp_pay, user_pay = session.determine_payments() session.update( status=status, status_code_1=status_code_1, diff --git a/tests/managers/thl/test_ledger/test_user_txs.py b/tests/managers/thl/test_ledger/test_user_txs.py index 881eb03..a6bfa79 100644 --- a/tests/managers/thl/test_ledger/test_user_txs.py +++ b/tests/managers/thl/test_ledger/test_user_txs.py @@ -4,6 +4,7 @@ from decimal import Decimal from typing import TYPE_CHECKING from uuid import uuid4 +from generalresearch.config import GRLBaseSettings from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager from generalresearch.managers.thl.user_compensate import user_compensate from generalresearch.models.thl.definitions import ( @@ -47,7 +48,7 @@ def test_user_txs( s: Session = session_with_tx_factory(user=user, wall_req_cpi=Decimal("1.00")) - bribe_uuid = user_compensate( + user_compensate( ledger_manager=thl_lm, user=user, amount_int=100, @@ -213,13 +214,13 @@ def test_user_txs_rolling_balance( user_factory: Callable[..., User], product_amt_true: Product, create_main_accounts, - thl_lm, - lm, + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, delete_ledger_db: Callable[..., None], session_with_tx_factory, adj_to_fail_with_tx_factory, user_payout_event_manager, - settings: GRLSettings, + settings: GRLBaseSettings, ): """ Creates 3 $1.00 bonuses (postive), @@ -232,11 +233,11 @@ def test_user_txs_rolling_balance( create_main_accounts() user: User = user_factory(product=product_amt_true) - account = thl_lm.get_account_or_create_user_wallet(user) + account = thl_ledger_manager.get_account_or_create_user_wallet(user) for _ in range(3): user_compensate( - ledger_manager=thl_lm, + ledger_manager=thl_ledger_manager, user=user, amount_int=100, skip_flag_check=True, @@ -250,19 +251,19 @@ def test_user_txs_rolling_balance( payout_type=PayoutType.AMT_BONUS, request_data={}, ) - thl_lm.create_tx_user_payout_request( + thl_ledger_manager.create_tx_user_payout_request( user=user, payout_event=pe, ) for _ in range(3): user_compensate( - ledger_manager=thl_lm, + ledger_manager=thl_ledger_manager, user=user, amount_int=100, skip_flag_check=True, ) - txs = thl_lm.get_user_txs(user, page=1, size=10) + txs = thl_ledger_manager.get_user_txs(user, page=1, size=10) assert txs.transactions[0].balance_after == 100 assert txs.transactions[1].balance_after == 200 assert txs.transactions[2].balance_after == 300 @@ -273,7 +274,7 @@ def test_user_txs_rolling_balance( # Ascending order, get 2nd page, make sure the balances include # the previous txs. (will return last 3 txs) - txs = thl_lm.get_user_txs(user, page=2, size=4) + txs = thl_ledger_manager.get_user_txs(user, page=2, size=4) assert len(txs.transactions) == 3 assert txs.transactions[0].balance_after == 250 assert txs.transactions[1].balance_after == 350 @@ -281,7 +282,7 @@ def test_user_txs_rolling_balance( # Descending order, get 1st page. Will # return most recent 3 txs in desc order - txs = thl_lm.get_user_txs(user, page=1, size=3, order_by="-created") + txs = thl_ledger_manager.get_user_txs(user, page=1, size=3, order_by="-created") assert len(txs.transactions) == 3 assert txs.transactions[0].balance_after == 450 assert txs.transactions[1].balance_after == 350 diff --git a/tests/managers/thl/test_payout.py b/tests/managers/thl/test_payout.py index c4f20cd..39bbe6b 100644 --- a/tests/managers/thl/test_payout.py +++ b/tests/managers/thl/test_payout.py @@ -1,6 +1,8 @@ import io import logging import os +from collections.abc import Callable +from dask.distributed import Client as DaskClient from datetime import UTC, datetime, timedelta from decimal import Decimal from random import choice as rand_choice @@ -11,6 +13,7 @@ import pandas as pd import pytest from generalresearch.currency import USDCent +from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager from generalresearch.managers.thl.payout import UserPayoutEventManager from generalresearch.models.thl.definitions import PayoutStatus @@ -21,8 +24,10 @@ from generalresearch.models.thl.payout import ( UserPayoutEvent, ) from generalresearch.models.thl.product import Product +from generalresearch.models.gr.business import Business from generalresearch.models.thl.user import User from generalresearch.models.thl.wallet import PayoutType +from generalresearch.pg_helper import PostgresConfig logger = logging.getLogger() @@ -32,13 +37,13 @@ cashout_method_uuid = uuid4().hex class TestPayout: def test_get_by_uuid_and_create( self, - user, + user: User, user_payout_event_manager: UserPayoutEventManager, - thl_lm, - utc_now, + thl_ledger_manager: ThlLedgerManager, + utc_now: datetime, ): - user_account: LedgerAccount = thl_lm.get_account_or_create_user_wallet( - user=user + user_account: LedgerAccount = ( + thl_ledger_manager.get_account_or_create_user_wallet(user=user) ) pe1: UserPayoutEvent = user_payout_event_manager.create( @@ -58,11 +63,18 @@ class TestPayout: assert pe1 == pe2 - def test_update(self, user, user_payout_event_manager, lm, thl_lm, utc_now): + def test_update( + self, + user: User, + user_payout_event_manager, + ledger_manager: LedgerManager, + thl_ledger_manager: ThlLedgerManager, + utc_now: datetime, + ): from generalresearch.models.thl.definitions import PayoutStatus from generalresearch.models.thl.wallet import PayoutType - user_account = thl_lm.get_account_or_create_user_wallet(user=user) + user_account = thl_ledger_manager.get_account_or_create_user_wallet(user=user) pe1 = user_payout_event_manager.create( status=PayoutStatus.PENDING, @@ -90,15 +102,6 @@ class TestPayout: def test_create_bp_payout( self, - user, - thl_web_rr, - user_payout_event_manager, - lm, - thl_lm, - product, - brokerage_product_payout_event_manager, - delete_ledger_db, - create_main_accounts, ): # create_bp_payout_event does not get called directly. We have tests # for the ledger methods already @@ -107,11 +110,11 @@ class TestPayout: @pytest.fixture def pending_bp_pe( self, - thl_web_rw, - product, + thl_web_rw: PostgresConfig, + product: Product, thl_lm: ThlLedgerManager, brokerage_product_payout_event_manager, - utc_now, + utc_now: datetime, ) -> BrokerageProductPayoutEvent: account = thl_lm.get_account_or_create_bp_wallet(product=product) bp_pe = BrokerageProductPayoutEvent( @@ -140,14 +143,11 @@ class TestPayout: def test_create_bp_payout_quick_dupe( self, - user, - product, - thl_web_rw, + product: Product, brokerage_product_payout_event_manager, - thl_lm, - lm, - utc_now, - create_main_accounts, + thl_lm: ThlLedgerManager, + ledger_manager, + utc_now: datetime, pending_bp_pe, ): thl_lm.get_account_or_create_bp_wallet(product=product) @@ -170,19 +170,18 @@ class TestPayout: def test_filter( self, - thl_web_rw, - thl_lm, - lm, - product, - user, + thl_ledger_manager: ThlLedgerManager, + ledger_manager, + product: Product, + user: User, user_payout_event_manager, - utc_now, + utc_now: datetime, ): from generalresearch.models.thl.definitions import PayoutStatus from generalresearch.models.thl.wallet import PayoutType - user_account = thl_lm.get_account_or_create_user_wallet(user=user) - bp_account = thl_lm.get_account_or_create_bp_wallet(product=product) + user_account = thl_ledger_manager.get_account_or_create_user_wallet(user=user) + bp_account = thl_ledger_manager.get_account_or_create_bp_wallet(product=product) user_payout_event_manager.create( status=PayoutStatus.PENDING, @@ -267,11 +266,10 @@ class TestBusinessPayoutEventManager: self, brokerage_product_payout_event_manager, business_payout_event_manager, - delete_ledger_db, + delete_ledger_db: Callable[..., None], create_main_accounts, - thl_lm, - thl_web_rr, - product_factory, + thl_ledger_manager: ThlLedgerManager, + product_factory: Callable[..., Product], bp_payout_factory, business, ): @@ -279,7 +277,7 @@ class TestBusinessPayoutEventManager: create_main_accounts() p1: Product = product_factory(business=business) - thl_lm.get_account_or_create_bp_wallet(product=p1) + thl_ledger_manager.get_account_or_create_bp_wallet(product=p1) ach_id1 = uuid4().hex ach_id2 = uuid4().hex @@ -315,24 +313,22 @@ class TestBusinessPayoutEventManager: def test_update_ext_reference_ids( self, - brokerage_product_payout_event_manager, business_payout_event_manager, - delete_ledger_db, - create_main_accounts, + delete_ledger_db: Callable[..., None],, + create_main_accounts: Callable[..., None], thl_ledger_manager, - thl_web_rr, - product_factory, - bp_payout_factory, + thl_web_rr: PostgresConfig, + product_factory: Callable[..., Product], delete_df_collection, user_factory, ledger_collection, session_with_tx_factory, pop_ledger_merge, client_no_amm, - mnt_filepath, + mnt_filepath: GRLDataset, product_manager, - start, - business, + start: datetime, + business: Business, ): delete_ledger_db() create_main_accounts() @@ -766,7 +762,6 @@ class TestBusinessPayoutEventManager: thl_lm.get_account_or_create_bp_wallet(product=p3) ach_id1 = uuid4().hex - ach_id2 = uuid4().hex # Product 1: Complete, Payout, Recon.. s1 = session_with_tx_factory( @@ -915,8 +910,8 @@ class TestBusinessPayoutEventManager: self, product, mnt_filepath, - thl_lm, - client_no_amm, + thl_ledger_manager: ThlLedgerManager, + client_no_amm: DaskClient, thl_redis_config, payout_event_manager, brokerage_product_payout_event_manager, @@ -934,7 +929,7 @@ class TestBusinessPayoutEventManager: bp_payout_factory, adj_to_fail_with_tx_factory, thl_web_rr, - lm, + ledger_manager, product_manager, rm_ledger_collection, rm_pop_ledger_merge, @@ -980,7 +975,7 @@ class TestBusinessPayoutEventManager: pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) business.prebuild_balance( thl_pg_config=thl_web_rr, - lm=lm, + lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, pop_ledger=pop_ledger_merge, @@ -1032,7 +1027,7 @@ class TestBusinessPayoutEventManager: # sent to the Business business.prebuild_balance( thl_pg_config=thl_web_rr, - lm=lm, + lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, pop_ledger=pop_ledger_merge, @@ -1050,11 +1045,9 @@ class TestBusinessPayoutEventManager: def test_ach_tx_id_reference( self, - product, mnt_filepath, - thl_lm, + thl_ledger_manager, client_no_amm, - thl_redis_config, payout_event_manager, brokerage_product_payout_event_manager, business_payout_event_manager, @@ -1062,12 +1055,12 @@ class TestBusinessPayoutEventManager: create_main_accounts, delete_df_collection, ledger_collection, - business, + business: Business, user_factory, product_factory, session_with_tx_factory, pop_ledger_merge, - start, + start: datetime, bp_payout_factory, adj_to_fail_with_tx_factory, thl_web_rr, @@ -1088,9 +1081,9 @@ class TestBusinessPayoutEventManager: u1: User = user_factory(product=p1) u2: User = user_factory(product=p2) u3: User = user_factory(product=p3) - thl_lm.get_account_or_create_bp_wallet(product=p1) - thl_lm.get_account_or_create_bp_wallet(product=p2) - thl_lm.get_account_or_create_bp_wallet(product=p3) + thl_ledger_manager.get_account_or_create_bp_wallet(product=p1) + thl_ledger_manager.get_account_or_create_bp_wallet(product=p2) + thl_ledger_manager.get_account_or_create_bp_wallet(product=p3) ach_id1 = uuid4().hex ach_id2 = uuid4().hex @@ -1102,7 +1095,7 @@ class TestBusinessPayoutEventManager: wall_req_cpi=Decimal("7.50"), started=start + timedelta(days=1, hours=1 + iidx, minutes=1 + idx), ) - payout_event_manager.set_account_lookup_table(thl_lm=thl_lm) + payout_event_manager.set_account_lookup_table(thl_lm=thl_ledger_manager) rm_ledger_collection() rm_pop_ledger_merge() @@ -1121,7 +1114,7 @@ class TestBusinessPayoutEventManager: amount=USDCent(100_01), transaction_id=ach_id1, pm=product_manager, - thl_lm=thl_lm, + thl_lm=thl_ledger_manager, created=start + timedelta(days=2, hours=1), ) @@ -1142,7 +1135,7 @@ class TestBusinessPayoutEventManager: amount=USDCent(100_02), transaction_id=ach_id2, pm=product_manager, - thl_lm=thl_lm, + thl_lm=thl_ledger_manager, created=start + timedelta(days=4, hours=1), ) @@ -1155,7 +1148,7 @@ class TestBusinessPayoutEventManager: pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) business.prebuild_payouts( thl_pg_config=thl_web_rr, - thl_lm=thl_lm, + thl_lm=thl_ledger_manager, bpem=business_payout_event_manager, ) business.prebuild_balance( diff --git a/tests/managers/thl/test_survey.py b/tests/managers/thl/test_survey.py index 117a6aa..eec7c3b 100644 --- a/tests/managers/thl/test_survey.py +++ b/tests/managers/thl/test_survey.py @@ -174,9 +174,9 @@ class TestSurvey: qs = sorted(qs, key=lambda x: x.importance.task_count if x.importance else 0) # qd = {q.id: q for q in qs} - q = [x for x in qs if x.ext_question_id == "i:adhoc_13126"][0] + q = next(x for x in qs if x.ext_question_id == "i:adhoc_13126") q.explanation_template = "You have been diagnosed with: {answer}." - q = [x for x in qs if x.ext_question_id == "gr:gender"][0] + q = next(x for x in qs if x.ext_question_id == "gr:gender") q.explanation_template = "Your gender is {answer}." ecs = [] @@ -251,9 +251,9 @@ class TestSurveyStat: survey_stats.append(ss) print(len(survey_stats)) print(survey_stats[12].natural_key, survey_stats[2000].natural_key) - print(f"----a-----: {datetime.now().isoformat()}") + print(f"----a-----: {datetime.now(tz=UTC).isoformat()}") res = surveystat_manager.update_or_create(survey_stats) - print(f"----b-----: {datetime.now().isoformat()}") + print(f"----b-----: {datetime.now(tz=UTC).isoformat()}") assert len(res) == 20_000 return diff --git a/tests/managers/thl/test_user_manager/test_base.py b/tests/managers/thl/test_user_manager/test_base.py index 7d83c11..5235a0f 100644 --- a/tests/managers/thl/test_user_manager/test_base.py +++ b/tests/managers/thl/test_user_manager/test_base.py @@ -5,6 +5,7 @@ from uuid import uuid4 import pytest +from generalresearch.managers.thl.product import ProductManager from generalresearch.managers.thl.user_manager import ( UserCreateNotAllowedError, get_bp_user_create_limit_hourly, @@ -15,8 +16,10 @@ from generalresearch.managers.thl.user_manager.rate_limit import ( from generalresearch.managers.thl.user_manager.user_manager import ( UserManager, ) +from generalresearch.managers.thl.userhealth import AuditLogManager from generalresearch.models.thl.product import Product, UserCreateConfig from generalresearch.models.thl.user import User +from generalresearch.pg_helper import PostgresConfig logger = logging.getLogger() @@ -83,7 +86,7 @@ class TestUserManager: class TestBlockUserManager: - def test_block_user(self, product, user_manager: UserManager): + def test_block_user(self, product: Product, user_manager: UserManager): product_user_id = f"user-{uuid4().hex[:10]}" # mysql_user_manager to skip user creation limit check @@ -109,7 +112,9 @@ class TestBlockUserManager: user = user_manager.get_user(user_id=user.user_id) assert user.blocked - def test_block_user_whitelist(self, product, user_manager, thl_web_rw): + def test_block_user_whitelist( + self, product: Product, user_manager: UserManager, thl_web_rw: PostgresConfig + ): product_user_id = f"user-{uuid4().hex[:10]}" # mysql_user_manager to skip user creation limit check @@ -135,7 +140,12 @@ class TestBlockUserManager: class TestCreateUserManager: - def test_create_user(self, product_manager, thl_web_rw, user_manager): + def test_create_user( + self, + product_manager: ProductManager, + thl_web_rw: PostgresConfig, + user_manager: UserManager, + ): product: Product = product_manager.create_dummy( user_create_config=UserCreateConfig( min_hourly_create_limit=10, max_hourly_create_limit=69 @@ -172,7 +182,9 @@ class TestCreateUserManager: assert u2.user_id == user.user_id assert u2.uuid == user.uuid - def test_create_user_integrity_error(self, product_manager, user_manager, caplog): + def test_create_user_integrity_error( + self, product_manager, user_manager: UserManager, caplog + ): product: Product = product_manager.create_dummy( product_id=uuid4().hex, team_id=uuid4().hex, @@ -213,7 +225,9 @@ class TestCreateUserManager: assert user1 == user2 - def test_raise_allow_user_create(self, product_manager, user_manager): + def test_raise_allow_user_create( + self, product_manager: ProductManager, user_manager: UserManager + ): rand_num = randint(25, 200) product: Product = product_manager.create_dummy( product_id=uuid4().hex, @@ -250,7 +264,7 @@ class TestCreateUserManager: user_manager.user_manager_limiter.storage.clear(key=key) n = 0 - with pytest.raises(expected_exception=UserCreateNotAllowedError) as cm: + with pytest.raises(expected_exception=UserCreateNotAllowedError): for n, _ in enumerate(range(rl_value + 5)): user_manager.user_manager_limiter.raise_allow_user_create( product=product @@ -260,7 +274,9 @@ class TestCreateUserManager: class TestUserManagerMethods: - def test_audit_log(self, user_manager, user, audit_log_manager): + def test_audit_log( + self, user_manager: UserManager, user: User, audit_log_manager: AuditLogManager + ): from generalresearch.models.thl.userhealth import AuditLog res = audit_log_manager.filter_by_user_id(user_id=user.user_id) diff --git a/tests/managers/thl/test_user_streak.py b/tests/managers/thl/test_user_streak.py index be0729c..e87869f 100644 --- a/tests/managers/thl/test_user_streak.py +++ b/tests/managers/thl/test_user_streak.py @@ -192,11 +192,11 @@ def test_user_streak_complete_active(user_streak_manager, user, session_manager) streaks = user_streak_manager.get_user_streaks( user_id=user.user_id, country_iso="us" ) - streak = [ + streak = next( s for s in streaks if s.fulfillment == StreakFulfillment.COMPLETE and s.period == StreakPeriod.DAY - ][0] + ) assert streak == expected_streak # And now they complete today diff --git a/tests/managers/thl/test_wall_manager.py b/tests/managers/thl/test_wall_manager.py index 067a29e..1071627 100644 --- a/tests/managers/thl/test_wall_manager.py +++ b/tests/managers/thl/test_wall_manager.py @@ -233,7 +233,7 @@ class TestWallCacheManager: attempts = wall_cache_manager.get_attempts(user_id=user.user_id) assert len(attempts) == 1 - wall2 = wall_manager.create_dummy( + wall_manager.create_dummy( session_id=session.id, user_id=session.user_id, started=start2, @@ -260,7 +260,7 @@ class TestWallCacheManager: wall_cache_manager.update_attempts_redis_(attempts10000, user_id=user.user_id) session = session_manager.create_dummy(started=start3, user=user) - wall3 = wall_manager.create_dummy( + wall_manager.create_dummy( session_id=session.id, user_id=session.user_id, started=start3, @@ -274,5 +274,5 @@ class TestWallCacheManager: redis_key = wall_cache_manager.get_cache_key_(user_id=user.user_id) assert wall_cache_manager.redis_client.llen(redis_key) == 5000 - assert len(attempts) == 5000 + assert len(attempts) == 5_000 assert attempts[0].req_survey_id == "33333" diff --git a/tests/models/custom_types/test_aware_datetime.py b/tests/models/custom_types/test_aware_datetime.py index 7c45710..e8a5aa3 100644 --- a/tests/models/custom_types/test_aware_datetime.py +++ b/tests/models/custom_types/test_aware_datetime.py @@ -42,7 +42,7 @@ class TestAwareDatetimeISO: AwareDatetimeISOModel.model_validate_json(t.model_dump_json()) def test_no_tz(self): - dt = datetime(2023, 10, 10, 1, 1, 1) + dt = datetime(2023, 10, 10, 1, 1, 1) # noqa with pytest.raises(expected_exception=ValidationError): AwareDatetimeISOModel(dt=dt, dt_optional=None) diff --git a/tests/models/gr/test_business.py b/tests/models/gr/test_business.py index 9554028..48a7bb0 100644 --- a/tests/models/gr/test_business.py +++ b/tests/models/gr/test_business.py @@ -4,6 +4,7 @@ import os from collections.abc import Callable from datetime import UTC, datetime, timedelta from decimal import Decimal +from pathlib import Path from uuid import uuid4 import pandas as pd @@ -17,6 +18,8 @@ from distributed.utils_test import ( from pytest import approx from generalresearch.currency import USDCent +from generalresearch.incite.base import GRLDatasets +from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge from generalresearch.managers.gr.business import BusinessBankAccountManager from generalresearch.managers.gr.team import TeamManager from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager @@ -31,7 +34,7 @@ from generalresearch.models.gr.business import ( BusinessBankAccount, BusinessContact, ) -from generalresearch.models.gr.team import Membership, Team +from generalresearch.models.gr.team import Team from generalresearch.models.thl.finance import ( BusinessBalances, ProductBalances, @@ -121,7 +124,7 @@ class TestBusiness: ledger_manager: LedgerManager, thl_ledger_manager: ThlLedgerManager, business_payout_event_manager: BusinessPayoutEventManager, - bp_payout_factory: Callable[..., Bus], + bp_payout_factory: Callable[..., BusinessPayoutEventManager], start: datetime, user_factory: Callable[..., User], session_with_tx_factory: Callable[..., Session], @@ -244,8 +247,8 @@ class TestBusiness: assert business.products[0].uuid == p1.uuid # Add two more, but list is still one until we prefetch - p2 = product_factory(business=business) - p3 = product_factory(business=business) + product_factory(business=business) + product_factory(business=business) assert len(business.products) == 1 business.prefetch_products(thl_pg_config=thl_web_rr) @@ -262,11 +265,11 @@ class TestBusiness: def test_balance( self, business: Business, - mnt_filepath, + mnt_filepath: GRLDatasets, client_no_amm: DaskClient, thl_web_rr: PostgresConfig, ledger_manager: LedgerManager, - pop_ledger_merge, + pop_ledger_merge: PopLedgerMerge, ): assert business.balance is None @@ -357,6 +360,7 @@ class TestBusiness: thl_lm=thl_ledger_manager, bpem=business_payout_event_manager, ) + assert isinstance(business.payouts, list) assert len(business.payouts) == 1 assert len(business.payouts[0].bp_payouts) == 2 assert sum([p.amount for p in business.payouts]) == 246 @@ -407,6 +411,7 @@ class TestBusiness: bpem=business_payout_event_manager, ) + assert isinstance(business.payouts, list) assert len(business.payouts) == 1 assert len(business.payouts[0].bp_payouts) == 3 assert business.payouts_total == USDCent(76) @@ -417,9 +422,9 @@ class TestBusiness: business: Business, thl_web_rr: PostgresConfig, thl_ledger_manager: ThlLedgerManager, - mnt_filepath, + mnt_filepath: GRLDatasets, client_no_amm: DaskClient, - pop_ledger_merge, + pop_ledger_merge: PopLedgerMerge, ): assert business.pop_financial is None business.prebuild_pop_financial( @@ -479,20 +484,15 @@ class TestBusinessBalance: product_factory: Callable[..., Product], user_factory: Callable[..., User], mnt_filepath, - bp_payout_factory: Callable[..., BrokerageProductPayoutEvent], - thl_ledger_manager: ThlLedgerManager, ledger_manager: LedgerManager, - duration: timedelta, - offset, start: datetime, thl_web_rr: PostgresConfig, - payout_event_manager: PayoutEventManager, session_with_tx_factory: Callable[..., Session], delete_ledger_db: Callable[..., None], create_main_accounts: Callable[..., None], client_no_amm: DaskClient, ledger_collection, - pop_ledger_merge, + pop_ledger_merge: PopLedgerMerge, delete_df_collection: Callable[..., None], ): delete_ledger_db() @@ -544,15 +544,10 @@ class TestBusinessBalance: business: Business, product_factory: Callable[..., Product], user_factory: Callable[..., User], - mnt_filepath, - bp_payout_factory: Callable[..., BrokerageProductPayoutEvent], + mnt_filepath: GRLDatasets, ledger_manager: LedgerManager, - thl_ledger_manager: ThlLedgerManager, - duration: timedelta, - offset, start: datetime, thl_web_rr: PostgresConfig, - payout_event_manager: PayoutEventManager, session_with_tx_factory: Callable[..., Session], delete_ledger_db: Callable[..., None], create_main_accounts: Callable[..., None], @@ -633,12 +628,10 @@ class TestBusinessBalance: business: Business, product_factory: Callable[..., Product], user_factory: Callable[..., User], - mnt_filepath, + mnt_filepath: GRLDatasets, bp_payout_factory: Callable[..., BrokerageProductPayoutEvent], thl_ledger_manager: ThlLedgerManager, ledger_manager: LedgerManager, - duration: timedelta, - offset, start: datetime, thl_web_rr: PostgresConfig, payout_event_manager: PayoutEventManager, @@ -709,12 +702,10 @@ class TestBusinessBalance: business: Business, product_factory: Callable[..., Product], user_factory: Callable[..., User], - mnt_filepath, + mnt_filepath: GRLDatasets, bp_payout_factory: Callable[..., BrokerageProductPayoutEvent], ledger_manager: LedgerManager, thl_ledger_manager: ThlLedgerManager, - duration: timedelta, - offset, start: datetime, thl_web_rr: PostgresConfig, payout_event_manager: PayoutEventManager, @@ -724,8 +715,6 @@ class TestBusinessBalance: ledger_collection, task_adj_collection, pop_ledger_merge: PopLedgerMerge, - wall_manager: WallManager, - session_manager: SessionManager, adj_to_fail_with_tx_factory: Callable[..., None], delete_df_collection: Callable[..., None], ): @@ -837,12 +826,9 @@ class TestBusinessBalance: def test_neg_balance_cache( self, - product: Product, - mnt_filepath, + mnt_filepath: GRLDatasets, thl_ledger_manager: ThlLedgerManager, client_no_amm: DaskClient, - thl_redis_config: RedisConfig, - brokerage_product_payout_event_manager, delete_ledger_db: Callable[..., None], create_main_accounts: Callable[..., None], delete_df_collection: Callable[..., None], @@ -913,13 +899,14 @@ class TestBusinessBalance: business.prebuild_balance( thl_pg_config=thl_web_rr, - lm=lm, + lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, pop_ledger=pop_ledger_merge, ) # Check Product 1 + assert isinstance(business.balance, BusinessBalances) pb1 = business.balance.product_balances[0] assert pb1.product_id == p1.uuid assert pb1.payout == 71 @@ -962,25 +949,21 @@ class TestBusinessBalance: business: Business, product_factory: Callable[..., Product], user_factory: Callable[..., User], - mnt_filepath, + mnt_filepath: GRLDatasets, bp_payout_factory: Callable[..., BrokerageProductPayoutEvent], thl_ledger_manager: ThlLedgerManager, ledger_manager: LedgerManager, - duration: timedelta, - offset, start: datetime, thl_web_rr: PostgresConfig, payout_event_manager, session_with_tx_factory, - delete_ledger_db, + delete_ledger_db: Callable[..., None], create_main_accounts: Callable[..., None], client_no_amm: DaskClient, ledger_collection, task_adj_collection, - pop_ledger_merge, - wall_manager: WallManager, - session_manager: SessionManager, - adj_to_fail_with_tx_factory, + pop_ledger_merge: PopLedgerMerge, + adj_to_fail_with_tx_factory: Callable[..., None], delete_df_collection: Callable[..., None], ): """ @@ -1205,22 +1188,21 @@ class TestBusinessMethods: gr_db: PostgresConfig, thl_web_rr: PostgresConfig, client_no_amm: DaskClient, - mnt_filepath, + mnt_filepath: GRLDatasets, ledger_manager: LedgerManager, thl_ledger_manager: ThlLedgerManager, business_payout_event_manager, product_factory: Callable[..., Product], - membership_factory: Callable[..., Membership], team: Team, session_with_tx_factory: Callable[..., Session], user_factory: Callable[..., User], ledger_collection, - pop_ledger_merge, + pop_ledger_merge: PopLedgerMerge, utc_60days_ago: datetime, delete_ledger_db: Callable[..., None], create_main_accounts: Callable[..., None], gr_redis_config: RedisConfig, - mnt_gr_api_dir, + mnt_gr_api_dir: Path, ): assert gr_redis.get(name=business.cache_key) is None @@ -1263,17 +1245,13 @@ class TestBusinessMethods: def test_set_cache_business( self, - gr_user: GRUser, business: Business, - gr_user_token: GRUserToken, - gr_redis: RedisConfig, gr_db: PostgresConfig, thl_web_rr: PostgresConfig, product_factory: Callable[..., Product], team: Team, - membership_factory: Callable[..., Membership], client_no_amm: DaskClient, - mnt_filepath, + mnt_filepath: GRLDatasets, ledger_manager: LedgerManager, thl_ledger_manager: ThlLedgerManager, business_payout_event_manager, @@ -1283,10 +1261,10 @@ class TestBusinessMethods: session_with_tx_factory: Callable[..., Session], ledger_collection, team_manager: TeamManager, - pop_ledger_merge, + pop_ledger_merge: PopLedgerMerge, gr_redis_config: RedisConfig, utc_60days_ago: datetime, - mnt_gr_api_dir, + mnt_gr_api_dir: Path, ): from generalresearch.models.gr.business import Business @@ -1336,12 +1314,16 @@ class TestBusinessMethods: gr_redis_config=gr_redis_config, ) + assert isinstance(business2, Business) assert business.model_dump_json() == business2.model_dump_json() # assert isinstance(business2.balance, BusinessBalances) + assert isinstance(business2.products, list) + assert isinstance(business2.teams, list) assert p1.uuid in [p.uuid for p in business2.products] assert len(business2.teams) == 1 assert team.uuid in [t.uuid for t in business2.teams] + assert isinstance(business2.balance, BusinessBalances) assert business2.balance.payout == 48 assert business2.balance.balance == 48 assert business2.balance.net == 48 @@ -1354,27 +1336,26 @@ class TestBusinessMethods: assert len(business2.bp_accounts) == 1 assert len(business2.bp_accounts) == len(business2.product_uuids) + assert isinstance(business2.pop_financial, list) assert len(business2.pop_financial) == 1 assert business2.pop_financial[0].payout == business2.balance.payout assert business2.pop_financial[0].net == business2.balance.net def test_prebuild_enriched_session_parquet( self, - event_report_request, enriched_session_merge, client_no_amm: DaskClient, wall_collection, session_collection, thl_web_rr: PostgresConfig, - session_report_request, user_factory: Callable[..., User], start: datetime, session_factory: Callable[..., Session], product_factory: Callable[..., Product], delete_df_collection: Callable[..., None], business: Business, - mnt_filepath, - mnt_gr_api_dir, + mnt_filepath: GRLDatasets, + mnt_gr_api_dir: Path, ): delete_df_collection(coll=wall_collection) @@ -1418,22 +1399,19 @@ class TestBusinessMethods: def test_prebuild_enriched_wall_parquet( self, - event_report_request, - enriched_session_merge, enriched_wall_merge, client_no_amm: DaskClient, wall_collection, session_collection, thl_web_rr: PostgresConfig, - session_report_request, user_factory: Callable[..., User], start: datetime, session_factory: Callable[..., Session], product_factory: Callable[..., Product], delete_df_collection: Callable[..., None], business: Business, - mnt_filepath, - mnt_gr_api_dir, + mnt_filepath: GRLDatasets, + mnt_gr_api_dir: Path, ): delete_df_collection(coll=wall_collection) diff --git a/tests/models/test_currency.py b/tests/models/test_currency.py index 1102717..9bc2216 100644 --- a/tests/models/test_currency.py +++ b/tests/models/test_currency.py @@ -8,11 +8,12 @@ from random import randint import pytest +from generalresearch.currency import USDCent, USDMill, format_usd_cent + class TestUSDCentModel: def test_construct_int(self): - from generalresearch.currency import USDCent for _ in range(100): int_val = randint(0, 999_999) @@ -20,10 +21,9 @@ class TestUSDCentModel: assert int_val == instance def test_construct_float(self): - from generalresearch.currency import USDCent + float_val: float = 10.6789 with pytest.warns(expected_warning=Warning) as record: - float_val: float = 10.6789 instance = USDCent(float_val) assert len(record) == 1 @@ -34,10 +34,9 @@ class TestUSDCentModel: assert instance == 10 def test_construct_decimal(self): - from generalresearch.currency import USDCent + decimal_val: Decimal = Decimal("10.0") with pytest.warns(expected_warning=Warning) as record: - decimal_val: Decimal = Decimal("10.0") instance = USDCent(decimal_val) assert len(record) == 1 @@ -50,8 +49,8 @@ class TestUSDCentModel: assert instance == 10 # Now with rounding + decimal_val: Decimal = Decimal("10.6789") with pytest.warns(Warning) as record: - decimal_val: Decimal = Decimal("10.6789") instance = USDCent(decimal_val) assert len(record) == 1 @@ -64,16 +63,12 @@ class TestUSDCentModel: assert instance == 10 def test_construct_negative(self): - from generalresearch.currency import USDCent - with pytest.raises(expected_exception=ValueError) as cm: USDCent(-1) assert "USDCent not be less than zero" in str(cm.value) def test_operation_add(self): - from generalresearch.currency import USDCent - - for i in range(100): + for _ in range(100): int_val1 = randint(0, 999_999) int_val2 = randint(0, 999_999) @@ -83,9 +78,7 @@ class TestUSDCentModel: assert int_val1 + int_val2 == instance1 + instance2 def test_operation_subtract(self): - from generalresearch.currency import USDCent - - for i in range(100): + for _ in range(100): int_val1 = randint(500_000, 999_999) int_val2 = randint(0, 499_999) @@ -95,9 +88,7 @@ class TestUSDCentModel: assert int_val1 - int_val2 == instance1 - instance2 def test_operation_subtract_to_neg(self): - from generalresearch.currency import USDCent - - for i in range(100): + for _ in range(100): int_val = randint(0, 999_999) instance = USDCent(int_val) @@ -107,9 +98,7 @@ class TestUSDCentModel: assert "USDCent not be less than zero" in str(cm.value) def test_operation_multiply(self): - from generalresearch.currency import USDCent - - for i in range(100): + for _ in range(100): int_val1 = randint(0, 999_999) int_val2 = randint(0, 999_999) @@ -119,15 +108,11 @@ class TestUSDCentModel: assert int_val1 * int_val2 == instance1 * instance2 def test_operation_div(self): - from generalresearch.currency import USDCent - with pytest.raises(ValueError) as cm: - USDCent(10) / 2 + _ = USDCent(10) / 2 assert "Division not allowed for USDCent" in str(cm.value) def test_operation_result_type(self): - from generalresearch.currency import USDCent - int_val = randint(1, 999_999) instance = USDCent(int_val) @@ -141,8 +126,6 @@ class TestUSDCentModel: assert isinstance(res_multipy, USDCent) def test_operation_partner_add(self): - from generalresearch.currency import USDCent - int_val = randint(1, 999_999) instance = USDCent(int_val) @@ -159,18 +142,14 @@ class TestUSDCentModel: _ = instance + True def test_abs(self): - from generalresearch.currency import USDCent - - for i in range(100): + for _ in range(100): int_val = abs(randint(0, 999_999)) instance = abs(USDCent(int_val)) assert int_val == instance def test_str(self): - from generalresearch.currency import USDCent - - for i in range(100): + for _ in range(100): int_val = randint(0, 999_999) instance = USDCent(int_val) @@ -180,8 +159,6 @@ class TestUSDCentModel: """There is no correct answer here, but we at least want to make sure that a USDCent is returned """ - from generalresearch.currency import USDCent - res = USDCent(10) // 1.2 assert not isinstance(res, USDCent) assert isinstance(res, float) @@ -206,18 +183,14 @@ class TestUSDCentModel: class TestUSDMillModel: def test_construct_int(self): - from generalresearch.currency import USDMill - - for i in range(100): + for _ in range(100): int_val = randint(0, 999_999) instance = USDMill(int_val) assert int_val == instance def test_construct_float(self): - from generalresearch.currency import USDMill - + float_val: float = 10.6789 with pytest.warns(expected_warning=Warning) as record: - float_val: float = 10.6789 instance = USDMill(float_val) assert len(record) == 1 @@ -228,10 +201,8 @@ class TestUSDMillModel: assert instance == 10 def test_construct_decimal(self): - from generalresearch.currency import USDMill - + decimal_val: Decimal = Decimal("10.0") with pytest.warns(expected_warning=Warning) as record: - decimal_val: Decimal = Decimal("10.0") instance = USDMill(decimal_val) assert len(record) == 1 @@ -244,12 +215,11 @@ class TestUSDMillModel: assert instance == 10 # Now with rounding + decimal_val: Decimal = Decimal("10.6789") with pytest.warns(expected_warning=Warning) as record: - decimal_val: Decimal = Decimal("10.6789") instance = USDMill(decimal_val) - assert isinstance(instance, USDMill) - + assert isinstance(instance, USDMill) assert len(record) == 1 assert ( "USDMill init with a Decimal. Rounding behavior may be unexpected" @@ -260,16 +230,12 @@ class TestUSDMillModel: assert instance == 10 def test_construct_negative(self): - from generalresearch.currency import USDMill - with pytest.raises(expected_exception=ValueError) as cm: USDMill(-1) assert "USDMill not be less than zero" in str(cm.value) def test_operation_add(self): - from generalresearch.currency import USDMill - - for i in range(100): + for _ in range(100): int_val1 = randint(0, 999_999) int_val2 = randint(0, 999_999) @@ -279,9 +245,7 @@ class TestUSDMillModel: assert int_val1 + int_val2 == instance1 + instance2 def test_operation_subtract(self): - from generalresearch.currency import USDMill - - for i in range(100): + for _ in range(100): int_val1 = randint(500_000, 999_999) int_val2 = randint(0, 499_999) @@ -291,21 +255,17 @@ class TestUSDMillModel: assert int_val1 - int_val2 == instance1 - instance2 def test_operation_subtract_to_neg(self): - from generalresearch.currency import USDMill - - for i in range(100): + for _ in range(100): int_val = randint(0, 999_999) instance = USDMill(int_val) with pytest.raises(expected_exception=ValueError) as cm: - instance - USDMill(1_000_000) + _ = instance - USDMill(1_000_000) assert "USDMill not be less than zero" in str(cm.value) def test_operation_multiply(self): - from generalresearch.currency import USDMill - - for i in range(100): + for _ in range(100): int_val1 = randint(0, 999_999) int_val2 = randint(0, 999_999) @@ -315,15 +275,11 @@ class TestUSDMillModel: assert int_val1 * int_val2 == instance1 * instance2 def test_operation_div(self): - from generalresearch.currency import USDMill - with pytest.raises(ValueError) as cm: - USDMill(10) / 2 + _ = USDMill(10) / 2 assert "Division not allowed for USDMill" in str(cm.value) def test_operation_result_type(self): - from generalresearch.currency import USDMill - int_val = randint(1, 999_999) instance = USDMill(int_val) @@ -337,36 +293,30 @@ class TestUSDMillModel: assert isinstance(res_multipy, USDMill) def test_operation_partner_add(self): - from generalresearch.currency import USDMill - int_val = randint(1, 999_999) instance = USDMill(int_val) with pytest.raises(expected_exception=AssertionError): - instance + 0.10 + _ = instance + 0.10 with pytest.raises(expected_exception=AssertionError): - instance + Decimal(".10") + _ = instance + Decimal(".10") with pytest.raises(expected_exception=AssertionError): - instance + "9.9" + _ = instance + "9.9" with pytest.raises(expected_exception=AssertionError): - instance + True + _ = instance + True def test_abs(self): - from generalresearch.currency import USDMill - - for i in range(100): + for _ in range(100): int_val = abs(randint(0, 999_999)) instance = abs(USDMill(int_val)) assert int_val == instance def test_str(self): - from generalresearch.currency import USDMill - - for i in range(100): + for _ in range(100): int_val = randint(0, 999_999) instance = USDMill(int_val) @@ -376,8 +326,6 @@ class TestUSDMillModel: """There is no correct answer here, but we at least want to make sure that a USDMill is returned """ - from generalresearch.currency import USDCent, USDMill - res = USDMill(10) // 1.2 assert not isinstance(res, USDMill) assert isinstance(res, float) @@ -402,11 +350,7 @@ class TestUSDMillModel: class TestNegativeFormatting: def test_pos(self): - from generalresearch.currency import format_usd_cent - assert "-$987.65" == format_usd_cent(-98765) def test_neg(self): - from generalresearch.currency import format_usd_cent - assert "-$123.45" == format_usd_cent(-12345) |
