diff options
42 files changed, 1194 insertions, 691 deletions
diff --git a/generalresearch/models/thl/user.py b/generalresearch/models/thl/user.py index 11e0d67..d3ffb0d 100644 --- a/generalresearch/models/thl/user.py +++ b/generalresearch/models/thl/user.py @@ -4,7 +4,7 @@ import json import logging import re from datetime import UTC, datetime -from typing import TYPE_CHECKING, Annotated, Self +from typing import TYPE_CHECKING, Annotated, Any, Self from uuid import UUID, uuid4 from pydantic import ( @@ -104,7 +104,7 @@ class User(BaseModel): # --- Validation --- @field_validator("product_user_id") - def check_product_user_id(cls, v: str) -> str: + def check_product_user_id(cls, v: str | None) -> str: if v is not None: if " " in v: raise ValueError("String cannot contain spaces") @@ -122,23 +122,27 @@ class User(BaseModel): # noinspection PyNestedDecoratorsk @field_validator("created", "last_seen") @classmethod - def check_not_in_future(cls, v: AwareDatetime) -> AwareDatetime: + def check_not_in_future(cls, v: AwareDatetime | None) -> AwareDatetime: if v is not None: try: assert v < datetime.now(tz=UTC) - except Exception: + except AssertionError: raise ValueError("Input is in the future") + + assert isinstance(v, AwareDatetime) return v # noinspection PyNestedDecorators @field_validator("created", "last_seen") @classmethod - def check_after_anno_domini(cls, v: AwareDatetime) -> AwareDatetime: + def check_after_anno_domini(cls, v: AwareDatetime | None) -> AwareDatetime: if v is not None: try: assert v > datetime(year=2016, month=7, day=13, tzinfo=UTC) - except Exception: + except AssertionError: raise ValueError("Input is before Anno Domini") + + assert isinstance(v, AwareDatetime) return v @model_validator(mode="after") @@ -168,7 +172,7 @@ class User(BaseModel): ) @classmethod - def is_valid_ubp(cls, *, product_id, product_user_id) -> bool: + def is_valid_ubp(cls, *, product_id: str, product_user_id: str) -> bool: # Attempt to create common_struct solely for validation purposes, # using the product_id and product_user_id try: @@ -178,7 +182,7 @@ class User(BaseModel): product_id=product_id, product_user_id=product_user_id, ) - except Exception as e: + except ValueError as e: logger.info(e) return False else: @@ -186,7 +190,7 @@ class User(BaseModel): # --- Methods --- @staticmethod - def check_bpuid_is_not_bpid(product_id, product_user_id): + def check_bpuid_is_not_bpid(product_id: str | None, product_user_id: str | None): """Unfortunately users were already created failing this constraint, so only check for new users! """ @@ -198,7 +202,7 @@ class User(BaseModel): raise ValueError("product_user_id must not equal the product_id") return True - def to_dict(self) -> dict: + def to_dict(self) -> dict[str, Any]: return self.model_dump(mode="python", exclude={"product"}) def to_json(self) -> str: @@ -291,7 +295,7 @@ class User(BaseModel): # --- Prebuild --- @classmethod - def from_db(cls, res) -> Self: + def from_db(cls, res: dict[str, Any]) -> Self: if res["created"]: res["created"] = res["created"].replace(tzinfo=UTC) if res["last_seen"]: diff --git a/test_utils/managers/conftest.py b/test_utils/managers/conftest.py index 2a9ea00..e5d6015 100644 --- a/test_utils/managers/conftest.py +++ b/test_utils/managers/conftest.py @@ -15,11 +15,17 @@ from generalresearch.managers.gr.team import ( ) from generalresearch.managers.spectrum.survey import SpectrumSurveyManager from generalresearch.managers.thl.buyer import BuyerManager +from generalresearch.managers.thl.cashout_method import ( + CashoutMethodManager, +) from generalresearch.managers.thl.ipinfo import ( GeoIpInfoManager, IPGeonameManager, IPInformationManager, ) +from generalresearch.managers.thl.user_streak import ( + UserStreakManager, +) from generalresearch.managers.thl.userhealth import ( AuditLogManager, IPRecordManager, @@ -114,12 +120,9 @@ def geoipinfo_manager( @pytest.fixture(scope="session") -def cashout_method_manager(thl_web_rw: PostgresConfig): +def cashout_method_manager(thl_web_rw: PostgresConfig) -> CashoutMethodManager: assert thl_web_rw.dsn.path assert "/unittest-" in thl_web_rw.dsn.path - from generalresearch.managers.thl.cashout_method import ( - CashoutMethodManager, - ) return CashoutMethodManager(pg_config=thl_web_rw) @@ -132,12 +135,9 @@ def event_manager(thl_redis_config: RedisConfig): @pytest.fixture(scope="session") -def user_streak_manager(thl_web_rw: PostgresConfig): +def user_streak_manager(thl_web_rw: PostgresConfig) -> UserStreakManager: assert thl_web_rw.dsn.path assert "/unittest-" in thl_web_rw.dsn.path - from generalresearch.managers.thl.user_streak import ( - UserStreakManager, - ) return UserStreakManager(pg_config=thl_web_rw) @@ -169,10 +169,16 @@ def delete_cashoutmethod_db(thl_web_rw: PostgresConfig) -> Callable[..., None]: @pytest.fixture(scope="session") -def setup_cashoutmethod_db(cashout_method_manager, delete_cashoutmethod_db): - delete_cashoutmethod_db() - for x in EXAMPLE_TANGO_CASHOUT_METHODS: - cashout_method_manager.create(x) +def setup_cashoutmethod_db( + cashout_method_manager: CashoutMethodManager, + delete_cashoutmethod_db: Callable[..., None], +) -> Callable[..., None]: + + def _inner(): + delete_cashoutmethod_db() + + for x in EXAMPLE_TANGO_CASHOUT_METHODS: + cashout_method_manager.create(x) # TODO: convert these ids into instances to use. # settings.amt_bonus_cashout_method_id @@ -180,7 +186,9 @@ def setup_cashoutmethod_db(cashout_method_manager, delete_cashoutmethod_db): # cashout_method_manager.create(AMT_ASSIGNMENT_CASHOUT_METHOD) # cashout_method_manager.create(AMT_BONUS_CASHOUT_METHOD) - raise NotImplementedError("Need to implement setup_cashoutmethod_db") + # raise NotImplementedError("Need to implement setup_cashoutmethod_db") + + return _inner # === THL: Marketplaces === @@ -259,33 +267,39 @@ def membership_manager(gr_db: PostgresConfig) -> MembershipManager: @pytest.fixture(scope="session") -def delete_buyers_surveys(thl_web_rw: PostgresConfig, buyer_manager: BuyerManager): - # assert "/unittest-" in thl_web_rw.dsn.path - thl_web_rw.execute_write( - """ - DELETE FROM marketplace_surveystat - WHERE survey_id IN ( - SELECT id - FROM marketplace_survey - WHERE source = %(source)s - );""", - params={"source": Source.TESTING.value}, - ) - thl_web_rw.execute_write( - """ - DELETE FROM marketplace_survey - WHERE buyer_id IN ( - SELECT id - FROM marketplace_buyer - WHERE source = %(source)s - );""", - params={"source": Source.TESTING.value}, - ) - thl_web_rw.execute_write( - """ - DELETE from marketplace_buyer - WHERE source=%(source)s; - """, - params={"source": Source.TESTING.value}, - ) - buyer_manager.populate_caches() +def delete_buyers_surveys( + thl_web_rw: PostgresConfig, buyer_manager: BuyerManager +) -> Callable[..., None]: + + def _inner(): + # assert "/unittest-" in thl_web_rw.dsn.path + thl_web_rw.execute_write( + """ + DELETE FROM marketplace_surveystat + WHERE survey_id IN ( + SELECT id + FROM marketplace_survey + WHERE source = %(source)s + );""", + params={"source": Source.TESTING.value}, + ) + thl_web_rw.execute_write( + """ + DELETE FROM marketplace_survey + WHERE buyer_id IN ( + SELECT id + FROM marketplace_buyer + WHERE source = %(source)s + );""", + params={"source": Source.TESTING.value}, + ) + thl_web_rw.execute_write( + """ + DELETE from marketplace_buyer + WHERE source=%(source)s; + """, + params={"source": Source.TESTING.value}, + ) + buyer_manager.populate_caches() + + return _inner diff --git a/test_utils/managers/thl/conftest.py b/test_utils/managers/thl/conftest.py index 21b2007..d40b7d2 100644 --- a/test_utils/managers/thl/conftest.py +++ b/test_utils/managers/thl/conftest.py @@ -20,6 +20,12 @@ from generalresearch.managers.thl.session import SessionManager from generalresearch.managers.thl.task_adjustment import ( TaskAdjustmentManager, ) +from generalresearch.managers.thl.user_manager.mysql_user_manager import ( + MysqlUserManager, +) +from generalresearch.managers.thl.user_manager.redis_user_manager import ( + RedisUserManager, +) from generalresearch.managers.thl.user_manager.user_manager import ( UserManager, ) @@ -161,6 +167,16 @@ def user_manager( @pytest.fixture(scope="session") +def mysql_user_manager(thl_web_rw: PostgresConfig) -> MysqlUserManager: + return MysqlUserManager(pg_config=thl_web_rw, is_read_replica=False) + + +@pytest.fixture(scope="session") +def redis_user_manager(thl_redis_config: RedisConfig) -> RedisUserManager: + return RedisUserManager(redis_dsn=thl_redis_config) + + +@pytest.fixture(scope="session") def user_metadata_manager(thl_web_rw: PostgresConfig) -> UserMetadataManager: assert thl_web_rw.dsn assert thl_web_rw.dsn.path diff --git a/test_utils/models/upk/conftest.py b/test_utils/models/upk/conftest.py index c8855da..ef77dd6 100644 --- a/test_utils/models/upk/conftest.py +++ b/test_utils/models/upk/conftest.py @@ -2,6 +2,7 @@ from __future__ import annotations import os import time +from collections.abc import Callable from typing import TYPE_CHECKING from uuid import UUID @@ -169,9 +170,13 @@ def upk_data( propertymarketplaceassociation_data, propertyitemrange_data, question_data, -) -> None: - # Wait a second to make sure the HarmonizerCache refresh loop pulls these in - time.sleep(2) +) -> Callable[..., None]: + + def _inner(): + # Wait a second to make sure the HarmonizerCache refresh loop pulls these in + time.sleep(2) + + return _inner def test_fixtures(upk_data): 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, |
