aboutsummaryrefslogtreecommitdiff
path: root/tests
diff options
context:
space:
mode:
authorMax Nanis2026-08-26 09:17:39 -0700
committerMax Nanis2026-08-26 09:17:39 -0700
commit47ea200eac0eaa7bef02f6ebb05de9afad5ee0d7 (patch)
tree0721a0145bf9db795ef442ca211caf9db2fad420 /tests
parent3b4059135be47f7752a08e4277a85f9e57ceaa9d (diff)
downloadgeneralresearch-47ea200eac0eaa7bef02f6ebb05de9afad5ee0d7.tar.gz
generalresearch-47ea200eac0eaa7bef02f6ebb05de9afad5ee0d7.zip
Ruff morning!
Diffstat (limited to 'tests')
-rw-r--r--tests/managers/leaderboard.py158
-rw-r--r--tests/managers/test_events.py65
-rw-r--r--tests/managers/test_lucid.py9
-rw-r--r--tests/managers/test_userpid.py2
-rw-r--r--tests/managers/thl/test_buyer.py11
-rw-r--r--tests/managers/thl/test_cashout_method.py67
-rw-r--r--tests/managers/thl/test_category.py109
-rw-r--r--tests/managers/thl/test_harmonized_uqa.py2
-rw-r--r--tests/managers/thl/test_ipinfo.py57
-rw-r--r--tests/managers/thl/test_ledger/test_lm_tx_locks.py37
-rw-r--r--tests/managers/thl/test_ledger/test_wallet.py2
-rw-r--r--tests/managers/thl/test_payout.py13
-rw-r--r--tests/managers/thl/test_product.py45
-rw-r--r--tests/managers/thl/test_product_prod.py18
-rw-r--r--tests/managers/thl/test_profiling/test_question.py25
-rw-r--r--tests/managers/thl/test_profiling/test_schema.py15
-rw-r--r--tests/managers/thl/test_profiling/test_uqa.py1
-rw-r--r--tests/managers/thl/test_profiling/test_user_upk.py20
-rw-r--r--tests/managers/thl/test_session_manager.py51
-rw-r--r--tests/managers/thl/test_survey.py64
-rw-r--r--tests/managers/thl/test_survey_penalty.py15
-rw-r--r--tests/managers/thl/test_task_adjustment.py195
-rw-r--r--tests/managers/thl/test_task_status.py111
-rw-r--r--tests/managers/thl/test_user_manager/test_base.py15
-rw-r--r--tests/managers/thl/test_user_manager/test_mysql.py27
-rw-r--r--tests/managers/thl/test_user_manager/test_redis.py46
-rw-r--r--tests/managers/thl/test_user_manager/test_user_fetch.py12
-rw-r--r--tests/managers/thl/test_user_manager/test_user_metadata.py33
-rw-r--r--tests/managers/thl/test_user_streak.py32
-rw-r--r--tests/managers/thl/test_userhealth.py135
-rw-r--r--tests/managers/thl/test_wall_manager.py58
-rw-r--r--tests/models/custom_types/test_dsn.py2
-rw-r--r--tests/models/custom_types/test_therest.py2
-rw-r--r--tests/models/dynata/test_eligbility.py2
-rw-r--r--tests/models/gr/test_authentication.py126
-rw-r--r--tests/models/gr/test_base.py2
-rw-r--r--tests/models/gr/test_business.py134
-rw-r--r--tests/models/gr/test_team.py14
38 files changed, 1098 insertions, 634 deletions
diff --git a/tests/managers/leaderboard.py b/tests/managers/leaderboard.py
index 3d1818b..d97714d 100644
--- a/tests/managers/leaderboard.py
+++ b/tests/managers/leaderboard.py
@@ -1,6 +1,9 @@
+from __future__ import annotations
+
import os
import time
import zoneinfo
+from collections.abc import Callable
from datetime import UTC, datetime
from decimal import Decimal
from uuid import uuid4
@@ -19,10 +22,11 @@ from generalresearch.models.thl.product import (
PayoutConfig,
PayoutTransformation,
PayoutTransformationPercentArgs,
- product: Product,
+ Product,
)
from generalresearch.models.thl.session import Session
from generalresearch.models.thl.user import User
+from generalresearch.redis_helper import RedisConfig
# random uuid for leaderboard tests
product_id = uuid4().hex
@@ -44,7 +48,9 @@ def session_factory():
def _create_session(
- product_user_id="aaa", country_iso="us", user_payout=Decimal("1.00")
+ product_user_id: str = "aaa",
+ country_iso: str = "us",
+ user_payout: Decimal = Decimal("1.00"),
):
user = User(
product_id=product_id,
@@ -74,59 +80,67 @@ def _create_session(
@pytest.fixture(scope="function")
-def setup_leaderboards(thl_redis):
- complete_count = {
- "aaa": 10,
- "bbb": 6,
- "ccc": 6,
- "ddd": 6,
- "eee": 2,
- "fff": 1,
- "ggg": 1,
- }
- sum_payout = {"aaa": 345, "bbb": 100, "ccc": 100}
- max_payout = sum_payout
- country_iso = "us"
- for freq in [
- LeaderboardFrequency.DAILY,
- LeaderboardFrequency.WEEKLY,
- LeaderboardFrequency.MONTHLY,
- ]:
- m = LeaderboardManager(
- redis_client=thl_redis,
- board_code=LeaderboardCode.COMPLETE_COUNT,
- freq=freq,
- product_id=product_id,
- country_iso=country_iso,
- within_time=datetime(2025, 2, 5, 12, 12, 12),
- )
- thl_redis.delete(m.key)
- thl_redis.zadd(m.key, complete_count)
- m = LeaderboardManager(
- redis_client=thl_redis,
- board_code=LeaderboardCode.SUM_PAYOUTS,
- freq=freq,
- product_id=product_id,
- country_iso=country_iso,
- within_time=datetime(2025, 2, 5, 12, 12, 12),
- )
- thl_redis.delete(m.key)
- thl_redis.zadd(m.key, sum_payout)
- m = LeaderboardManager(
- redis_client=thl_redis,
- board_code=LeaderboardCode.LARGEST_PAYOUT,
- freq=freq,
- product_id=product_id,
- country_iso=country_iso,
- within_time=datetime(2025, 2, 5, 12, 12, 12),
- )
- thl_redis.delete(m.key)
- thl_redis.zadd(m.key, max_payout)
+def setup_leaderboards(thl_redis: RedisConfig) -> Callable[..., None]:
+
+ def _inner():
+ complete_count = {
+ "aaa": 10,
+ "bbb": 6,
+ "ccc": 6,
+ "ddd": 6,
+ "eee": 2,
+ "fff": 1,
+ "ggg": 1,
+ }
+ sum_payout = {"aaa": 345, "bbb": 100, "ccc": 100}
+ max_payout = sum_payout
+ country_iso = "us"
+ for freq in [
+ LeaderboardFrequency.DAILY,
+ LeaderboardFrequency.WEEKLY,
+ LeaderboardFrequency.MONTHLY,
+ ]:
+ m = LeaderboardManager(
+ redis_client=thl_redis,
+ board_code=LeaderboardCode.COMPLETE_COUNT,
+ freq=freq,
+ product_id=product_id,
+ country_iso=country_iso,
+ within_time=datetime(2025, 2, 5, 12, 12, 12, tzinfo=UTC),
+ )
+ thl_redis.delete(m.key)
+ thl_redis.zadd(m.key, complete_count)
+ m = LeaderboardManager(
+ redis_client=thl_redis,
+ board_code=LeaderboardCode.SUM_PAYOUTS,
+ freq=freq,
+ product_id=product_id,
+ country_iso=country_iso,
+ within_time=datetime(2025, 2, 5, 12, 12, 12, tzinfo=UTC),
+ )
+ thl_redis.delete(m.key)
+ thl_redis.zadd(m.key, sum_payout)
+ m = LeaderboardManager(
+ redis_client=thl_redis,
+ board_code=LeaderboardCode.LARGEST_PAYOUT,
+ freq=freq,
+ product_id=product_id,
+ country_iso=country_iso,
+ within_time=datetime(2025, 2, 5, 12, 12, 12, tzinfo=UTC),
+ )
+ thl_redis.delete(m.key)
+ thl_redis.zadd(m.key, max_payout)
+
+ return _inner
class TestLeaderboards:
- def test_leaderboard_manager(self, setup_leaderboards, thl_redis):
+ def test_leaderboard_manager(
+ self, setup_leaderboards: Callable[..., None], thl_redis: RedisConfig
+ ):
+ setup_leaderboards()
+
country_iso = "us"
board_code = LeaderboardCode.COMPLETE_COUNT
freq = LeaderboardFrequency.DAILY
@@ -136,7 +150,7 @@ class TestLeaderboards:
freq=freq,
product_id=product_id,
country_iso=country_iso,
- within_time=datetime(2025, 2, 5, 0, 0, 0),
+ within_time=datetime(2025, 2, 5, 0, 0, 0, tzinfo=UTC),
)
lb = m.get_leaderboard()
assert lb.period_start_local == datetime(
@@ -164,7 +178,11 @@ class TestLeaderboards:
LeaderboardRow(bpuid="ggg", rank=6, value=1),
]
- def test_leaderboard_manager_bpuid(self, setup_leaderboards, thl_redis):
+ def test_leaderboard_manager_bpuid(
+ self, setup_leaderboards: Callable[..., None], thl_redis: RedisConfig
+ ):
+ setup_leaderboards()
+
country_iso = "us"
board_code = LeaderboardCode.COMPLETE_COUNT
freq = LeaderboardFrequency.DAILY
@@ -174,7 +192,7 @@ class TestLeaderboards:
freq=freq,
product_id=product_id,
country_iso=country_iso,
- within_time=datetime(2025, 2, 5, 12, 12, 12),
+ within_time=datetime(2025, 2, 5, 12, 12, 12, tzinfo=UTC),
)
lb = m.get_leaderboard(bp_user_id="fff", limit=1)
@@ -191,7 +209,14 @@ class TestLeaderboards:
lb.censor()
assert lb.rows[0].bpuid == "ee*"
- def test_leaderboard_hit(self, setup_leaderboards, session_factory, thl_redis):
+ def test_leaderboard_hit(
+ self,
+ setup_leaderboards: Callable[..., None],
+ session_factory: Callable[..., Session],
+ thl_redis: RedisConfig,
+ ):
+ setup_leaderboards()
+
hit_leaderboards(redis_client=thl_redis, session=session_factory())
for freq in [
@@ -205,7 +230,7 @@ class TestLeaderboards:
freq=freq,
product_id=product_id,
country_iso="us",
- within_time=datetime(2025, 2, 5, 12, 12, 12),
+ within_time=datetime(2025, 2, 5, 12, 12, 12, tzinfo=UTC),
)
lb = m.get_leaderboard(limit=1)
assert lb.row_count == 7
@@ -216,7 +241,7 @@ class TestLeaderboards:
freq=freq,
product_id=product_id,
country_iso="us",
- within_time=datetime(2025, 2, 5, 12, 12, 12),
+ within_time=datetime(2025, 2, 5, 12, 12, 12, tzinfo=UTC),
)
lb = m.get_leaderboard(limit=1)
assert lb.row_count == 3
@@ -227,15 +252,20 @@ class TestLeaderboards:
freq=freq,
product_id=product_id,
country_iso="us",
- within_time=datetime(2025, 2, 5, 12, 12, 12),
+ within_time=datetime(2025, 2, 5, 12, 12, 12, tzinfo=UTC),
)
lb = m.get_leaderboard(limit=1)
assert lb.row_count == 3
assert lb.rows == [LeaderboardRow(bpuid="aaa", rank=1, value=345 + 100)]
def test_leaderboard_hit_new_row(
- self, setup_leaderboards, session_factory, thl_redis
+ self,
+ setup_leaderboards: Callable[..., None],
+ session_factory: Callable[..., None],
+ thl_redis: RedisConfig,
):
+ setup_leaderboards()
+
session = session_factory(product_user_id="zzz")
hit_leaderboards(redis_client=thl_redis, session=session)
m = LeaderboardManager(
@@ -244,24 +274,20 @@ class TestLeaderboards:
freq=LeaderboardFrequency.DAILY,
product_id=product_id,
country_iso="us",
- within_time=datetime(2025, 2, 5, 12, 12, 12),
+ within_time=datetime(2025, 2, 5, 12, 12, 12, tzinfo=UTC),
)
lb = m.get_leaderboard()
assert lb.row_count == 8
assert LeaderboardRow(bpuid="zzz", value=1, rank=6) in lb.rows
- def test_leaderboard_country(self, thl_redis):
+ def test_leaderboard_country(self, thl_redis: RedisConfig):
m = LeaderboardManager(
redis_client=thl_redis,
board_code=LeaderboardCode.COMPLETE_COUNT,
freq=LeaderboardFrequency.DAILY,
product_id=product_id,
country_iso="jp",
- within_time=datetime(
- 2025,
- 2,
- 1,
- ),
+ within_time=datetime(2025, 2, 1, tzinfo=UTC),
)
lb = m.get_leaderboard()
assert lb.row_count == 0
diff --git a/tests/managers/test_events.py b/tests/managers/test_events.py
index a6d3a6b..cb32275 100644
--- a/tests/managers/test_events.py
+++ b/tests/managers/test_events.py
@@ -1,6 +1,9 @@
+from __future__ import annotations
+
import math
import random
import time
+from collections.abc import Callable
from datetime import UTC, datetime, timedelta
from decimal import Decimal
from functools import partial
@@ -9,7 +12,8 @@ from uuid import uuid4
import pytest
-from generalresearch.managers.events import EventSubscriber
+from generalresearch.managers.events import EventManager, EventSubscriber
+from generalresearch.managers.thl.product import ProductManager
from generalresearch.models import Source
from generalresearch.models.events import (
AggregateBySource,
@@ -21,22 +25,23 @@ from generalresearch.models.legacy.bucket import Bucket
from generalresearch.models.thl.definitions import Status, StatusCode1
from generalresearch.models.thl.session import Session, Wall
from generalresearch.models.thl.user import User
+from generalresearch.redis_helper import RedisConfig
# We don't need anything in the db, so not using the db fixtures
@pytest.fixture(scope="function")
-def product_id(product_manager):
+def product_id(product_manager: ProductManager) -> str:
return uuid4().hex
@pytest.fixture(scope="function")
-def user_factory(product_id):
+def user_factory(product_id: str):
return partial(create_dummy, product_id=product_id)
@pytest.fixture(scope="function")
-def event_subscriber(thl_redis_config: RedisConfig, product_id):
- return EventSubscriber(redis_config=thl_redis_config: RedisConfig, product_id=product_id)
+def event_subscriber(thl_redis_config: RedisConfig, product_id: str) -> EventSubscriber:
+ return EventSubscriber(redis_config=thl_redis_config, product_id=product_id)
def create_dummy(
@@ -53,7 +58,7 @@ def create_dummy(
class TestActiveUsers:
- def test_run_empty(self, event_manager, product_id):
+ def test_run_empty(self, event_manager: EventManager, product_id: str):
res = event_manager.get_user_stats(product_id)
assert res == {
"active_users_last_1h": 0,
@@ -62,7 +67,12 @@ class TestActiveUsers:
"in_progress_users": 0,
}
- def test_run(self, event_manager, product_id, user_factory):
+ def test_run(
+ self,
+ event_manager: EventManager,
+ product_id: str,
+ user_factory: Callable[..., User],
+ ):
event_manager.clear_global_user_stats()
user1: User = user_factory()
@@ -89,6 +99,8 @@ class TestActiveUsers:
# Create a 2nd user in another product
product_id2 = uuid4().hex
user2: User = user_factory(product_id=product_id2)
+ assert isinstance(user2, User)
+ assert isinstance(user2.created, datetime)
# Change to say user was created >24 hrs ago
user2.created = user2.created - timedelta(hours=25)
event_manager.handle_user(user2)
@@ -115,7 +127,12 @@ class TestActiveUsers:
"in_progress_users": 0,
}
- def test_inprogress(self, event_manager, product_id, user_factory):
+ def test_inprogress(
+ self,
+ event_manager: EventSubscriber,
+ product_id: str,
+ user_factory: Callable[..., User],
+ ):
event_manager.clear_global_user_stats()
user1: User = user_factory()
user2: User = user_factory()
@@ -138,7 +155,12 @@ class TestActiveUsers:
res = event_manager.get_user_stats(product_id)
assert res["in_progress_users"] == 1
- def test_expiry(self, event_manager, product_id, user_factory):
+ def test_expiry(
+ self,
+ event_manager: EventManager,
+ product_id: str,
+ user_factory: Callable[..., User],
+ ):
event_manager.clear_global_user_stats()
user1: User = user_factory()
event_manager.handle_user(user1)
@@ -166,7 +188,7 @@ class TestActiveUsers:
class TestSessionStats:
- def test_run_empty(self, event_manager, product_id):
+ def test_run_empty(self, event_manager: EventManager, product_id: str):
res = event_manager.get_session_stats(product_id)
assert res == {
"session_enters_last_1h": 0,
@@ -185,7 +207,14 @@ class TestSessionStats:
"session_fail_avg_loi_last_24h": None,
}
- def test_run(self, event_manager, product_id, user_factory: Callable[..., User], utc_now, utc_hour_ago):
+ def test_run(
+ self,
+ event_manager: EventManager,
+ product_id: str,
+ user_factory: Callable[..., User],
+ utc_now: datetime,
+ utc_hour_ago: datetime,
+ ):
event_manager.clear_global_session_stats()
user: User = user_factory()
@@ -306,7 +335,7 @@ class TestSessionStats:
class TestTaskStatsManager:
- def test_empty(self, event_manager):
+ def test_empty(self, event_manager: EventManager):
event_manager.clear_task_stats()
assert event_manager.get_task_stats_raw() == {
"live_task_count": AggregateBySource(total=0),
@@ -320,7 +349,7 @@ class TestTaskStatsManager:
assert sm.data.task_created_count_last_24h.total == 0
assert sm.data.live_tasks_max_payout.value is None
- def test(self, event_manager):
+ def test(self, event_manager: EventManager):
event_manager.clear_task_stats()
event_manager.set_source_task_stats(
source=Source.TESTING,
@@ -445,12 +474,12 @@ class TestTaskStatsManager:
class TestChannelsSubscriptions:
def test_stats_worker(
self,
- event_manager,
- event_subscriber,
- product_id,
+ event_manager: EventManager,
+ event_subscriber: EventSubscriber,
+ product_id: str,
user_factory: Callable[..., User],
- utc_hour_ago,
- utc_now,
+ utc_hour_ago: datetime,
+ utc_now: datetime,
):
event_manager.clear_stats()
assert event_subscriber.pubsub
diff --git a/tests/managers/test_lucid.py b/tests/managers/test_lucid.py
index 654b58d..20dca22 100644
--- a/tests/managers/test_lucid.py
+++ b/tests/managers/test_lucid.py
@@ -1,6 +1,9 @@
+from __future__ import annotations
+
import pytest
from generalresearch.managers.lucid.profiling import get_profiling_library
+from generalresearch.pg_helper import PostgresConfig
qids = ["42", "43", "45", "97", "120", "639", "15297"]
@@ -8,9 +11,9 @@ qids = ["42", "43", "45", "97", "120", "639", "15297"]
class TestLucidProfiling:
@pytest.mark.skip
- def test_get_library(self, thl_web_rr):
+ def test_get_library(self, thl_web_rr: PostgresConfig):
pks = [(qid, "us", "eng") for qid in qids]
- qs = get_profiling_library(thl_web_rr: PostgresConfig, pks=pks)
+ qs = get_profiling_library(thl_web_rr, pks=pks)
assert len(qids) == len(qs)
# just making sure this doesn't raise errors
@@ -19,5 +22,5 @@ class TestLucidProfiling:
# a lot will fail parsing because they have no options or the options are blank
# just asserting that we get some back
- qs = get_profiling_library(thl_web_rr: PostgresConfig, country_iso="mx", language_iso="spa")
+ qs = get_profiling_library(thl_web_rr, country_iso="mx", language_iso="spa")
assert len(qs) > 100
diff --git a/tests/managers/test_userpid.py b/tests/managers/test_userpid.py
index 36c2de9..e74e40b 100644
--- a/tests/managers/test_userpid.py
+++ b/tests/managers/test_userpid.py
@@ -1,3 +1,5 @@
+from __future__ import annotations
+
import pytest
from pydantic import MySQLDsn
diff --git a/tests/managers/thl/test_buyer.py b/tests/managers/thl/test_buyer.py
index 69ea105..6776ab3 100644
--- a/tests/managers/thl/test_buyer.py
+++ b/tests/managers/thl/test_buyer.py
@@ -1,3 +1,8 @@
+from __future__ import annotations
+
+from collections.abc import Callable
+
+from generalresearch.managers.thl.buyer import BuyerManager
from generalresearch.models import Source
@@ -5,10 +10,12 @@ class TestBuyer:
def test(
self,
- delete_buyers_surveys,
- buyer_manager,
+ delete_buyers_surveys: Callable[..., None],
+ buyer_manager: BuyerManager,
):
+ delete_buyers_surveys()
+
bs = buyer_manager.bulk_get_or_create(source=Source.TESTING, codes=["a", "b"])
assert len(bs) == 2
buyer_a = bs[0]
diff --git a/tests/managers/thl/test_cashout_method.py b/tests/managers/thl/test_cashout_method.py
index ee52188..451d3e0 100644
--- a/tests/managers/thl/test_cashout_method.py
+++ b/tests/managers/thl/test_cashout_method.py
@@ -1,5 +1,14 @@
+from __future__ import annotations
+
+from collections.abc import Callable
+
import pytest
+from generalresearch.config import GRLBaseSettings
+from generalresearch.managers.thl.cashout_method import (
+ CashoutMethodManager,
+)
+from generalresearch.models.thl.user import User
from generalresearch.models.thl.wallet import PayoutType
from generalresearch.models.thl.wallet.cashout_method import (
CashMailCashoutMethodData,
@@ -13,15 +22,26 @@ from test_utils.managers.cashout_methods import (
class TestTangoCashoutMethods:
- def test_create_and_get(self, cashout_method_manager, setup_cashoutmethod_db):
+ def test_create_and_get(
+ self,
+ cashout_method_manager: CashoutMethodManager,
+ setup_cashoutmethod_db: Callable[..., None],
+ ):
+ setup_cashoutmethod_db()
+
res = cashout_method_manager.filter(payout_types=[PayoutType.TANGO])
assert len(res) == 2
- cm = [x for x in res if x.ext_id == "U025035"][0]
+ cm = next(x for x in res if x.ext_id == "U025035")
assert EXAMPLE_TANGO_CASHOUT_METHODS[0] == cm
def test_user(
- self, cashout_method_manager, user_with_wallet, setup_cashoutmethod_db
+ self,
+ cashout_method_manager: CashoutMethodManager,
+ user_with_wallet: User,
+ setup_cashoutmethod_db: Callable[..., None],
):
+ setup_cashoutmethod_db()
+
res = cashout_method_manager.get_cashout_methods(user_with_wallet)
# This user ONLY has the two tango cashout methods, no AMT
assert len(res) == 2
@@ -29,19 +49,31 @@ class TestTangoCashoutMethods:
class TestAMTCashoutMethods:
- def test_create_and_get(self, cashout_method_manager, setup_cashoutmethod_db):
+ def test_create_and_get(
+ self,
+ settings: GRLBaseSettings,
+ cashout_method_manager: CashoutMethodManager,
+ setup_cashoutmethod_db: Callable[..., None],
+ ):
+ setup_cashoutmethod_db()
+
res = cashout_method_manager.filter(payout_types=[PayoutType.AMT])
assert len(res) == 2
- cm = [x for x in res if x.name == "AMT Assignment"][0]
- assert AMT_ASSIGNMENT_CASHOUT_METHOD == cm
+ cm = next(x for x in res if x.name == "AMT Assignment")
+ assert settings.amt_assignment_cashout_method_id == cm
- cm = [x for x in res if x.name == "AMT Bonus"][0]
- assert AMT_BONUS_CASHOUT_METHOD == cm
+ cm = next(x for x in res if x.name == "AMT Bonus")
+ assert settings.amt_bonus_cashout_method_id == cm
def test_user(
- self, cashout_method_manager, user_with_wallet_amt, setup_cashoutmethod_db
+ self,
+ cashout_method_manager: CashoutMethodManager,
+ user_with_wallet_amt: User,
+ setup_cashoutmethod_db: Callable[..., None],
):
+ setup_cashoutmethod_db()
+
res = cashout_method_manager.get_cashout_methods(user_with_wallet_amt)
# This user has the 2 tango, plus amt bonus & assignment
assert len(res) == 4
@@ -49,14 +81,22 @@ class TestAMTCashoutMethods:
class TestUserCashoutMethods:
- def test(self, cashout_method_manager, user_with_wallet, delete_cashoutmethod_db):
+ def test(
+ self,
+ cashout_method_manager: CashoutMethodManager,
+ user_with_wallet: User,
+ delete_cashoutmethod_db: Callable[..., None],
+ ):
delete_cashoutmethod_db()
res = cashout_method_manager.get_cashout_methods(user_with_wallet)
assert len(res) == 0
def test_cash_in_mail(
- self, cashout_method_manager, user_with_wallet, delete_cashoutmethod_db
+ self,
+ cashout_method_manager: CashoutMethodManager,
+ user_with_wallet: User,
+ delete_cashoutmethod_db: Callable[..., None],
):
delete_cashoutmethod_db()
@@ -95,7 +135,10 @@ class TestUserCashoutMethods:
assert len(res) == 2
def test_paypal(
- self, cashout_method_manager, user_with_wallet, delete_cashoutmethod_db
+ self,
+ cashout_method_manager: CashoutMethodManager,
+ user_with_wallet: User,
+ delete_cashoutmethod_db: Callable[..., None],
):
delete_cashoutmethod_db()
diff --git a/tests/managers/thl/test_category.py b/tests/managers/thl/test_category.py
index ad0f07b..ec52aae 100644
--- a/tests/managers/thl/test_category.py
+++ b/tests/managers/thl/test_category.py
@@ -1,12 +1,18 @@
+from __future__ import annotations
+
+from collections.abc import Callable
+
import pytest
+from generalresearch.managers.thl.category import CategoryManager
from generalresearch.models.thl.category import Category
+from generalresearch.pg_helper import PostgresConfig
class TestCategory:
@pytest.fixture
- def beauty_fitness(self, thl_web_rw):
+ def beauty_fitness(self, thl_web_rw: PostgresConfig) -> Category:
return Category(
uuid="12c1e96be82c4642a07a12a90ce6f59e",
@@ -16,72 +22,83 @@ class TestCategory:
)
@pytest.fixture
- def hair_care(self, beauty_fitness):
+ def hair_care(self, beauty_fitness: Category) -> Category:
return Category(
uuid="dd76c4b565d34f198dad3687326503d6",
adwords_vertical_id="146",
label="Hair Care",
- path="/Beauty & Fitness/Hair Care",
+ path=f"{beauty_fitness.path}/Hair Care",
)
@pytest.fixture
- def hair_loss(self, hair_care):
+ def hair_loss(self, hair_care: Category) -> Category:
return Category(
uuid="aacff523c8e246888215611ec3b823c0",
adwords_vertical_id="235",
label="Hair Loss",
- path="/Beauty & Fitness/Hair Care/Hair Loss",
+ path=f"{hair_care}/Hair Loss",
)
@pytest.fixture
def category_data(
- self, category_manager, thl_web_rw, beauty_fitness, hair_care, hair_loss
- ):
- cats = [beauty_fitness, hair_care, hair_loss]
- data = [x.model_dump(mode="json") for x in cats]
- # We need the parent pk's to set the parent_id. So insert all without a parent,
- # then pull back all pks and map to the parents as parsed by the parent_path
- query = """
- INSERT INTO marketplace_category
- (uuid, adwords_vertical_id, label, path)
- VALUES
- (%(uuid)s, %(adwords_vertical_id)s, %(label)s, %(path)s)
- ON CONFLICT (uuid) DO NOTHING;
- """
- with thl_web_rw.make_connection() as conn:
- with conn.cursor() as c:
- c.executemany(query=query, params_seq=data)
- conn.commit()
-
- res = thl_web_rw.execute_sql_query("SELECT id, path FROM marketplace_category")
- path_id = {x["path"]: x["id"] for x in res}
- data = [
- {"id": path_id[c.path], "parent_id": path_id[c.parent_path]}
- for c in cats
- if c.parent_path
- ]
- query = """
- UPDATE marketplace_category
- SET parent_id = %(parent_id)s
- WHERE id = %(id)s;
- """
- with thl_web_rw.make_connection() as conn:
- with conn.cursor() as c:
- c.executemany(query=query, params_seq=data)
- conn.commit()
-
- category_manager.populate_caches()
+ self,
+ category_manager: CategoryManager,
+ thl_web_rw: PostgresConfig,
+ beauty_fitness: Category,
+ hair_care: Category,
+ hair_loss: Category,
+ ) -> Callable[..., None]:
+
+ def _inner():
+ cats = [beauty_fitness, hair_care, hair_loss]
+ data = [x.model_dump(mode="json") for x in cats]
+ # We need the parent pk's to set the parent_id. So insert all without a parent,
+ # then pull back all pks and map to the parents as parsed by the parent_path
+ query = """
+ INSERT INTO marketplace_category
+ (uuid, adwords_vertical_id, label, path)
+ VALUES
+ (%(uuid)s, %(adwords_vertical_id)s, %(label)s, %(path)s)
+ ON CONFLICT (uuid) DO NOTHING;
+ """
+ with thl_web_rw.make_connection() as conn:
+ with conn.cursor() as c:
+ c.executemany(query=query, params_seq=data)
+ conn.commit()
+
+ res = thl_web_rw.execute_sql_query(
+ "SELECT id, path FROM marketplace_category"
+ )
+ path_id = {x["path"]: x["id"] for x in res}
+ data = [
+ {"id": path_id[c.path], "parent_id": path_id[c.parent_path]}
+ for c in cats
+ if c.parent_path
+ ]
+ query = """
+ UPDATE marketplace_category
+ SET parent_id = %(parent_id)s
+ WHERE id = %(id)s;
+ """
+ with thl_web_rw.make_connection() as conn:
+ with conn.cursor() as c:
+ c.executemany(query=query, params_seq=data)
+ conn.commit()
+
+ category_manager.populate_caches()
+
+ return _inner
def test(
self,
- category_data,
- category_manager,
- beauty_fitness,
- hair_care,
- hair_loss,
+ category_data: Callable[..., None],
+ category_manager: CategoryManager,
+ beauty_fitness: Category,
):
+ category_data()
+
# category_manager on init caches the category info. This rarely/never changes so this is fine,
# but now that tests get run on a new db each time, the category_manager is inited before
# the fixtures run. so category_manager's cache needs to be rerun
diff --git a/tests/managers/thl/test_harmonized_uqa.py b/tests/managers/thl/test_harmonized_uqa.py
index 81ac080..84eeb56 100644
--- a/tests/managers/thl/test_harmonized_uqa.py
+++ b/tests/managers/thl/test_harmonized_uqa.py
@@ -1,3 +1,5 @@
+from __future__ import annotations
+
from datetime import UTC, datetime
import pytest
diff --git a/tests/managers/thl/test_ipinfo.py b/tests/managers/thl/test_ipinfo.py
index c89312b..48b9efd 100644
--- a/tests/managers/thl/test_ipinfo.py
+++ b/tests/managers/thl/test_ipinfo.py
@@ -1,3 +1,5 @@
+from collections.abc import Callable
+
import faker
from generalresearch.managers.thl.ipinfo import (
@@ -5,14 +7,22 @@ from generalresearch.managers.thl.ipinfo import (
IPGeonameManager,
IPInformationManager,
)
-from generalresearch.models.thl.ipinfo import IPGeoname, IPInformation
+from generalresearch.models.thl.ipinfo import (
+ GeoIPInformation,
+ IPGeoname,
+ IPInformation,
+)
+from generalresearch.pg_helper import PostgresConfig
+from generalresearch.redis_helper import RedisConfig
fake = faker.Faker()
class TestIPGeonameManager:
- def test_init(self, thl_web_rr: PostgresConfig, ip_geoname_manager: IPGeonameManager):
+ def test_init(
+ self, thl_web_rr: PostgresConfig, ip_geoname_manager: IPGeonameManager
+ ):
instance = IPGeonameManager(pg_config=thl_web_rr)
assert isinstance(instance, IPGeonameManager)
@@ -31,7 +41,9 @@ class TestIPGeonameManager:
class TestIPInformationManager:
- def test_init(self, thl_web_rr: PostgresConfig, ip_information_manager: IPInformationManager):
+ def test_init(
+ self, thl_web_rr: PostgresConfig, ip_information_manager: IPInformationManager
+ ):
instance = IPInformationManager(pg_config=thl_web_rr)
assert isinstance(instance, IPInformationManager)
assert isinstance(ip_information_manager, IPInformationManager)
@@ -45,7 +57,12 @@ class TestIPInformationManager:
assert res[0].model_dump_json() == instance.model_dump_json()
- def test_prefetch_geoname(self, ip_information, ip_geoname, thl_web_rr):
+ def test_prefetch_geoname(
+ self,
+ ip_information: IPInformation,
+ ip_geoname: IPGeoname,
+ thl_web_rr: PostgresConfig,
+ ):
assert isinstance(ip_information, IPInformation)
assert ip_information.geoname_id == ip_geoname.geoname_id
@@ -62,11 +79,16 @@ class TestGeoIpInfoManager:
thl_redis_config: RedisConfig,
geoipinfo_manager: GeoIpInfoManager,
):
- instance = GeoIpInfoManager(pg_config=thl_web_rr: PostgresConfig, redis_config=thl_redis_config)
+ instance = GeoIpInfoManager(pg_config=thl_web_rr, redis_config=thl_redis_config)
assert isinstance(instance, GeoIpInfoManager)
assert isinstance(geoipinfo_manager, GeoIpInfoManager)
- def test_multi(self, ip_information_factory, ip_geoname, geoipinfo_manager):
+ def test_multi(
+ self,
+ ip_information_factory: Callable[..., IPInformation],
+ ip_geoname: IPGeoname,
+ geoipinfo_manager: GeoIpInfoManager,
+ ):
ip = fake.ipv4_public()
ip_information_factory(ip=ip, geoname=ip_geoname)
ips = [ip]
@@ -93,7 +115,12 @@ class TestGeoIpInfoManager:
assert res[ip] is not None
assert res[ip2] is not None
- def test_multi_ipv6(self, ip_information_factory, ip_geoname, geoipinfo_manager):
+ def test_multi_ipv6(
+ self,
+ ip_information_factory: Callable[..., IPInformation],
+ ip_geoname: IPGeoname,
+ geoipinfo_manager: GeoIpInfoManager,
+ ):
ip = fake.ipv6()
# Make another IP that will be in the same /64 block.
ip2 = ip[:-1] + "a" if ip[-1] != "a" else ip[:-1] + "b"
@@ -108,13 +135,19 @@ class TestGeoIpInfoManager:
# Looks up in redis, if not exists, looks in mysql, then sets
# the caches that didn't exist.
res = geoipinfo_manager.get_multi(ip_addresses=ips)
- assert res[ip].ip == ip
- assert res[ip].lookup_prefix == "/64"
- assert res[ip2].ip == ip2
- assert res[ip2].lookup_prefix == "/64"
+
+ res1 = res[ip]
+ assert isinstance(res1, GeoIPInformation)
+ assert res1.ip == ip
+ assert res1.lookup_prefix == "/64"
+
+ res2 = res[ip2]
+ assert isinstance(res2, GeoIPInformation)
+ assert res2.ip == ip2
+ assert res2.lookup_prefix == "/64"
# they should be the same basically, except for the ip
- def test_doesnt_exist(self, geoipinfo_manager):
+ def test_doesnt_exist(self, geoipinfo_manager: GeoIpInfoManager):
ip = fake.ipv4_public()
res = geoipinfo_manager.get_multi(ip_addresses=[ip])
assert res == {ip: None}
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 9158e15..e603632 100644
--- a/tests/managers/thl/test_ledger/test_lm_tx_locks.py
+++ b/tests/managers/thl/test_ledger/test_lm_tx_locks.py
@@ -1,9 +1,10 @@
from __future__ import annotations
import logging
-from collections.abc import Callable
+from collections.abc import Callable, Generator
from datetime import UTC, datetime, timedelta
from decimal import Decimal
+from logging import LogCaptureFixture
import pytest
@@ -41,7 +42,7 @@ class TestLedgerLocks:
session_factory: Callable[..., Session],
product_user_wallet_no: Product,
create_main_accounts: Callable[..., None],
- caplog,
+ caplog: Generator[LogCaptureFixture],
thl_ledger_manager: ThlLedgerManager,
ledger_manager: LedgerManager,
utc_hour_ago: datetime,
@@ -126,14 +127,15 @@ class TestLedgerLocks:
# purposely hold the lock open
tx = None
ledger_manager.redis_client.set(lock_name, "1")
- with caplog.at_level(logging.ERROR):
- with pytest.raises(expected_exception=LedgerTransactionCreateLockError):
- tx = thl_ledger_manager.create_tx_protected(
- lock_key=lock_key,
- condition=condition,
- create_tx_func=create_tx_func,
- )
- assert tx is None
+ with caplog.at_level(logging.ERROR), pytest.raises(
+ expected_exception=LedgerTransactionCreateLockError
+ ):
+ tx = thl_ledger_manager.create_tx_protected(
+ lock_key=lock_key,
+ condition=condition,
+ create_tx_func=create_tx_func,
+ )
+ assert tx is None
assert "Unable to acquire lock within the time specified" in caplog.text
ledger_manager.redis_client.delete(lock_name)
@@ -143,7 +145,7 @@ class TestLedgerLocks:
product_user_wallet_no: Product,
create_main_accounts: Callable[..., None],
delete_ledger_db: Callable[..., None],
- caplog,
+ caplog: Generator[LogCaptureFixture],
thl_ledger_manager: ThlLedgerManager,
ledger_manager: LedgerManager,
):
@@ -226,12 +228,13 @@ class TestLedgerLocks:
# Purposely hold the lock open
ledger_manager.redis_client.set(name=lock_name, value="1")
- with caplog.at_level(logging.DEBUG):
- with pytest.raises(expected_exception=LedgerTransactionCreateLockError):
- tx = thl_ledger_manager.create_tx_task_complete(
- wall=wall3, user=user, created=wall3.started
- )
- assert isinstance(tx, LedgerTransaction)
+ with caplog.at_level(logging.DEBUG), pytest.raises(
+ expected_exception=LedgerTransactionCreateLockError
+ ):
+ tx = thl_ledger_manager.create_tx_task_complete(
+ wall=wall3, user=user, created=wall3.started
+ )
+ assert isinstance(tx, LedgerTransaction)
assert "Unable to acquire lock within the time specified" in caplog.text
# Release the lock
diff --git a/tests/managers/thl/test_ledger/test_wallet.py b/tests/managers/thl/test_ledger/test_wallet.py
index 9e886db..cad3ea4 100644
--- a/tests/managers/thl/test_ledger/test_wallet.py
+++ b/tests/managers/thl/test_ledger/test_wallet.py
@@ -6,7 +6,6 @@ from uuid import uuid4
import pytest
-from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager
from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager
from generalresearch.managers.thl.product import ProductManager
from generalresearch.models.thl.product import (
@@ -55,6 +54,7 @@ class TestGetUserWalletBalance:
user: User = user_factory(schrute_product)
balance = thl_ledger_manager.get_user_wallet_balance(user=user)
assert balance == 0
+ assert isinstance(user.product, Product)
balance_string = user.product.format_payout_format(Decimal(balance) / 100)
assert balance_string == "0 Schrute Bucks"
redeemable_balance = thl_ledger_manager.get_user_redeemable_wallet_balance(
diff --git a/tests/managers/thl/test_payout.py b/tests/managers/thl/test_payout.py
index 153bee9..e6c597b 100644
--- a/tests/managers/thl/test_payout.py
+++ b/tests/managers/thl/test_payout.py
@@ -513,7 +513,7 @@ class TestBusinessPayoutEventManager:
assert int(res.deduction.sum()) == available_balance - 1
# Slightly less
- with pytest.raises(expected_exception=Exception) as cm:
+ with pytest.raises(expected_exception=ValueError):
res = business_payout_event_manager.recoup_proportional(
df=df, target_amount=available_balance + 1
)
@@ -731,12 +731,9 @@ class TestBusinessPayoutEventManager:
def test_ach_payment(
self,
- product: Product,
mnt_filepath: GRLDatasets,
thl_ledger_manager: ThlLedgerManager,
client_no_amm: DaskClient,
- thl_redis_config: RedisConfig,
- payout_event_manager: PayoutEventManager,
brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager,
business_payout_event_manager: BusinessPayoutEventManager,
delete_ledger_db: Callable[..., None],
@@ -899,8 +896,6 @@ class TestBusinessPayoutEventManager:
pop_ledger=pop_ledger_merge,
)
business.prebuild_payouts(
- thl_pg_config=thl_web_rr,
- thl_lm=thl_ledger_manager,
bpem=business_payout_event_manager,
)
assert isinstance(business.payouts, list)
@@ -930,13 +925,10 @@ class TestBusinessPayoutEventManager:
def test_ach_payment_partial_amount(
self,
- product: Product,
mnt_filepath: GRLDatasets,
thl_ledger_manager: ThlLedgerManager,
client_no_amm: DaskClient,
- thl_redis_config: RedisConfig,
payout_event_manager: PayoutEventManager,
- brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager,
business_payout_event_manager: BusinessPayoutEventManager,
delete_ledger_db: Callable[..., None],
create_main_accounts: Callable[..., None],
@@ -948,8 +940,6 @@ class TestBusinessPayoutEventManager:
session_with_tx_factory: Callable[..., None],
pop_ledger_merge: PopLedgerMerge,
start: datetime,
- bp_payout_factory: Callable[..., BrokerageProductPayoutEvent],
- adj_to_fail_with_tx_factory: Callable[..., None],
thl_web_rr: PostgresConfig,
ledger_manager: LedgerManager,
product_manager: ProductManager,
@@ -1076,7 +1066,6 @@ class TestBusinessPayoutEventManager:
thl_ledger_manager: ThlLedgerManager,
client_no_amm: DaskClient,
payout_event_manager: PayoutEventManager,
- brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager,
business_payout_event_manager: BusinessPayoutEventManager,
delete_ledger_db: Callable[..., None],
create_main_accounts: Callable[..., None],
diff --git a/tests/managers/thl/test_product.py b/tests/managers/thl/test_product.py
index 31b7b73..f93ac36 100644
--- a/tests/managers/thl/test_product.py
+++ b/tests/managers/thl/test_product.py
@@ -1,10 +1,15 @@
+from __future__ import annotations
+
+from collections.abc import Callable
from uuid import uuid4
import pytest
+from generalresearch.managers.thl.product import ProductManager
from generalresearch.models import Source
+from generalresearch.models.gr.team import Team
from generalresearch.models.thl.product import (
- product: Product,
+ Product,
ProfilingConfig,
SourceConfig,
SourcesConfig,
@@ -16,7 +21,7 @@ from generalresearch.models.thl.product import (
class TestProductManagerGetMethods:
- def test_get_by_uuid(self, product_manager):
+ def test_get_by_uuid(self, product_manager: ProductManager):
product: Product = product_manager.create_dummy(
product_id=uuid4().hex,
team_id=uuid4().hex,
@@ -36,10 +41,10 @@ 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):
+ def test_get_by_uuids(self, product_manager: ProductManager):
cnt = 5
- product_uuids = [uuid4().hex for idx in range(cnt)]
+ product_uuids = [uuid4().hex for _ in range(cnt)]
for product_id in product_uuids:
product_manager.create_dummy(
product_id=product_id,
@@ -61,7 +66,7 @@ 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):
+ def test_get_by_uuid_if_exists(self, product_manager: ProductManager):
product: Product = product_manager.create_dummy(
product_id=uuid4().hex,
team_id=uuid4().hex,
@@ -73,7 +78,7 @@ 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):
+ def test_get_by_uuids_if_exists(self, product_manager: ProductManager):
product_uuids = [uuid4().hex for _ in range(2)]
for product_id in product_uuids:
product_manager.create_dummy(
@@ -105,8 +110,8 @@ class TestProductManagerGetMethods:
# for instance in res:
# assert isinstance(instance, Product)
- def test_get_by_business_ids(self, product_manager):
- business_ids = [uuid4().hex for i in range(5)]
+ def test_get_by_business_ids(self, product_manager: ProductManager):
+ business_ids = [uuid4().hex for _ in range(5)]
product_manager.fetch_uuids(business_uuids=business_ids)
@@ -123,7 +128,7 @@ class TestProductManagerGetMethods:
class TestProductManagerCreation:
- def test_base(self, product_manager):
+ def test_base(self, product_manager: ProductManager):
instance = product_manager.create_dummy(
product_id=uuid4().hex,
team_id=uuid4().hex,
@@ -135,7 +140,7 @@ class TestProductManagerCreation:
class TestProductManagerCreate:
- def test_create_simple(self, product_manager):
+ def test_create_simple(self, product_manager: ProductManager):
# Always required: product_id, team_id, name, redirect_url
# Required internally - if not passed use default: harmonizer_domain,
# commission_pct, sources
@@ -181,9 +186,9 @@ class TestProductManager:
def test_get_by_uuid1(
self,
product_manager: ProductManager,
- team,
+ team: Team,
product: Product,
- product_factory,
+ product_factory: Callable[..., Product],
):
p1 = product_factory(team=team)
instance = product_manager.get_by_uuid(product_uuid=p1.uuid)
@@ -209,7 +214,9 @@ class TestProductManager:
assert 0 == instance.user_create_config.min_hourly_create_limit
assert instance.user_create_config.max_hourly_create_limit is None
- def test_get_by_uuid3(self, product_manager: ProductManager, product_factory):
+ def test_get_by_uuid3(
+ self, product_manager: ProductManager, product_factory: Callable[..., Product]
+ ):
p3 = product_factory()
instance = product_manager.get_by_uuid(p3.id)
assert instance.id == p3.id
@@ -225,7 +232,7 @@ 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):
+ def test_sources(self, 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)
@@ -240,7 +247,7 @@ 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):
+ def test_global_sources(self, product_manager: ProductManager):
sources_config = SupplyConfig(
policies=[
SupplyPolicy(
@@ -267,7 +274,7 @@ class TestProductManager:
p2 = product_manager.get_by_uuid(p1.id)
assert p1 == p2
- def test_user_health_config(self, product_manager):
+ def test_user_health_config(self, product_manager: ProductManager):
p = product_manager.create_dummy(
user_health_config=UserHealthConfig(banned_countries=["ng", "in"])
)
@@ -278,7 +285,7 @@ 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):
+ def test_profiling_config(self, product_manager: ProductManager):
p = product_manager.create_dummy(
profiling_config=ProfilingConfig(max_questions=1)
)
@@ -325,7 +332,7 @@ class TestProductManager:
class TestProductManagerUpdate:
- def test_update(self, product_manager):
+ def test_update(self, product_manager: ProductManager):
p = product_manager.create_dummy()
p.name = "new name"
p.enabled = False
@@ -346,7 +353,7 @@ class TestProductManagerUpdate:
class TestProductManagerCacheClear:
- def test_cache_clear(self, product_manager):
+ def test_cache_clear(self, product_manager: ProductManager):
p = product_manager.create_dummy()
product_manager.get_by_uuid(product_uuid=p.id)
product_manager.get_by_uuid(product_uuid=p.id)
diff --git a/tests/managers/thl/test_product_prod.py b/tests/managers/thl/test_product_prod.py
index 0f622b6..8734210 100644
--- a/tests/managers/thl/test_product_prod.py
+++ b/tests/managers/thl/test_product_prod.py
@@ -1,8 +1,12 @@
+from __future__ import annotations
+
import logging
+from collections.abc import Callable
from uuid import uuid4
import pytest
+from generalresearch.managers.thl.product import ProductManager
from generalresearch.models.thl.product import Product
logger = logging.getLogger()
@@ -10,7 +14,9 @@ logger = logging.getLogger()
class TestProductManagerGetMethods:
- def test_get_by_uuid(self, product_manager: ProductManager, product_factory):
+ def test_get_by_uuid(
+ self, product_manager: ProductManager, product_factory: Callable[..., Product]
+ ):
# Just test that we load properly
for p in [product_factory(), product_factory(), product_factory()]:
instance = product_manager.get_by_uuid(product_uuid=p.id)
@@ -22,7 +28,9 @@ 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, product_factory):
+ def test_get_by_uuids(
+ self, product_manager: ProductManager, product_factory: Callable[..., Product]
+ ):
products = [product_factory(), product_factory(), product_factory()]
cnt = len(products)
res = product_manager.get_by_uuids(product_uuids=[p.id for p in products])
@@ -43,7 +51,7 @@ class TestProductManagerGetMethods:
assert "invalid uuid passed" in str(cm.value)
def test_get_by_uuid_if_exists(
- self, product_factory: Callable[..., Product], product_manager
+ self, product_factory: Callable[..., Product], product_manager: ProductManager
):
products = [product_factory(), product_factory(), product_factory()]
@@ -54,7 +62,7 @@ class TestProductManagerGetMethods:
assert instance is None
def test_get_by_uuids_if_exists(
- self, product_manager: ProductManager, product_factory
+ self, product_manager: ProductManager, product_factory: Callable[..., Product]
):
products = [product_factory(), product_factory(), product_factory()]
@@ -78,7 +86,7 @@ class TestProductManagerGetMethods:
class TestProductManagerGetAll:
@pytest.mark.skip(reason="TODO")
- def test_get_ALL_by_ids(self, product_manager):
+ def test_get_ALL_by_ids(self, product_manager: ProductManager):
products = product_manager.get_all(rand_limit=50)
logger.info(f"Fetching {len(products)} product uuids")
# todo: once timebucks stops spamming broken accounts, fetch more
diff --git a/tests/managers/thl/test_profiling/test_question.py b/tests/managers/thl/test_profiling/test_question.py
index 998466e..97e7365 100644
--- a/tests/managers/thl/test_profiling/test_question.py
+++ b/tests/managers/thl/test_profiling/test_question.py
@@ -1,3 +1,4 @@
+from collections.abc import Callable
from uuid import uuid4
from generalresearch.managers.thl.profiling.question import QuestionManager
@@ -6,7 +7,11 @@ from generalresearch.models import Source
class TestQuestionManager:
- def test_get_multi_upk(self, question_manager: QuestionManager, upk_data):
+ def test_get_multi_upk(
+ self, question_manager: QuestionManager, upk_data: Callable[..., None]
+ ):
+ upk_data()
+
qs = question_manager.get_multi_upk(
question_ids=[
"8a22de34f985476aac85e15547100db8",
@@ -17,13 +22,21 @@ class TestQuestionManager:
)
assert len(qs) == 3
- def test_get_questions_ranked(self, question_manager: QuestionManager, upk_data):
+ def test_get_questions_ranked(
+ self, question_manager: QuestionManager, upk_data: Callable[..., None]
+ ):
+ upk_data()
+
qs = question_manager.get_questions_ranked(country_iso="mx", language_iso="spa")
assert len(qs) >= 40
assert qs[0].importance.task_score > qs[40].importance.task_score
assert all(q.country_iso == "mx" and q.language_iso == "spa" for q in qs)
- def test_lookup_by_property(self, question_manager: QuestionManager, upk_data):
+ def test_lookup_by_property(
+ self, question_manager: QuestionManager, upk_data: Callable[..., None]
+ ):
+ upk_data()
+
q = question_manager.lookup_by_property(
property_code="i:industry", country_iso="us", language_iso="eng"
)
@@ -38,7 +51,11 @@ class TestQuestionManager:
)
assert q.explanation_template
- def test_filter_by_property(self, question_manager: QuestionManager, upk_data):
+ def test_filter_by_property(
+ self, question_manager: QuestionManager, upk_data: Callable[..., None]
+ ):
+ upk_data()
+
lookup = [
("i:industry", "us", "eng"),
("i:industry", "mx", "eng"),
diff --git a/tests/managers/thl/test_profiling/test_schema.py b/tests/managers/thl/test_profiling/test_schema.py
index ae61527..b0eae31 100644
--- a/tests/managers/thl/test_profiling/test_schema.py
+++ b/tests/managers/thl/test_profiling/test_schema.py
@@ -1,9 +1,18 @@
+from collections.abc import Callable
+
+from generalresearch.managers.thl.profiling.schema import (
+ UpkSchemaManager,
+)
from generalresearch.models.thl.profiling.upk_property import PropertyType
class TestUpkSchemaManager:
- def test_get_props_info(self, upk_schema_manager, upk_data):
+ def test_get_props_info(
+ self, upk_schema_manager: UpkSchemaManager, upk_data: Callable[..., None]
+ ):
+ upk_data()
+
props = upk_schema_manager.get_props_info()
assert (
len(props) == 16955
@@ -35,10 +44,10 @@ class TestUpkSchemaManager:
assert age.prop_type == PropertyType.UPK_NUMERICAL
assert age.gold_standard
- cars = [
+ cars = next(
x
for x in props
if x.country_iso == "us" and x.property_label == "household_auto_type"
- ][0]
+ )
assert not cars.gold_standard
assert cars.categories[0].label == "Autos & Vehicles"
diff --git a/tests/managers/thl/test_profiling/test_uqa.py b/tests/managers/thl/test_profiling/test_uqa.py
deleted file mode 100644
index 8b13789..0000000
--- a/tests/managers/thl/test_profiling/test_uqa.py
+++ /dev/null
@@ -1 +0,0 @@
-
diff --git a/tests/managers/thl/test_profiling/test_user_upk.py b/tests/managers/thl/test_profiling/test_user_upk.py
index 8b995b1..fa10b67 100644
--- a/tests/managers/thl/test_profiling/test_user_upk.py
+++ b/tests/managers/thl/test_profiling/test_user_upk.py
@@ -1,6 +1,8 @@
+from collections.abc import Callable
from datetime import UTC, datetime
from generalresearch.managers.thl.profiling.user_upk import UserUpkManager
+from generalresearch.models.thl.user import User
now = datetime.now(tz=UTC)
base = {
@@ -21,11 +23,25 @@ for a in upk_ans_dict:
class TestUserUpkManager:
- def test_user_upk_empty(self, user_upk_manager: UserUpkManager, upk_data, user):
+ def test_user_upk_empty(
+ self,
+ user_upk_manager: UserUpkManager,
+ upk_data: Callable[..., None],
+ user: User,
+ ):
+ upk_data()
+
res = user_upk_manager.get_user_upk_mysql(user_id=user.user_id)
assert len(res) == 0
- def test_user_upk(self, user_upk_manager: UserUpkManager, upk_data, user):
+ def test_user_upk(
+ self,
+ user_upk_manager: UserUpkManager,
+ upk_data: Callable[..., None],
+ user: User,
+ ):
+ upk_data()
+
for x in upk_ans_dict:
x["user_id"] = user.user_id
user_upk = user_upk_manager.populate_user_upk_from_dict(upk_ans_dict)
diff --git a/tests/managers/thl/test_session_manager.py b/tests/managers/thl/test_session_manager.py
index 6c5f820..05a49c1 100644
--- a/tests/managers/thl/test_session_manager.py
+++ b/tests/managers/thl/test_session_manager.py
@@ -1,22 +1,34 @@
-from datetime import timedelta
+from __future__ import annotations
+
+from collections.abc import Callable
+from datetime import datetime, timedelta
from decimal import Decimal
from uuid import uuid4
from faker import Faker
+from generalresearch.managers.thl.session import SessionManager
from generalresearch.models import DeviceType
+from generalresearch.models.gr.business import Business
+from generalresearch.models.gr.team import Team
from generalresearch.models.legacy.bucket import Bucket
from generalresearch.models.thl.definitions import (
SessionStatusCode2,
Status,
StatusCode1,
)
+from generalresearch.models.thl.product import Product
+from generalresearch.models.thl.session import Session
+from generalresearch.models.thl.user import User
+from generalresearch.pg_helper import PostgresConfig
fake = Faker()
class TestSessionManager:
- def test_create_session(self, session_manager, user, utc_hour_ago):
+ def test_create_session(
+ self, session_manager: SessionManager, user: User, utc_hour_ago: datetime
+ ):
bucket = Bucket(
loi_min=timedelta(seconds=60),
loi_max=timedelta(seconds=120),
@@ -39,7 +51,9 @@ class TestSessionManager:
s2 = session_manager.get_from_uuid(session_uuid=s1.uuid)
assert s1 == s2
- def test_finish_with_status(self, session_manager, user, utc_hour_ago):
+ def test_finish_with_status(
+ self, session_manager: SessionManager, user: User, utc_hour_ago: datetime
+ ):
uuid_1 = uuid4().hex
session = session_manager.create(
started=utc_hour_ago, user=user, uuid_id=uuid_1
@@ -59,7 +73,7 @@ class TestSessionManager:
class TestSessionManagerFilter:
- def test_base(self, session_manager, user, utc_now):
+ def test_base(self, session_manager: SessionManager, user: User, utc_now: datetime):
uuid_id = uuid4().hex
session_manager.create(started=utc_now, user=user, uuid_id=uuid_id)
res = session_manager.filter(limit=1)
@@ -67,7 +81,9 @@ class TestSessionManagerFilter:
assert isinstance(res, list)
assert res[0].uuid == uuid_id
- def test_user(self, session_manager, user, utc_hour_ago):
+ def test_user(
+ self, session_manager: SessionManager, user: User, utc_hour_ago: datetime
+ ):
session_manager.create(started=utc_hour_ago, user=user, uuid_id=uuid4().hex)
session_manager.create(started=utc_hour_ago, user=user, uuid_id=uuid4().hex)
@@ -78,16 +94,13 @@ class TestSessionManagerFilter:
self,
product_factory: Callable[..., Product],
user_factory: Callable[..., User],
- session_manager,
- user,
- utc_hour_ago,
+ session_manager: SessionManager,
+ utc_hour_ago: datetime,
):
- from generalresearch.models.thl.session import Session
- from generalresearch.models.thl.user import User
p1 = product_factory()
- for n in range(5):
+ for _ in range(5):
u = user_factory(product=p1)
session_manager.create(started=utc_hour_ago, user=u, uuid_id=uuid4().hex)
@@ -102,15 +115,14 @@ class TestSessionManagerFilter:
self,
product_factory: Callable[..., Product],
user_factory: Callable[..., User],
- team,
- session_manager,
- user,
- utc_hour_ago,
+ team: Team,
+ session_manager: SessionManager,
+ utc_hour_ago: datetime,
thl_web_rr: PostgresConfig,
):
p1 = product_factory(team=team)
- for n in range(5):
+ for _ in range(5):
u = user_factory(product=p1)
session_manager.create(started=utc_hour_ago, user=u, uuid_id=uuid4().hex)
@@ -124,14 +136,13 @@ class TestSessionManagerFilter:
product_factory: Callable[..., Product],
business: Business,
user_factory: Callable[..., User],
- session_manager,
- user,
- utc_hour_ago,
+ session_manager: SessionManager,
+ utc_hour_ago: datetime,
thl_web_rr: PostgresConfig,
):
p1 = product_factory(business=business)
- for n in range(5):
+ for _ in range(5):
u = user_factory(product=p1)
session_manager.create(started=utc_hour_ago, user=u, uuid_id=uuid4().hex)
diff --git a/tests/managers/thl/test_survey.py b/tests/managers/thl/test_survey.py
index eec7c3b..2c2bf9d 100644
--- a/tests/managers/thl/test_survey.py
+++ b/tests/managers/thl/test_survey.py
@@ -1,9 +1,21 @@
+from __future__ import annotations
+
import uuid
+from collections.abc import Callable
from datetime import UTC, datetime
from decimal import Decimal
import pytest
+from generalresearch.managers.thl.buyer import BuyerManager
+from generalresearch.managers.thl.profiling.question import (
+ QuestionManager,
+)
+from generalresearch.managers.thl.profiling.schema import (
+ UpkSchemaManager,
+)
+from generalresearch.managers.thl.profiling.uqa import UQAManager
+from generalresearch.managers.thl.survey import SurveyManager, SurveyStatManager
from generalresearch.models import Source
from generalresearch.models.legacy.bucket import (
DurationSummary,
@@ -23,7 +35,7 @@ from generalresearch.models.thl.survey.model import (
@pytest.fixture(scope="session")
-def surveys_fixture():
+def surveys_fixture() -> list[Survey]:
return [
Survey(source=Source.TESTING, survey_id="a", buyer_code="buyer1"),
Survey(source=Source.TESTING, survey_id="b", buyer_code="buyer2"),
@@ -73,11 +85,13 @@ class TestSurvey:
def test(
self,
- delete_buyers_surveys,
- buyer_manager,
- survey_manager,
- surveys_fixture,
+ delete_buyers_surveys: Callable[..., None],
+ buyer_manager: BuyerManager,
+ survey_manager: SurveyManager,
+ surveys_fixture: list[Survey],
):
+ delete_buyers_surveys()
+
survey_manager.create_or_update(surveys_fixture)
survey_ids = {s.survey_id for s in surveys_fixture}
res = survey_manager.filter_by_natural_key(
@@ -98,7 +112,7 @@ class TestSurvey:
assert res2[0] == res[0]
assert len(res2) == len(surveys2)
- def test_category(self, survey_manager):
+ def test_category(self, survey_manager: SurveyManager):
survey1 = Survey(id=562289, survey_id="a", source=Source.TESTING)
survey2 = Survey(id=562290, survey_id="a", source=Source.TESTING)
categories = list(survey_manager.category_manager.categories.values())
@@ -110,8 +124,14 @@ class TestSurvey:
survey_manager.update_surveys_categories(surveys)
def test_survey_eligibility(
- self, survey_manager, upk_data, question_manager, uqa_manager
+ self,
+ survey_manager: SurveyManager,
+ upk_data: Callable[..., None],
+ question_manager: QuestionManager,
+ uqa_manager: UQAManager,
):
+ upk_data()
+
bucket = TopNPlusBucket(
id="c82cf98c578a43218334544ab376b00e",
contents=[],
@@ -205,10 +225,10 @@ class TestSurvey:
class TestSurveyStat:
def test(
self,
- delete_buyers_surveys,
+ delete_buyers_surveys: Callable[..., None],
surveystat_manager,
- survey_manager,
- surveys_fixture,
+ survey_manager: SurveyManager,
+ surveys_fixture: list[Survey],
):
survey_manager.create_or_update(surveys_fixture)
ss = [ssa, ssb]
@@ -276,15 +296,17 @@ class TestSurveyStat:
def test_ymsp(
self,
- delete_buyers_surveys,
- surveys_fixture,
- survey_manager,
- surveystat_manager,
+ delete_buyers_surveys: Callable[..., None],
+ surveys_fixture: list[Survey],
+ survey_manager: SurveyManager,
+ surveystat_manager: SurveyStatManager,
):
+ delete_buyers_surveys()
+
source = Source.TESTING
survey = surveys_fixture[0].model_copy()
surveys = []
- for idx in range(100):
+ for _ in range(100):
s = survey.model_copy()
s.survey_id = uuid.uuid4().hex
surveys.append(s)
@@ -305,7 +327,7 @@ class TestSurveyStat:
surveys = surveys[10:]
# and 2 new ones are created
- for idx in range(2):
+ for _ in range(2):
s = survey.model_copy()
s.survey_id = uuid.uuid4().hex
surveys.append(s)
@@ -329,11 +351,13 @@ class TestSurveyStat:
def test_filter(
self,
- delete_buyers_surveys,
- surveys_fixture,
- survey_manager,
- surveystat_manager,
+ delete_buyers_surveys: Callable[..., None],
+ surveys_fixture: list[Survey],
+ survey_manager: SurveyManager,
+ surveystat_manager: SurveyStatManager,
):
+ delete_buyers_surveys()
+
surveys = []
survey = surveys_fixture[0].model_copy()
survey.source = Source.TESTING
diff --git a/tests/managers/thl/test_survey_penalty.py b/tests/managers/thl/test_survey_penalty.py
index c7862bb..9c29a0a 100644
--- a/tests/managers/thl/test_survey_penalty.py
+++ b/tests/managers/thl/test_survey_penalty.py
@@ -1,7 +1,10 @@
+from __future__ import annotations
+
import uuid
import pytest
+from generalresearch.managers.thl.survey_penalty import SurveyPenaltyManager
from generalresearch.models import Source
from generalresearch.models.thl.survey.penalty import (
BPSurveyPenalty,
@@ -22,7 +25,9 @@ def team_uuid() -> str:
@pytest.fixture
-def penalties(product_uuid, team_uuid):
+def penalties(
+ product_uuid: str, team_uuid: str
+) -> list[BPSurveyPenalty | TeamSurveyPenalty]:
return [
BPSurveyPenalty(
source=Source.TESTING, survey_id="a", penalty=0.1, product_id=product_uuid
@@ -48,7 +53,13 @@ def penalties(product_uuid, team_uuid):
class TestSurveyPenalty:
- def test(self, surveypenalty_manager, penalties, product_uuid, team_uuid):
+ def test(
+ self,
+ surveypenalty_manager: SurveyPenaltyManager,
+ penalties: list[BPSurveyPenalty | TeamSurveyPenalty],
+ product_uuid: str,
+ team_uuid: str,
+ ):
surveypenalty_manager.set_penalties(penalties)
res = surveypenalty_manager.get_penalties_for(
diff --git a/tests/managers/thl/test_task_adjustment.py b/tests/managers/thl/test_task_adjustment.py
index 7b77a68..a7324c3 100644
--- a/tests/managers/thl/test_task_adjustment.py
+++ b/tests/managers/thl/test_task_adjustment.py
@@ -1,20 +1,31 @@
+from __future__ import annotations
+
import logging
+from collections.abc import Callable
from datetime import UTC, datetime, timedelta
from decimal import Decimal
from random import randint
import pytest
+from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager
+from generalresearch.managers.thl.session import SessionManager
+from generalresearch.managers.thl.task_adjustment import (
+ TaskAdjustmentManager,
+)
+from generalresearch.managers.thl.wall import WallManager
from generalresearch.models import Source
from generalresearch.models.thl.definitions import (
Status,
StatusCode1,
WallAdjustedStatus,
)
+from generalresearch.models.thl.session import Session
+from generalresearch.models.thl.user import User
@pytest.fixture()
-def session_complete(session_with_tx_factory: Callable[..., None], user):
+def session_complete(session_with_tx_factory: Callable[..., Session], user: User):
return session_with_tx_factory(
user=user, final_status=Status.COMPLETE, wall_req_cpi=Decimal("1.23")
)
@@ -22,7 +33,7 @@ def session_complete(session_with_tx_factory: Callable[..., None], user):
@pytest.fixture()
def session_complete_with_wallet(
- session_with_tx_factory: Callable[..., None], user_with_wallet
+ session_with_tx_factory: Callable[..., None], user_with_wallet: User
):
return session_with_tx_factory(
user=user_with_wallet,
@@ -32,7 +43,9 @@ def session_complete_with_wallet(
@pytest.fixture()
-def session_fail(user, session_manager, wall_manager):
+def session_fail(
+ user: User, session_manager: SessionManager, wall_manager: WallManager
+) -> Session:
session = session_manager.create_dummy(started=datetime.now(UTC), user=user)
wall1 = wall_manager.create_dummy(
session_id=session.id,
@@ -56,38 +69,37 @@ class TestHandleRecons:
def test_complete_to_recon(
self,
- session_complete,
- thl_lm,
- task_adjustment_manager,
- wall_manager,
- session_manager,
+ session_complete: Session,
+ thl_ledger_manager: ThlLedgerManager,
+ task_adjustment_manager: TaskAdjustmentManager,
+ wall_manager: WallManager,
+ session_manager: SessionManager,
caplog,
):
print(wall_manager.pg_config.dsn)
mid = session_complete.uuid
wall_uuid = session_complete.wall_events[-1].uuid
s = session_complete
- ledger_manager = thl_lm
- revenue_account = ledger_manager.get_account_task_complete_revenue()
- current_amount = ledger_manager.get_account_filtered_balance(
+ revenue_account = thl_ledger_manager.get_account_task_complete_revenue()
+ current_amount = thl_ledger_manager.get_account_filtered_balance(
revenue_account, "thl_wall", wall_uuid
)
assert (
current_amount == 123
), "this is the amount of revenue from this task complete"
- bp_wallet_account = ledger_manager.get_account_or_create_bp_wallet(
+ bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet(
s.user.product
)
- current_bp_payout = ledger_manager.get_account_filtered_balance(
+ current_bp_payout = thl_ledger_manager.get_account_filtered_balance(
bp_wallet_account, "thl_session", mid
)
assert current_bp_payout == 117, "this is the amount paid to the BP"
# Do the work here !! ----v
task_adjustment_manager.handle_single_recon(
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
wall_uuid=wall_uuid,
adjusted_status=WallAdjustedStatus.ADJUSTED_TO_FAIL,
)
@@ -95,18 +107,18 @@ class TestHandleRecons:
len(task_adjustment_manager.filter_by_wall_uuid(wall_uuid=wall_uuid)) == 1
)
- current_amount = ledger_manager.get_account_filtered_balance(
+ current_amount = thl_ledger_manager.get_account_filtered_balance(
revenue_account, "thl_wall", wall_uuid
)
assert current_amount == 0, "after recon, it should be zeroed"
- current_bp_payout = ledger_manager.get_account_filtered_balance(
+ current_bp_payout = thl_ledger_manager.get_account_filtered_balance(
bp_wallet_account, "thl_session", mid
)
assert current_bp_payout == 0, "this is the amount paid to the BP"
- commission_account = ledger_manager.get_account_or_create_bp_commission(
+ commission_account = thl_ledger_manager.get_account_or_create_bp_commission(
s.user.product
)
- assert ledger_manager.get_account_balance(commission_account) == 0
+ assert thl_ledger_manager.get_account_balance(commission_account) == 0
# Now, say we get the exact same *adjust to incomplete* msg again. It should do nothing!
adjusted_timestamp = datetime.now(tz=UTC)
@@ -122,7 +134,7 @@ class TestHandleRecons:
session = session_manager.get_from_id(wall.session_id)
user = session.user
with caplog.at_level(logging.INFO):
- ledger_manager.create_tx_task_adjustment(
+ thl_ledger_manager.create_tx_task_adjustment(
wall, user=user, created=adjusted_timestamp
)
assert "No transactions needed" in caplog.text
@@ -135,212 +147,219 @@ class TestHandleRecons:
assert "is already f" in caplog.text or "is already Status.FAIL" in caplog.text
with caplog.at_level(logging.INFO, logger="LedgerManager"):
- ledger_manager.create_tx_bp_adjustment(session, created=adjusted_timestamp)
+ thl_ledger_manager.create_tx_bp_adjustment(
+ session, created=adjusted_timestamp
+ )
assert "No transactions needed" in caplog.text
- current_amount = ledger_manager.get_account_filtered_balance(
+ current_amount = thl_ledger_manager.get_account_filtered_balance(
revenue_account, "thl_wall", wall_uuid
)
assert current_amount == 0, "after recon, it should be zeroed"
- current_bp_payout = ledger_manager.get_account_filtered_balance(
+ current_bp_payout = thl_ledger_manager.get_account_filtered_balance(
bp_wallet_account, "thl_session", mid
)
assert current_bp_payout == 0, "this is the amount paid to the BP"
# And if we get an adj to fail, and handle it, it should do nothing at all
task_adjustment_manager.handle_single_recon(
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
wall_uuid=wall_uuid,
adjusted_status=WallAdjustedStatus.ADJUSTED_TO_FAIL,
)
assert (
len(task_adjustment_manager.filter_by_wall_uuid(wall_uuid=wall_uuid)) == 1
)
- current_amount = ledger_manager.get_account_filtered_balance(
+ current_amount = thl_ledger_manager.get_account_filtered_balance(
revenue_account, "thl_wall", wall_uuid
)
assert current_amount == 0, "after recon, it should be zeroed"
- def test_fail_to_complete(self, session_fail, thl_lm, task_adjustment_manager):
- s = session_fail
+ def test_fail_to_complete(
+ self,
+ session_fail: Session,
+ thl_ledger_manager: ThlLedgerManager,
+ task_adjustment_manager: TaskAdjustmentManager,
+ ):
mid = session_fail.uuid
wall_uuid = session_fail.wall_events[-1].uuid
- ledger_manager = thl_lm
- revenue_account = ledger_manager.get_account_task_complete_revenue()
- current_amount = ledger_manager.get_account_filtered_balance(
+ revenue_account = thl_ledger_manager.get_account_task_complete_revenue()
+ current_amount = thl_ledger_manager.get_account_filtered_balance(
revenue_account, "thl_wall", mid
)
assert (
current_amount == 0
), "this is the amount of revenue from this task complete"
- bp_wallet_account = ledger_manager.get_account_or_create_bp_wallet(
- s.user.product
+ bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet(
+ session_fail.user.product
)
- current_bp_payout = ledger_manager.get_account_filtered_balance(
+ current_bp_payout = thl_ledger_manager.get_account_filtered_balance(
bp_wallet_account, "thl_session", mid
)
assert current_bp_payout == 0, "this is the amount paid to the BP"
task_adjustment_manager.handle_single_recon(
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
wall_uuid=wall_uuid,
adjusted_status=WallAdjustedStatus.ADJUSTED_TO_COMPLETE,
)
- current_amount = ledger_manager.get_account_filtered_balance(
+ current_amount = thl_ledger_manager.get_account_filtered_balance(
revenue_account, "thl_wall", wall_uuid
)
assert current_amount == 322, "after recon, we should be paid"
- current_bp_payout = ledger_manager.get_account_filtered_balance(
+ current_bp_payout = thl_ledger_manager.get_account_filtered_balance(
bp_wallet_account, "thl_session", mid
)
assert current_bp_payout == 306, "this is the amount paid to the BP"
# Now reverse it back to fail
task_adjustment_manager.handle_single_recon(
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
wall_uuid=wall_uuid,
adjusted_status=WallAdjustedStatus.ADJUSTED_TO_FAIL,
)
- current_amount = ledger_manager.get_account_filtered_balance(
+ current_amount = thl_ledger_manager.get_account_filtered_balance(
revenue_account, "thl_wall", wall_uuid
)
assert current_amount == 0
- current_bp_payout = ledger_manager.get_account_filtered_balance(
+ current_bp_payout = thl_ledger_manager.get_account_filtered_balance(
bp_wallet_account, "thl_session", mid
)
assert current_bp_payout == 0
- commission_account = ledger_manager.get_account_or_create_bp_commission(
- s.user.product
+ commission_account = thl_ledger_manager.get_account_or_create_bp_commission(
+ session_fail.user.product
)
- assert ledger_manager.get_account_balance(commission_account) == 0
+ assert thl_ledger_manager.get_account_balance(commission_account) == 0
def test_complete_already_complete(
- self, session_complete, thl_lm, task_adjustment_manager
+ self,
+ session_complete: Session,
+ thl_ledger_manager: ThlLedgerManager,
+ task_adjustment_manager: TaskAdjustmentManager,
):
- s = session_complete
mid = session_complete.uuid
wall_uuid = session_complete.wall_events[-1].uuid
- ledger_manager = thl_lm
for _ in range(4):
# just run it 4 times to make sure nothing happens 4 times
task_adjustment_manager.handle_single_recon(
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
wall_uuid=wall_uuid,
adjusted_status=WallAdjustedStatus.ADJUSTED_TO_COMPLETE,
)
- revenue_account = ledger_manager.get_account_task_complete_revenue()
- bp_wallet_account = ledger_manager.get_account_or_create_bp_wallet(
- s.user.product
+ revenue_account = thl_ledger_manager.get_account_task_complete_revenue()
+ bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet(
+ session_complete.user.product
)
- commission_account = ledger_manager.get_account_or_create_bp_commission(
- s.user.product
+ commission_account = thl_ledger_manager.get_account_or_create_bp_commission(
+ session_complete.user.product
)
- current_amount = ledger_manager.get_account_filtered_balance(
+ current_amount = thl_ledger_manager.get_account_filtered_balance(
revenue_account, "thl_wall", wall_uuid
)
assert current_amount == 123
- assert ledger_manager.get_account_balance(commission_account) == 6
+ assert thl_ledger_manager.get_account_balance(commission_account) == 6
- current_bp_payout = ledger_manager.get_account_filtered_balance(
+ current_bp_payout = thl_ledger_manager.get_account_filtered_balance(
bp_wallet_account, "thl_session", mid
)
assert current_bp_payout == 117
def test_incomplete_already_incomplete(
- self, session_fail, thl_lm, task_adjustment_manager
+ self,
+ session_fail: Session,
+ thl_ledger_manager: ThlLedgerManager,
+ task_adjustment_manager: TaskAdjustmentManager,
):
- s = session_fail
mid = session_fail.uuid
wall_uuid = session_fail.wall_events[-1].uuid
- ledger_manager = thl_lm
for _ in range(4):
task_adjustment_manager.handle_single_recon(
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
wall_uuid=wall_uuid,
adjusted_status=WallAdjustedStatus.ADJUSTED_TO_FAIL,
)
- revenue_account = ledger_manager.get_account_task_complete_revenue()
- bp_wallet_account = ledger_manager.get_account_or_create_bp_wallet(
- s.user.product
+ revenue_account = thl_ledger_manager.get_account_task_complete_revenue()
+ bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet(
+ session_fail.user.product
)
- commission_account = ledger_manager.get_account_or_create_bp_commission(
- s.user.product
+ commission_account = thl_ledger_manager.get_account_or_create_bp_commission(
+ session_fail.user.product
)
- current_amount = ledger_manager.get_account_filtered_balance(
+ current_amount = thl_ledger_manager.get_account_filtered_balance(
revenue_account, "thl_wall", mid
)
assert current_amount == 0
- assert ledger_manager.get_account_balance(commission_account) == 0
+ assert thl_ledger_manager.get_account_balance(commission_account) == 0
- current_bp_payout = ledger_manager.get_account_filtered_balance(
+ current_bp_payout = thl_ledger_manager.get_account_filtered_balance(
bp_wallet_account, "thl_session", mid
)
assert current_bp_payout == 0
def test_complete_to_recon_user_wallet(
self,
- session_complete_with_wallet,
- user_with_wallet,
- thl_lm,
- task_adjustment_manager,
+ session_complete_with_wallet: Session,
+ # user_with_wallet: User,
+ thl_ledger_manager: ThlLedgerManager,
+ task_adjustment_manager: TaskAdjustmentManager,
):
- s = session_complete_with_wallet
- mid = s.uuid
- wall_uuid = s.wall_events[-1].uuid
- ledger_manager = thl_lm
+ mid = session_complete_with_wallet.uuid
+ wall_uuid = session_complete_with_wallet.wall_events[-1].uuid
- revenue_account = ledger_manager.get_account_task_complete_revenue()
- amount = ledger_manager.get_account_filtered_balance(
+ revenue_account = thl_ledger_manager.get_account_task_complete_revenue()
+ amount = thl_ledger_manager.get_account_filtered_balance(
revenue_account, "thl_wall", wall_uuid
)
assert amount == 123, "this is the amount of revenue from this task complete"
- bp_wallet_account = ledger_manager.get_account_or_create_bp_wallet(
- s.user.product
+ bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet(
+ session_complete_with_wallet.user.product
)
- user_wallet_account = ledger_manager.get_account_or_create_user_wallet(s.user)
- commission_account = ledger_manager.get_account_or_create_bp_commission(
- s.user.product
+ user_wallet_account = thl_ledger_manager.get_account_or_create_user_wallet(
+ session_complete_with_wallet.user
+ )
+ commission_account = thl_ledger_manager.get_account_or_create_bp_commission(
+ session_complete_with_wallet.user.product
)
- amount = ledger_manager.get_account_filtered_balance(
+ amount = thl_ledger_manager.get_account_filtered_balance(
bp_wallet_account, "thl_session", mid
)
assert amount == 70, "this is the amount paid to the BP"
- amount = ledger_manager.get_account_filtered_balance(
+ amount = thl_ledger_manager.get_account_filtered_balance(
user_wallet_account, "thl_session", mid
)
assert amount == 47, "this is the amount paid to the user"
assert (
- ledger_manager.get_account_balance(commission_account) == 6
+ thl_ledger_manager.get_account_balance(commission_account) == 6
), "earned commission"
task_adjustment_manager.handle_single_recon(
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
wall_uuid=wall_uuid,
adjusted_status=WallAdjustedStatus.ADJUSTED_TO_FAIL,
)
- amount = ledger_manager.get_account_filtered_balance(
+ amount = thl_ledger_manager.get_account_filtered_balance(
revenue_account, "thl_wall", wall_uuid
)
assert amount == 0
- amount = ledger_manager.get_account_filtered_balance(
+ amount = thl_ledger_manager.get_account_filtered_balance(
bp_wallet_account, "thl_session", mid
)
assert amount == 0
- amount = ledger_manager.get_account_filtered_balance(
+ amount = thl_ledger_manager.get_account_filtered_balance(
user_wallet_account, "thl_session", mid
)
assert amount == 0
assert (
- ledger_manager.get_account_balance(commission_account) == 0
+ thl_ledger_manager.get_account_balance(commission_account) == 0
), "earned commission"
diff --git a/tests/managers/thl/test_task_status.py b/tests/managers/thl/test_task_status.py
index 44938a6..b47f650 100644
--- a/tests/managers/thl/test_task_status.py
+++ b/tests/managers/thl/test_task_status.py
@@ -1,9 +1,14 @@
+from __future__ import annotations
+
+from collections.abc import Callable
from datetime import UTC, datetime, timedelta
from decimal import Decimal
import pytest
+from generalresearch.managers.thl.product import ProductManager
from generalresearch.managers.thl.session import SessionManager
+from generalresearch.managers.thl.wall import WallManager
from generalresearch.models import Source
from generalresearch.models.thl.definitions import (
Status,
@@ -14,6 +19,7 @@ from generalresearch.models.thl.product import (
PayoutConfig,
PayoutTransformation,
PayoutTransformationPercentArgs,
+ Product,
UserWalletConfig,
)
from generalresearch.models.thl.session import Session, WallOut
@@ -30,7 +36,7 @@ finish3 = start3 + timedelta(minutes=5)
@pytest.fixture(scope="session")
-def bp1(product_manager):
+def bp1(product_manager: ProductManager) -> Product:
# user wallet disabled, payout xform NULL
return product_manager.create_dummy(
user_wallet_config=UserWalletConfig(enabled=False),
@@ -39,7 +45,7 @@ def bp1(product_manager):
@pytest.fixture(scope="session")
-def bp2(product_manager):
+def bp2(product_manager: ProductManager) -> Product:
# user wallet disabled, payout xform 40%
return product_manager.create_dummy(
user_wallet_config=UserWalletConfig(enabled=False),
@@ -53,7 +59,7 @@ def bp2(product_manager):
@pytest.fixture(scope="session")
-def bp3(product_manager):
+def bp3(product_manager: ProductManager) -> Product:
# user wallet enabled, payout xform 50%
return product_manager.create_dummy(
user_wallet_config=UserWalletConfig(enabled=True),
@@ -70,9 +76,9 @@ class TestTaskStatus:
def test_task_status_complete_1(
self,
- bp1,
+ bp1: Product,
user_factory: Callable[..., User],
- finished_session_factory,
+ finished_session_factory: Callable[..., Session],
session_manager: SessionManager,
):
# User Payout xform NULL
@@ -131,10 +137,10 @@ class TestTaskStatus:
def test_task_status_complete_2(
self,
- bp2,
+ bp2: Product,
user_factory: Callable[..., User],
- finished_session_factory,
- session_manager,
+ finished_session_factory: Callable[..., Session],
+ session_manager: SessionManager,
):
# User Payout xform 40%
user2: User = user_factory(product=bp2)
@@ -202,10 +208,10 @@ class TestTaskStatus:
def test_task_status_complete_3(
self,
- bp3,
+ bp3: Product,
user_factory: Callable[..., User],
- finished_session_factory,
- session_manager,
+ finished_session_factory: Callable[..., Session],
+ session_manager: SessionManager,
):
# Wallet enabled User Payout xform 50% (the response is identical
# to the user wallet disabled w same xform)
@@ -235,16 +241,17 @@ class TestTaskStatus:
}
)
tsr = session_manager.get_task_status_response(s3.uuid)
+ assert isinstance(tsr, TaskStatusResponse)
# Not bothering with wall events ...
tsr.wall_events = None
assert tsr == expected_tsr
def test_task_status_fail(
self,
- bp1,
+ bp1: Product,
user_factory: Callable[..., User],
- finished_session_factory,
- session_manager,
+ finished_session_factory: Callable[..., Session],
+ session_manager: SessionManager,
):
# User Payout xform NULL: user payout is None always
user1: User = user_factory(product=bp1)
@@ -275,16 +282,17 @@ class TestTaskStatus:
}
)
tsr = session_manager.get_task_status_response(s1.uuid)
+ assert isinstance(tsr, TaskStatusResponse)
# Not bothering with wall events ...
tsr.wall_events = None
assert tsr == expected_tsr
def test_task_status_fail_xform(
self,
- bp2,
+ bp2: Product,
user_factory: Callable[..., User],
- finished_session_factory,
- session_manager,
+ finished_session_factory: Callable[..., Session],
+ session_manager: SessionManager,
):
# User Payout xform 40%: user_payout is 0 (not None)
@@ -314,16 +322,17 @@ class TestTaskStatus:
}
)
tsr = session_manager.get_task_status_response(s.uuid)
+ assert isinstance(tsr, TaskStatusResponse)
# Not bothering with wall events ...
tsr.wall_events = None
assert tsr == expected_tsr
def test_task_status_abandon(
self,
- bp1,
+ bp1: Product,
user_factory: Callable[..., User],
- session_factory,
- session_manager,
+ session_factory: Callable[..., Session],
+ session_manager: SessionManager,
):
# User Payout xform NULL: all payout fields are None
user: User = user_factory(product=bp1)
@@ -352,16 +361,17 @@ class TestTaskStatus:
}
)
tsr = session_manager.get_task_status_response(s.uuid)
+ assert isinstance(tsr, TaskStatusResponse)
# Not bothering with wall events ...
tsr.wall_events = None
assert tsr == expected_tsr
def test_task_status_abandon_xform(
self,
- bp2,
+ bp2: Product,
user_factory: Callable[..., User],
- session_factory,
- session_manager,
+ session_factory: Callable[..., Session],
+ session_manager: SessionManager,
):
# User Payout xform 40%: all payout fields are None (same as when payout xform is null)
user: User = user_factory(product=bp2)
@@ -393,17 +403,18 @@ class TestTaskStatus:
}
)
tsr = session_manager.get_task_status_response(s.uuid)
+ assert isinstance(tsr, TaskStatusResponse)
# Not bothering with wall events ...
tsr.wall_events = None
assert tsr == expected_tsr
def test_task_status_adj_fail(
self,
- bp1,
+ bp1: Product,
user_factory: Callable[..., User],
- finished_session_factory,
- wall_manager,
- session_manager,
+ finished_session_factory: Callable[..., Session],
+ wall_manager: WallManager,
+ session_manager: SessionManager,
):
# Complete -> Fail
# User Payout xform NULL: adjusted_user_* and user_* is still all None
@@ -442,17 +453,18 @@ class TestTaskStatus:
}
)
tsr = session_manager.get_task_status_response(s.uuid)
+ assert isinstance(tsr, TaskStatusResponse)
# Not bothering with wall events ...
tsr.wall_events = None
assert tsr == expected_tsr
def test_task_status_adj_fail_xform(
self,
- bp2,
+ bp2: Product,
user_factory: Callable[..., User],
- finished_session_factory,
- wall_manager,
- session_manager,
+ finished_session_factory: Callable[..., Session],
+ wall_manager: WallManager,
+ session_manager: SessionManager,
):
# Complete -> Fail
# User Payout xform 40%: adjusted_user_payout is 0 (not null)
@@ -494,17 +506,18 @@ class TestTaskStatus:
}
)
tsr = session_manager.get_task_status_response(s.uuid)
+ assert isinstance(tsr, TaskStatusResponse)
# Not bothering with wall events ...
tsr.wall_events = None
assert tsr == expected_tsr
def test_task_status_adj_complete_from_abandon(
self,
- bp1,
+ bp1: Product,
user_factory: Callable[..., User],
- session_factory,
- wall_manager,
- session_manager,
+ session_factory: Callable[..., Session],
+ wall_manager: WallManager,
+ session_manager: SessionManager,
):
# User Payout xform NULL
user: User = user_factory(product=bp1)
@@ -548,17 +561,18 @@ class TestTaskStatus:
}
)
tsr = session_manager.get_task_status_response(s.uuid)
+ assert isinstance(tsr, TaskStatusResponse)
# Not bothering with wall events ...
tsr.wall_events = None
assert tsr == expected_tsr
def test_task_status_adj_complete_from_abandon_xform(
self,
- bp2,
+ bp2: Product,
user_factory: Callable[..., User],
- session_factory,
- wall_manager,
- session_manager,
+ session_factory: Callable[..., Session],
+ wall_manager: WallManager,
+ session_manager: SessionManager,
):
# User Payout xform 40%
user: User = user_factory(product=bp2)
@@ -605,17 +619,18 @@ class TestTaskStatus:
}
)
tsr = session_manager.get_task_status_response(s.uuid)
+ assert isinstance(tsr, TaskStatusResponse)
# Not bothering with wall events ...
tsr.wall_events = None
assert tsr == expected_tsr
def test_task_status_adj_complete_from_fail(
self,
- bp1,
+ bp1: Product,
user_factory: Callable[..., User],
- finished_session_factory,
- wall_manager,
- session_manager,
+ finished_session_factory: Callable[..., Session],
+ wall_manager: WallManager,
+ session_manager: SessionManager,
):
# User Payout xform NULL
user: User = user_factory(product=bp1)
@@ -659,17 +674,18 @@ class TestTaskStatus:
}
)
tsr = session_manager.get_task_status_response(s.uuid)
+ assert isinstance(tsr, TaskStatusResponse)
# Not bothering with wall events ...
tsr.wall_events = None
assert tsr == expected_tsr
def test_task_status_adj_complete_from_fail_xform(
self,
- bp2,
+ bp2: Product,
user_factory: Callable[..., User],
- finished_session_factory,
- wall_manager,
- session_manager,
+ finished_session_factory: Callable[..., Session],
+ wall_manager: WallManager,
+ session_manager: SessionManager,
):
# User Payout xform 40%
user: User = user_factory(product=bp2)
@@ -715,6 +731,7 @@ class TestTaskStatus:
}
)
tsr = session_manager.get_task_status_response(s.uuid)
+ assert isinstance(tsr, TaskStatusResponse)
# Not bothering with wall events ...
tsr.wall_events = None
assert tsr == expected_tsr
diff --git a/tests/managers/thl/test_user_manager/test_base.py b/tests/managers/thl/test_user_manager/test_base.py
index 6b259ff..5822207 100644
--- a/tests/managers/thl/test_user_manager/test_base.py
+++ b/tests/managers/thl/test_user_manager/test_base.py
@@ -10,14 +10,18 @@ from generalresearch.managers.thl.user_manager import (
UserCreateNotAllowedError,
get_bp_user_create_limit_hourly,
)
+from generalresearch.managers.thl.user_manager.mysql_user_manager import (
+ MysqlUserManager,
+)
from generalresearch.managers.thl.user_manager.rate_limit import (
RateLimitItemPerHourConstantKey,
+ UserManagerLimiter,
)
from generalresearch.managers.thl.user_manager.user_manager import (
UserManager,
)
from generalresearch.managers.thl.userhealth import AuditLogManager
-from generalresearch.models.thl.product import product: Product, UserCreateConfig
+from generalresearch.models.thl.product import Product, UserCreateConfig, product
from generalresearch.models.thl.user import User
from generalresearch.pg_helper import PostgresConfig
@@ -86,10 +90,11 @@ class TestUserManager:
class TestBlockUserManager:
- def test_block_user(self, product: product: 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
+ assert isinstance(user_manager.mysql_user_manager, MysqlUserManager)
user: User = user_manager.mysql_user_manager.create_user(
product_id=product.id, product_user_id=product_user_id
)
@@ -113,11 +118,12 @@ class TestBlockUserManager:
assert user.blocked
def test_block_user_whitelist(
- self, product: product: Product, user_manager: UserManager, thl_web_rw: PostgresConfig
+ 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
+ assert isinstance(user_manager.mysql_user_manager, MysqlUserManager)
user: User = user_manager.mysql_user_manager.create_user(
product_id=product.id, product_user_id=product_user_id
)
@@ -154,6 +160,7 @@ class TestCreateUserManager:
product_user_id = f"user-{uuid4().hex[:10]}"
+ assert isinstance(user_manager.mysql_user_manager, MysqlUserManager)
user: User = user_manager.mysql_user_manager.create_user(
product_id=product.id, product_user_id=product_user_id
)
@@ -200,6 +207,7 @@ class TestCreateUserManager:
product_user_id = f"user-{uuid4().hex[:10]}"
rand_msg = f"log-{uuid4().hex}"
+ assert isinstance(user_manager.mysql_user_manager, MysqlUserManager)
with caplog.at_level(logging.INFO):
logger.info(rand_msg)
user1 = user_manager.mysql_user_manager.create_user(
@@ -264,6 +272,7 @@ class TestCreateUserManager:
assert key == f"LIMITER/thl-grpc/allow_user_create/{instance.id}"
# make sure we clear the key or subsequent tests will fail
+ assert isinstance(user_manager.user_manager_limiter, UserManagerLimiter)
user_manager.user_manager_limiter.storage.clear(key=key)
n = 0
diff --git a/tests/managers/thl/test_user_manager/test_mysql.py b/tests/managers/thl/test_user_manager/test_mysql.py
index d414a13..e6f43ef 100644
--- a/tests/managers/thl/test_user_manager/test_mysql.py
+++ b/tests/managers/thl/test_user_manager/test_mysql.py
@@ -1,24 +1,25 @@
+from __future__ import annotations
+
+from generalresearch.managers.thl.user_manager.mysql_user_manager import (
+ MysqlUserManager,
+)
+from generalresearch.models.thl.user import User
class TestUserManagerMysqlNew:
- def test_get_notset(self, user_manager):
- assert (
- user_manager.mysql_user_manager.get_user_from_mysql(user_id=-3105) is None
- )
+ def test_get_notset(self, mysql_user_manager: MysqlUserManager):
+ assert mysql_user_manager.get_user_from_mysql(user_id=-3105) is None
- def test_get_user_id(self, user, user_manager):
- assert (
- user_manager.mysql_user_manager.get_user_from_mysql(user_id=user.user_id)
- == user
- )
+ def test_get_user_id(self, user: User, mysql_user_manager: MysqlUserManager):
+ assert mysql_user_manager.get_user_from_mysql(user_id=user.user_id) == user
- def test_get_uuid(self, user, user_manager):
- u = user_manager.mysql_user_manager.get_user_from_mysql(user_uuid=user.uuid)
+ def test_get_uuid(self, user: User, mysql_user_manager: MysqlUserManager):
+ u = mysql_user_manager.get_user_from_mysql(user_uuid=user.uuid)
assert u == user
- def test_get_ubp(self, user, user_manager):
- u = user_manager.mysql_user_manager.get_user_from_mysql(
+ def test_get_ubp(self, user: User, mysql_user_manager: MysqlUserManager):
+ u = mysql_user_manager.get_user_from_mysql(
product_id=user.product_id, product_user_id=user.product_user_id
)
assert u == user
diff --git a/tests/managers/thl/test_user_manager/test_redis.py b/tests/managers/thl/test_user_manager/test_redis.py
index 0731438..04071ee 100644
--- a/tests/managers/thl/test_user_manager/test_redis.py
+++ b/tests/managers/thl/test_user_manager/test_redis.py
@@ -1,29 +1,37 @@
+from __future__ import annotations
+
import pytest
+from generalresearch.config import GRLBaseSettings
from generalresearch.managers.base import Permission
+from generalresearch.managers.thl.user_manager.redis_user_manager import (
+ RedisUserManager,
+)
+from generalresearch.models.thl.user import User
+from generalresearch.pg_helper import PostgresConfig
class TestUserManagerRedis:
- def test_get_notset(self, user_manager, user):
- user_manager.clear_user_inmemory_cache(user=user)
- assert user_manager.redis_user_manager.get_user(user_id=user.user_id) is None
+ def test_get_notset(self, redis_user_manager: RedisUserManager, user: User):
+ redis_user_manager.clear_user_inmemory_cache(user=user)
+ assert redis_user_manager.get_user(user_id=user.user_id) is None
- def test_get_user_id(self, user_manager, user):
- user_manager.redis_user_manager.set_user(user=user)
+ def test_get_user_id(self, redis_user_manager: RedisUserManager, user: User):
+ redis_user_manager.set_user(user=user)
- assert user_manager.redis_user_manager.get_user(user_id=user.user_id) == user
+ assert redis_user_manager.get_user(user_id=user.user_id) == user
- def test_get_uuid(self, user_manager, user):
- user_manager.redis_user_manager.set_user(user=user)
+ def test_get_uuid(self, redis_user_manager: RedisUserManager, user: User):
+ redis_user_manager.set_user(user=user)
- assert user_manager.redis_user_manager.get_user(user_uuid=user.uuid) == user
+ assert redis_user_manager.get_user(user_uuid=user.uuid) == user
- def test_get_ubp(self, user_manager, user):
- user_manager.redis_user_manager.set_user(user=user)
+ def test_get_ubp(self, redis_user_manager: RedisUserManager, user: User):
+ redis_user_manager.set_user(user=user)
assert (
- user_manager.redis_user_manager.get_user(
+ redis_user_manager.get_user(
product_id=user.product_id, product_user_id=user.product_user_id
)
== user
@@ -34,7 +42,13 @@ class TestUserManagerRedis:
# I mean, the sets are implicitly tested by the get tests above. no point
pass
- def test_get_with_cache_prefix(self, settings, user, thl_web_rw, thl_web_rr):
+ def test_get_with_cache_prefix(
+ self,
+ settings: GRLBaseSettings,
+ user: User,
+ thl_web_rw: PostgresConfig,
+ thl_web_rr: PostgresConfig,
+ ):
"""
Confirm the prefix functionality is working; we do this so it
is easier to migrate between any potentially breaking versions
@@ -47,7 +61,7 @@ class TestUserManagerRedis:
um1 = UserManager(
pg_config=thl_web_rw,
- pg_config_rr=thl_web_rr: PostgresConfig,
+ pg_config_rr=thl_web_rr,
sql_permissions=[Permission.UPDATE, Permission.CREATE],
redis=settings.redis,
redis_timeout=settings.redis_timeout,
@@ -55,7 +69,7 @@ class TestUserManagerRedis:
um2 = UserManager(
pg_config=thl_web_rw,
- pg_config_rr=thl_web_rr: PostgresConfig,
+ pg_config_rr=thl_web_rr,
sql_permissions=[Permission.UPDATE, Permission.CREATE],
redis=settings.redis,
redis_timeout=settings.redis_timeout,
@@ -69,9 +83,11 @@ class TestUserManagerRedis:
product_id=user.product_id, product_user_id=user.product_user_id
)
+ assert isinstance(um1.redis_user_manager, RedisUserManager)
res1 = um1.redis_user_manager.client.get(f"user-lookup:user_id:{user.user_id}")
assert res1 is not None
+ assert isinstance(um2.redis_user_manager, RedisUserManager)
res2 = um2.redis_user_manager.client.get(
f"user-lookup-v2:user_id:{user.user_id}"
)
diff --git a/tests/managers/thl/test_user_manager/test_user_fetch.py b/tests/managers/thl/test_user_manager/test_user_fetch.py
index 5c608b3..87d010a 100644
--- a/tests/managers/thl/test_user_manager/test_user_fetch.py
+++ b/tests/managers/thl/test_user_manager/test_user_fetch.py
@@ -1,14 +1,22 @@
+from __future__ import annotations
+
+from collections.abc import Callable
from uuid import uuid4
import pytest
+from generalresearch.managers.thl.user_manager.user_manager import UserManager
+from generalresearch.models.thl.product import Product
from generalresearch.models.thl.user import User
class TestUserManagerFetch:
def test_fetch(
- self, user_factory: Callable[..., User], product: Product, user_manager
+ self,
+ user_factory: Callable[..., User],
+ product: Product,
+ user_manager: UserManager,
):
user1: User = user_factory(product=product)
user2: User = user_factory(product=product)
@@ -31,7 +39,7 @@ class TestUserManagerFetch:
res = user_manager.fetch(user_uuids=[uuid4().hex])
assert len(res) == 0
- def test_fetch_invalid(self, user_manager):
+ def test_fetch_invalid(self, user_manager: UserManager):
with pytest.raises(AssertionError) as e:
user_manager.fetch(user_uuids=[], user_ids=None)
assert "Must pass ONE of user_ids, user_uuids" in str(e.value)
diff --git a/tests/managers/thl/test_user_manager/test_user_metadata.py b/tests/managers/thl/test_user_manager/test_user_metadata.py
index 0b99afe..670e38a 100644
--- a/tests/managers/thl/test_user_manager/test_user_metadata.py
+++ b/tests/managers/thl/test_user_manager/test_user_metadata.py
@@ -1,21 +1,35 @@
+from __future__ import annotations
+
+from collections.abc import Callable
from uuid import uuid4
import pytest
+from generalresearch.managers.thl.user_manager.user_metadata_manager import (
+ UserMetadataManager,
+)
+from generalresearch.models.thl.product import Product
+from generalresearch.models.thl.user import User
from generalresearch.models.thl.user_profile import UserMetadata
class TestUserMetadataManager:
- def test_get_notset(self, user, user_manager, user_metadata_manager):
+ def test_get_notset(
+ self,
+ user: User,
+ user_metadata_manager: UserMetadataManager,
+ ):
# The row in the db won't exist. It just returns the default obj with everything None (except for the user_id)
um1 = user_metadata_manager.get(user_id=user.user_id)
assert um1 == UserMetadata(user_id=user.user_id)
def test_create(
- self, user_factory: Callable[..., User], product: Product, user_metadata_manager
+ self,
+ user_factory: Callable[..., User],
+ product: Product,
+ user_metadata_manager: UserMetadataManager,
):
- from generalresearch.models.thl.user import User
u1: User = user_factory(product=product)
@@ -29,9 +43,11 @@ class TestUserMetadataManager:
assert um == um2
def test_create_no_email(
- self, product: Product, user_factory: Callable[..., User], user_metadata_manager
+ self,
+ product: Product,
+ user_factory: Callable[..., User],
+ user_metadata_manager: UserMetadataManager,
):
- from generalresearch.models.thl.user import User
u1: User = user_factory(product=product)
um = UserMetadata(user_id=u1.user_id)
@@ -42,9 +58,11 @@ class TestUserMetadataManager:
assert um == um2
def test_update(
- self, product: Product, user_factory: Callable[..., User], user_metadata_manager
+ self,
+ product: Product,
+ user_factory: Callable[..., User],
+ user_metadata_manager: UserMetadataManager,
):
- from generalresearch.models.thl.user import User
u: User = user_factory(product=product)
@@ -66,7 +84,6 @@ class TestUserMetadataManager:
def test_filter(
self, user_factory: Callable[..., User], product: Product, user_metadata_manager
):
- from generalresearch.models.thl.user import User
user1: User = user_factory(product=product)
user2: User = user_factory(product=product)
diff --git a/tests/managers/thl/test_user_streak.py b/tests/managers/thl/test_user_streak.py
index e87869f..61e2947 100644
--- a/tests/managers/thl/test_user_streak.py
+++ b/tests/managers/thl/test_user_streak.py
@@ -1,3 +1,5 @@
+from __future__ import annotations
+
import copy
from datetime import UTC, date, datetime, timedelta
from decimal import Decimal
@@ -5,8 +7,13 @@ from zoneinfo import ZoneInfo
import pytest
-from generalresearch.managers.thl.user_streak import compute_streaks_from_days
+from generalresearch.managers.thl.session import SessionManager
+from generalresearch.managers.thl.user_streak import (
+ UserStreakManager,
+ compute_streaks_from_days,
+)
from generalresearch.models.thl.definitions import Status, StatusCode1
+from generalresearch.models.thl.user import User
from generalresearch.models.thl.user_streak import (
StreakFulfillment,
StreakPeriod,
@@ -59,7 +66,7 @@ def test_compute_streaks_from_days():
@pytest.fixture
-def broken_active_streak(user):
+def broken_active_streak(user: User) -> list[UserStreak]:
return [
UserStreak(
period=StreakPeriod.DAY,
@@ -94,7 +101,7 @@ def broken_active_streak(user):
]
-def create_session_fail(session_manager, start, user):
+def create_session_fail(session_manager: SessionManager, start: datetime, user: User):
session = session_manager.create_dummy(started=start, country_iso="us", user=user)
session_manager.finish_with_status(
session,
@@ -104,7 +111,9 @@ def create_session_fail(session_manager, start, user):
)
-def create_session_complete(session_manager, start, user):
+def create_session_complete(
+ session_manager: SessionManager, start: datetime, user: User
+):
session = session_manager.create_dummy(started=start, country_iso="us", user=user)
session_manager.finish_with_status(
session,
@@ -115,7 +124,7 @@ def create_session_complete(session_manager, start, user):
)
-def test_user_streak_empty(user_streak_manager, user):
+def test_user_streak_empty(user_streak_manager: UserStreakManager, user: User):
streaks = user_streak_manager.get_user_streaks(
user_id=user.user_id, country_iso="us"
)
@@ -123,7 +132,10 @@ def test_user_streak_empty(user_streak_manager, user):
def test_user_streaks_active_broken(
- user_streak_manager, user, session_manager, broken_active_streak
+ user_streak_manager: UserStreakManager,
+ user: User,
+ session_manager: SessionManager,
+ broken_active_streak: list[UserStreak],
):
# Testing active streak, but broken (not today or yesterday)
start1 = datetime(2025, 2, 12, tzinfo=UTC)
@@ -171,7 +183,9 @@ def test_user_streaks_active_broken(
assert streaks == expected_streaks
-def test_user_streak_complete_active(user_streak_manager, user, session_manager):
+def test_user_streak_complete_active(
+ user_streak_manager: UserStreakManager, user: User, session_manager: SessionManager
+):
"""Testing active streak that is today"""
# They completed yesterday NY time. Today isn't over so streak is pending
@@ -217,9 +231,9 @@ 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
diff --git a/tests/managers/thl/test_userhealth.py b/tests/managers/thl/test_userhealth.py
index 98d8f25..ea54359 100644
--- a/tests/managers/thl/test_userhealth.py
+++ b/tests/managers/thl/test_userhealth.py
@@ -1,3 +1,6 @@
+from __future__ import annotations
+
+from collections.abc import Callable
from datetime import UTC, datetime
from uuid import uuid4
@@ -5,23 +8,27 @@ import faker
import pytest
from generalresearch.managers.thl.userhealth import (
+ AuditLogManager,
IPRecordManager,
UserIpHistoryManager,
)
-from generalresearch.models.thl.ipinfo import GeoIPInformation
+from generalresearch.models.thl.ipinfo import GeoIPInformation, IPGeoname, IPInformation
+from generalresearch.models.thl.product import Product
+from generalresearch.models.thl.user import User
from generalresearch.models.thl.user_iphistory import (
IPRecord,
+ UserIPHistory,
)
from generalresearch.models.thl.userhealth import AuditLog, AuditLogLevel
+from generalresearch.pg_helper import PostgresConfig
+from generalresearch.redis_helper import RedisConfig
fake = faker.Faker()
class TestAuditLog:
- def test_init(self, thl_web_rr: PostgresConfig, audit_log_manager):
- from generalresearch.managers.thl.userhealth import AuditLogManager
-
+ def test_init(self, thl_web_rr: PostgresConfig, audit_log_manager: AuditLogManager):
alm = AuditLogManager(pg_config=thl_web_rr)
assert isinstance(alm, AuditLogManager)
@@ -33,15 +40,16 @@ class TestAuditLog:
argnames="level",
argvalues=list(AuditLogLevel),
)
- def test_create(self, audit_log_manager, user, level):
+ def test_create(
+ self, audit_log_manager: AuditLogManager, user: User, level: AuditLogLevel
+ ):
instance = audit_log_manager.create(
user_id=user.user_id, level=level, event_type=uuid4().hex
)
assert isinstance(instance, AuditLog)
assert instance.id != 1
- def test_get_by_id(self, audit_log, audit_log_manager):
- from generalresearch.models.thl.userhealth import AuditLog
+ def test_get_by_id(self, audit_log: AuditLog, audit_log_manager: AuditLogManager):
with pytest.raises(expected_exception=Exception) as cm:
audit_log_manager.get_by_id(auditlog_id=999_999_999_999)
@@ -57,8 +65,8 @@ class TestAuditLog:
self,
user_factory: Callable[..., User],
product_factory: Callable[..., Product],
- audit_log_factory,
- audit_log_manager,
+ audit_log_factory: Callable[..., AuditLog],
+ audit_log_manager: AuditLogManager,
):
p1 = product_factory()
p2 = product_factory()
@@ -82,7 +90,11 @@ class TestAuditLog:
assert len(res) == 1
def test_filter_by_user_id(
- self, user_factory: Callable[..., User], product: Product, audit_log_factory, audit_log_manager
+ self,
+ user_factory: Callable[..., User],
+ product: Product,
+ audit_log_factory: Callable[..., AuditLog],
+ audit_log_manager: AuditLogManager,
):
u1 = user_factory(product=product)
u2 = user_factory(product=product)
@@ -110,8 +122,8 @@ class TestAuditLog:
self,
user_factory: Callable[..., User],
product_factory: Callable[..., Product],
- audit_log_factory,
- audit_log_manager,
+ audit_log_factory: Callable[..., AuditLog],
+ audit_log_manager: AuditLogManager,
):
p1 = product_factory()
p2 = product_factory()
@@ -144,8 +156,8 @@ class TestAuditLog:
self,
user_factory: Callable[..., User],
product_factory: Callable[..., Product],
- audit_log_factory,
- audit_log_manager,
+ audit_log_factory: Callable[..., AuditLog],
+ audit_log_manager: AuditLogManager,
):
p1 = product_factory()
p2 = product_factory()
@@ -205,18 +217,29 @@ class TestAuditLog:
class TestIPRecordManager:
- def test_init(self, thl_web_rr: PostgresConfig, thl_redis_config: RedisConfig, ip_record_manager):
- instance = IPRecordManager(pg_config=thl_web_rr: PostgresConfig, redis_config=thl_redis_config)
+ def test_init(
+ self,
+ thl_web_rr: PostgresConfig,
+ thl_redis_config: RedisConfig,
+ ip_record_manager: IPRecordManager,
+ ):
+ instance = IPRecordManager(pg_config=thl_web_rr, redis_config=thl_redis_config)
assert isinstance(instance, IPRecordManager)
assert isinstance(ip_record_manager, IPRecordManager)
- def test_create(self, ip_record_manager, user, ip_information):
+ def test_create(
+ self,
+ ip_record_manager: IPRecordManager,
+ user: User,
+ ip_information: IPInformation,
+ ):
instance = ip_record_manager.create_dummy(
user_id=user.user_id, ip=ip_information.ip
)
assert isinstance(instance, IPRecord)
assert isinstance(instance.forwarded_ips, list)
+ assert isinstance(instance.forwarded_ip_records, list)
assert isinstance(instance.forwarded_ip_records[0], IPRecord)
assert isinstance(instance.forwarded_ips[0], str)
@@ -228,10 +251,10 @@ class TestIPRecordManager:
def test_prefetch_info(
self,
- ip_record_factory,
- ip_information_factory,
- ip_geoname,
- user,
+ ip_record_factory: Callable[..., IPRecord],
+ ip_information_factory: Callable[..., IPInformation],
+ ip_geoname: IPGeoname,
+ user: User,
thl_web_rr: PostgresConfig,
thl_redis_config: RedisConfig,
):
@@ -239,15 +262,17 @@ class TestIPRecordManager:
ip = fake.ipv4_public()
ip_information_factory(ip=ip, geoname=ip_geoname)
ipr: IPRecord = ip_record_factory(user_id=user.user_id, ip=ip)
+ assert isinstance(ipr, IPRecord)
assert ipr.information is None
assert len(ipr.forwarded_ip_records) >= 1
+ assert isinstance(ipr.forwarded_ip_records, list)
fipr = ipr.forwarded_ip_records[0]
assert fipr.information is None
ipr.prefetch_ipinfo(
- pg_config=thl_web_rr: PostgresConfig,
- redis_config=thl_redis_config: RedisConfig,
+ pg_config=thl_web_rr,
+ redis_config=thl_redis_config,
include_forwarded=True,
)
assert isinstance(ipr.information, GeoIPInformation)
@@ -256,8 +281,8 @@ class TestIPRecordManager:
ip_information_factory(ip=fipr.ip, geoname=ip_geoname)
ipr.prefetch_ipinfo(
- pg_config=thl_web_rr: PostgresConfig,
- redis_config=thl_redis_config: RedisConfig,
+ pg_config=thl_web_rr,
+ redis_config=thl_redis_config,
include_forwarded=True,
)
assert fipr.information is not None
@@ -265,28 +290,35 @@ class TestIPRecordManager:
@pytest.mark.usefixtures("user_iphistory_manager_clear_cache")
class TestUserIpHistoryManager:
- def test_init(self, thl_web_rr: PostgresConfig, thl_redis_config: RedisConfig, user_iphistory_manager):
+ def test_init(
+ self,
+ thl_web_rr: PostgresConfig,
+ thl_redis_config: RedisConfig,
+ user_iphistory_manager: UserIpHistoryManager,
+ ):
instance = UserIpHistoryManager(
- pg_config=thl_web_rr: PostgresConfig, redis_config=thl_redis_config
+ pg_config=thl_web_rr, redis_config=thl_redis_config
)
assert isinstance(instance, UserIpHistoryManager)
assert isinstance(user_iphistory_manager, UserIpHistoryManager)
def test_latest_record(
self,
- user_iphistory_manager,
- user,
- ip_record_factory,
- ip_information_factory,
- ip_geoname,
+ user_iphistory_manager: UserIpHistoryManager,
+ user: User,
+ ip_record_factory: Callable[..., IPRecord],
+ ip_information_factory: Callable[..., IPInformation],
+ ip_geoname: IPGeoname,
):
ip = fake.ipv4_public()
ip_information_factory(ip=ip, geoname=ip_geoname, is_anonymous=True)
ipr1: IPRecord = ip_record_factory(user_id=user.user_id, ip=ip)
ipr = user_iphistory_manager.get_user_latest_ip_record(user=user)
+ assert isinstance(ipr, IPRecord)
assert ipr.ip == ipr1.ip
assert ipr.is_anonymous
+ assert isinstance(ipr.information, GeoIPInformation)
assert ipr.information.lookup_prefix == "/32"
ip = fake.ipv6()
@@ -294,7 +326,9 @@ class TestUserIpHistoryManager:
ipr2: IPRecord = ip_record_factory(user_id=user.user_id, ip=ip)
ipr = user_iphistory_manager.get_user_latest_ip_record(user=user)
+ assert isinstance(ipr, IPRecord)
assert ipr.ip == ipr2.ip
+ assert isinstance(ipr.information, GeoIPInformation)
assert ipr.information.lookup_prefix == "/64"
assert ipr.information is not None
assert not ipr.is_anonymous
@@ -303,6 +337,8 @@ class TestUserIpHistoryManager:
assert country_iso == ip_geoname.country_iso
iph = user_iphistory_manager.get_user_ip_history(user_id=user.user_id)
+ assert isinstance(iph, UserIPHistory)
+ assert isinstance(iph.ips, list)
assert iph.ips[0].information is not None
assert iph.ips[1].information is not None
assert iph.ips[0].country_iso == country_iso
@@ -310,7 +346,12 @@ class TestUserIpHistoryManager:
assert iph.ips[0].ip == ipr1.ip
assert iph.ips[1].ip == ipr2.ip
- def test_virgin(self, user, user_iphistory_manager, ip_record_factory):
+ def test_virgin(
+ self,
+ user: User,
+ user_iphistory_manager: UserIpHistoryManager,
+ ip_record_factory: Callable[..., IPRecord],
+ ):
iph = user_iphistory_manager.get_user_ip_history(user_id=user.user_id)
assert len(iph.ips) == 0
@@ -320,16 +361,18 @@ class TestUserIpHistoryManager:
def test_out_of_order(
self,
- ip_record_factory,
- user,
- user_iphistory_manager,
- ip_information_factory,
- ip_geoname,
+ ip_record_factory: Callable[..., IPRecord],
+ user: User,
+ user_iphistory_manager: UserIpHistoryManager,
+ ip_information_factory: Callable[..., IPInformation],
+ ip_geoname: IPGeoname,
):
# Create the user-ip association BEFORE the ip even exists in the ipinfo table
ip = fake.ipv4_public()
ip_record_factory(user_id=user.user_id, ip=ip)
iph = user_iphistory_manager.get_user_ip_history(user_id=user.user_id)
+ assert isinstance(iph, UserIPHistory)
+ assert isinstance(iph.ips, list)
assert len(iph.ips) == 1
ipr = iph.ips[0]
assert ipr.information is None
@@ -337,6 +380,8 @@ class TestUserIpHistoryManager:
ip_information_factory(ip=ip, geoname=ip_geoname, is_anonymous=True)
iph = user_iphistory_manager.get_user_ip_history(user_id=user.user_id)
+ assert isinstance(iph, UserIPHistory)
+ assert isinstance(iph.ips, list)
assert len(iph.ips) == 1
ipr = iph.ips[0]
assert ipr.information is not None
@@ -344,16 +389,18 @@ class TestUserIpHistoryManager:
def test_out_of_order_ipv6(
self,
- ip_record_factory,
- user,
- user_iphistory_manager,
- ip_information_factory,
- ip_geoname,
+ ip_record_factory: Callable[..., IPRecord],
+ user: User,
+ user_iphistory_manager: UserIpHistoryManager,
+ ip_information_factory: Callable[..., IPInformation],
+ ip_geoname: IPGeoname,
):
# Create the user-ip association BEFORE the ip even exists in the ipinfo table
ip = fake.ipv6()
ip_record_factory(user_id=user.user_id, ip=ip)
iph = user_iphistory_manager.get_user_ip_history(user_id=user.user_id)
+ assert isinstance(iph, UserIPHistory)
+ assert isinstance(iph.ips, list)
assert len(iph.ips) == 1
ipr = iph.ips[0]
assert ipr.information is None
@@ -361,6 +408,8 @@ class TestUserIpHistoryManager:
ip_information_factory(ip=ip, geoname=ip_geoname, is_anonymous=True)
iph = user_iphistory_manager.get_user_ip_history(user_id=user.user_id)
+ assert isinstance(iph, UserIPHistory)
+ assert isinstance(iph.ips, list)
assert len(iph.ips) == 1
ipr = iph.ips[0]
assert ipr.information is not None
diff --git a/tests/managers/thl/test_wall_manager.py b/tests/managers/thl/test_wall_manager.py
index 1071627..b8a636f 100644
--- a/tests/managers/thl/test_wall_manager.py
+++ b/tests/managers/thl/test_wall_manager.py
@@ -1,24 +1,36 @@
+from __future__ import annotations
+
+from collections.abc import Callable
from datetime import UTC, datetime, timedelta
from decimal import Decimal
from uuid import uuid4
import pytest
+from pydantic import PositiveInt
+from generalresearch.managers.thl.session import SessionManager
+from generalresearch.managers.thl.wall import WallCacheManager, WallManager
from generalresearch.models import Source
from generalresearch.models.thl.session import (
ReportValue,
+ Session,
Status,
StatusCode1,
)
+from generalresearch.models.thl.user import User
class TestWallManager:
@pytest.mark.parametrize("wall_count", [1, 2, 5, 10, 50, 99])
def test_get_wall_events(
- self, wall_manager, session_factory, user, wall_count, utc_hour_ago
+ self,
+ wall_manager: WallManager,
+ session_factory: Callable[..., Session],
+ user: User,
+ wall_count: PositiveInt,
+ utc_hour_ago: datetime,
):
- from generalresearch.models.thl.session import Session
s1: Session = session_factory(
user=user, wall_count=wall_count, started=utc_hour_ago
@@ -62,12 +74,15 @@ class TestWallManager:
]
def test_get_wall_events_list_input(
- self, wall_manager, session_factory, user, utc_hour_ago
+ self,
+ wall_manager: WallManager,
+ session_factory: Callable[..., Session],
+ user: User,
+ utc_hour_ago: datetime,
):
- from generalresearch.models.thl.session import Session
session_ids = []
- for idx in range(10):
+ for _ in range(10):
s: Session = session_factory(user=user, wall_count=5, started=utc_hour_ago)
session_ids.append(s.id)
@@ -82,7 +97,7 @@ class TestWallManager:
assert session_ids == res1
- def test_create_wall(self, wall_manager, user, session):
+ def test_create_wall(self, wall_manager: WallManager, user: User, session: Session):
w = wall_manager.create(
session_id=session.id,
user_id=user.user_id,
@@ -98,7 +113,13 @@ class TestWallManager:
w2 = wall_manager.get_from_uuid(wall_uuid=w.uuid)
assert w == w2
- def test_report_wall_abandon(self, wall_manager, user, session, utc_hour_ago):
+ def test_report_wall_abandon(
+ self,
+ wall_manager: WallManager,
+ user: User,
+ session: Session,
+ utc_hour_ago: datetime,
+ ):
w1 = wall_manager.create(
session_id=session.id,
user_id=user.user_id,
@@ -138,7 +159,12 @@ class TestWallManager:
# the status and finished get updated
def test_report_wall(
- self, wall_manager, session_manager, user, session, utc_hour_ago
+ self,
+ wall_manager: WallManager,
+ session_manager: SessionManager,
+ user: User,
+ session: Session,
+ utc_hour_ago: datetime,
):
w1 = wall_manager.create(
session_id=session.id,
@@ -174,7 +200,13 @@ class TestWallManager:
assert Status.COMPLETE == w2.status
assert "This survey blows!" == w2.report_notes
- def test_filter_wall_attempts(self, wall_manager, user, session, utc_hour_ago):
+ def test_filter_wall_attempts(
+ self,
+ wall_manager: WallManager,
+ user: User,
+ session: Session,
+ utc_hour_ago: datetime,
+ ):
res = wall_manager.filter_wall_attempts(user_id=user.user_id)
assert len(res) == 0
wall_manager.create(
@@ -205,12 +237,16 @@ class TestWallManager:
class TestWallCacheManager:
- def test_get_attempts_none(self, wall_cache_manager, user):
+ def test_get_attempts_none(self, wall_cache_manager: WallCacheManager, user: User):
attempts = wall_cache_manager.get_attempts(user.user_id)
assert len(attempts) == 0
def test_get_wall_events(
- self, wall_cache_manager, wall_manager, session_manager, user
+ self,
+ wall_cache_manager: WallCacheManager,
+ wall_manager: WallManager,
+ session_manager: SessionManager,
+ user: User,
):
start1 = datetime.now(UTC) - timedelta(hours=3)
start2 = datetime.now(UTC) - timedelta(hours=2)
diff --git a/tests/models/custom_types/test_dsn.py b/tests/models/custom_types/test_dsn.py
index eb27526..d8c7c53 100644
--- a/tests/models/custom_types/test_dsn.py
+++ b/tests/models/custom_types/test_dsn.py
@@ -1,3 +1,5 @@
+from __future__ import annotations
+
from uuid import uuid4
import pytest
diff --git a/tests/models/custom_types/test_therest.py b/tests/models/custom_types/test_therest.py
index 13e9bae..01bc644 100644
--- a/tests/models/custom_types/test_therest.py
+++ b/tests/models/custom_types/test_therest.py
@@ -1,3 +1,5 @@
+from __future__ import annotations
+
import json
from uuid import UUID
diff --git a/tests/models/dynata/test_eligbility.py b/tests/models/dynata/test_eligbility.py
index 27de5b3..b3a9f13 100644
--- a/tests/models/dynata/test_eligbility.py
+++ b/tests/models/dynata/test_eligbility.py
@@ -1,3 +1,5 @@
+from __future__ import annotations
+
from datetime import UTC, datetime
diff --git a/tests/models/gr/test_authentication.py b/tests/models/gr/test_authentication.py
index d4db112..d2a7054 100644
--- a/tests/models/gr/test_authentication.py
+++ b/tests/models/gr/test_authentication.py
@@ -1,3 +1,5 @@
+from __future__ import annotations
+
import binascii
import json
import os
@@ -7,9 +9,14 @@ from random import randint
from uuid import uuid4
import pytest
+from redis import Redis
-from generalresearch.models.gr.authentication import GRUser
+from generalresearch.models.gr.authentication import Claims, GRToken, GRUser
+from generalresearch.models.gr.business import Business
from generalresearch.models.gr.team import Membership, Team
+from generalresearch.models.thl.product import Product
+from generalresearch.pg_helper import PostgresConfig
+from generalresearch.redis_helper import RedisConfig
SSO_ISSUER = ""
@@ -29,7 +36,13 @@ class TestGRUser:
def test_businesses(self):
pass
- def test_teams(self, gr_user: GRUser, membership, gr_db, gr_redis_config):
+ def test_teams(
+ self,
+ gr_user: GRUser,
+ membership: Membership,
+ gr_db: PostgresConfig,
+ gr_redis_config: RedisConfig,
+ ):
assert gr_user.teams is None
@@ -41,15 +54,15 @@ class TestGRUser:
def test_prefetch_team_duplicates(
self,
- gr_user_token,
+ gr_user_token: GRToken,
gr_user: GRUser,
membership: Membership,
product_factory: Callable[..., Product],
- membership_factory,
+ membership_factory: Callable[..., Membership],
team: Team,
thl_web_rr: PostgresConfig,
- gr_redis_config,
- gr_db,
+ gr_redis_config: RedisConfig,
+ gr_db: PostgresConfig,
):
product_factory(team=team)
membership_factory(team=team, gr_user=gr_user)
@@ -67,9 +80,9 @@ class TestGRUser:
product_factory: Callable[..., Product],
team: Team,
membership: Membership,
- gr_db,
+ gr_db: PostgresConfig,
thl_web_rr: PostgresConfig,
- gr_redis_config,
+ gr_redis_config: RedisConfig,
):
from generalresearch.models.thl.product import Product
@@ -78,6 +91,8 @@ class TestGRUser:
# Create a new Team membership, and then create a Product that
# is part of that team
membership.prefetch_team(pg_config=gr_db, redis_config=gr_redis_config)
+ assert isinstance(membership.team, Team)
+
p: Product = product_factory(team=team)
assert p.id_int
assert team.uuid == membership.team.uuid
@@ -87,7 +102,7 @@ class TestGRUser:
gr_user.prefetch_products(
pg_config=gr_db,
- thl_pg_config=thl_web_rr: PostgresConfig,
+ thl_pg_config=thl_web_rr,
redis_config=gr_redis_config,
)
assert isinstance(gr_user.products, list)
@@ -97,7 +112,7 @@ class TestGRUser:
class TestGRUserMethods:
- def test_cache_key(self, gr_user, gr_redis):
+ def test_cache_key(self, gr_user: GRUser, gr_redis: RedisConfig):
assert isinstance(gr_user.cache_key, str)
assert ":" in gr_user.cache_key
assert str(gr_user.id) in gr_user.cache_key
@@ -105,11 +120,11 @@ class TestGRUserMethods:
def test_to_redis(
self,
gr_user: GRUser,
- gr_redis,
+ gr_redis: Redis,
team: Team,
business: Business,
product_factory: Callable[..., Product],
- membership_factory: Callable[Membership],
+ membership_factory: Callable[..., Membership],
):
product_factory(team=team, business=business)
membership_factory(team=team, gr_user=gr_user)
@@ -125,11 +140,11 @@ class TestGRUserMethods:
def test_set_cache(
self,
gr_user: GRUser,
- gr_user_token,
- gr_redis,
- gr_db,
+ gr_user_token: GRToken,
+ gr_redis: Redis,
+ gr_db: PostgresConfig,
thl_web_rr: PostgresConfig,
- gr_redis_config,
+ gr_redis_config: RedisConfig,
):
assert gr_redis.get(name=gr_user.cache_key) is None
assert gr_redis.get(name=f"{gr_user.cache_key}:team_uuids") is None
@@ -137,7 +152,7 @@ class TestGRUserMethods:
assert gr_redis.get(name=f"{gr_user.cache_key}:product_uuids") is None
gr_user.set_cache(
- pg_config=gr_db, thl_web_rr=thl_web_rr: PostgresConfig, redis_config=gr_redis_config
+ pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config
)
assert gr_redis.get(name=gr_user.cache_key) is not None
@@ -148,14 +163,14 @@ class TestGRUserMethods:
def test_set_cache_gr_user(
self,
gr_user: GRUser,
- gr_user_token,
- gr_redis,
- gr_redis_config,
- gr_db,
+ gr_user_token: GRToken,
+ gr_redis: RedisConfig,
+ gr_redis_config: RedisConfig,
+ gr_db: PostgresConfig,
thl_web_rr: PostgresConfig,
product_factory: Callable[..., Product],
- team,
- membership_factory,
+ team: Team,
+ membership_factory: Callable[..., Membership],
thl_redis_config: RedisConfig,
):
from generalresearch.models.gr.authentication import GRUser
@@ -164,7 +179,7 @@ class TestGRUserMethods:
membership_factory(team=team, gr_user=gr_user)
gr_user.set_cache(
- pg_config=gr_db, thl_web_rr=thl_web_rr: PostgresConfig, redis_config=gr_redis_config
+ pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config
)
res: str = gr_redis.get(name=gr_user.cache_key)
@@ -176,27 +191,27 @@ class TestGRUserMethods:
gru2.prefetch_products(
pg_config=gr_db,
- thl_pg_config=thl_web_rr: PostgresConfig,
- redis_config=thl_redis_config: RedisConfig,
+ thl_pg_config=thl_web_rr,
+ redis_config=thl_redis_config,
)
assert gru2.product_uuids == [p1.uuid]
def test_set_cache_team_uuids(
self,
- gr_user,
- membership,
- gr_user_token,
- gr_redis,
- gr_db,
+ gr_user: GRUser,
+ membership: Membership,
+ gr_user_token: GRToken,
+ gr_redis: Redis,
+ gr_db: PostgresConfig,
thl_web_rr: PostgresConfig,
product_factory: Callable[..., Product],
- team,
- gr_redis_config,
+ team: Team,
+ gr_redis_config: RedisConfig,
):
product_factory(team=team)
gr_user.set_cache(
- pg_config=gr_db, thl_web_rr=thl_web_rr: PostgresConfig, redis_config=gr_redis_config
+ pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config
)
res = json.loads(gr_redis.get(name=f"{gr_user.cache_key}:team_uuids"))
assert len(res) == 1
@@ -206,18 +221,18 @@ class TestGRUserMethods:
def test_set_cache_business_uuids(
self,
gr_user: GRUser,
- gr_redis,
- gr_db,
+ gr_redis: Redis,
+ gr_db: PostgresConfig,
thl_web_rr: PostgresConfig,
product_factory: Callable[..., Product],
business: Business,
- team,
- gr_redis_config,
+ team: Team,
+ gr_redis_config: RedisConfig,
):
product_factory(team=team, business=business)
gr_user.set_cache(
- pg_config=gr_db, thl_web_rr=thl_web_rr: PostgresConfig, redis_config=gr_redis_config
+ pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config
)
res = json.loads(gr_redis.get(name=f"{gr_user.cache_key}:business_uuids"))
assert len(res) == 1
@@ -225,20 +240,20 @@ class TestGRUserMethods:
def test_set_cache_product_uuids(
self,
- gr_user,
- membership,
- gr_user_token,
- gr_redis,
- gr_db,
+ gr_user: GRUser,
+ membership: Membership,
+ gr_user_token: GRToken,
+ gr_redis: Redis,
+ gr_db: PostgresConfig,
thl_web_rr: PostgresConfig,
product_factory: Callable[..., Product],
- team,
- gr_redis_config,
+ team: Team,
+ gr_redis_config: RedisConfig,
):
product_factory(team=team)
gr_user.set_cache(
- pg_config=gr_db, thl_web_rr=thl_web_rr: PostgresConfig, redis_config=gr_redis_config
+ pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config
)
res = json.loads(gr_redis.get(name=f"{gr_user.cache_key}:product_uuids"))
assert len(res) == 1
@@ -248,9 +263,7 @@ class TestGRUserMethods:
class TestGRToken:
@pytest.fixture
- def gr_token(self, gr_user):
- from generalresearch.models.gr.authentication import GRToken
-
+ def gr_token(self, gr_user: GRUser):
now = datetime.now(tz=UTC)
token = binascii.hexlify(os.urandom(20)).decode()
@@ -258,29 +271,26 @@ class TestGRToken:
return gr_token
- def test_init(self, gr_token):
- from generalresearch.models.gr.authentication import GRToken
-
+ def test_init(self, gr_token: GRToken):
assert isinstance(gr_token, GRToken)
assert gr_token.created
- def test_user(self, gr_token, gr_db, gr_redis_config):
- from generalresearch.models.gr.authentication import GRUser
-
+ def test_user(
+ self, gr_token: GRToken, gr_db: PostgresConfig, gr_redis_config: RedisConfig
+ ):
assert gr_token.user is None
gr_token.prefetch_user(pg_config=gr_db, redis_config=gr_redis_config)
assert isinstance(gr_token.user, GRUser)
- def test_auth_header(self, gr_token):
+ def test_auth_header(self, gr_token: GRToken):
assert isinstance(gr_token.auth_header, dict)
class TestClaims:
def test_init(self):
- from generalresearch.models.gr.authentication import Claims
d = {
"iss": SSO_ISSUER,
diff --git a/tests/models/gr/test_base.py b/tests/models/gr/test_base.py
index 8da28d3..412fa52 100644
--- a/tests/models/gr/test_base.py
+++ b/tests/models/gr/test_base.py
@@ -1,3 +1,5 @@
+from __future__ import annotations
+
import subprocess
from collections.abc import Callable
from pathlib import Path
diff --git a/tests/models/gr/test_business.py b/tests/models/gr/test_business.py
index 5239ac2..f0107de 100644
--- a/tests/models/gr/test_business.py
+++ b/tests/models/gr/test_business.py
@@ -19,6 +19,10 @@ from pytest import approx
from generalresearch.currency import USDCent
from generalresearch.incite.base import GRLDatasets
+from generalresearch.incite.collections.thl_web import (
+ SessionDFCollection,
+ WallDFCollection,
+)
from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge
from generalresearch.managers.gr.business import BusinessBankAccountManager
from generalresearch.managers.gr.team import TeamManager
@@ -29,7 +33,7 @@ from generalresearch.managers.thl.payout import (
PayoutEventManager,
)
from generalresearch.models.gr.business import (
- business: Business,
+ Business,
BusinessAddress,
BusinessBankAccount,
BusinessContact,
@@ -50,7 +54,7 @@ class TestBusinessBankAccount:
def test_init(
self,
- business: business: Business,
+ business: Business,
business_bank_account_manager: BusinessBankAccountManager,
):
from generalresearch.models.gr.business import (
@@ -68,7 +72,7 @@ class TestBusinessBankAccount:
def test_business(
self,
business_bank_account: BusinessBankAccount,
- business: business: Business,
+ business: Business,
gr_db: PostgresConfig,
gr_redis_config: RedisConfig,
):
@@ -79,7 +83,7 @@ class TestBusinessBankAccount:
business_bank_account.prefetch_business(
pg_config=gr_db, redis_config=gr_redis_config
)
- assert isinstance(business_bank_account.business: Business, Business)
+ assert isinstance(business_bank_account.business, Business)
assert business_bank_account.business.uuid == business.uuid
@@ -112,13 +116,13 @@ class TestBusiness:
def test_init(self, business: Business):
- assert isinstance(business: Business, Business)
+ assert isinstance(business, Business)
assert isinstance(business.id, int)
assert isinstance(business.uuid, str)
def test_str_and_repr(
self,
- business: business: Business,
+ business: Business,
product_factory: Callable[..., Product],
thl_web_rr: PostgresConfig,
ledger_manager: LedgerManager,
@@ -181,12 +185,12 @@ class TestBusiness:
pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection)
business.prebuild_payouts(
- thl_pg_config=thl_web_rr: PostgresConfig,
+ thl_pg_config=thl_web_rr,
thl_lm=thl_ledger_manager,
bpem=business_payout_event_manager,
)
business.prebuild_balance(
- thl_pg_config=thl_web_rr: PostgresConfig,
+ thl_pg_config=thl_web_rr,
lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
@@ -198,7 +202,7 @@ class TestBusiness:
def test_addresses(
self,
- business: business: Business,
+ business: Business,
business_address: BusinessAddress,
gr_db: PostgresConfig,
):
@@ -213,7 +217,7 @@ class TestBusiness:
def test_teams(
self,
- business: business: Business,
+ business: Business,
team: Team,
team_manager: TeamManager,
gr_db: PostgresConfig,
@@ -231,7 +235,7 @@ class TestBusiness:
def test_products(
self,
- business: business: Business,
+ business: Business,
product_factory: Callable[..., Product],
thl_web_rr: PostgresConfig,
):
@@ -254,7 +258,7 @@ class TestBusiness:
business.prefetch_products(thl_pg_config=thl_web_rr)
assert len(business.products) == 3
- def test_bank_accounts(self, business: business: Business, gr_db: PostgresConfig):
+ def test_bank_accounts(self, business: Business, gr_db: PostgresConfig):
assert business.products is None
# It's an empty list after prefetch
@@ -264,7 +268,7 @@ class TestBusiness:
def test_balance(
self,
- business: business: Business,
+ business: Business,
mnt_filepath: GRLDatasets,
client_no_amm: DaskClient,
thl_web_rr: PostgresConfig,
@@ -275,7 +279,7 @@ class TestBusiness:
with pytest.raises(expected_exception=AssertionError) as cm:
business.prebuild_balance(
- thl_pg_config=thl_web_rr: PostgresConfig,
+ thl_pg_config=thl_web_rr,
lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
@@ -289,7 +293,7 @@ class TestBusiness:
def test_payouts_no_accounts(
self,
- business: business: Business,
+ business: Business,
product_factory: Callable[..., Product],
thl_web_rr: PostgresConfig,
thl_ledger_manager: ThlLedgerManager,
@@ -299,7 +303,7 @@ class TestBusiness:
with pytest.raises(expected_exception=AssertionError) as cm:
business.prebuild_payouts(
- thl_pg_config=thl_web_rr: PostgresConfig,
+ thl_pg_config=thl_web_rr,
thl_lm=thl_ledger_manager,
bpem=business_payout_event_manager,
)
@@ -309,7 +313,7 @@ class TestBusiness:
thl_ledger_manager.get_account_or_create_bp_wallet(product=p)
business.prebuild_payouts(
- thl_pg_config=thl_web_rr: PostgresConfig,
+ thl_pg_config=thl_web_rr,
thl_lm=thl_ledger_manager,
bpem=business_payout_event_manager,
)
@@ -318,7 +322,7 @@ class TestBusiness:
def test_payouts(
self,
- business: business: Business,
+ business: Business,
product_factory: Callable[..., Product],
bp_payout_factory: Callable[..., BrokerageProductPayoutEvent],
thl_ledger_manager: ThlLedgerManager,
@@ -338,7 +342,7 @@ class TestBusiness:
)
business.prebuild_payouts(
- thl_pg_config=thl_web_rr: PostgresConfig,
+ thl_pg_config=thl_web_rr,
thl_lm=thl_ledger_manager,
bpem=business_payout_event_manager,
)
@@ -356,7 +360,7 @@ class TestBusiness:
thl_lm=thl_ledger_manager
)
business.prebuild_payouts(
- thl_pg_config=thl_web_rr: PostgresConfig,
+ thl_pg_config=thl_web_rr,
thl_lm=thl_ledger_manager,
bpem=business_payout_event_manager,
)
@@ -367,12 +371,12 @@ class TestBusiness:
def test_payouts_totals(
self,
- business: business: Business,
+ business: Business,
product_factory: Callable[..., Product],
bp_payout_factory: Callable[..., BrokerageProductPayoutEvent],
thl_ledger_manager: ThlLedgerManager,
thl_web_rr: PostgresConfig,
- business_payout_event_manager,
+ business_payout_event_manager: BusinessPayoutEventManager,
create_main_accounts: Callable[..., None],
):
@@ -406,7 +410,7 @@ class TestBusiness:
)
business.prebuild_payouts(
- thl_pg_config=thl_web_rr: PostgresConfig,
+ thl_pg_config=thl_web_rr,
thl_lm=thl_ledger_manager,
bpem=business_payout_event_manager,
)
@@ -419,7 +423,7 @@ class TestBusiness:
def test_pop_financial(
self,
- business: business: Business,
+ business: Business,
thl_web_rr: PostgresConfig,
thl_ledger_manager: ThlLedgerManager,
mnt_filepath: GRLDatasets,
@@ -428,7 +432,7 @@ class TestBusiness:
):
assert business.pop_financial is None
business.prebuild_pop_financial(
- thl_pg_config=thl_web_rr: PostgresConfig,
+ thl_pg_config=thl_web_rr,
thl_lm=thl_ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
@@ -438,7 +442,7 @@ class TestBusiness:
def test_bp_accounts(
self,
- business: business: Business,
+ business: Business,
thl_web_rr: PostgresConfig,
product_factory: Callable[..., Product],
thl_ledger_manager: ThlLedgerManager,
@@ -480,7 +484,7 @@ class TestBusinessBalance:
def test_single_product(
self,
- business: business: Business,
+ business: Business,
product_factory: Callable[..., Product],
user_factory: Callable[..., User],
mnt_filepath,
@@ -519,7 +523,7 @@ class TestBusinessBalance:
pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection)
business.prebuild_balance(
- thl_pg_config=thl_web_rr: PostgresConfig,
+ thl_pg_config=thl_web_rr,
lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
@@ -541,7 +545,7 @@ class TestBusinessBalance:
def test_multi_product(
self,
- business: business: Business,
+ business: Business,
product_factory: Callable[..., Product],
user_factory: Callable[..., User],
mnt_filepath: GRLDatasets,
@@ -579,7 +583,7 @@ class TestBusinessBalance:
pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection)
business.prebuild_balance(
- thl_pg_config=thl_web_rr: PostgresConfig,
+ thl_pg_config=thl_web_rr,
lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
@@ -625,7 +629,7 @@ class TestBusinessBalance:
def test_multi_product_multi_payout(
self,
- business: business: Business,
+ business: Business,
product_factory: Callable[..., Product],
user_factory: Callable[..., User],
mnt_filepath: GRLDatasets,
@@ -665,7 +669,7 @@ class TestBusinessBalance:
payout_event_manager.set_account_lookup_table(thl_lm=thl_ledger_manager)
bp_payout_factory(
- product=u1.product: Product,
+ product=u1.product,
amount=USDCent(5),
created=start + timedelta(days=4),
skip_wallet_balance_check=True,
@@ -673,7 +677,7 @@ class TestBusinessBalance:
)
bp_payout_factory(
- product=u2.product: Product,
+ product=u2.product,
amount=USDCent(50),
created=start + timedelta(days=4),
skip_wallet_balance_check=True,
@@ -684,7 +688,7 @@ class TestBusinessBalance:
pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection)
business.prebuild_balance(
- thl_pg_config=thl_web_rr: PostgresConfig,
+ thl_pg_config=thl_web_rr,
lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
@@ -699,7 +703,7 @@ class TestBusinessBalance:
def test_multi_product_multi_payout_adjustment(
self,
- business: business: Business,
+ business: Business,
product_factory: Callable[..., Product],
user_factory: Callable[..., User],
mnt_filepath: GRLDatasets,
@@ -758,7 +762,7 @@ class TestBusinessBalance:
payout_event_manager.set_account_lookup_table(thl_lm=thl_ledger_manager)
bp_payout_factory(
- product=u1.product: Product,
+ product=u1.product,
amount=USDCent(250),
created=start + timedelta(days=3),
skip_wallet_balance_check=True,
@@ -766,7 +770,7 @@ class TestBusinessBalance:
)
bp_payout_factory(
- product=u2.product: Product,
+ product=u2.product,
amount=USDCent(50),
created=start + timedelta(days=4),
skip_wallet_balance_check=True,
@@ -796,7 +800,7 @@ class TestBusinessBalance:
assert df.shape == (20, 28)
business.prebuild_balance(
- thl_pg_config=thl_web_rr: PostgresConfig,
+ thl_pg_config=thl_web_rr,
lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
@@ -833,7 +837,7 @@ class TestBusinessBalance:
create_main_accounts: Callable[..., None],
delete_df_collection: Callable[..., None],
ledger_collection,
- business: business: Business,
+ business: Business,
user_factory: Callable[..., User],
product_factory: Callable[..., Product],
session_with_tx_factory: Callable[..., Session],
@@ -869,7 +873,7 @@ class TestBusinessBalance:
)
payout_event_manager.set_account_lookup_table(thl_lm=thl_ledger_manager)
bp_payout_factory(
- product=u1.product: Product,
+ product=u1.product,
amount=USDCent(71),
ext_ref_id=uuid4().hex,
created=start + timedelta(days=1, minutes=1),
@@ -898,7 +902,7 @@ class TestBusinessBalance:
pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection)
business.prebuild_balance(
- thl_pg_config=thl_web_rr: PostgresConfig,
+ thl_pg_config=thl_web_rr,
lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
@@ -946,7 +950,7 @@ class TestBusinessBalance:
def test_multi_product_multi_payout_adjustment_at_timestamp(
self,
- business: business: Business,
+ business: Business,
product_factory: Callable[..., Product],
user_factory: Callable[..., User],
mnt_filepath: GRLDatasets,
@@ -1022,7 +1026,7 @@ class TestBusinessBalance:
payout_event_manager.set_account_lookup_table(thl_lm=thl_ledger_manager)
bp_payout_factory(
- product=u1.product: Product,
+ product=u1.product,
amount=USDCent(250),
created=start + timedelta(days=3),
skip_wallet_balance_check=True,
@@ -1030,7 +1034,7 @@ class TestBusinessBalance:
)
bp_payout_factory(
- product=u2.product: Product,
+ product=u2.product,
amount=USDCent(50),
created=start + timedelta(days=4),
skip_wallet_balance_check=True,
@@ -1060,7 +1064,7 @@ class TestBusinessBalance:
assert df.shape == (20, 28)
business.prebuild_balance(
- thl_pg_config=thl_web_rr: PostgresConfig,
+ thl_pg_config=thl_web_rr,
lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
@@ -1068,7 +1072,7 @@ class TestBusinessBalance:
)
business.prebuild_balance(
- thl_pg_config=thl_web_rr: PostgresConfig,
+ thl_pg_config=thl_web_rr,
lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
@@ -1078,7 +1082,7 @@ class TestBusinessBalance:
day1_bal = business.balance
business.prebuild_balance(
- thl_pg_config=thl_web_rr: PostgresConfig,
+ thl_pg_config=thl_web_rr,
lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
@@ -1088,7 +1092,7 @@ class TestBusinessBalance:
day2_bal = business.balance
business.prebuild_balance(
- thl_pg_config=thl_web_rr: PostgresConfig,
+ thl_pg_config=thl_web_rr,
lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
@@ -1098,7 +1102,7 @@ class TestBusinessBalance:
day3_bal = business.balance
business.prebuild_balance(
- thl_pg_config=thl_web_rr: PostgresConfig,
+ thl_pg_config=thl_web_rr,
lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
@@ -1108,7 +1112,7 @@ class TestBusinessBalance:
day4_bal = business.balance
business.prebuild_balance(
- thl_pg_config=thl_web_rr: PostgresConfig,
+ thl_pg_config=thl_web_rr,
lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
@@ -1118,7 +1122,7 @@ class TestBusinessBalance:
day5_bal = business.balance
business.prebuild_balance(
- thl_pg_config=thl_web_rr: PostgresConfig,
+ thl_pg_config=thl_web_rr,
lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
@@ -1183,7 +1187,7 @@ class TestBusinessMethods:
def test_set_cache(
self,
- business: business: Business,
+ business: Business,
gr_redis: RedisConfig,
gr_db: PostgresConfig,
thl_web_rr: PostgresConfig,
@@ -1219,7 +1223,7 @@ class TestBusinessMethods:
business.set_cache(
pg_config=gr_db,
- thl_web_rr=thl_web_rr: PostgresConfig,
+ thl_web_rr=thl_web_rr,
redis_config=gr_redis_config,
client=client_no_amm,
ds=mnt_filepath,
@@ -1245,7 +1249,7 @@ class TestBusinessMethods:
def test_set_cache_business(
self,
- business: business: Business,
+ business: Business,
gr_db: PostgresConfig,
thl_web_rr: PostgresConfig,
product_factory: Callable[..., Product],
@@ -1282,7 +1286,7 @@ class TestBusinessMethods:
business.set_cache(
pg_config=gr_db,
- thl_web_rr=thl_web_rr: PostgresConfig,
+ thl_web_rr=thl_web_rr,
redis_config=gr_redis_config,
client=client_no_amm,
ds=mnt_filepath,
@@ -1345,15 +1349,15 @@ class TestBusinessMethods:
self,
enriched_session_merge,
client_no_amm: DaskClient,
- wall_collection,
- session_collection,
+ wall_collection: WallDFCollection,
+ session_collection: SessionDFCollection,
thl_web_rr: PostgresConfig,
user_factory: Callable[..., User],
start: datetime,
session_factory: Callable[..., Session],
product_factory: Callable[..., Product],
delete_df_collection: Callable[..., None],
- business: business: Business,
+ business: Business,
mnt_filepath: GRLDatasets,
mnt_gr_api_dir: Path,
):
@@ -1380,11 +1384,11 @@ class TestBusinessMethods:
client=client_no_amm,
session_coll=session_collection,
wall_coll=wall_collection,
- pg_config=thl_web_rr: PostgresConfig,
+ pg_config=thl_web_rr,
)
business.prebuild_enriched_session_parquet(
- thl_pg_config=thl_web_rr: PostgresConfig,
+ thl_pg_config=thl_web_rr,
ds=mnt_filepath,
client=client_no_amm,
mnt_gr_api=mnt_gr_api_dir,
@@ -1401,15 +1405,15 @@ class TestBusinessMethods:
self,
enriched_wall_merge,
client_no_amm: DaskClient,
- wall_collection,
- session_collection,
+ wall_collection: WallDFCollection,
+ session_collection: SessionDFCollection,
thl_web_rr: PostgresConfig,
user_factory: Callable[..., User],
start: datetime,
session_factory: Callable[..., Session],
product_factory: Callable[..., Product],
delete_df_collection: Callable[..., None],
- business: business: Business,
+ business: Business,
mnt_filepath: GRLDatasets,
mnt_gr_api_dir: Path,
):
@@ -1436,11 +1440,11 @@ class TestBusinessMethods:
client=client_no_amm,
session_coll=session_collection,
wall_coll=wall_collection,
- pg_config=thl_web_rr: PostgresConfig,
+ pg_config=thl_web_rr,
)
business.prebuild_enriched_wall_parquet(
- thl_pg_config=thl_web_rr: PostgresConfig,
+ thl_pg_config=thl_web_rr,
ds=mnt_filepath,
client=client_no_amm,
mnt_gr_api=mnt_gr_api_dir,
diff --git a/tests/models/gr/test_team.py b/tests/models/gr/test_team.py
index 26300b9..dc7d4b9 100644
--- a/tests/models/gr/test_team.py
+++ b/tests/models/gr/test_team.py
@@ -97,7 +97,7 @@ class TestTeam:
def test_businesses(
self,
team: Team,
- business: business: Business,
+ business: Business,
team_manager: TeamManager,
gr_db: PostgresConfig,
gr_redis_config: RedisConfig,
@@ -160,7 +160,7 @@ class TestTeamMethods:
team.set_cache(
pg_config=gr_db,
- thl_web_rr=thl_web_rr: PostgresConfig,
+ thl_web_rr=thl_web_rr,
redis_config=gr_redis_config,
client=client_no_amm,
ds=mnt_filepath,
@@ -192,7 +192,7 @@ class TestTeamMethods:
team.set_cache(
pg_config=gr_db,
- thl_web_rr=thl_web_rr: PostgresConfig,
+ thl_web_rr=thl_web_rr,
redis_config=gr_redis_config,
client=client_no_amm,
ds=mnt_filepath,
@@ -254,11 +254,11 @@ class TestTeamMethods:
client=client_no_amm,
session_coll=session_collection,
wall_coll=wall_collection,
- pg_config=thl_web_rr: PostgresConfig,
+ pg_config=thl_web_rr,
)
team.prebuild_enriched_session_parquet(
- thl_pg_config=thl_web_rr: PostgresConfig,
+ thl_pg_config=thl_web_rr,
ds=mnt_filepath,
client=client_no_amm,
mnt_gr_api=mnt_gr_api_dir,
@@ -310,11 +310,11 @@ class TestTeamMethods:
client=client_no_amm,
session_coll=session_collection,
wall_coll=wall_collection,
- pg_config=thl_web_rr: PostgresConfig,
+ pg_config=thl_web_rr,
)
team.prebuild_enriched_wall_parquet(
- thl_pg_config=thl_web_rr: PostgresConfig,
+ thl_pg_config=thl_web_rr,
ds=mnt_filepath,
client=client_no_amm,
mnt_gr_api=mnt_gr_api_dir,