aboutsummaryrefslogtreecommitdiff
path: root/tests
diff options
context:
space:
mode:
authorMax Nanis2026-08-24 17:12:54 -0700
committerMax Nanis2026-08-24 17:12:54 -0700
commit5682f48a96b0713929a4bb52e72eec5907d5dd32 (patch)
tree3e5ba728145cf69f4ea0fde48e630b341c164cd0 /tests
parent979efd01b0c61d493d1fef3dff9885f5d38a8975 (diff)
downloadgeneralresearch-5682f48a96b0713929a4bb52e72eec5907d5dd32.tar.gz
generalresearch-5682f48a96b0713929a4bb52e72eec5907d5dd32.zip
pytest fixture annotations, ruf manual edits
Diffstat (limited to 'tests')
-rw-r--r--tests/managers/thl/test_ledger/test_thl_lm_tx.py14
-rw-r--r--tests/managers/thl/test_ledger/test_user_txs.py23
-rw-r--r--tests/managers/thl/test_payout.py125
-rw-r--r--tests/managers/thl/test_survey.py8
-rw-r--r--tests/managers/thl/test_user_manager/test_base.py30
-rw-r--r--tests/managers/thl/test_user_streak.py4
-rw-r--r--tests/managers/thl/test_wall_manager.py6
-rw-r--r--tests/models/custom_types/test_aware_datetime.py2
-rw-r--r--tests/models/gr/test_business.py100
-rw-r--r--tests/models/test_currency.py114
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)