From 6cf7ccbaa8306700e64ada19d6f99807743b2865 Mon Sep 17 00:00:00 2001 From: Max Nanis Date: Fri, 21 Aug 2026 17:19:41 -0700 Subject: Ruff auto updates to 3.14 --- .../incite/collections/test_df_collection_base.py | 24 +++++++++++----------- .../collections/test_df_collection_item_base.py | 20 +++++++++--------- .../collections/test_df_collection_item_thl_web.py | 10 ++++----- .../test_df_collection_thl_marketplaces.py | 6 +++--- .../collections/test_df_collection_thl_web.py | 2 +- 5 files changed, 31 insertions(+), 31 deletions(-) (limited to 'tests/incite/collections') diff --git a/tests/incite/collections/test_df_collection_base.py b/tests/incite/collections/test_df_collection_base.py index 31d1720..c3c64e4 100644 --- a/tests/incite/collections/test_df_collection_base.py +++ b/tests/incite/collections/test_df_collection_base.py @@ -1,4 +1,4 @@ -from datetime import datetime, timezone +from datetime import UTC, datetime, timezone from typing import TYPE_CHECKING import pandas as pd @@ -24,7 +24,7 @@ class TestDFCollectionBase: """ - def test_init(self, mnt_filepath: "GRLDatasets", df_coll_type: DFCollectionType): + def test_init(self, mnt_filepath: GRLDatasets, df_coll_type: DFCollectionType): """Try to initialize the DFCollection with various invalid parameters""" with pytest.raises(expected_exception=ValueError) as cm: DFCollection(archive_path=mnt_filepath.data_src) @@ -46,11 +46,11 @@ class TestDFCollectionBase: class TestDFCollectionBaseProperties: @pytest.mark.skip - def test_df_collection_items(self, mnt_filepath: "GRLDatasets", df_coll_type): + def test_df_collection_items(self, mnt_filepath: GRLDatasets, df_coll_type): instance = DFCollection( data_type=df_coll_type, - start=datetime(year=1800, month=1, day=1, tzinfo=timezone.utc), - finished=datetime(year=1900, month=1, day=1, tzinfo=timezone.utc), + start=datetime(year=1800, month=1, day=1, tzinfo=UTC), + finished=datetime(year=1900, month=1, day=1, tzinfo=UTC), offset="100d", archive_path=mnt_filepath.archive_path(enum_type=df_coll_type), ) @@ -58,11 +58,11 @@ class TestDFCollectionBaseProperties: assert len(instance.interval_range) == len(instance.items) assert len(instance.items) == 366 - def test_df_collection_progress(self, mnt_filepath: "GRLDatasets", df_coll_type): + def test_df_collection_progress(self, mnt_filepath: GRLDatasets, df_coll_type): instance = DFCollection( data_type=df_coll_type, - start=datetime(year=1800, month=1, day=1, tzinfo=timezone.utc), - finished=datetime(year=1900, month=1, day=1, tzinfo=timezone.utc), + start=datetime(year=1800, month=1, day=1, tzinfo=UTC), + finished=datetime(year=1900, month=1, day=1, tzinfo=UTC), offset="100d", archive_path=mnt_filepath.archive_path(enum_type=df_coll_type), ) @@ -71,7 +71,7 @@ class TestDFCollectionBaseProperties: assert isinstance(instance.progress, pd.DataFrame) assert instance.progress.shape == (366, 6) - def test_df_collection_schema(self, mnt_filepath: "GRLDatasets", df_coll_type): + def test_df_collection_schema(self, mnt_filepath: GRLDatasets, df_coll_type): instance1 = DFCollection( data_type=DFCollectionType.WALL, archive_path=mnt_filepath.data_src ) @@ -88,12 +88,12 @@ class TestDFCollectionBaseProperties: class TestDFCollectionBaseMethods: @pytest.mark.skip - def test_initial_load(self, mnt_filepath: "GRLDatasets", thl_web_rr): + def test_initial_load(self, mnt_filepath: GRLDatasets, thl_web_rr): instance = DFCollection( pg_config=thl_web_rr, data_type=DFCollectionType.USER, - start=datetime(year=2022, month=1, day=1, minute=0, tzinfo=timezone.utc), - finished=datetime(year=2022, month=1, day=1, minute=5, tzinfo=timezone.utc), + start=datetime(year=2022, month=1, day=1, minute=0, tzinfo=UTC), + finished=datetime(year=2022, month=1, day=1, minute=5, tzinfo=UTC), offset="2min", archive_path=mnt_filepath.data_src, ) diff --git a/tests/incite/collections/test_df_collection_item_base.py b/tests/incite/collections/test_df_collection_item_base.py index 136d234..8cf719d 100644 --- a/tests/incite/collections/test_df_collection_item_base.py +++ b/tests/incite/collections/test_df_collection_item_base.py @@ -1,4 +1,4 @@ -from datetime import datetime, timezone +from datetime import UTC, datetime, timezone from typing import TYPE_CHECKING import pytest @@ -19,12 +19,12 @@ df_collection_types = [e for e in DFCollectionType if e is not DFCollectionType. @pytest.mark.parametrize("df_coll_type", df_collection_types) class TestDFCollectionItemBase: - def test_init(self, mnt_filepath: "GRLDatasets", df_coll_type): + def test_init(self, mnt_filepath: GRLDatasets, df_coll_type): collection = DFCollection( data_type=df_coll_type, offset="100d", - start=datetime(year=1800, month=1, day=1, tzinfo=timezone.utc), - finished=datetime(year=1900, month=1, day=1, tzinfo=timezone.utc), + start=datetime(year=1800, month=1, day=1, tzinfo=UTC), + finished=datetime(year=1900, month=1, day=1, tzinfo=UTC), archive_path=mnt_filepath.archive_path(enum_type=df_coll_type), ) @@ -45,12 +45,12 @@ class TestDFCollectionItemProperties: @pytest.mark.parametrize("df_coll_type", df_collection_types) class TestDFCollectionItemMethods: - def test_has_mysql_false(self, mnt_filepath: "GRLDatasets", df_coll_type): + def test_has_mysql_false(self, mnt_filepath: GRLDatasets, df_coll_type): collection = DFCollection( data_type=df_coll_type, offset="100d", - start=datetime(year=1800, month=1, day=1, tzinfo=timezone.utc), - finished=datetime(year=1900, month=1, day=1, tzinfo=timezone.utc), + start=datetime(year=1800, month=1, day=1, tzinfo=UTC), + finished=datetime(year=1900, month=1, day=1, tzinfo=UTC), archive_path=mnt_filepath.archive_path(enum_type=df_coll_type), ) @@ -58,13 +58,13 @@ class TestDFCollectionItemMethods: assert not instance1.has_mysql() def test_has_mysql_true( - self, thl_web_rr: PostgresConfig, mnt_filepath: "GRLDatasets", df_coll_type + self, thl_web_rr: PostgresConfig, mnt_filepath: GRLDatasets, df_coll_type ): collection = DFCollection( data_type=df_coll_type, offset="100d", - start=datetime(year=1800, month=1, day=1, tzinfo=timezone.utc), - finished=datetime(year=1900, month=1, day=1, tzinfo=timezone.utc), + start=datetime(year=1800, month=1, day=1, tzinfo=UTC), + finished=datetime(year=1900, month=1, day=1, tzinfo=UTC), archive_path=mnt_filepath.archive_path(enum_type=df_coll_type), pg_config=thl_web_rr, ) diff --git a/tests/incite/collections/test_df_collection_item_thl_web.py b/tests/incite/collections/test_df_collection_item_thl_web.py index 8b8bcbe..062171d 100644 --- a/tests/incite/collections/test_df_collection_item_thl_web.py +++ b/tests/incite/collections/test_df_collection_item_thl_web.py @@ -1,11 +1,11 @@ from __future__ import annotations -from collections.abc import Generator -from datetime import datetime, timedelta, timezone +from collections.abc import Callable, Generator +from datetime import UTC, datetime, timedelta, timezone from itertools import product as iter_product from os.path import join as pjoin from pathlib import Path, PurePath -from typing import TYPE_CHECKING, Callable +from typing import TYPE_CHECKING from uuid import uuid4 import dask.dataframe as dd @@ -52,7 +52,7 @@ unsupported_mock_types = { } -def combo_object() -> Generator[str, None, None]: +def combo_object() -> Generator[str]: for x in iter_product( df_collections, ["15min", "45min", "1H"], @@ -632,7 +632,7 @@ class TestDFCollectionItemMethodBase: for item in df_collection.items: item: DFCollectionItem - if datetime.now(tz=timezone.utc) > item.finish + aa: + if datetime.now(tz=UTC) > item.finish + aa: assert item.should_archive() else: assert not item.should_archive() diff --git a/tests/incite/collections/test_df_collection_thl_marketplaces.py b/tests/incite/collections/test_df_collection_thl_marketplaces.py index 981f62e..2597d38 100644 --- a/tests/incite/collections/test_df_collection_thl_marketplaces.py +++ b/tests/incite/collections/test_df_collection_thl_marketplaces.py @@ -1,4 +1,4 @@ -from datetime import datetime, timezone +from datetime import UTC, datetime, timezone from itertools import product from typing import TYPE_CHECKING @@ -57,8 +57,8 @@ class TestDFCollection_thl_marketplaces: archive_path=mnt_filepath.archive_path(enum_type=data_type), sql_helper=spectrum_rw, offset=offset, - start=datetime(year=2023, month=6, day=1, minute=0, tzinfo=timezone.utc), - finished=datetime(year=2023, month=6, day=1, minute=5, tzinfo=timezone.utc), + start=datetime(year=2023, month=6, day=1, minute=0, tzinfo=UTC), + finished=datetime(year=2023, month=6, day=1, minute=5, tzinfo=UTC), ) assert isinstance(instance, DFCollection) diff --git a/tests/incite/collections/test_df_collection_thl_web.py b/tests/incite/collections/test_df_collection_thl_web.py index b09d44c..2cb0ba0 100644 --- a/tests/incite/collections/test_df_collection_thl_web.py +++ b/tests/incite/collections/test_df_collection_thl_web.py @@ -20,7 +20,7 @@ if TYPE_CHECKING: ) -def combo_object() -> Generator[tuple, None, None]: +def combo_object() -> Generator[tuple]: for x in product( [ DFCollectionType.USER, -- cgit v1.2.3 From e2c5de703be45746bacaea4136f24440ff5a291c Mon Sep 17 00:00:00 2001 From: Max Nanis Date: Mon, 24 Aug 2026 12:35:35 -0700 Subject: Ruff std replacements --- generalresearch/__init__.py | 2 +- generalresearch/config.py | 2 +- generalresearch/grliq/managers/event_plotter.py | 4 +- generalresearch/grliq/managers/forensic_data.py | 4 +- generalresearch/grliq/managers/forensic_events.py | 76 ++++++++++------------ generalresearch/grliq/managers/forensic_results.py | 7 +- generalresearch/grliq/models/custom_types.py | 3 +- generalresearch/grliq/models/forensic_data.py | 1 - generalresearch/grliq/models/forensic_summary.py | 3 +- generalresearch/grliq/utils.py | 2 +- generalresearch/grpc.py | 2 +- generalresearch/incite/base.py | 10 +-- generalresearch/incite/collections/__init__.py | 18 ++--- generalresearch/incite/defaults.py | 2 +- .../incite/mergers/foundations/__init__.py | 2 +- .../incite/mergers/foundations/enriched_session.py | 10 +-- .../mergers/foundations/enriched_task_adjust.py | 5 +- .../incite/mergers/foundations/enriched_wall.py | 8 +-- generalresearch/incite/mergers/ym_survey_wall.py | 10 ++- generalresearch/incite/mergers/ym_wall_summary.py | 3 +- generalresearch/incite/schemas/thl_web.py | 2 +- generalresearch/locales/setup_json.py | 1 + generalresearch/locales/timezone.py | 1 - generalresearch/managers/cint/survey.py | 2 +- generalresearch/managers/criteria.py | 4 +- generalresearch/managers/dynata/survey.py | 2 +- generalresearch/managers/events.py | 4 +- generalresearch/managers/gr/authentication.py | 5 +- generalresearch/managers/gr/team.py | 2 +- generalresearch/managers/innovate/survey.py | 2 +- generalresearch/managers/leaderboard/__init__.py | 3 +- generalresearch/managers/leaderboard/manager.py | 2 +- generalresearch/managers/morning/survey.py | 2 +- generalresearch/managers/network/label.py | 9 ++- generalresearch/managers/network/tool_run.py | 22 +++---- generalresearch/managers/precision/survey.py | 10 +-- generalresearch/managers/prodege/survey.py | 2 +- generalresearch/managers/repdata/survey.py | 4 +- generalresearch/managers/sago/survey.py | 2 +- generalresearch/managers/spectrum/survey.py | 2 +- generalresearch/managers/survey.py | 2 - generalresearch/managers/thl/buyer.py | 2 +- generalresearch/managers/thl/cashout_method.py | 2 +- generalresearch/managers/thl/contest_manager.py | 3 +- .../managers/thl/ledger_manager/conditions.py | 4 +- .../managers/thl/ledger_manager/exceptions.py | 5 -- .../managers/thl/ledger_manager/ledger.py | 15 ++--- .../managers/thl/ledger_manager/thl_ledger.py | 22 +++---- generalresearch/managers/thl/product.py | 10 +-- generalresearch/managers/thl/profiling/question.py | 2 +- generalresearch/managers/thl/profiling/uqa.py | 2 +- generalresearch/managers/thl/profiling/user_upk.py | 2 +- generalresearch/managers/thl/session.py | 11 +--- generalresearch/managers/thl/survey.py | 25 ++++--- generalresearch/managers/thl/task_adjustment.py | 7 +- generalresearch/managers/thl/user_compensate.py | 2 +- .../thl/user_manager/mysql_user_manager.py | 2 +- .../managers/thl/user_manager/user_manager.py | 1 - generalresearch/managers/thl/userhealth.py | 2 +- generalresearch/managers/thl/wall.py | 2 +- generalresearch/managers/thl/wallet/tango.py | 2 +- generalresearch/models/admin/request.py | 2 +- generalresearch/models/custom_types.py | 3 +- generalresearch/models/dynata/survey.py | 4 +- generalresearch/models/gr/authentication.py | 5 +- generalresearch/models/gr/business.py | 6 +- generalresearch/models/gr/team.py | 7 +- generalresearch/models/legacy/questions.py | 3 - generalresearch/models/network/mtr/execute.py | 2 +- generalresearch/models/network/nmap/parser.py | 2 +- generalresearch/models/network/rdns/execute.py | 2 +- generalresearch/models/spectrum/survey.py | 2 +- generalresearch/models/string_utils.py | 1 - generalresearch/models/thl/contest/contest.py | 4 +- .../models/thl/contest/contest_entry.py | 2 +- generalresearch/models/thl/contest/examples.py | 7 -- generalresearch/models/thl/contest/io.py | 2 +- generalresearch/models/thl/contest/leaderboard.py | 2 +- generalresearch/models/thl/contest/milestone.py | 3 +- generalresearch/models/thl/contest/raffle.py | 2 +- generalresearch/models/thl/finance.py | 5 +- generalresearch/models/thl/ipinfo.py | 2 +- generalresearch/models/thl/ledger_example.py | 2 +- generalresearch/models/thl/offerwall/cache.py | 2 +- generalresearch/models/thl/payout_format.py | 12 ++-- generalresearch/models/thl/product.py | 2 - .../models/thl/profiling/marketplace.py | 2 +- .../models/thl/profiling/upk_question.py | 2 +- .../models/thl/profiling/upk_question_answer.py | 2 +- .../models/thl/profiling/user_question_answer.py | 2 +- generalresearch/models/thl/session.py | 2 +- generalresearch/models/thl/survey/__init__.py | 3 - generalresearch/models/thl/survey/buyer.py | 2 +- generalresearch/models/thl/survey/model.py | 2 +- generalresearch/models/thl/survey/penalty.py | 2 +- generalresearch/models/thl/task_adjustment.py | 2 +- generalresearch/models/thl/user.py | 2 +- generalresearch/models/thl/user_iphistory.py | 3 +- generalresearch/models/thl/wallet/payout.py | 2 +- generalresearch/pg_helper.py | 12 ++-- generalresearch/schemas/survey_stats.py | 2 +- generalresearch/sql_helper.py | 19 ++---- generalresearch/thl_django/apps.py | 14 ++-- generalresearch/thl_django/fields.py | 3 +- .../thl_django/migrations/0001_initial.py | 3 +- ..._live_alter_surveycategory_strength_and_more.py | 2 +- ...rveystat_surveystat_live_survey_idx_and_more.py | 2 +- ...ssion_thl_session_status_d578b7_idx_and_more.py | 2 +- ...p_portscanport_iplabel_mtr_portscan_and_more.py | 5 +- generalresearch/thl_django/network/models.py | 5 +- generalresearch/utils/enum.py | 4 +- generalresearch/wall_status_codes/lucid.py | 2 +- generalresearch/wall_status_codes/morning.py | 2 +- generalresearch/wall_status_codes/pollfish.py | 2 +- generalresearch/wall_status_codes/precision.py | 2 +- generalresearch/wall_status_codes/repdata.py | 2 +- test_utils/conftest.py | 3 +- test_utils/grliq/conftest.py | 2 +- test_utils/incite/conftest.py | 2 +- test_utils/managers/conftest.py | 2 - test_utils/models/conftest.py | 9 +-- test_utils/models/contest/conftest.py | 2 +- test_utils/models/network/conftest.py | 2 +- test_utils/models/thl/conftest.py | 2 +- test_utils/spectrum/conftest.py | 2 +- .../incite/collections/test_df_collection_base.py | 3 +- .../collections/test_df_collection_item_base.py | 2 +- .../collections/test_df_collection_item_thl_web.py | 11 +--- .../test_df_collection_thl_marketplaces.py | 7 +- .../collections/test_df_collection_thl_web.py | 2 - .../mergers/foundations/test_enriched_session.py | 7 +- .../foundations/test_enriched_task_adjust.py | 7 -- .../mergers/foundations/test_enriched_wall.py | 23 +------ .../mergers/foundations/test_user_id_product.py | 15 +---- tests/incite/mergers/test_merge_collection.py | 3 +- tests/incite/mergers/test_merge_collection_item.py | 9 +-- tests/incite/mergers/test_pop_ledger.py | 8 +-- tests/incite/mergers/test_ym_survey_merge.py | 17 +---- tests/incite/test_collection_base.py | 3 +- tests/incite/test_collection_base_item.py | 2 +- tests/incite/test_grl_flow.py | 11 ++-- tests/incite/test_interval_idx.py | 3 +- tests/managers/gr/test_authentication.py | 3 - tests/managers/gr/test_business.py | 4 +- tests/managers/gr/test_team.py | 2 - tests/managers/leaderboard.py | 2 +- tests/managers/network/test_label.py | 2 +- tests/managers/test_events.py | 15 ++--- tests/managers/test_userpid.py | 4 +- .../managers/thl/test_contest/test_leaderboard.py | 8 +-- tests/managers/thl/test_contest/test_milestone.py | 14 +--- tests/managers/thl/test_contest/test_raffle.py | 14 +--- tests/managers/thl/test_harmonized_uqa.py | 2 +- tests/managers/thl/test_ipinfo.py | 2 +- tests/managers/thl/test_ledger/test_lm_accounts.py | 6 -- tests/managers/thl/test_ledger/test_lm_tx_locks.py | 23 ++----- .../thl/test_ledger/test_thl_lm_accounts.py | 32 ++++----- .../thl/test_ledger/test_thl_lm_bp_payout.py | 10 +-- tests/managers/thl/test_ledger/test_thl_lm_tx.py | 64 +++--------------- .../test_ledger/test_thl_lm_tx__user_payouts.py | 3 +- tests/managers/thl/test_ledger/test_thl_pem.py | 7 +- tests/managers/thl/test_ledger/test_user_txs.py | 3 +- tests/managers/thl/test_ledger/test_wallet.py | 2 +- tests/managers/thl/test_product.py | 9 ++- tests/managers/thl/test_product_prod.py | 2 - tests/managers/thl/test_profiling/test_user_upk.py | 2 +- tests/managers/thl/test_session_manager.py | 7 +- tests/managers/thl/test_survey.py | 2 +- tests/managers/thl/test_survey_penalty.py | 1 - tests/managers/thl/test_task_adjustment.py | 2 +- tests/managers/thl/test_task_status.py | 2 +- tests/managers/thl/test_user_manager/test_base.py | 4 +- tests/managers/thl/test_user_manager/test_mysql.py | 1 - .../thl/test_user_manager/test_user_fetch.py | 1 - .../thl/test_user_manager/test_user_metadata.py | 1 - tests/managers/thl/test_user_streak.py | 2 +- tests/managers/thl/test_userhealth.py | 2 +- tests/managers/thl/test_wall_manager.py | 13 ++-- tests/models/admin/test_report_request.py | 2 +- tests/models/custom_types/test_aware_datetime.py | 2 +- tests/models/custom_types/test_dsn.py | 1 - tests/models/dynata/test_eligbility.py | 2 +- tests/models/gr/test_authentication.py | 2 +- tests/models/gr/test_business.py | 3 +- tests/models/innovate/test_question.py | 6 +- .../models/legacy/test_user_question_answer_in.py | 4 +- tests/models/morning/test.py | 2 +- tests/models/network/test_mtr.py | 4 +- tests/models/network/test_nmap_parser.py | 1 + tests/models/prodege/test_survey_participation.py | 2 +- tests/models/spectrum/test_question.py | 2 +- tests/models/spectrum/test_survey.py | 2 +- tests/models/spectrum/test_survey_manager.py | 2 +- tests/models/test_finance.py | 2 +- tests/models/thl/question/test_question_info.py | 2 +- tests/models/thl/test_adjustments.py | 30 ++------- .../thl/test_contest/test_leaderboard_contest.py | 2 +- tests/models/thl/test_ledger.py | 2 +- tests/models/thl/test_payout.py | 5 +- tests/models/thl/test_product.py | 2 +- tests/models/thl/test_upkquestion.py | 12 +--- tests/models/thl/test_user.py | 6 +- tests/models/thl/test_user_iphistory.py | 2 +- tests/models/thl/test_user_streak.py | 2 +- tests/models/thl/test_wall.py | 2 +- tests/models/thl/test_wall_session.py | 2 +- tests/sql_helper.py | 4 +- 207 files changed, 398 insertions(+), 725 deletions(-) (limited to 'tests/incite/collections') diff --git a/generalresearch/__init__.py b/generalresearch/__init__.py index 604b7e2..3b2ec3d 100644 --- a/generalresearch/__init__.py +++ b/generalresearch/__init__.py @@ -129,7 +129,7 @@ def synchronized(wrapped): if lock is None: lock = threading.RLock() - setattr(context, "_synchronized_lock", lock) + context._synchronized_lock = lock return lock diff --git a/generalresearch/config.py b/generalresearch/config.py index 76e3995..c6f41e8 100644 --- a/generalresearch/config.py +++ b/generalresearch/config.py @@ -1,7 +1,7 @@ from __future__ import annotations import os -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from pathlib import Path from pydantic import DirectoryPath, Field, MariaDBDsn, PostgresDsn, RedisDsn diff --git a/generalresearch/grliq/managers/event_plotter.py b/generalresearch/grliq/managers/event_plotter.py index 0bba7c5..94b70ef 100644 --- a/generalresearch/grliq/managers/event_plotter.py +++ b/generalresearch/grliq/managers/event_plotter.py @@ -13,7 +13,7 @@ def make_events_svg( mouse_events: list[MouseEvent], keyboard_events: list[KeyboardEvent] ) -> str: if len(mouse_events) + len(keyboard_events) == 0: - return f'\n' + "\n" + return '\n' + "\n" t = np.array([pm.timeStamp for pm in mouse_events]) t_diff = t.max() - t.min() @@ -88,7 +88,7 @@ def make_events_svg( svg_elements.append(svg_multiline_text(text, cx + 5, cy - 5, font_size)) svg = ( - f'' + '' + "\n".join(svg_elements) + "\n" ) diff --git a/generalresearch/grliq/managers/forensic_data.py b/generalresearch/grliq/managers/forensic_data.py index 7567552..093f7ae 100644 --- a/generalresearch/grliq/managers/forensic_data.py +++ b/generalresearch/grliq/managers/forensic_data.py @@ -1,8 +1,8 @@ from __future__ import annotations +from collections.abc import Collection from datetime import datetime from typing import Any -from collections.abc import Collection from psycopg import sql from pydantic import NonNegativeInt, PositiveInt @@ -611,7 +611,7 @@ class GrlIqDataManager: if res and res["c"] >= 0: return int(res["c"]) - except (Exception,) as e: + except Exception: pass query = f""" diff --git a/generalresearch/grliq/managers/forensic_events.py b/generalresearch/grliq/managers/forensic_events.py index c847d4d..93da481 100644 --- a/generalresearch/grliq/managers/forensic_events.py +++ b/generalresearch/grliq/managers/forensic_events.py @@ -36,36 +36,35 @@ class GrlIqEventManager: "uuid": uuid4().hex, } - with self.postgres_config.make_connection() as conn: - with conn.cursor() as c: - c.execute("SELECT pg_advisory_xact_lock(hashtext(%s))", (session_uuid,)) - # Try to update first - update_query = sql.SQL(""" + with self.postgres_config.make_connection() as conn, conn.cursor() as c: + c.execute("SELECT pg_advisory_xact_lock(hashtext(%s))", (session_uuid,)) + # Try to update first + update_query = sql.SQL(""" UPDATE grliq_forensicevents SET timing_data = %(timing_data)s WHERE session_uuid = %(session_uuid)s AND timing_data IS NULL RETURNING id """) - c.execute(update_query, data) - result = c.fetchone() + c.execute(update_query, data) + result = c.fetchone() - if result: - pk = result["id"] - conn.commit() - return pk + if result: + pk = result["id"] + conn.commit() + return pk - # No matching row to update. Do an insert - insert_query = sql.SQL(""" + # No matching row to update. Do an insert + insert_query = sql.SQL(""" INSERT INTO grliq_forensicevents (uuid, session_uuid, timing_data) VALUES (%(uuid)s, %(session_uuid)s, %(timing_data)s) RETURNING id """) - c.execute(insert_query, data) - pk = c.fetchone()["id"] - conn.commit() + c.execute(insert_query, data) + pk = c.fetchone()["id"] + conn.commit() return int(pk) @@ -88,11 +87,10 @@ class GrlIqEventManager: "event_end": event_end, } - with self.postgres_config.make_connection() as conn: - with conn.cursor() as c: - c.execute("SELECT pg_advisory_xact_lock(hashtext(%s))", (session_uuid,)) - # Try to update first - update_query = sql.SQL(""" + with self.postgres_config.make_connection() as conn, conn.cursor() as c: + c.execute("SELECT pg_advisory_xact_lock(hashtext(%s))", (session_uuid,)) + # Try to update first + update_query = sql.SQL(""" UPDATE grliq_forensicevents SET events = %(events)s, mouse_events = %(mouse_events)s, @@ -102,16 +100,16 @@ class GrlIqEventManager: AND events IS NULL RETURNING id """) - c.execute(update_query, data) - result = c.fetchone() + c.execute(update_query, data) + result = c.fetchone() - if result: - pk = result["id"] - conn.commit() - return pk + if result: + pk = result["id"] + conn.commit() + return pk - # No matching row to update. Do an insert - insert_query = sql.SQL(""" + # No matching row to update. Do an insert + insert_query = sql.SQL(""" INSERT INTO grliq_forensicevents (uuid, session_uuid, events, mouse_events, event_start, event_end) @@ -120,9 +118,9 @@ class GrlIqEventManager: %(event_start)s, %(event_end)s) RETURNING id """) - c.execute(insert_query, data) - pk = c.fetchone()["id"] - conn.commit() + c.execute(insert_query, data) + pk = c.fetchone()["id"] + conn.commit() return int(pk) @@ -167,10 +165,9 @@ class GrlIqEventManager: {filter_str} ORDER BY {order_by} LIMIT {limit} """ - with self.postgres_config.make_connection() as conn: - with conn.cursor() as c: - c.execute(query=query, params=params) - res = c.fetchall() + with self.postgres_config.make_connection() as conn, conn.cursor() as c: + c.execute(query=query, params=params) + res = c.fetchall() for x in res: if x.get("mouse_events"): @@ -206,10 +203,9 @@ class GrlIqEventManager: AND timing_data IS NOT NULL ORDER BY session_uuid, fe.id DESC; """) - with self.postgres_config.make_connection() as conn: - with conn.cursor() as c: - c.execute(query, params) - res = c.fetchall() + with self.postgres_config.make_connection() as conn, conn.cursor() as c: + c.execute(query, params) + res = c.fetchall() for x in res: x["timing_data"] = TimingData.model_validate(x["timing_data"]) diff --git a/generalresearch/grliq/managers/forensic_results.py b/generalresearch/grliq/managers/forensic_results.py index 706a7db..158e582 100644 --- a/generalresearch/grliq/managers/forensic_results.py +++ b/generalresearch/grliq/managers/forensic_results.py @@ -90,10 +90,9 @@ class GrlIqCategoryResultsReader: ORDER BY created_at DESC LIMIT {limit} """ - with self.postgres_config.make_connection() as conn: - with conn.cursor() as c: - c.execute(query, params) - res = c.fetchall() + with self.postgres_config.make_connection() as conn, conn.cursor() as c: + c.execute(query, params) + res = c.fetchall() for x in res: x["client_ip"] = str(x["client_ip"]) diff --git a/generalresearch/grliq/models/custom_types.py b/generalresearch/grliq/models/custom_types.py index 5c7f155..903e230 100644 --- a/generalresearch/grliq/models/custom_types.py +++ b/generalresearch/grliq/models/custom_types.py @@ -1,6 +1,7 @@ -import annotated_types from typing import Annotated +import annotated_types + GrlIqScore = Annotated[int, annotated_types.Ge(0), annotated_types.Le(100)] GrlIqAvgScore = Annotated[float, annotated_types.Ge(0), annotated_types.Le(100)] GrlIqRate = Annotated[float, annotated_types.Ge(0), annotated_types.Le(1)] diff --git a/generalresearch/grliq/models/forensic_data.py b/generalresearch/grliq/models/forensic_data.py index 8d65696..666cb81 100644 --- a/generalresearch/grliq/models/forensic_data.py +++ b/generalresearch/grliq/models/forensic_data.py @@ -779,7 +779,6 @@ class GrlIqData(BaseModel): minutes=90 ), "expired session" - return None def model_dump_sql(self, **kwargs) -> dict[str, Any]: d = dict() diff --git a/generalresearch/grliq/models/forensic_summary.py b/generalresearch/grliq/models/forensic_summary.py index 6f80768..f5b0f25 100644 --- a/generalresearch/grliq/models/forensic_summary.py +++ b/generalresearch/grliq/models/forensic_summary.py @@ -3,7 +3,6 @@ from __future__ import annotations import random from typing import ( Literal, - Optional, Union, get_args, get_origin, @@ -125,7 +124,7 @@ def generate_GrlIqCheckerResultsSummary(): if base_type == GrlIqCheckerResult: if is_opt: fields[f"{field_name}_avg"] = ( - Optional[GrlIqAvgScore], + GrlIqAvgScore | None, Field(default=None, examples=[random.randint(0, 100)]), ) fields[f"{field_name}_pct_none"] = ( diff --git a/generalresearch/grliq/utils.py b/generalresearch/grliq/utils.py index ca8c6a1..bceaa30 100644 --- a/generalresearch/grliq/utils.py +++ b/generalresearch/grliq/utils.py @@ -1,7 +1,7 @@ from __future__ import annotations import os -from datetime import datetime, timezone, UTC +from datetime import UTC, datetime from pathlib import Path from uuid import UUID diff --git a/generalresearch/grpc.py b/generalresearch/grpc.py index 178521e..f1b5611 100644 --- a/generalresearch/grpc.py +++ b/generalresearch/grpc.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from google.protobuf.duration_pb2 import Duration from google.protobuf.timestamp_pb2 import Timestamp diff --git a/generalresearch/incite/base.py b/generalresearch/incite/base.py index 6888ca2..44647fc 100644 --- a/generalresearch/incite/base.py +++ b/generalresearch/incite/base.py @@ -7,8 +7,9 @@ import re import shutil import subprocess import warnings +from collections.abc import Callable, Sequence from concurrent.futures import Future -from datetime import datetime, timedelta, timezone, UTC +from datetime import UTC, datetime, timedelta from os import R_OK, access, listdir from os.path import isdir from os.path import join as pjoin @@ -17,8 +18,8 @@ from sys import platform from typing import ( TYPE_CHECKING, Any, + Self, ) -from collections.abc import Callable, Sequence from uuid import uuid4 import dask @@ -42,7 +43,6 @@ from pydantic import ( ) from pydantic.json_schema import SkipJsonSchema from sentry_sdk import capture_exception -from typing import Self from generalresearch.config import is_debug from generalresearch.incite.schemas import ( @@ -481,7 +481,7 @@ class CollectionBase(BaseModel): item.cleanup_partials() def clear_tmp_archives(self) -> None: - regex = re.compile(r"\.parquet\.[0-9a-f]{32}", re.I) + regex = re.compile(r"\.parquet\.[0-9a-f]{32}", re.IGNORECASE) for fn in os.listdir(self.archive_path): if regex.search(fn): @@ -554,7 +554,7 @@ class CollectionBase(BaseModel): try: pq.ParquetDataset(highest_version).read().to_pandas() - except (Exception,): + except Exception: # If the most recent version isn't valid, we don't want to # create a symlink to it. # TODO: We could try to be smart and iterate down the most recent diff --git a/generalresearch/incite/collections/__init__.py b/generalresearch/incite/collections/__init__.py index 38749b3..42c3d31 100644 --- a/generalresearch/incite/collections/__init__.py +++ b/generalresearch/incite/collections/__init__.py @@ -214,13 +214,13 @@ class DFCollectionItem(CollectionItemBase): """, params=[start, finish], ) - except (Exception,) as e: + except Exception as e: capture_exception(error=e) LOG.error(f"_from_mysql Exception: {e}") return None if not res: - LOG.warning(f"_from_mysql query returned nothing") + LOG.warning("_from_mysql query returned nothing") # Return an empty df.DataFrame with the correct columns return empty_dataframe_from_schema(coll._schema) @@ -228,7 +228,7 @@ class DFCollectionItem(CollectionItemBase): df = self.validate_df(df=df) if df is None: - LOG.warning(f"_from_mysql query results failed validation") + LOG.warning("_from_mysql query results failed validation") # Schema validation can fail... return None @@ -265,13 +265,13 @@ class DFCollectionItem(CollectionItemBase): """, params=[start, finish], ) - except (Exception,) as e: + except Exception as e: capture_exception(error=e) LOG.error(f"_from_postgres Exception: {e}") return None if not res: - LOG.warning(f"_from_postgres query returned nothing") + LOG.warning("_from_postgres query returned nothing") # Return an empty df.DataFrame with the correct columns return empty_dataframe_from_schema(coll._schema) @@ -279,7 +279,7 @@ class DFCollectionItem(CollectionItemBase): df = self.validate_df(df=df) if df is None: - LOG.warning(f"_from_postgres query results failed validation") + LOG.warning("_from_postgres query results failed validation") # Schema validation can fail... return None @@ -351,7 +351,7 @@ class DFCollectionItem(CollectionItemBase): c: Cursor = conn.cursor() for chunk in chunked(tx_ids, n=5_000): c.execute( - query=f""" + query=""" SELECT ltm.transaction_id AS tx_id, ltm.id AS tx_metadata_id, ltm.key, ltm.value @@ -466,7 +466,7 @@ class DFCollectionItem(CollectionItemBase): compression="brotli", ) - except (Exception,) as e: + except Exception as e: LOG.exception(e) self.delete_archive(tmp_path) return False @@ -553,7 +553,7 @@ class DFCollectionItem(CollectionItemBase): write_metadata_file=True, compression="brotli", ) - except (Exception,) as e: + except Exception as e: LOG.exception(e) self.delete_archive(next_numbered_path) return False diff --git a/generalresearch/incite/defaults.py b/generalresearch/incite/defaults.py index 5a95607..d4025fc 100644 --- a/generalresearch/incite/defaults.py +++ b/generalresearch/incite/defaults.py @@ -1,4 +1,4 @@ -from datetime import datetime, timezone, UTC +from datetime import UTC, datetime from generalresearch.incite.base import GRLDatasets from generalresearch.incite.collections import DFCollectionType diff --git a/generalresearch/incite/mergers/foundations/__init__.py b/generalresearch/incite/mergers/foundations/__init__.py index f7a45a8..d3100fa 100644 --- a/generalresearch/incite/mergers/foundations/__init__.py +++ b/generalresearch/incite/mergers/foundations/__init__.py @@ -117,7 +117,7 @@ def annotate_product_and_team_id( try: with conn.cursor() as c: c.execute( - query=f""" + query=""" SELECT u.id AS user_id, u.product_id, bp.team_id FROM thl_user u diff --git a/generalresearch/incite/mergers/foundations/enriched_session.py b/generalresearch/incite/mergers/foundations/enriched_session.py index a368e6c..8a707ca 100644 --- a/generalresearch/incite/mergers/foundations/enriched_session.py +++ b/generalresearch/incite/mergers/foundations/enriched_session.py @@ -60,17 +60,17 @@ class EnrichedSessionMergeItem(MergeCollectionItem): return # --- Session --- - LOG.warning(f"EnrichedSessionMergeItem: get session_collection") + LOG.warning("EnrichedSessionMergeItem: get session_collection") session_items = [w for w in session_coll.items if w.interval.overlaps(ir)] if len(session_items) == 0: - LOG.warning(f"EnrichedSessionMergeItem: no session items. set_empty.") + LOG.warning("EnrichedSessionMergeItem: no session items. set_empty.") if self.should_archive(): self.set_empty() return if not ( session_items[-1].has_partial_archive() or session_items[-1].has_archive() ): - LOG.warning(f"EnrichedSessionMergeItem: session isn't updated!") + LOG.warning("EnrichedSessionMergeItem: session isn't updated!") return sddf = session_coll.ddf( @@ -81,7 +81,7 @@ class EnrichedSessionMergeItem(MergeCollectionItem): ) # --- Walls --- - LOG.warning(f"EnrichedSessionMergeItem: merge wall_collection") + LOG.warning("EnrichedSessionMergeItem: merge wall_collection") wall_items = [ w for w in wall_coll.items @@ -95,7 +95,7 @@ class EnrichedSessionMergeItem(MergeCollectionItem): ] if len(wall_items) == 0: - LOG.error(f"EnrichedSessionMergeItem: no wall items") + LOG.error("EnrichedSessionMergeItem: no wall items") return wddf = wall_coll.ddf( diff --git a/generalresearch/incite/mergers/foundations/enriched_task_adjust.py b/generalresearch/incite/mergers/foundations/enriched_task_adjust.py index f3ab8d8..234cd5b 100644 --- a/generalresearch/incite/mergers/foundations/enriched_task_adjust.py +++ b/generalresearch/incite/mergers/foundations/enriched_task_adjust.py @@ -55,7 +55,7 @@ class EnrichedTaskAdjustMergeItem(MergeCollectionItem): LOG.warning(f"EnrichedReconMergeItem.build({ir})") # --- Task Adjustments --- - LOG.warning(f"EnrichedReconMergeItem: get session_collection") + LOG.warning("EnrichedReconMergeItem: get session_collection") task_adj_coll_items = [ w for w in task_adj_coll.items if w.interval.overlaps(ir) ] @@ -209,6 +209,5 @@ class EnrichedTaskAdjustMerge(MergeCollection): enriched_wall=enriched_wall, pg_config=pg_config, ) - except (Exception,) as e: + except Exception as e: capture_exception(error=e) - pass diff --git a/generalresearch/incite/mergers/foundations/enriched_wall.py b/generalresearch/incite/mergers/foundations/enriched_wall.py index 5a7dd2b..b2ac7bb 100644 --- a/generalresearch/incite/mergers/foundations/enriched_wall.py +++ b/generalresearch/incite/mergers/foundations/enriched_wall.py @@ -55,10 +55,10 @@ class EnrichedWallMergeItem(MergeCollectionItem): return # --- Wall --- - LOG.warning(f"EnrichedWallMergeItem: get wall_collection") + LOG.warning("EnrichedWallMergeItem: get wall_collection") wall_items = [w for w in wall_coll.items if w.interval.overlaps(ir)] if len(wall_items) == 0: - LOG.warning(f"EnrichedWallMergeItem: no wall items. set_empty.") + LOG.warning("EnrichedWallMergeItem: no wall items. set_empty.") if self.should_archive(): self.set_empty() return @@ -93,7 +93,7 @@ class EnrichedWallMergeItem(MergeCollectionItem): wdf = wdf.reset_index(drop=False) # --- Sessions --- - LOG.warning(f"EnrichedWallMergeItem: merge session_collection") + LOG.warning("EnrichedWallMergeItem: merge session_collection") session_items = [ s for s in session_coll.items @@ -107,7 +107,7 @@ class EnrichedWallMergeItem(MergeCollectionItem): ] if len(session_items) == 0: - LOG.error(f"EnrichedWallMergeItem: no session items. breaking early.") + LOG.error("EnrichedWallMergeItem: no session items. breaking early.") return sdf = session_coll.ddf( diff --git a/generalresearch/incite/mergers/ym_survey_wall.py b/generalresearch/incite/mergers/ym_survey_wall.py index c060aae..9750b57 100644 --- a/generalresearch/incite/mergers/ym_survey_wall.py +++ b/generalresearch/incite/mergers/ym_survey_wall.py @@ -62,7 +62,7 @@ class YMSurveyWallMergeCollectionItem(MergeCollectionItem): ) ddf = ddf[ddf["started"] > start] - LOG.warning(f"YMSurveyWallMerge: merge session_collection") + LOG.warning("YMSurveyWallMerge: merge session_collection") session_items = [ s for s in enriched_session.items @@ -98,17 +98,16 @@ class YMSurveyWallMergeCollectionItem(MergeCollectionItem): df.dropna(subset="product_id", how="any", inplace=True) df.sort_values(by="started", inplace=True) - LOG.debug(f"YMSurveyWallMerge.build() validation") + LOG.debug("YMSurveyWallMerge.build() validation") df = self.validate_df(df=df) if df is not None: ddf = dd.from_pandas(df, npartitions=4) - LOG.info(f"YMSurveyWallMerge.build() saving") + LOG.info("YMSurveyWallMerge.build() saving") self.to_archive_symlink(client=client, ddf=ddf) else: LOG.warning("YMSurveyWallMerge failed validation") - return None class YMSurveyWallMerge(MergeCollection): @@ -144,8 +143,7 @@ class YMSurveyWallMerge(MergeCollection): wall_coll=wall_coll, enriched_session=enriched_session, ) - except (Exception,) as e: + except Exception as e: capture_exception(error=e) - pass item.delete_dangling_partials(keep_latest=2, target_path=item.path) diff --git a/generalresearch/incite/mergers/ym_wall_summary.py b/generalresearch/incite/mergers/ym_wall_summary.py index 618b810..37fc3b9 100644 --- a/generalresearch/incite/mergers/ym_wall_summary.py +++ b/generalresearch/incite/mergers/ym_wall_summary.py @@ -117,9 +117,8 @@ class YMWallSummaryMerge(MergeCollection): # item every time build is run even if it isn't closed # if item.should_archive(): item.fetch(wall_collection, session_collection, user_id_product) - except (Exception,) as e: + except Exception as e: capture_exception(e) - pass @staticmethod def build_groupbys(df: pd.DataFrame) -> pd.DataFrame: diff --git a/generalresearch/incite/schemas/thl_web.py b/generalresearch/incite/schemas/thl_web.py index b831b9a..c1be202 100644 --- a/generalresearch/incite/schemas/thl_web.py +++ b/generalresearch/incite/schemas/thl_web.py @@ -1,4 +1,4 @@ -from datetime import datetime, timedelta, timezone, UTC +from datetime import UTC, datetime, timedelta import pandas as pd from pandera.pandas import Check, Column, DataFrameSchema, Index, MultiIndex diff --git a/generalresearch/locales/setup_json.py b/generalresearch/locales/setup_json.py index 356084a..57caa00 100644 --- a/generalresearch/locales/setup_json.py +++ b/generalresearch/locales/setup_json.py @@ -11,6 +11,7 @@ def country_default_lang(): """ raise ValueError("no need to run this, I already ran it.") import pandas as pd + from generalresearch.locales import Localelator l = Localelator() diff --git a/generalresearch/locales/timezone.py b/generalresearch/locales/timezone.py index fce6e0e..810dba3 100644 --- a/generalresearch/locales/timezone.py +++ b/generalresearch/locales/timezone.py @@ -1,4 +1,3 @@ -from typing import Optional from pytz import country_timezones diff --git a/generalresearch/managers/cint/survey.py b/generalresearch/managers/cint/survey.py index da1ecd9..819ae3d 100644 --- a/generalresearch/managers/cint/survey.py +++ b/generalresearch/managers/cint/survey.py @@ -2,7 +2,7 @@ from __future__ import annotations import logging from collections.abc import Collection -from datetime import datetime, timezone, UTC +from datetime import UTC, datetime import pymysql from pymysql import IntegrityError diff --git a/generalresearch/managers/criteria.py b/generalresearch/managers/criteria.py index b5d9830..c13b8ac 100644 --- a/generalresearch/managers/criteria.py +++ b/generalresearch/managers/criteria.py @@ -2,7 +2,7 @@ from __future__ import annotations from abc import ABC from collections.abc import Collection -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from more_itertools import chunked @@ -30,7 +30,6 @@ class CriteriaManager(SqlManager, ABC): """ Create a single criterion """ - ... def filter(self, hashes: Collection[str]) -> dict[str, MarketplaceCondition]: """ @@ -96,7 +95,6 @@ class CriteriaManager(SqlManager, ABC): ) conn.commit() - return None @property def mysql_fields(self) -> str: diff --git a/generalresearch/managers/dynata/survey.py b/generalresearch/managers/dynata/survey.py index 3a15c4d..7643dc4 100644 --- a/generalresearch/managers/dynata/survey.py +++ b/generalresearch/managers/dynata/survey.py @@ -2,7 +2,7 @@ from __future__ import annotations import logging from collections.abc import Collection -from datetime import datetime, timezone, UTC +from datetime import UTC, datetime import pymysql from pymysql import IntegrityError diff --git a/generalresearch/managers/events.py b/generalresearch/managers/events.py index f3c6a04..f6c429e 100644 --- a/generalresearch/managers/events.py +++ b/generalresearch/managers/events.py @@ -5,7 +5,7 @@ import math import socket import threading import time -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal from typing import TYPE_CHECKING @@ -724,7 +724,6 @@ class EventManager(StatsManager): ) ) self.publish_event(msg, product_id=user.product_id) - return def handle_task_finish(self, wall: Wall, session: Session, user: User): self.mark_user_active(user=user) @@ -818,7 +817,6 @@ class EventSubscriber(RedisManager): p.subscribe(self.get_channel_name()) self.pubsub_client = r self.pubsub = p - return def get_channel_name(self): return f"{self.cache_prefix}:event-channel:{self.product_id}" diff --git a/generalresearch/managers/gr/authentication.py b/generalresearch/managers/gr/authentication.py index 409cb10..80bee4b 100644 --- a/generalresearch/managers/gr/authentication.py +++ b/generalresearch/managers/gr/authentication.py @@ -3,7 +3,7 @@ from __future__ import annotations import binascii import logging import os -from datetime import datetime, timezone, UTC +from datetime import UTC, datetime from typing import TYPE_CHECKING, Any from psycopg import sql @@ -146,7 +146,7 @@ class GRUserManager(PostgresManagerWithRedis): for item in res: for k, v in item.items(): - if isinstance(item[k], datetime): + if isinstance(v, datetime): item[k] = item[k].replace(tzinfo=UTC) return [GRUser.model_validate(item) for item in res] @@ -270,7 +270,6 @@ class GRTokenManager(PostgresManager): ) conn.commit() - return def get_by_user_id(self, user_id: PositiveInt) -> GRToken | None: # django authtoken_token table has (user_id) UNIQUE constraint diff --git a/generalresearch/managers/gr/team.py b/generalresearch/managers/gr/team.py index 393f446..d04370b 100644 --- a/generalresearch/managers/gr/team.py +++ b/generalresearch/managers/gr/team.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import datetime, timezone, UTC +from datetime import UTC, datetime from typing import TYPE_CHECKING from uuid import uuid4 diff --git a/generalresearch/managers/innovate/survey.py b/generalresearch/managers/innovate/survey.py index c65b100..cddfba2 100644 --- a/generalresearch/managers/innovate/survey.py +++ b/generalresearch/managers/innovate/survey.py @@ -2,7 +2,7 @@ from __future__ import annotations import logging from collections.abc import Collection -from datetime import datetime, timezone, UTC +from datetime import UTC, datetime import pymysql from pymysql import IntegrityError diff --git a/generalresearch/managers/leaderboard/__init__.py b/generalresearch/managers/leaderboard/__init__.py index d5138cd..048d2cb 100644 --- a/generalresearch/managers/leaderboard/__init__.py +++ b/generalresearch/managers/leaderboard/__init__.py @@ -1,8 +1,9 @@ from __future__ import annotations +from zoneinfo import ZoneInfo + import pytz from cachetools import LRUCache, cached -from zoneinfo import ZoneInfo @cached(cache=LRUCache(maxsize=1)) diff --git a/generalresearch/managers/leaderboard/manager.py b/generalresearch/managers/leaderboard/manager.py index 71c3a73..07e3e2c 100644 --- a/generalresearch/managers/leaderboard/manager.py +++ b/generalresearch/managers/leaderboard/manager.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import datetime, timedelta, timezone, UTC +from datetime import UTC, datetime, timedelta from decimal import Decimal from functools import cached_property from typing import TYPE_CHECKING diff --git a/generalresearch/managers/morning/survey.py b/generalresearch/managers/morning/survey.py index 5fba70d..0e29010 100644 --- a/generalresearch/managers/morning/survey.py +++ b/generalresearch/managers/morning/survey.py @@ -3,7 +3,7 @@ from __future__ import annotations import json import logging from collections.abc import Collection -from datetime import datetime, timezone, UTC +from datetime import UTC, datetime import pymysql from pymysql import IntegrityError diff --git a/generalresearch/managers/network/label.py b/generalresearch/managers/network/label.py index c5306a6..1f44862 100644 --- a/generalresearch/managers/network/label.py +++ b/generalresearch/managers/network/label.py @@ -1,7 +1,7 @@ from __future__ import annotations from collections.abc import Collection -from datetime import datetime, timedelta, timezone, UTC +from datetime import UTC, datetime, timedelta from psycopg import sql from pydantic import IPvAnyNetwork, TypeAdapter @@ -28,10 +28,9 @@ class IPLabelManager(PostgresManager): %(provider)s, %(metadata)s ) RETURNING id;""") params = ip_label.model_dump_postgres() - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute(query, params) - pk = c.fetchone()["id"] + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute(query, params) + pk = c.fetchone()["id"] return ip_label def make_filter_str( diff --git a/generalresearch/managers/network/tool_run.py b/generalresearch/managers/network/tool_run.py index 026b3d3..73afc13 100644 --- a/generalresearch/managers/network/tool_run.py +++ b/generalresearch/managers/network/tool_run.py @@ -49,7 +49,6 @@ class ToolRunManager(PostgresManager): c.execute(query, params) run_id = c.fetchone()["id"] run.id = run_id - return None def create_tool_run(self, run: NmapRun | RDNSRun | MTRRun): if type(run) is NmapRun: @@ -77,10 +76,9 @@ class ToolRunManager(PostgresManager): """ Insert a PortScan + PortScanPorts from a Pydantic NmapResult. """ - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - self._create_tool_run(run, c) - self.nmap_manager._create(run, c=c) + with self.pg_config.make_connection() as conn, conn.cursor() as c: + self._create_tool_run(run, c) + self.nmap_manager._create(run, c=c) return run def get_nmap_run(self, id: int) -> NmapRun: @@ -98,10 +96,9 @@ class ToolRunManager(PostgresManager): """ Insert a RDnsRun + RDNSResult """ - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - self._create_tool_run(run, c) - self.rdns_manager._create(run, c=c) + with self.pg_config.make_connection() as conn, conn.cursor() as c: + self._create_tool_run(run, c) + self.rdns_manager._create(run, c=c) return run def get_rdns_run(self, id: int) -> RDNSRun: @@ -120,10 +117,9 @@ class ToolRunManager(PostgresManager): return RDNSRun.model_validate(res) def create_mtr_run(self, run: MTRRun) -> MTRRun: - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - self._create_tool_run(run, c) - self.mtr_manager._create(run, c=c) + with self.pg_config.make_connection() as conn, conn.cursor() as c: + self._create_tool_run(run, c) + self.mtr_manager._create(run, c=c) return run def get_mtr_run(self, id: int) -> MTRRun: diff --git a/generalresearch/managers/precision/survey.py b/generalresearch/managers/precision/survey.py index 833cb28..c13dca8 100644 --- a/generalresearch/managers/precision/survey.py +++ b/generalresearch/managers/precision/survey.py @@ -2,7 +2,7 @@ from __future__ import annotations import logging from collections.abc import Collection -from datetime import datetime, timezone, UTC +from datetime import UTC, datetime import pymysql from pymysql import IntegrityError @@ -125,7 +125,7 @@ class PrecisionSurveyManager(SurveyManager): country_data = [(survey.survey_id, c) for c in survey.country_isos] c.executemany( - f""" + """ INSERT INTO `thl-precision`.`precision_survey_country` (survey_id, country_iso, is_active) VALUES (%s, %s, TRUE) @@ -134,7 +134,7 @@ class PrecisionSurveyManager(SurveyManager): ) lang_data = [(survey.survey_id, c) for c in survey.language_isos] c.executemany( - f""" + """ INSERT INTO `thl-precision`.`precision_survey_language` (survey_id, language_iso, is_active) VALUES (%s, %s, TRUE) @@ -188,7 +188,7 @@ class PrecisionSurveyManager(SurveyManager): country_data = [(survey.survey_id, c) for c in survey.country_isos] # Turn ON countries in this survey's list of countries, insert row, if already exists, set active. c.executemany( - query=f""" + query=""" INSERT INTO `thl-precision`.`precision_survey_country` (survey_id, country_iso, is_active) VALUES (%s, %s, TRUE) ON DUPLICATE KEY UPDATE is_active = TRUE; @@ -207,7 +207,7 @@ class PrecisionSurveyManager(SurveyManager): ) language_data = [(survey.survey_id, c) for c in survey.language_isos] c.executemany( - query=f""" + query=""" INSERT INTO `thl-precision`.`precision_survey_language` (survey_id, language_iso, is_active) VALUES (%s, %s, TRUE) ON DUPLICATE KEY UPDATE is_active = TRUE; diff --git a/generalresearch/managers/prodege/survey.py b/generalresearch/managers/prodege/survey.py index 750383f..983572c 100644 --- a/generalresearch/managers/prodege/survey.py +++ b/generalresearch/managers/prodege/survey.py @@ -1,7 +1,7 @@ from __future__ import annotations from collections.abc import Collection -from datetime import datetime, timezone, UTC +from datetime import UTC, datetime import pymysql diff --git a/generalresearch/managers/repdata/survey.py b/generalresearch/managers/repdata/survey.py index 1e1c3c6..fe6f621 100644 --- a/generalresearch/managers/repdata/survey.py +++ b/generalresearch/managers/repdata/survey.py @@ -2,7 +2,7 @@ from __future__ import annotations import json from collections.abc import Collection -from datetime import datetime, timezone, UTC +from datetime import UTC, datetime import pymysql @@ -105,7 +105,7 @@ class RepDataSurveyManager(SurveyManager): surveys = {s.survey_id: s for s in surveys} if surveys: res = self.sql_helper.execute_sql_query( - query=f""" + query=""" SELECT * FROM `thl-repdata`.`repdata_surveystream` WHERE survey_id IN %s diff --git a/generalresearch/managers/sago/survey.py b/generalresearch/managers/sago/survey.py index 2582902..462d2ef 100644 --- a/generalresearch/managers/sago/survey.py +++ b/generalresearch/managers/sago/survey.py @@ -2,7 +2,7 @@ from __future__ import annotations import logging from collections.abc import Collection -from datetime import datetime, timezone, UTC +from datetime import UTC, datetime import pymysql from pymysql import IntegrityError diff --git a/generalresearch/managers/spectrum/survey.py b/generalresearch/managers/spectrum/survey.py index 58f8a1a..9b58d43 100644 --- a/generalresearch/managers/spectrum/survey.py +++ b/generalresearch/managers/spectrum/survey.py @@ -2,7 +2,7 @@ from __future__ import annotations import logging from collections.abc import Collection -from datetime import datetime, timezone, UTC +from datetime import UTC, datetime import pymysql from pymysql import IntegrityError diff --git a/generalresearch/managers/survey.py b/generalresearch/managers/survey.py index 964344f..1f057fb 100644 --- a/generalresearch/managers/survey.py +++ b/generalresearch/managers/survey.py @@ -12,14 +12,12 @@ class SurveyManager(SqlManager, ABC): """ Create a single survey """ - ... def update(self, surveys: list[MarketplaceTask]) -> bool: """ Update a list of surveys. Depending on the implementation, this may operate one by one or as a bulk update. """ - ... def update_field(self, survey: MarketplaceTask, field: str) -> bool: """ diff --git a/generalresearch/managers/thl/buyer.py b/generalresearch/managers/thl/buyer.py index ae40bc8..2cb582f 100644 --- a/generalresearch/managers/thl/buyer.py +++ b/generalresearch/managers/thl/buyer.py @@ -1,7 +1,7 @@ from __future__ import annotations from collections.abc import Collection -from datetime import datetime, timezone, UTC +from datetime import UTC, datetime from generalresearch.managers.base import Permission, PostgresManager from generalresearch.models import Source diff --git a/generalresearch/managers/thl/cashout_method.py b/generalresearch/managers/thl/cashout_method.py index 94617f7..10282d4 100644 --- a/generalresearch/managers/thl/cashout_method.py +++ b/generalresearch/managers/thl/cashout_method.py @@ -2,7 +2,7 @@ from __future__ import annotations from collections.abc import Collection from copy import copy -from datetime import datetime, timezone, UTC +from datetime import UTC, datetime from typing import Any from uuid import UUID, uuid4 diff --git a/generalresearch/managers/thl/contest_manager.py b/generalresearch/managers/thl/contest_manager.py index 286de3d..62146d7 100644 --- a/generalresearch/managers/thl/contest_manager.py +++ b/generalresearch/managers/thl/contest_manager.py @@ -1,7 +1,7 @@ from __future__ import annotations from collections.abc import Collection -from datetime import datetime, timezone, UTC +from datetime import UTC, datetime from typing import Any, Literal, cast from uuid import UUID @@ -827,7 +827,6 @@ class MilestoneContestManager(ContestBaseManager): ) self.end_milestone_contest(contest) - return None def enter_contest_db_work_milestone( self, contest: MilestoneUserView, user: User, incr: PositiveInt diff --git a/generalresearch/managers/thl/ledger_manager/conditions.py b/generalresearch/managers/thl/ledger_manager/conditions.py index f79a0f0..b2fd465 100644 --- a/generalresearch/managers/thl/ledger_manager/conditions.py +++ b/generalresearch/managers/thl/ledger_manager/conditions.py @@ -1,9 +1,9 @@ from __future__ import annotations import logging -from datetime import datetime, timedelta, timezone, UTC -from typing import TYPE_CHECKING from collections.abc import Callable +from datetime import UTC, datetime, timedelta +from typing import TYPE_CHECKING from generalresearch.config import JAMES_BILLINGS_BPID, JAMES_BILLINGS_TX_CUTOFF from generalresearch.currency import USDCent diff --git a/generalresearch/managers/thl/ledger_manager/exceptions.py b/generalresearch/managers/thl/ledger_manager/exceptions.py index b79c153..48c102b 100644 --- a/generalresearch/managers/thl/ledger_manager/exceptions.py +++ b/generalresearch/managers/thl/ledger_manager/exceptions.py @@ -11,7 +11,6 @@ class LedgerTransactionCreateError(Exception): Ledger transaction creation failed """ - pass class LedgerTransactionCreateLockError(LedgerTransactionCreateError): @@ -19,7 +18,6 @@ class LedgerTransactionCreateLockError(LedgerTransactionCreateError): Ledger transaction creation failed because we could not acquire a lock """ - pass class LedgerTransactionReleaseLockError(LedgerTransactionCreateError): @@ -29,7 +27,6 @@ class LedgerTransactionReleaseLockError(LedgerTransactionCreateError): back-populate as in sentry I see this very rarely. """ - pass class LedgerTransactionFlagAlreadyExistsError(LedgerTransactionCreateError): @@ -38,7 +35,6 @@ class LedgerTransactionFlagAlreadyExistsError(LedgerTransactionCreateError): tx was already set """ - pass class LedgerTransactionConditionFailedError(LedgerTransactionCreateError): @@ -46,4 +42,3 @@ class LedgerTransactionConditionFailedError(LedgerTransactionCreateError): We tried to create a transaction but the condition check failed. """ - pass diff --git a/generalresearch/managers/thl/ledger_manager/ledger.py b/generalresearch/managers/thl/ledger_manager/ledger.py index 00fac27..410f6ca 100644 --- a/generalresearch/managers/thl/ledger_manager/ledger.py +++ b/generalresearch/managers/thl/ledger_manager/ledger.py @@ -2,10 +2,9 @@ from __future__ import annotations import logging from collections import defaultdict -from collections.abc import Collection -from datetime import datetime, timedelta, timezone, UTC +from collections.abc import Callable, Collection +from datetime import UTC, datetime, timedelta from typing import Any -from collections.abc import Callable from uuid import UUID import redis @@ -343,7 +342,7 @@ class LedgerTransactionManager(LedgerManagerBasePostgres): assert len(tag) > 6, "Please confirm the tag is valid" res = self.pg_config.execute_sql_query( - query=f""" + query=""" SELECT lt.id FROM ledger_transaction AS lt WHERE tag = %s @@ -361,7 +360,7 @@ class LedgerTransactionManager(LedgerManagerBasePostgres): def get_tx_ids_by_tags(self, tags: list[str]) -> set[PositiveInt]: res = self.pg_config.execute_sql_query( - query=f""" + query=""" SELECT lt.id, lt.tag, lt.created, lt.ext_description FROM ledger_transaction AS lt WHERE tag = ANY(%s) @@ -869,7 +868,7 @@ class LedgerAccountManager(LedgerManagerBasePostgres): # qualified_name has a unique index so there can only be 0 or 1 match. res = self.pg_config.execute_sql_query( - query=f""" + query=""" SELECT uuid, display_name, qualified_name, account_type, normal_balance, reference_type, @@ -928,7 +927,7 @@ class LedgerAccountManager(LedgerManagerBasePostgres): # TODO: Move to RR with long timeout (2min+), it causes problems res = self.pg_config.execute_sql_query( - query=f""" + query=""" SELECT SUM(amount * direction) AS total FROM ledger_entry WHERE account_id = %s @@ -1046,7 +1045,7 @@ class LedgerManager( """This is for testing only, as it'll take forever to run this if the ledger_manager is huge """ - res = self.pg_config.execute_sql_query(f""" + res = self.pg_config.execute_sql_query(""" SELECT SUM(CASE WHEN normal_balance = -1 THEN total ELSE 0 END) AS credit_total, SUM(CASE WHEN normal_balance = 1 THEN total ELSE 0 END) AS debit_total diff --git a/generalresearch/managers/thl/ledger_manager/thl_ledger.py b/generalresearch/managers/thl/ledger_manager/thl_ledger.py index c9203f6..ded6518 100644 --- a/generalresearch/managers/thl/ledger_manager/thl_ledger.py +++ b/generalresearch/managers/thl/ledger_manager/thl_ledger.py @@ -1,11 +1,10 @@ from __future__ import annotations import logging -from collections.abc import Collection -from datetime import datetime, timedelta, timezone, UTC +from collections.abc import Callable, Collection +from datetime import UTC, datetime, timedelta from decimal import Decimal from typing import TYPE_CHECKING -from collections.abc import Callable from uuid import UUID import numpy as np @@ -407,8 +406,7 @@ class ThlLedgerManager(LedgerManager): f"bp_pay {bp_pay} > thl_net {thl_net}. Capping bp_pay to thl_net." ) bp_pay = thl_net - if user_pay > bp_pay: - user_pay = bp_pay + user_pay = min(user_pay, bp_pay) commission_amount = round(thl_net - bp_pay) @@ -538,7 +536,7 @@ class ThlLedgerManager(LedgerManager): ] else: - logger.info(f"create_transaction_task_adjustment. No transactions needed.") + logger.info("create_transaction_task_adjustment. No transactions needed.") return None amt_str = f"${abs(change_amount) / 100:,.2f}" @@ -666,7 +664,7 @@ class ThlLedgerManager(LedgerManager): else: logger.info( - f"create_transaction_bp_adjustment. No transactions needed." + "create_transaction_bp_adjustment. No transactions needed." ) return None else: @@ -738,7 +736,7 @@ class ThlLedgerManager(LedgerManager): else: logger.info( - f"create_transaction_bp_adjustment. No transactions needed." + "create_transaction_bp_adjustment. No transactions needed." ) return None @@ -856,7 +854,7 @@ class ThlLedgerManager(LedgerManager): ), ] - ext_description = f"BP Payout" + ext_description = "BP Payout" t = self.create_tx( entries=entries, metadata=metadata, @@ -983,7 +981,7 @@ class ThlLedgerManager(LedgerManager): raise ValueError("Invalid Direction") if description is None: - description = f"BP Plug" + description = "BP Plug" t = self.create_tx( entries=entries, @@ -1191,7 +1189,7 @@ class ThlLedgerManager(LedgerManager): f"Trying to cancel user payout {payout_event.uuid} with no request tx found." ) - description = f"User Payout Cancelled" + description = "User Payout Cancelled" f = lambda: self.create_tx_user_payout_cancelled_( user=user, payout_event=payout_event, @@ -1900,7 +1898,7 @@ class ThlLedgerManager(LedgerManager): reserve = round(wall["user_payout_int"].sum() - wall["redeemable"].sum()) redeemable_balance = user_wallet_balance - reserve - redeemable_balance = 0 if redeemable_balance < 0 else redeemable_balance + redeemable_balance = max(redeemable_balance, 0) if redeemable_balance > 0: # it is possible the user_wallet_balance is negative, in which case diff --git a/generalresearch/managers/thl/product.py b/generalresearch/managers/thl/product.py index 3b92361..d924e17 100644 --- a/generalresearch/managers/thl/product.py +++ b/generalresearch/managers/thl/product.py @@ -4,7 +4,7 @@ import json import logging import operator from collections.abc import Collection -from datetime import datetime, timezone, UTC +from datetime import UTC, datetime from decimal import Decimal from threading import Lock from typing import TYPE_CHECKING @@ -94,9 +94,9 @@ class ProductManager(PostgresManager): return self.fetch_uuids( product_uuids=[product_uuid], )[0] - except (AssertionError,): + except AssertionError: return None - except (IndexError,): + except IndexError: return None def get_by_uuids_if_exists( @@ -397,11 +397,11 @@ class ProductManager(PostgresManager): # # from pymysql import IntegrityError # except IntegrityError as e: - except (Exception,) as e: + except Exception as e: try: return self.get_by_uuid(product_uuid=instance.id) - except (Exception,) as e2: + except Exception: pass finally: self.cache_clear(instance.id) diff --git a/generalresearch/managers/thl/profiling/question.py b/generalresearch/managers/thl/profiling/question.py index 10a9e32..078894a 100644 --- a/generalresearch/managers/thl/profiling/question.py +++ b/generalresearch/managers/thl/profiling/question.py @@ -85,7 +85,7 @@ class QuestionManager(PostgresManager): def lookup_by_property( self, property_code: str, country_iso: str, language_iso: str ) -> UpkQuestion: - query = f""" + query = """ SELECT data, property_code, explanation_template, explanation_fragment_template FROM marketplace_question WHERE property_code = %(property_code)s diff --git a/generalresearch/managers/thl/profiling/uqa.py b/generalresearch/managers/thl/profiling/uqa.py index 3854333..fa1747b 100644 --- a/generalresearch/managers/thl/profiling/uqa.py +++ b/generalresearch/managers/thl/profiling/uqa.py @@ -2,7 +2,7 @@ from __future__ import annotations import logging from collections.abc import Collection -from datetime import datetime, timedelta, timezone, UTC +from datetime import UTC, datetime, timedelta from generalresearch.managers.base import PostgresManagerWithRedis from generalresearch.models.thl.profiling.user_question_answer import ( diff --git a/generalresearch/managers/thl/profiling/user_upk.py b/generalresearch/managers/thl/profiling/user_upk.py index ee36124..a2cddb3 100644 --- a/generalresearch/managers/thl/profiling/user_upk.py +++ b/generalresearch/managers/thl/profiling/user_upk.py @@ -3,7 +3,7 @@ from __future__ import annotations import json from collections import defaultdict from collections.abc import Collection -from datetime import datetime, timedelta, timezone, UTC +from datetime import UTC, datetime, timedelta from typing import Any from uuid import UUID diff --git a/generalresearch/managers/thl/session.py b/generalresearch/managers/thl/session.py index 771f882..959003c 100644 --- a/generalresearch/managers/thl/session.py +++ b/generalresearch/managers/thl/session.py @@ -1,7 +1,7 @@ from __future__ import annotations from collections.abc import Collection -from datetime import datetime, timedelta, timezone, UTC +from datetime import UTC, datetime, timedelta from decimal import Decimal from typing import Any from uuid import UUID, uuid4 @@ -190,14 +190,7 @@ class SessionManager(PostgresManager): # re-run model_validate after finished = finished if finished else datetime.now(tz=UTC) session.update( - **{ - "status": status, - "status_code_1": status_code_1, - "status_code_2": status_code_2, - "finished": finished, - "payout": payout, - "user_payout": user_payout, - } + status=status, status_code_1=status_code_1, status_code_2=status_code_2, finished=finished, payout=payout, user_payout=user_payout ) d = session.model_dump_mysql() self.pg_config.execute_write( diff --git a/generalresearch/managers/thl/survey.py b/generalresearch/managers/thl/survey.py index 966e96d..a9ec841 100644 --- a/generalresearch/managers/thl/survey.py +++ b/generalresearch/managers/thl/survey.py @@ -2,7 +2,7 @@ from __future__ import annotations from collections import defaultdict from collections.abc import Collection -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from typing import Any import pandas as pd @@ -280,7 +280,6 @@ class SurveyManager(PostgresManager): query=query, params={"survey_pks": survey_pks}, ) - return None def update_surveys_categories(self, surveys: list[Survey] | None = None) -> None: for chunk in chunked(surveys, 500): @@ -329,12 +328,11 @@ class SurveyManager(PostgresManager): ] with self.pg_config.make_connection() as conn: # noinspection PyArgumentList - with conn.transaction(): - with conn.cursor() as c: - c.execute(temp_table_sql) - c.executemany(insert_values_sql, rows) - c.execute(delete_sql) - c.execute(upsert_sql) + with conn.transaction(), conn.cursor() as c: + c.execute(temp_table_sql) + c.executemany(insert_values_sql, rows) + c.execute(delete_sql) + c.execute(upsert_sql) conn.commit() def get_survey_categories(self): @@ -760,12 +758,11 @@ class SurveyStatManager(PostgresManager): print(query) print(params) - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute("SET work_mem = '256MB';") - c.execute("SET statement_timeout = '10s';") - c.execute(query, params=params) - res = c.fetchall() + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute("SET work_mem = '256MB';") + c.execute("SET statement_timeout = '10s';") + c.execute(query, params=params) + res = c.fetchall() return [SurveyStat.model_validate(x) for x in res] diff --git a/generalresearch/managers/thl/task_adjustment.py b/generalresearch/managers/thl/task_adjustment.py index 3ec3d41..60bade7 100644 --- a/generalresearch/managers/thl/task_adjustment.py +++ b/generalresearch/managers/thl/task_adjustment.py @@ -1,7 +1,7 @@ from __future__ import annotations import logging -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from decimal import Decimal from functools import cached_property @@ -148,10 +148,7 @@ class TaskAdjustmentManager(PostgresManager): if ( wall.status == Status.COMPLETE and adjusted_status == WallAdjustedStatus.ADJUSTED_TO_COMPLETE - ): - new_adjusted_status = None - new_adjusted_cpi = None - elif ( + ) or ( wall.status != Status.COMPLETE and adjusted_status == WallAdjustedStatus.ADJUSTED_TO_FAIL ): diff --git a/generalresearch/managers/thl/user_compensate.py b/generalresearch/managers/thl/user_compensate.py index c6c0747..8338424 100644 --- a/generalresearch/managers/thl/user_compensate.py +++ b/generalresearch/managers/thl/user_compensate.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from decimal import Decimal from uuid import uuid4 diff --git a/generalresearch/managers/thl/user_manager/mysql_user_manager.py b/generalresearch/managers/thl/user_manager/mysql_user_manager.py index 7931ba4..d2d0ffc 100644 --- a/generalresearch/managers/thl/user_manager/mysql_user_manager.py +++ b/generalresearch/managers/thl/user_manager/mysql_user_manager.py @@ -2,7 +2,7 @@ from __future__ import annotations import logging from collections.abc import Collection -from datetime import datetime, timezone, UTC +from datetime import UTC, datetime from functools import lru_cache from uuid import uuid4 diff --git a/generalresearch/managers/thl/user_manager/user_manager.py b/generalresearch/managers/thl/user_manager/user_manager.py index 3794020..7486869 100644 --- a/generalresearch/managers/thl/user_manager/user_manager.py +++ b/generalresearch/managers/thl/user_manager/user_manager.py @@ -103,7 +103,6 @@ class UserManager: event_value=event_value, ) - return None def cache_clear(self): # Generally this is used in testing. This clears the .get_user's lru_cache. diff --git a/generalresearch/managers/thl/userhealth.py b/generalresearch/managers/thl/userhealth.py index 0bc60ec..babed04 100644 --- a/generalresearch/managers/thl/userhealth.py +++ b/generalresearch/managers/thl/userhealth.py @@ -2,7 +2,7 @@ from __future__ import annotations import ipaddress from collections.abc import Collection -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from itertools import zip_longest from typing import Any diff --git a/generalresearch/managers/thl/wall.py b/generalresearch/managers/thl/wall.py index 7e413d7..03ca1c6 100644 --- a/generalresearch/managers/thl/wall.py +++ b/generalresearch/managers/thl/wall.py @@ -3,7 +3,7 @@ from __future__ import annotations import logging from collections import defaultdict from collections.abc import Collection -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal from functools import cached_property from uuid import uuid4 diff --git a/generalresearch/managers/thl/wallet/tango.py b/generalresearch/managers/thl/wallet/tango.py index a4ebf22..4abfc70 100644 --- a/generalresearch/managers/thl/wallet/tango.py +++ b/generalresearch/managers/thl/wallet/tango.py @@ -44,7 +44,7 @@ def complete_tango_order( tango_client=tango_client, ) - except Exception as e: + except Exception: # todo: its possible the order went through, but something else was wrong # we should try to retrieve the order by its ref_id and confirm it really # failed... diff --git a/generalresearch/models/admin/request.py b/generalresearch/models/admin/request.py index 2d68de1..5fdc784 100644 --- a/generalresearch/models/admin/request.py +++ b/generalresearch/models/admin/request.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from enum import Enum from typing import Literal diff --git a/generalresearch/models/custom_types.py b/generalresearch/models/custom_types.py index 9346064..c200b34 100644 --- a/generalresearch/models/custom_types.py +++ b/generalresearch/models/custom_types.py @@ -2,8 +2,7 @@ from __future__ import annotations import json import re -import sys as _sys -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from typing import Annotated, Any, Literal from uuid import UUID diff --git a/generalresearch/models/dynata/survey.py b/generalresearch/models/dynata/survey.py index 9adac42..5a9f763 100644 --- a/generalresearch/models/dynata/survey.py +++ b/generalresearch/models/dynata/survey.py @@ -132,9 +132,7 @@ class DynataCondition(MarketplaceCondition): if cell["kind"] == "RANGE": d["values"] = [ - "{0}-{1}".format( - cell["range"]["from"] or "inf", cell["range"]["to"] or "inf" - ) + f"{cell["range"]["from"] or "inf"}-{cell["range"]["to"] or "inf"}" ] d["value_type"] = ConditionValueType.RANGE return cls.model_validate(d) diff --git a/generalresearch/models/gr/authentication.py b/generalresearch/models/gr/authentication.py index 67a8fc2..f9644fe 100644 --- a/generalresearch/models/gr/authentication.py +++ b/generalresearch/models/gr/authentication.py @@ -3,7 +3,7 @@ from __future__ import annotations import binascii import json import os -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from typing import TYPE_CHECKING, Any, Self from pydantic import ( @@ -181,7 +181,7 @@ class GRUser(BaseModel): team_products = pm.fetch_uuids(team_uuids=team_uuids) if team_uuids else [] products = {p.id: p for p in business_products + team_products} - self.products = sorted(products.values(), key=lambda x: getattr(x, "created")) + self.products = sorted(products.values(), key=lambda x: x.created) def prefetch_token(self, pg_config: PostgresConfig): from generalresearch.managers.gr.authentication import ( @@ -283,7 +283,6 @@ class GRUser(BaseModel): ex=ex_secs, ) - return None # --- ORM --- diff --git a/generalresearch/models/gr/business.py b/generalresearch/models/gr/business.py index 534b23f..74b5c29 100644 --- a/generalresearch/models/gr/business.py +++ b/generalresearch/models/gr/business.py @@ -420,7 +420,7 @@ class Business(BaseModel): # that is still valid. Don't attempt to build a balance, leave it # as None rather than all zeros LOG.warning(f"Business({self.uuid=}).prebuild_balance empty dataframe") - return None + return LOG.debug(f"Business.prebuild_balance.groupby() {df.head()}") df = df.groupby("account_id").sum() @@ -442,13 +442,11 @@ class Business(BaseModel): order_by=OrderBy.DESC, ) self.prebuild_payouts_total() - return None def prebuild_payouts_total(self): assert self.payouts is not None self.payouts_total = USDCent(sum([po.amount for po in self.payouts])) self.payouts_total_str = self.payouts_total.to_usd_str() - return None def prebuild_pop_financial( self, @@ -548,7 +546,6 @@ class Business(BaseModel): except Exception as e: raise OSError(f"Parquet verification failed: {e}") - return None def prebuild_enriched_wall_parquet( self, @@ -593,7 +590,6 @@ class Business(BaseModel): except Exception as e: raise OSError(f"Parquet verification failed: {e}") - return None @classmethod def required_fields(cls) -> list[str]: diff --git a/generalresearch/models/gr/team.py b/generalresearch/models/gr/team.py index 900062f..78a9ba9 100644 --- a/generalresearch/models/gr/team.py +++ b/generalresearch/models/gr/team.py @@ -2,7 +2,7 @@ from __future__ import annotations import json import os -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from enum import Enum from pathlib import Path from typing import TYPE_CHECKING, Self @@ -194,7 +194,6 @@ class Team(BaseModel): except Exception as e: raise OSError(f"Parquet verification failed: {e}") - return def prebuild_enriched_wall_parquet( self, @@ -239,7 +238,6 @@ class Team(BaseModel): except Exception as e: raise OSError(f"Parquet verification failed: {e}") - return None @classmethod def required_fields(cls) -> list[str]: @@ -325,7 +323,6 @@ class Team(BaseModel): enriched_wall=enriched_wall, ) - return # --- ORM --- @@ -344,5 +341,5 @@ class Team(BaseModel): d = {val: json.loads(res[idx]) for idx, val in enumerate(keys)} return Team.model_validate(d) - except (Exception,) as e: + except Exception: return None diff --git a/generalresearch/models/legacy/questions.py b/generalresearch/models/legacy/questions.py index 8e19e57..4651ab0 100644 --- a/generalresearch/models/legacy/questions.py +++ b/generalresearch/models/legacy/questions.py @@ -106,7 +106,6 @@ class UserQuestionAnswerIn(BaseModel): if self.question_id == user_agent_qid: val = self.answer[0] # assert val == request.user_agent.to_header(): - pass return self @@ -217,7 +216,6 @@ class UserQuestionAnswers(BaseModel): # --- Prefetch --- def prefetch_user(self, um: UserManager) -> None: - from generalresearch.models.thl.user import User res: User | None = um.get_user_if_exists( product_id=self.product_id, product_user_id=self.product_user_id @@ -230,7 +228,6 @@ class UserQuestionAnswers(BaseModel): def prefetch_wall(self, wm: WallManager) -> None: from generalresearch.models import Source - from generalresearch.models.thl.session import Wall res: Wall | None = wm.get_from_uuid_if_exists(wall_uuid=self.session_id) diff --git a/generalresearch/models/network/mtr/execute.py b/generalresearch/models/network/mtr/execute.py index 953124d..c5b3c5c 100644 --- a/generalresearch/models/network/mtr/execute.py +++ b/generalresearch/models/network/mtr/execute.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from uuid import uuid4 from generalresearch.models.custom_types import UUIDStr diff --git a/generalresearch/models/network/nmap/parser.py b/generalresearch/models/network/nmap/parser.py index 6ad4ab4..ecaf2d1 100644 --- a/generalresearch/models/network/nmap/parser.py +++ b/generalresearch/models/network/nmap/parser.py @@ -1,7 +1,7 @@ from __future__ import annotations import xml.etree.ElementTree as ET -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from typing import Any from generalresearch.models.network.definitions import IPProtocol diff --git a/generalresearch/models/network/rdns/execute.py b/generalresearch/models/network/rdns/execute.py index 1d74df2..d6de84b 100644 --- a/generalresearch/models/network/rdns/execute.py +++ b/generalresearch/models/network/rdns/execute.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from uuid import uuid4 from generalresearch.models.custom_types import UUIDStr diff --git a/generalresearch/models/spectrum/survey.py b/generalresearch/models/spectrum/survey.py index aaf182e..f9a5e27 100644 --- a/generalresearch/models/spectrum/survey.py +++ b/generalresearch/models/spectrum/survey.py @@ -75,7 +75,7 @@ class SpectrumCondition(MarketplaceCondition): rs["from"] = round(rs["from"] / 12) rs["to"] = round(rs["to"] / 12) d["values"] = [ - "{0}-{1}".format(rs["from"] or "inf", rs["to"] or "inf") + f"{rs["from"] or "inf"}-{rs["to"] or "inf"}" for rs in d["range_sets"] ] d["value_type"] = ConditionValueType.RANGE diff --git a/generalresearch/models/string_utils.py b/generalresearch/models/string_utils.py index d76456f..dff2f4d 100644 --- a/generalresearch/models/string_utils.py +++ b/generalresearch/models/string_utils.py @@ -1,5 +1,4 @@ import unicodedata -from typing import Optional def remove_nbsp(s: str | None) -> str | None: diff --git a/generalresearch/models/thl/contest/contest.py b/generalresearch/models/thl/contest/contest.py index 8c45f13..2a8853d 100644 --- a/generalresearch/models/thl/contest/contest.py +++ b/generalresearch/models/thl/contest/contest.py @@ -2,7 +2,7 @@ from __future__ import annotations import json from abc import ABC, abstractmethod -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from typing import Any, Self from uuid import uuid4 @@ -167,7 +167,7 @@ class Contest(ContestBase): ended_at=datetime.now(tz=UTC), end_reason=reason, ) - return None + return def model_dump_mysql(self, **kwargs) -> dict[str, Any]: d = self.model_dump(mode="json", **kwargs) diff --git a/generalresearch/models/thl/contest/contest_entry.py b/generalresearch/models/thl/contest/contest_entry.py index 31ef317..bb3aef4 100644 --- a/generalresearch/models/thl/contest/contest_entry.py +++ b/generalresearch/models/thl/contest/contest_entry.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from uuid import uuid4 from pydantic import ( diff --git a/generalresearch/models/thl/contest/examples.py b/generalresearch/models/thl/contest/examples.py index 1810e63..018f090 100644 --- a/generalresearch/models/thl/contest/examples.py +++ b/generalresearch/models/thl/contest/examples.py @@ -93,7 +93,6 @@ def _example_raffle(schema: dict) -> None: product_id=EXAMPLE_PRODUCT_ID, ).model_dump(mode="json") - return None def _example_raffle_user_view(schema: dict[str, Any]) -> None: @@ -144,7 +143,6 @@ def _example_raffle_user_view(schema: dict[str, Any]) -> None: product_user_id="test-user", ).model_dump(mode="json") - return None def _example_milestone_create(schema: dict[str, Any]) -> None: @@ -185,7 +183,6 @@ def _example_milestone_create(schema: dict[str, Any]) -> None: terms_and_conditions=HttpUrl("https://www.example.com"), ).model_dump(mode="json") - return None def _example_milestone(schema: dict[str, Any]) -> None: @@ -231,7 +228,6 @@ def _example_milestone(schema: dict[str, Any]) -> None: win_count=12, ).model_dump(mode="json") - return None def _example_milestone_user_view(schema: dict[str, Any]) -> None: @@ -277,7 +273,6 @@ def _example_milestone_user_view(schema: dict[str, Any]) -> None: product_user_id="test-user", ).model_dump(mode="json") - return None def _example_leaderboard_contest_create(schema: dict[str, Any]) -> None: @@ -322,7 +317,6 @@ def _example_leaderboard_contest_create(schema: dict[str, Any]) -> None: leaderboard_key=f"leaderboard:{EXAMPLE_PRODUCT_ID}:us:weekly:2025-05-26:complete_count", ).model_dump(mode="json") - return None def _example_leaderboard_contest(schema: dict[str, Any]) -> None: @@ -368,7 +362,6 @@ def _example_leaderboard_contest(schema: dict[str, Any]) -> None: product_id=EXAMPLE_PRODUCT_ID, ).model_dump(mode="json") - return None def _example_leaderboard_contest_user_view(schema: dict[str, Any]) -> None: diff --git a/generalresearch/models/thl/contest/io.py b/generalresearch/models/thl/contest/io.py index c6af719..d11080f 100644 --- a/generalresearch/models/thl/contest/io.py +++ b/generalresearch/models/thl/contest/io.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from uuid import uuid4 from generalresearch.models.thl.contest.definitions import ContestType diff --git a/generalresearch/models/thl/contest/leaderboard.py b/generalresearch/models/thl/contest/leaderboard.py index c5e0626..696cdea 100644 --- a/generalresearch/models/thl/contest/leaderboard.py +++ b/generalresearch/models/thl/contest/leaderboard.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from typing import Any, Literal, Self from pydantic import ( diff --git a/generalresearch/models/thl/contest/milestone.py b/generalresearch/models/thl/contest/milestone.py index 8b74d50..8d96fcb 100644 --- a/generalresearch/models/thl/contest/milestone.py +++ b/generalresearch/models/thl/contest/milestone.py @@ -2,7 +2,7 @@ from __future__ import annotations import logging from datetime import timedelta -from typing import Any, Literal +from typing import Any, Literal, Self from pydantic import ( BaseModel, @@ -10,7 +10,6 @@ from pydantic import ( Field, PositiveInt, ) -from typing import Self from generalresearch.models.custom_types import AwareDatetimeISO from generalresearch.models.thl.contest.contest import ( diff --git a/generalresearch/models/thl/contest/raffle.py b/generalresearch/models/thl/contest/raffle.py index b497a44..08243f4 100644 --- a/generalresearch/models/thl/contest/raffle.py +++ b/generalresearch/models/thl/contest/raffle.py @@ -3,7 +3,7 @@ from __future__ import annotations import logging import random from collections import defaultdict -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from typing import Any, Literal, Self from pydantic import ( diff --git a/generalresearch/models/thl/finance.py b/generalresearch/models/thl/finance.py index a992f78..0856825 100644 --- a/generalresearch/models/thl/finance.py +++ b/generalresearch/models/thl/finance.py @@ -1,7 +1,7 @@ from __future__ import annotations import random -from datetime import UTC, timezone +from datetime import UTC from typing import TYPE_CHECKING from uuid import uuid4 @@ -28,8 +28,7 @@ payout_example = random.randint(150, 750 * 100) adjustment_example = random.randint(-1_000, 50 * 100) if TYPE_CHECKING: - from generalresearch.managers.thl.product import ProductManager - from generalresearch.models.thl.ledger import AccountType, Direction, LedgerAccount + from generalresearch.models.thl.ledger import LedgerAccount class AdjustmentType(BaseModel): diff --git a/generalresearch/models/thl/ipinfo.py b/generalresearch/models/thl/ipinfo.py index 3f212cf..e327bae 100644 --- a/generalresearch/models/thl/ipinfo.py +++ b/generalresearch/models/thl/ipinfo.py @@ -1,7 +1,7 @@ from __future__ import annotations import ipaddress -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from typing import Any, Literal, Self from faker import Faker diff --git a/generalresearch/models/thl/ledger_example.py b/generalresearch/models/thl/ledger_example.py index 92ad83d..0291691 100644 --- a/generalresearch/models/thl/ledger_example.py +++ b/generalresearch/models/thl/ledger_example.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from typing import Any from uuid import uuid4 diff --git a/generalresearch/models/thl/offerwall/cache.py b/generalresearch/models/thl/offerwall/cache.py index 82ab36d..97546b2 100644 --- a/generalresearch/models/thl/offerwall/cache.py +++ b/generalresearch/models/thl/offerwall/cache.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from typing import Any from pydantic import BaseModel, Field diff --git a/generalresearch/models/thl/payout_format.py b/generalresearch/models/thl/payout_format.py index a7fc4fb..4d616b6 100644 --- a/generalresearch/models/thl/payout_format.py +++ b/generalresearch/models/thl/payout_format.py @@ -45,7 +45,7 @@ def format_payout_format(payout_format: str, payout_int: int) -> str: try: xform, formatstr = inside.split(":") - except ValueError as e: + except ValueError: raise ValueError( "Payout format string must contain ':' to distinguish between transformations and formatting." ) @@ -61,17 +61,17 @@ def format_payout_format(payout_format: str, payout_int: int) -> str: payout = decimal.Decimal(eval(xform, {"payout": payout_int})) - except NameError as e: + except NameError: raise ValueError("Payout format string must contain 'payout' variable.") - except ZeroDivisionError as e: + except ZeroDivisionError: raise ValueError("Cannot divide by zero.") - except TypeError as e: + except TypeError: # "{payout()*1:}" - TypeError: 'int' object is not callable raise ValueError("Invalid type reference.") - except Exception as e: - raise ValueError(f"Invalid payout transformation") + except Exception: + raise ValueError("Invalid payout transformation") formatstr = f"{{:{formatstr}}}" diff --git a/generalresearch/models/thl/product.py b/generalresearch/models/thl/product.py index 5b7ab9a..76a8e83 100644 --- a/generalresearch/models/thl/product.py +++ b/generalresearch/models/thl/product.py @@ -77,7 +77,6 @@ if TYPE_CHECKING: from generalresearch.models.thl.payout import ( BrokerageProductPayoutEvent, ) - from generalresearch.models.thl.user import User # fmt: off @@ -1088,7 +1087,6 @@ class Product(BaseModel, validate_assignment=True): from generalresearch.incite.schemas.mergers.pop_ledger import ( numerical_col_names, ) - from generalresearch.models.thl.ledger import LedgerAccount account: LedgerAccount = thl_lm.get_account_or_create_bp_wallet(product=self) assert self.id == account.reference_uuid diff --git a/generalresearch/models/thl/profiling/marketplace.py b/generalresearch/models/thl/profiling/marketplace.py index 9038cf6..ad4ce80 100644 --- a/generalresearch/models/thl/profiling/marketplace.py +++ b/generalresearch/models/thl/profiling/marketplace.py @@ -1,7 +1,7 @@ from __future__ import annotations from abc import ABC, abstractmethod -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from functools import cached_property from typing import Any diff --git a/generalresearch/models/thl/profiling/upk_question.py b/generalresearch/models/thl/profiling/upk_question.py index 08ea350..77bba6f 100644 --- a/generalresearch/models/thl/profiling/upk_question.py +++ b/generalresearch/models/thl/profiling/upk_question.py @@ -367,7 +367,7 @@ class UpkQuestion(BaseModel): self.choices is None ), f"No `choices` are allowed for type `{self.type}`" else: - assert self.choices is not None, f"`choices` must be set" + assert self.choices is not None, "`choices` must be set" return self @model_validator(mode="after") diff --git a/generalresearch/models/thl/profiling/upk_question_answer.py b/generalresearch/models/thl/profiling/upk_question_answer.py index c59d99d..d8323ad 100644 --- a/generalresearch/models/thl/profiling/upk_question_answer.py +++ b/generalresearch/models/thl/profiling/upk_question_answer.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from typing import Any, Self from uuid import uuid4 diff --git a/generalresearch/models/thl/profiling/user_question_answer.py b/generalresearch/models/thl/profiling/user_question_answer.py index a55b205..2db07b7 100644 --- a/generalresearch/models/thl/profiling/user_question_answer.py +++ b/generalresearch/models/thl/profiling/user_question_answer.py @@ -2,7 +2,7 @@ from __future__ import annotations import json from collections.abc import Iterator -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from typing import Any, Literal, Self from pydantic import ( diff --git a/generalresearch/models/thl/session.py b/generalresearch/models/thl/session.py index c8e681c..fe7194a 100644 --- a/generalresearch/models/thl/session.py +++ b/generalresearch/models/thl/session.py @@ -2,7 +2,7 @@ from __future__ import annotations import json import logging -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal from typing import TYPE_CHECKING, Annotated, Any, Self from uuid import uuid4 diff --git a/generalresearch/models/thl/survey/__init__.py b/generalresearch/models/thl/survey/__init__.py index bc3d699..6e6b475 100644 --- a/generalresearch/models/thl/survey/__init__.py +++ b/generalresearch/models/thl/survey/__init__.py @@ -111,7 +111,6 @@ class MarketplaceTask(BaseModel, ABC): """ The Condition Model for this survey class """ - pass @property @abstractmethod @@ -119,7 +118,6 @@ class MarketplaceTask(BaseModel, ABC): """ The age question ID """ - pass @property @abstractmethod @@ -129,7 +127,6 @@ class MarketplaceTask(BaseModel, ABC): """ Mapping of generic Gender to the marketplace condition for that gender """ - pass @property def marketplace_age_groups( diff --git a/generalresearch/models/thl/survey/buyer.py b/generalresearch/models/thl/survey/buyer.py index 384bab4..b888007 100644 --- a/generalresearch/models/thl/survey/buyer.py +++ b/generalresearch/models/thl/survey/buyer.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from decimal import Decimal from math import log from typing import Annotated diff --git a/generalresearch/models/thl/survey/model.py b/generalresearch/models/thl/survey/model.py index 57bcbe2..2eed8f7 100644 --- a/generalresearch/models/thl/survey/model.py +++ b/generalresearch/models/thl/survey/model.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from decimal import Decimal from typing import Annotated, Any diff --git a/generalresearch/models/thl/survey/penalty.py b/generalresearch/models/thl/survey/penalty.py index 05153fe..04f8e20 100644 --- a/generalresearch/models/thl/survey/penalty.py +++ b/generalresearch/models/thl/survey/penalty.py @@ -1,7 +1,7 @@ from __future__ import annotations import abc -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from typing import Annotated, Literal from pydantic import BaseModel, ConfigDict, Field, TypeAdapter diff --git a/generalresearch/models/thl/task_adjustment.py b/generalresearch/models/thl/task_adjustment.py index 1834898..27c47d4 100644 --- a/generalresearch/models/thl/task_adjustment.py +++ b/generalresearch/models/thl/task_adjustment.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from decimal import Decimal from uuid import uuid4 diff --git a/generalresearch/models/thl/user.py b/generalresearch/models/thl/user.py index 55bbd18..11e0d67 100644 --- a/generalresearch/models/thl/user.py +++ b/generalresearch/models/thl/user.py @@ -3,7 +3,7 @@ from __future__ import annotations import json import logging import re -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from typing import TYPE_CHECKING, Annotated, Self from uuid import UUID, uuid4 diff --git a/generalresearch/models/thl/user_iphistory.py b/generalresearch/models/thl/user_iphistory.py index 469f8ba..257c41a 100644 --- a/generalresearch/models/thl/user_iphistory.py +++ b/generalresearch/models/thl/user_iphistory.py @@ -1,7 +1,7 @@ from __future__ import annotations import ipaddress -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from typing import Self from faker import Faker @@ -204,7 +204,6 @@ class UserIPHistory(BaseModel): if res.get(x.ip): x.information = res[x.ip] - return None def collapse_ip_records(self): """ diff --git a/generalresearch/models/thl/wallet/payout.py b/generalresearch/models/thl/wallet/payout.py index cbb37fe..42530b3 100644 --- a/generalresearch/models/thl/wallet/payout.py +++ b/generalresearch/models/thl/wallet/payout.py @@ -2,7 +2,7 @@ from __future__ import annotations import json from collections.abc import Collection -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from typing import Any from uuid import uuid4 diff --git a/generalresearch/pg_helper.py b/generalresearch/pg_helper.py index b5e124a..1d5d30b 100644 --- a/generalresearch/pg_helper.py +++ b/generalresearch/pg_helper.py @@ -1,12 +1,11 @@ from __future__ import annotations -from datetime import UTC, timezone +from datetime import UTC import psycopg -from psycopg.adapt import Buffer from psycopg.rows import RowFactory, dict_row from psycopg.types.datetime import TimestampLoader -from psycopg.types.net import Address, InetLoader, Interface +from psycopg.types.net import InetLoader from psycopg.types.string import TextLoader from psycopg.types.uuid import UUIDLoader from pydantic import PostgresDsn @@ -103,10 +102,9 @@ class PostgresConfig: # This is only intended for SELECT queries assert "SELECT" in query.upper(), "Supports SELECTs only" - with self.make_connection() as conn: - with conn.cursor() as c: - c.execute(query=query, params=params) - return c.fetchall() + with self.make_connection() as conn, conn.cursor() as c: + c.execute(query=query, params=params) + return c.fetchall() def execute_write(self, query, params=None) -> int: cmd = query.lstrip().upper() diff --git a/generalresearch/schemas/survey_stats.py b/generalresearch/schemas/survey_stats.py index 0d509e1..b3acf34 100644 --- a/generalresearch/schemas/survey_stats.py +++ b/generalresearch/schemas/survey_stats.py @@ -70,7 +70,7 @@ UnitInterval = Column( SID_CHECKS = [ Check.str_length(min_value=3, max_value=67), - Check.str_matches("^[a-z]{1,2}\:[A-Za-z0-9]+"), + Check.str_matches(r"^[a-z]{1,2}\:[A-Za-z0-9]+"), Check( lambda x: len(set(x.str.split(":").str[0])) == 1, error="the sources must all be the same", diff --git a/generalresearch/sql_helper.py b/generalresearch/sql_helper.py index 2bbdbe5..ea45305 100644 --- a/generalresearch/sql_helper.py +++ b/generalresearch/sql_helper.py @@ -1,7 +1,7 @@ from __future__ import annotations import logging -from typing import Any, Optional +from typing import Any from uuid import UUID from pydantic import MariaDBDsn, MySQLDsn, PostgresDsn @@ -39,14 +39,7 @@ class SqlConnector: # I'm intentionally doing a match case here so that we'll make sure # we can NOT use this on old versions of python 😈 - if "mysql" in self.dsn.scheme: - import pymysql as engine_module - - self.engine_module = engine_module - self.cursor_class = engine_module.cursors.DictCursor - self.quote_char = "`" - - elif "maria" in self.dsn.scheme: + if "mysql" in self.dsn.scheme or "maria" in self.dsn.scheme: import pymysql as engine_module self.engine_module = engine_module @@ -199,7 +192,6 @@ class SqlHelper(SqlConnector): if cursor is None: c.connection.commit() - return None def bulk_update( self, @@ -233,7 +225,7 @@ class SqlHelper(SqlConnector): if cursor is None: c.connection.commit() - return None + return def get_or_create( self, @@ -253,7 +245,7 @@ class SqlHelper(SqlConnector): lookup_fns = ",".join( ["`" + x + "`" for x in set(lookup_dict.keys()) | {primary_key}] ) - lookup_vals = " AND ".join([f"`{fn}`=%({fn})s" for fn in lookup_dict.keys()]) + lookup_vals = " AND ".join([f"`{fn}`=%({fn})s" for fn in lookup_dict]) table_name_str = self._quote(table_name) query = f"SELECT {lookup_fns} FROM {table_name_str} WHERE {lookup_vals} LIMIT 2" if cursor is None: @@ -293,7 +285,7 @@ class SqlHelper(SqlConnector): else: c = cursor field_names = ",".join(map(self._quote, create_dict)) - vals = ",".join([f"%({fn})s" for fn in create_dict.keys()]) + vals = ",".join([f"%({fn})s" for fn in create_dict]) table_name_str = self._quote(table_name) query = f"INSERT INTO {table_name_str} ({field_names}) VALUES ({vals})" c.execute(query, create_dict) @@ -352,4 +344,3 @@ class SqlHelper(SqlConnector): if cursor is None: c.connection.commit() - return None diff --git a/generalresearch/thl_django/apps.py b/generalresearch/thl_django/apps.py index 2813947..bd87110 100644 --- a/generalresearch/thl_django/apps.py +++ b/generalresearch/thl_django/apps.py @@ -6,11 +6,11 @@ class THLSchemaConfig(AppConfig): label = "thl_django" def ready(self): - from .accounting import models # noqa: F401 # pycharm: keep - from .common import models # noqa: F401 # pycharm: keep - from .contest import models # noqa: F401 # pycharm: keep - from .event import models # noqa: F401 # pycharm: keep - from .marketplace import models # noqa: F401 # pycharm: keep - from .network import models # noqa: F401 # pycharm: keep - from .userhealth import models # noqa: F401 # pycharm: keep + from .accounting import models # pycharm: keep + from .common import models # pycharm: keep + from .contest import models # pycharm: keep + from .event import models # pycharm: keep + from .marketplace import models # pycharm: keep + from .network import models # pycharm: keep + from .userhealth import models # pycharm: keep from .userprofile import models # noqa: F401 # pycharm: keep diff --git a/generalresearch/thl_django/fields.py b/generalresearch/thl_django/fields.py index 5e40ef0..251faa5 100644 --- a/generalresearch/thl_django/fields.py +++ b/generalresearch/thl_django/fields.py @@ -1,6 +1,7 @@ -from django.db import models import ipaddress +from django.db import models + class CIDRField(models.Field): description = "PostgreSQL CIDR network" diff --git a/generalresearch/thl_django/migrations/0001_initial.py b/generalresearch/thl_django/migrations/0001_initial.py index ecae35a..cf147e3 100644 --- a/generalresearch/thl_django/migrations/0001_initial.py +++ b/generalresearch/thl_django/migrations/0001_initial.py @@ -1,7 +1,8 @@ # Generated by Django 6.0 on 2025-12-26 20:53 -import django.db.models.deletion import uuid + +import django.db.models.deletion from django.db import migrations, models diff --git a/generalresearch/thl_django/migrations/0002_surveystat_is_live_alter_surveycategory_strength_and_more.py b/generalresearch/thl_django/migrations/0002_surveystat_is_live_alter_surveycategory_strength_and_more.py index 211c48a..f767afc 100644 --- a/generalresearch/thl_django/migrations/0002_surveystat_is_live_alter_surveycategory_strength_and_more.py +++ b/generalresearch/thl_django/migrations/0002_surveystat_is_live_alter_surveycategory_strength_and_more.py @@ -1,7 +1,7 @@ # Generated by Django 6.0 on 2025-12-28 16:49 -from django.db import migrations, models from django.contrib.postgres.operations import AddIndexConcurrently +from django.db import migrations, models class Migration(migrations.Migration): diff --git a/generalresearch/thl_django/migrations/0003_remove_surveystat_surveystat_live_survey_idx_and_more.py b/generalresearch/thl_django/migrations/0003_remove_surveystat_surveystat_live_survey_idx_and_more.py index ecaf0a9..dcf9ef2 100644 --- a/generalresearch/thl_django/migrations/0003_remove_surveystat_surveystat_live_survey_idx_and_more.py +++ b/generalresearch/thl_django/migrations/0003_remove_surveystat_surveystat_live_survey_idx_and_more.py @@ -1,10 +1,10 @@ # Generated by Django 6.0 on 2025-12-29 21:22 -from django.db import migrations, models from django.contrib.postgres.operations import ( AddIndexConcurrently, RemoveIndexConcurrently, ) +from django.db import migrations, models class Migration(migrations.Migration): diff --git a/generalresearch/thl_django/migrations/0006_remove_thlsession_thl_session_status_d578b7_idx_and_more.py b/generalresearch/thl_django/migrations/0006_remove_thlsession_thl_session_status_d578b7_idx_and_more.py index e2492ab..64338c8 100644 --- a/generalresearch/thl_django/migrations/0006_remove_thlsession_thl_session_status_d578b7_idx_and_more.py +++ b/generalresearch/thl_django/migrations/0006_remove_thlsession_thl_session_status_d578b7_idx_and_more.py @@ -1,7 +1,7 @@ # Generated by Django 6.0 on 2026-01-02 17:38 -from django.db import migrations from django.contrib.postgres.operations import RemoveIndexConcurrently +from django.db import migrations class Migration(migrations.Migration): diff --git a/generalresearch/thl_django/migrations/0009_toolrun_mtrhop_portscanport_iplabel_mtr_portscan_and_more.py b/generalresearch/thl_django/migrations/0009_toolrun_mtrhop_portscanport_iplabel_mtr_portscan_and_more.py index e19a353..e8ac2c2 100644 --- a/generalresearch/thl_django/migrations/0009_toolrun_mtrhop_portscanport_iplabel_mtr_portscan_and_more.py +++ b/generalresearch/thl_django/migrations/0009_toolrun_mtrhop_portscanport_iplabel_mtr_portscan_and_more.py @@ -1,13 +1,14 @@ # Generated by Django 6.0 on 2026-03-15 20:17 +import uuid + import django.contrib.postgres.indexes import django.db.models.deletion import django.utils.timezone from django.contrib.postgres.operations import CreateExtension +from django.db import migrations, models import generalresearch.thl_django.fields -import uuid -from django.db import migrations, models class Migration(migrations.Migration): diff --git a/generalresearch/thl_django/network/models.py b/generalresearch/thl_django/network/models.py index 167af02..733c0ab 100644 --- a/generalresearch/thl_django/network/models.py +++ b/generalresearch/thl_django/network/models.py @@ -1,12 +1,11 @@ from uuid import uuid4 -from django.utils import timezone -from django.contrib.postgres.indexes import GistIndex, GinIndex +from django.contrib.postgres.indexes import GinIndex, GistIndex from django.db import models +from django.utils import timezone from generalresearch.thl_django.fields import CIDRField - ####### # ** Signals ** # ToolRun diff --git a/generalresearch/utils/enum.py b/generalresearch/utils/enum.py index e59a383..b4620b3 100644 --- a/generalresearch/utils/enum.py +++ b/generalresearch/utils/enum.py @@ -19,7 +19,7 @@ class ReprEnumMeta(EnumMeta): [f" - __{e.value}__ *({e.name})*: {descriptions[e.name]}" for e in self] ) else: - return f"\nAllowed values: \n" + "\n".join( + return "\nAllowed values: \n" + "\n".join( [f" - __{e.value}__ *({e.name})*: {descriptions[e.name]}" for e in self] ) @@ -35,7 +35,7 @@ class ReprEnumMeta(EnumMeta): [f" - __{e.name}__: {descriptions[e.name]}" for e in self] ) else: - return f"\nAllowed values: \n" + "\n".join( + return "\nAllowed values: \n" + "\n".join( [f" - __{e.name}__: {descriptions[e.name]}" for e in self] ) diff --git a/generalresearch/wall_status_codes/lucid.py b/generalresearch/wall_status_codes/lucid.py index 05b75fb..3cc1b5e 100644 --- a/generalresearch/wall_status_codes/lucid.py +++ b/generalresearch/wall_status_codes/lucid.py @@ -61,7 +61,7 @@ client_status_map: dict[str, StatusCode1] = { "35": StatusCode1.BUYER_QUALITY_FAIL, } -status_map = defaultdict(lambda: Status.FAIL, **{"s": Status.COMPLETE}) +status_map = defaultdict(lambda: Status.FAIL, s=Status.COMPLETE) status_codes_ext_map: dict[StatusCode1, list[str]] = { StatusCode1.COMPLETE: [], StatusCode1.BUYER_FAIL: ["3"], diff --git a/generalresearch/wall_status_codes/morning.py b/generalresearch/wall_status_codes/morning.py index bac9669..ffa6be2 100644 --- a/generalresearch/wall_status_codes/morning.py +++ b/generalresearch/wall_status_codes/morning.py @@ -52,7 +52,7 @@ short_code_to_status_codes_morning: dict[str, str] = { "sur_tim": "survey_timeout", "tem_ban": "temporarily_banned", } -status_map = defaultdict(lambda: Status.FAIL, **{"complete": Status.COMPLETE}) +status_map = defaultdict(lambda: Status.FAIL, complete=Status.COMPLETE) status_codes_ext_map: dict[StatusCode1, list[str]] = { StatusCode1.COMPLETE: ["complete"], diff --git a/generalresearch/wall_status_codes/pollfish.py b/generalresearch/wall_status_codes/pollfish.py index 1ae732d..a5c6e25 100644 --- a/generalresearch/wall_status_codes/pollfish.py +++ b/generalresearch/wall_status_codes/pollfish.py @@ -28,7 +28,7 @@ status_codes_map: dict[str, str] = { "su_al_ta": "survey_already_taken", "complete": "complete", } -status_map = defaultdict(lambda: Status.FAIL, **{"complete": Status.COMPLETE}) +status_map = defaultdict(lambda: Status.FAIL, complete=Status.COMPLETE) status_codes_ext_map: dict[StatusCode1, list[str]] = { StatusCode1.COMPLETE: ["complete"], StatusCode1.BUYER_FAIL: ["third_party_termination", "screenout"], diff --git a/generalresearch/wall_status_codes/precision.py b/generalresearch/wall_status_codes/precision.py index b033cf2..c7593cf 100644 --- a/generalresearch/wall_status_codes/precision.py +++ b/generalresearch/wall_status_codes/precision.py @@ -45,7 +45,7 @@ status_codes_precision: dict[str, str] = { "60": "Client Reject", "80": "Final Complete", } -status_map = defaultdict(lambda: Status.FAIL, **{"s": Status.COMPLETE}) +status_map = defaultdict(lambda: Status.FAIL, s=Status.COMPLETE) status_codes_ext_map: dict[StatusCode1, list[str]] = { StatusCode1.COMPLETE: ["10"], StatusCode1.BUYER_FAIL: ["20", "30"], diff --git a/generalresearch/wall_status_codes/repdata.py b/generalresearch/wall_status_codes/repdata.py index ad64338..24aab14 100644 --- a/generalresearch/wall_status_codes/repdata.py +++ b/generalresearch/wall_status_codes/repdata.py @@ -46,7 +46,7 @@ rd_threat_name: dict[str, str] = { "18": "MaxMind Failure", } -status_map = defaultdict(lambda: Status.FAIL, **{"complete": Status.COMPLETE}) +status_map = defaultdict(lambda: Status.FAIL, complete=Status.COMPLETE) status_code_map: dict[StatusCode1, list[str]] = { StatusCode1.COMPLETE: ["1000"], StatusCode1.BUYER_FAIL: ["2000", "4000"], diff --git a/test_utils/conftest.py b/test_utils/conftest.py index ffe458c..03e9305 100644 --- a/test_utils/conftest.py +++ b/test_utils/conftest.py @@ -5,9 +5,8 @@ import shutil import stat import subprocess import sys -import tempfile from collections.abc import Callable, Generator -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from os.path import join as pjoin from pathlib import Path from uuid import uuid4 diff --git a/test_utils/grliq/conftest.py b/test_utils/grliq/conftest.py index 7665b52..9a3bc56 100644 --- a/test_utils/grliq/conftest.py +++ b/test_utils/grliq/conftest.py @@ -1,7 +1,7 @@ from __future__ import annotations from collections.abc import Callable -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from uuid import uuid4 import pytest diff --git a/test_utils/incite/conftest.py b/test_utils/incite/conftest.py index 87ea7ae..2968d18 100644 --- a/test_utils/incite/conftest.py +++ b/test_utils/incite/conftest.py @@ -1,7 +1,7 @@ from __future__ import annotations from collections.abc import Callable -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from os.path import join as pjoin from pathlib import Path from random import choice as randchoice diff --git a/test_utils/managers/conftest.py b/test_utils/managers/conftest.py index 4dacb29..2a9ea00 100644 --- a/test_utils/managers/conftest.py +++ b/test_utils/managers/conftest.py @@ -20,14 +20,12 @@ from generalresearch.managers.thl.ipinfo import ( IPGeonameManager, IPInformationManager, ) -from generalresearch.managers.thl.profiling.uqa import UQAManager from generalresearch.managers.thl.userhealth import ( AuditLogManager, IPRecordManager, UserIpHistoryManager, ) from generalresearch.models import Source -from generalresearch.models.thl.user import User from generalresearch.pg_helper import PostgresConfig from generalresearch.redis_helper import RedisConfig from generalresearch.sql_helper import SqlHelper diff --git a/test_utils/models/conftest.py b/test_utils/models/conftest.py index 3a9e45c..5570b40 100644 --- a/test_utils/models/conftest.py +++ b/test_utils/models/conftest.py @@ -1,7 +1,7 @@ from __future__ import annotations from collections.abc import Callable -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal from random import choice as randchoice from random import randint @@ -19,17 +19,15 @@ from generalresearch.models.thl.definitions import ( ) from generalresearch.models.thl.survey.model import Buyer, Survey from generalresearch.pg_helper import PostgresConfig -from generalresearch.redis_helper import RedisConfig if TYPE_CHECKING: from generalresearch.currency import USDCent - from generalresearch.managers.gr.authentication import GRTokenManager, GRUserManager from generalresearch.managers.gr.business import ( BusinessAddressManager, BusinessBankAccountManager, BusinessManager, ) - from generalresearch.managers.gr.team import MembershipManager, TeamManager + from generalresearch.managers.gr.team import TeamManager from generalresearch.managers.thl.buyer import BuyerManager from generalresearch.managers.thl.ipinfo import ( IPGeonameManager, @@ -45,13 +43,12 @@ if TYPE_CHECKING: from generalresearch.managers.thl.user_manager.user_manager import UserManager from generalresearch.managers.thl.userhealth import AuditLogManager, IPRecordManager from generalresearch.managers.thl.wall import WallManager - from generalresearch.models.gr.authentication import GRToken, GRUser from generalresearch.models.gr.business import ( Business, BusinessAddress, BusinessBankAccount, ) - from generalresearch.models.gr.team import Membership, Team + from generalresearch.models.gr.team import Team from generalresearch.models.thl.ipinfo import IPGeoname, IPInformation from generalresearch.models.thl.payout import ( BrokerageProductPayoutEvent, diff --git a/test_utils/models/contest/conftest.py b/test_utils/models/contest/conftest.py index 0060946..b14e126 100644 --- a/test_utils/models/contest/conftest.py +++ b/test_utils/models/contest/conftest.py @@ -1,7 +1,7 @@ from __future__ import annotations from collections.abc import Callable -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from decimal import Decimal from uuid import uuid4 diff --git a/test_utils/models/network/conftest.py b/test_utils/models/network/conftest.py index cabd8dc..6ba37a3 100644 --- a/test_utils/models/network/conftest.py +++ b/test_utils/models/network/conftest.py @@ -1,5 +1,5 @@ import os -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from uuid import uuid4 import pytest diff --git a/test_utils/models/thl/conftest.py b/test_utils/models/thl/conftest.py index a2adcce..5b21c6b 100644 --- a/test_utils/models/thl/conftest.py +++ b/test_utils/models/thl/conftest.py @@ -1,7 +1,7 @@ from __future__ import annotations from collections.abc import Callable -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from decimal import ROUND_DOWN, Decimal from random import choice as rand_choice from random import choice as rchoice diff --git a/test_utils/spectrum/conftest.py b/test_utils/spectrum/conftest.py index 9c067d3..eb2e289 100644 --- a/test_utils/spectrum/conftest.py +++ b/test_utils/spectrum/conftest.py @@ -1,6 +1,6 @@ import logging import time -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from typing import TYPE_CHECKING import pytest diff --git a/tests/incite/collections/test_df_collection_base.py b/tests/incite/collections/test_df_collection_base.py index c3c64e4..176bf4b 100644 --- a/tests/incite/collections/test_df_collection_base.py +++ b/tests/incite/collections/test_df_collection_base.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from typing import TYPE_CHECKING import pandas as pd @@ -9,7 +9,6 @@ from generalresearch.incite.collections import ( DFCollection, DFCollectionType, ) -from test_utils.incite.conftest import mnt_filepath if TYPE_CHECKING: from generalresearch.incite.base import GRLDatasets diff --git a/tests/incite/collections/test_df_collection_item_base.py b/tests/incite/collections/test_df_collection_item_base.py index 8cf719d..0218f30 100644 --- a/tests/incite/collections/test_df_collection_item_base.py +++ b/tests/incite/collections/test_df_collection_item_base.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from typing import TYPE_CHECKING import pytest diff --git a/tests/incite/collections/test_df_collection_item_thl_web.py b/tests/incite/collections/test_df_collection_item_thl_web.py index 062171d..3d70e56 100644 --- a/tests/incite/collections/test_df_collection_item_thl_web.py +++ b/tests/incite/collections/test_df_collection_item_thl_web.py @@ -1,7 +1,7 @@ from __future__ import annotations from collections.abc import Callable, Generator -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from itertools import product as iter_product from os.path import join as pjoin from pathlib import Path, PurePath @@ -236,7 +236,6 @@ class TestDFCollectionItemMethod: incite_item_factory, delete_df_collection, ): - from generalresearch.models.thl.user import User if df_collection.data_type in unsupported_mock_types: return @@ -280,7 +279,6 @@ class TestDFCollectionItemMethod: incite_item_factory, delete_df_collection, ): - from generalresearch.models.thl.user import User if df_collection.data_type in unsupported_mock_types: return @@ -333,7 +331,6 @@ class TestDFCollectionItemMethod: delete_df_collection, mnt_filepath, ): - from generalresearch.models.thl.user import User if df_collection.data_type != DFCollectionType.LEDGER: return @@ -385,7 +382,6 @@ class TestDFCollectionItemMethod: delete_df_collection, mnt_filepath, ): - from generalresearch.models.thl.user import User if df_collection.data_type in unsupported_mock_types: return @@ -766,7 +762,6 @@ class TestDFCollectionItemFunctionalTest: delete_df_collection, mnt_filepath: GRLDatasets, ): - from generalresearch.models.thl.user import User if df_collection.data_type in unsupported_mock_types: return @@ -828,7 +823,6 @@ class TestDFCollectionItemFunctionalTest: import pyarrow.parquet as pq - from generalresearch.models.thl.user import User if df_collection.data_type in unsupported_mock_types: return @@ -864,7 +858,6 @@ class TestDFCollectionItemFunctionalTest: delete_df_collection, mnt_filepath: GRLDatasets, ): - from generalresearch.models.thl.user import User delete_df_collection(coll=df_collection) df_collection._client = client_no_amm @@ -920,7 +913,6 @@ class TestDFCollectionItemFunctionalTest: """Don't allow creating an archive for data that will likely be overwritten or updated """ - from generalresearch.models.thl.user import User if df_collection.data_type in unsupported_mock_types: return @@ -987,7 +979,6 @@ class TestDFCollectionItemFunctionalTest: duration: timedelta, mnt_filepath, ): - from generalresearch.models.thl.user import User delete_df_collection(coll=df_collection) df_collection._client = client_no_amm diff --git a/tests/incite/collections/test_df_collection_thl_marketplaces.py b/tests/incite/collections/test_df_collection_thl_marketplaces.py index 2597d38..d2d3ce4 100644 --- a/tests/incite/collections/test_df_collection_thl_marketplaces.py +++ b/tests/incite/collections/test_df_collection_thl_marketplaces.py @@ -1,6 +1,5 @@ -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from itertools import product -from typing import TYPE_CHECKING import pytest from pandera.pandas import Column, DataFrameSchema, Index @@ -12,10 +11,6 @@ from generalresearch.incite.collections.thl_marketplaces import ( SagoSurveyHistoryCollection, SpectrumSurveyTimeseriesCollection, ) -from test_utils.incite.conftest import mnt_filepath - -if TYPE_CHECKING: - from generalresearch.incite.base import GRLDatasets def combo_object(): diff --git a/tests/incite/collections/test_df_collection_thl_web.py b/tests/incite/collections/test_df_collection_thl_web.py index 2cb0ba0..bcdeb83 100644 --- a/tests/incite/collections/test_df_collection_thl_web.py +++ b/tests/incite/collections/test_df_collection_thl_web.py @@ -13,9 +13,7 @@ from pandera.pandas import DataFrameSchema from generalresearch.incite.collections import DFCollection, DFCollectionType if TYPE_CHECKING: - from generalresearch.incite.base import GRLDatasets from generalresearch.incite.collections import ( - DFCollectionItem, DFCollectionType, ) diff --git a/tests/incite/mergers/foundations/test_enriched_session.py b/tests/incite/mergers/foundations/test_enriched_session.py index a0ae01e..ba11725 100644 --- a/tests/incite/mergers/foundations/test_enriched_session.py +++ b/tests/incite/mergers/foundations/test_enriched_session.py @@ -1,7 +1,6 @@ -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal from itertools import product -from typing import Optional import dask.dataframe as dd import pandas as pd @@ -11,10 +10,6 @@ from generalresearch.incite.schemas.admin_responses import ( AdminPOPSessionSchema, ) from generalresearch.pg_helper import PostgresConfig -from test_utils.incite.collections.conftest import ( - session_collection, - wall_collection, -) @pytest.mark.parametrize( diff --git a/tests/incite/mergers/foundations/test_enriched_task_adjust.py b/tests/incite/mergers/foundations/test_enriched_task_adjust.py index 96c214f..8c3a647 100644 --- a/tests/incite/mergers/foundations/test_enriched_task_adjust.py +++ b/tests/incite/mergers/foundations/test_enriched_task_adjust.py @@ -5,13 +5,6 @@ import dask.dataframe as dd import pandas as pd import pytest -from test_utils.incite.collections.conftest import ( - wall_collection, - task_adj_collection, - session_collection, -) -from test_utils.incite.mergers.conftest import enriched_wall_merge - @pytest.mark.parametrize( argnames="offset, duration,", diff --git a/tests/incite/mergers/foundations/test_enriched_wall.py b/tests/incite/mergers/foundations/test_enriched_wall.py index b421df8..0e28bce 100644 --- a/tests/incite/mergers/foundations/test_enriched_wall.py +++ b/tests/incite/mergers/foundations/test_enriched_wall.py @@ -1,34 +1,15 @@ -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal from itertools import product as iter_product -from typing import Optional import dask.dataframe as dd import pandas as pd import pytest # noinspection PyUnresolvedReferences -from distributed.utils_test import ( - cleanup, - client, - client_no_amm, - cluster_fixture, - gen_cluster, - loop, - loop_in_thread, -) - from generalresearch.incite.mergers.foundations.enriched_wall import ( EnrichedWallMergeItem, ) -from test_utils.incite.collections.conftest import ( - session_collection, - wall_collection, -) -from test_utils.incite.conftest import incite_item_factory -from test_utils.incite.mergers.conftest import ( - enriched_wall_merge, -) @pytest.mark.parametrize( @@ -118,7 +99,7 @@ class TestEnrichedWall: try: modified_time1 = path.stat().st_mtime - except (Exception,): + except Exception: modified_time1 = 0 item.build( diff --git a/tests/incite/mergers/foundations/test_user_id_product.py b/tests/incite/mergers/foundations/test_user_id_product.py index a696b45..10802e5 100644 --- a/tests/incite/mergers/foundations/test_user_id_product.py +++ b/tests/incite/mergers/foundations/test_user_id_product.py @@ -1,24 +1,13 @@ -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from itertools import product import pandas as pd import pytest # noinspection PyUnresolvedReferences -from distributed.utils_test import ( - cleanup, - client, - client_no_amm, - cluster_fixture, - gen_cluster, - loop, - loop_in_thread, -) - from generalresearch.incite.mergers.foundations.user_id_product import ( UserIdProductMergeItem, ) -from test_utils.incite.mergers.conftest import user_id_product_merge @pytest.mark.parametrize( @@ -51,7 +40,7 @@ class TestUserIDProduct: try: modified_time1 = path.stat().st_mtime - except (Exception,): + except Exception: modified_time1 = 0 user_id_product_merge.build(client=client_no_amm, user_coll=user_collection) diff --git a/tests/incite/mergers/test_merge_collection.py b/tests/incite/mergers/test_merge_collection.py index 77fa8c7..15fa4db 100644 --- a/tests/incite/mergers/test_merge_collection.py +++ b/tests/incite/mergers/test_merge_collection.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from itertools import product import pandas as pd @@ -9,7 +9,6 @@ from generalresearch.incite.mergers import ( MergeCollection, MergeType, ) -from test_utils.incite.conftest import mnt_filepath merge_types = list(e for e in MergeType if e != MergeType.TEST) diff --git a/tests/incite/mergers/test_merge_collection_item.py b/tests/incite/mergers/test_merge_collection_item.py index 96f8789..3d0b644 100644 --- a/tests/incite/mergers/test_merge_collection_item.py +++ b/tests/incite/mergers/test_merge_collection_item.py @@ -1,17 +1,10 @@ -from datetime import datetime, timezone, timedelta +from datetime import timedelta from itertools import product from pathlib import PurePath import pytest from generalresearch.incite.mergers import MergeCollectionItem, MergeType -from generalresearch.incite.mergers.foundations.enriched_session import ( - EnrichedSessionMerge, -) -from generalresearch.incite.mergers.foundations.enriched_wall import ( - EnrichedWallMerge, -) -from test_utils.incite.mergers.conftest import merge_collection @pytest.mark.parametrize( diff --git a/tests/incite/mergers/test_pop_ledger.py b/tests/incite/mergers/test_pop_ledger.py index 7583faf..dc01179 100644 --- a/tests/incite/mergers/test_pop_ledger.py +++ b/tests/incite/mergers/test_pop_ledger.py @@ -1,18 +1,12 @@ -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from itertools import product as iter_product -from typing import Optional import pandas as pd import pytest -from distributed.utils_test import client_no_amm from generalresearch.incite.schemas.mergers.pop_ledger import ( numerical_col_names, ) -from test_utils.incite.collections.conftest import ledger_collection -from test_utils.incite.conftest import incite_item_factory, mnt_filepath -from test_utils.incite.mergers.conftest import pop_ledger_merge -from test_utils.managers.ledger.conftest import create_main_accounts @pytest.mark.parametrize( diff --git a/tests/incite/mergers/test_ym_survey_merge.py b/tests/incite/mergers/test_ym_survey_merge.py index 9107f21..850df8a 100644 --- a/tests/incite/mergers/test_ym_survey_merge.py +++ b/tests/incite/mergers/test_ym_survey_merge.py @@ -1,25 +1,10 @@ -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from itertools import product import pandas as pd import pytest # noinspection PyUnresolvedReferences -from distributed.utils_test import ( - cleanup, - client, - client_no_amm, - cluster_fixture, - gen_cluster, - loop, - loop_in_thread, -) - -from test_utils.incite.collections.conftest import session_collection, wall_collection -from test_utils.incite.mergers.conftest import ( - enriched_session_merge, - ym_survey_wall_merge, -) @pytest.mark.parametrize( diff --git a/tests/incite/test_collection_base.py b/tests/incite/test_collection_base.py index 5a63019..7e1577a 100644 --- a/tests/incite/test_collection_base.py +++ b/tests/incite/test_collection_base.py @@ -10,7 +10,6 @@ import pytest from _pytest._code.code import ExceptionInfo from generalresearch.incite.base import CollectionBase -from test_utils.incite.conftest import mnt_filepath AGO_15min = (datetime.now(tz=UTC) - timedelta(minutes=15)).replace(microsecond=0) AGO_1HR = (datetime.now(tz=UTC) - timedelta(hours=1)).replace(microsecond=0) @@ -108,7 +107,7 @@ class TestCollectionBase: with pytest.raises(expected_exception=ValueError) as cm: cm: ExceptionInfo - CollectionBase(offset=f"59sec", archive_path=mnt_filepath.data_src) + CollectionBase(offset="59sec", archive_path=mnt_filepath.data_src) assert "Must be equal to, or longer than 1 min" in str(cm.value) with pytest.raises(expected_exception=ValueError) as cm: diff --git a/tests/incite/test_collection_base_item.py b/tests/incite/test_collection_base_item.py index 3f4d023..7a0a581 100644 --- a/tests/incite/test_collection_base_item.py +++ b/tests/incite/test_collection_base_item.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from os.path import join as pjoin from pathlib import Path from uuid import uuid4 diff --git a/tests/incite/test_grl_flow.py b/tests/incite/test_grl_flow.py index c632f9a..6aea182 100644 --- a/tests/incite/test_grl_flow.py +++ b/tests/incite/test_grl_flow.py @@ -1,15 +1,16 @@ class TestGRLFlow: def test_init(self, mnt_filepath, thl_web_rr): + from generalresearch.incite.collections.thl_web import ( + LedgerDFCollection, + TaskAdjustmentDFCollection, + ) from generalresearch.incite.defaults import ( ledger_df_collection, task_df_collection, - pop_ledger as plm, ) - - from generalresearch.incite.collections.thl_web import ( - LedgerDFCollection, - TaskAdjustmentDFCollection, + from generalresearch.incite.defaults import ( + pop_ledger as plm, ) from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge diff --git a/tests/incite/test_interval_idx.py b/tests/incite/test_interval_idx.py index ea2bced..3034c21 100644 --- a/tests/incite/test_interval_idx.py +++ b/tests/incite/test_interval_idx.py @@ -1,5 +1,6 @@ +from datetime import datetime + import pandas as pd -from datetime import datetime, timezone, timedelta class TestIntervalIndex: diff --git a/tests/managers/gr/test_authentication.py b/tests/managers/gr/test_authentication.py index 53b6931..b9f43a6 100644 --- a/tests/managers/gr/test_authentication.py +++ b/tests/managers/gr/test_authentication.py @@ -1,11 +1,9 @@ import logging -from random import randint from uuid import uuid4 import pytest from generalresearch.models.gr.authentication import GRUser -from test_utils.models.conftest import gr_user SSO_ISSUER = "" @@ -13,7 +11,6 @@ SSO_ISSUER = "" class TestGRUserManager: def test_create(self, gr_um): - from generalresearch.models.gr.authentication import GRUser user: GRUser = gr_um.create_dummy() instance = gr_um.get_by_id(user.id) diff --git a/tests/managers/gr/test_business.py b/tests/managers/gr/test_business.py index 7eb77f8..74a5450 100644 --- a/tests/managers/gr/test_business.py +++ b/tests/managers/gr/test_business.py @@ -2,8 +2,6 @@ from uuid import uuid4 import pytest -from test_utils.models.conftest import business - class TestBusinessBankAccountManager: @@ -12,8 +10,8 @@ class TestBusinessBankAccountManager: def test_create(self, business, business_bank_account_manager): from generalresearch.models.gr.business import ( - TransferMethod, BusinessBankAccount, + TransferMethod, ) instance = business_bank_account_manager.create( diff --git a/tests/managers/gr/test_team.py b/tests/managers/gr/test_team.py index 9215da4..0918ab8 100644 --- a/tests/managers/gr/test_team.py +++ b/tests/managers/gr/test_team.py @@ -1,7 +1,5 @@ from uuid import uuid4 -from test_utils.models.conftest import team - class TestMembershipManager: diff --git a/tests/managers/leaderboard.py b/tests/managers/leaderboard.py index 149bdbb..7773030 100644 --- a/tests/managers/leaderboard.py +++ b/tests/managers/leaderboard.py @@ -1,7 +1,7 @@ import os import time import zoneinfo -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from decimal import Decimal from uuid import uuid4 diff --git a/tests/managers/network/test_label.py b/tests/managers/network/test_label.py index 5b9a790..bfc7518 100644 --- a/tests/managers/network/test_label.py +++ b/tests/managers/network/test_label.py @@ -9,8 +9,8 @@ from generalresearch.managers.network.label import IPLabelManager from generalresearch.models.network.label import ( IPLabel, IPLabelKind, - IPLabelSource, IPLabelMetadata, + IPLabelSource, ) from generalresearch.models.thl.ipinfo import normalize_ip diff --git a/tests/managers/test_events.py b/tests/managers/test_events.py index 6941c00..cc9a1bf 100644 --- a/tests/managers/test_events.py +++ b/tests/managers/test_events.py @@ -1,11 +1,10 @@ import math import random import time -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal from functools import partial from math import floor -from typing import Optional from uuid import uuid4 import pytest @@ -384,7 +383,7 @@ class TestTaskStatsManager: "task_created_count_last_24h": AggregateBySource(total=0), } event_manager.set_source_task_stats( - source=Source.TESTING, live_task_count=0, live_tasks_max_payout=Decimal("0") + source=Source.TESTING, live_task_count=0, live_tasks_max_payout=Decimal(0) ) assert event_manager.get_task_stats_raw() == { "live_task_count": AggregateBySource( @@ -400,7 +399,7 @@ class TestTaskStatsManager: event_manager.set_source_task_stats( source=Source.TESTING, live_task_count=0, - live_tasks_max_payout=Decimal("0"), + live_tasks_max_payout=Decimal(0), created_count=10, ) res = event_manager.get_task_stats_raw() @@ -414,7 +413,7 @@ class TestTaskStatsManager: event_manager.set_source_task_stats( source=Source.TESTING, live_task_count=0, - live_tasks_max_payout=Decimal("0"), + live_tasks_max_payout=Decimal(0), created_count=10, ) res = event_manager.get_task_stats_raw() @@ -428,7 +427,7 @@ class TestTaskStatsManager: event_manager.set_source_task_stats( source=Source.TESTING2, live_task_count=0, - live_tasks_max_payout=Decimal("0"), + live_tasks_max_payout=Decimal(0), created_count=1, ) res = event_manager.get_task_stats_raw() @@ -481,7 +480,7 @@ class TestChannelsSubscriptions: wall = Wall( req_survey_id="a", - req_cpi=Decimal("1"), + req_cpi=Decimal(1), source=Source.TESTING, session_id=session.id, user_id=user.user_id, @@ -497,7 +496,7 @@ class TestChannelsSubscriptions: status=Status.COMPLETE, status_code_1=StatusCode1.COMPLETE, finished=datetime.now(tz=UTC), - cpi=Decimal("1"), + cpi=Decimal(1), ) event_manager.handle_task_finish(wall, session, user) msg = event_subscriber.get_next_message() diff --git a/tests/managers/test_userpid.py b/tests/managers/test_userpid.py index 4a3f699..36c2de9 100644 --- a/tests/managers/test_userpid.py +++ b/tests/managers/test_userpid.py @@ -1,11 +1,10 @@ import pytest from pydantic import MySQLDsn -from generalresearch.managers.marketplace.user_pid import UserPidMultiManager -from generalresearch.sql_helper import SqlHelper from generalresearch.managers.cint.user_pid import CintUserPidManager from generalresearch.managers.dynata.user_pid import DynataUserPidManager from generalresearch.managers.innovate.user_pid import InnovateUserPidManager +from generalresearch.managers.marketplace.user_pid import UserPidMultiManager from generalresearch.managers.morning.user_pid import MorningUserPidManager # from generalresearch.managers.precision import PrecisionUserPidManager @@ -13,6 +12,7 @@ from generalresearch.managers.prodege.user_pid import ProdegeUserPidManager from generalresearch.managers.repdata.user_pid import RepdataUserPidManager from generalresearch.managers.sago.user_pid import SagoUserPidManager from generalresearch.managers.spectrum.user_pid import SpectrumUserPidManager +from generalresearch.sql_helper import SqlHelper dsn = "" diff --git a/tests/managers/thl/test_contest/test_leaderboard.py b/tests/managers/thl/test_contest/test_leaderboard.py index 1a52f83..7adea9c 100644 --- a/tests/managers/thl/test_contest/test_leaderboard.py +++ b/tests/managers/thl/test_contest/test_leaderboard.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from zoneinfo import ZoneInfo from generalresearch.currency import USDCent @@ -12,12 +12,6 @@ from generalresearch.models.thl.contest.leaderboard import ( ) from generalresearch.models.thl.product import Product from generalresearch.models.thl.user import User -from test_utils.managers.contest.conftest import ( - leaderboard_contest_create as contest_create, -) -from test_utils.managers.contest.conftest import ( - leaderboard_contest_in_db as contest_in_db, -) class TestLeaderboardContestCRUD: diff --git a/tests/managers/thl/test_contest/test_milestone.py b/tests/managers/thl/test_contest/test_milestone.py index 66c5dc4..ed0bbb5 100644 --- a/tests/managers/thl/test_contest/test_milestone.py +++ b/tests/managers/thl/test_contest/test_milestone.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from generalresearch.models.thl.contest.definitions import ( ContestEndReason, @@ -12,18 +12,6 @@ from generalresearch.models.thl.contest.milestone import ( ) from generalresearch.models.thl.product import Product from generalresearch.models.thl.user import User -from test_utils.managers.contest.conftest import ( - milestone_contest as contest, -) -from test_utils.managers.contest.conftest import ( - milestone_contest_create as contest_create, -) -from test_utils.managers.contest.conftest import ( - milestone_contest_factory as contest_factory, -) -from test_utils.managers.contest.conftest import ( - milestone_contest_in_db as contest_in_db, -) class TestMilestoneContest: diff --git a/tests/managers/thl/test_contest/test_raffle.py b/tests/managers/thl/test_contest/test_raffle.py index 5804ea3..736a5e9 100644 --- a/tests/managers/thl/test_contest/test_raffle.py +++ b/tests/managers/thl/test_contest/test_raffle.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime import pytest from pydantic import ValidationError @@ -28,18 +28,6 @@ from generalresearch.models.thl.contest.raffle import ( ) from generalresearch.models.thl.product import Product from generalresearch.models.thl.user import User -from test_utils.managers.contest.conftest import ( - raffle_contest as contest, -) -from test_utils.managers.contest.conftest import ( - raffle_contest_create as contest_create, -) -from test_utils.managers.contest.conftest import ( - raffle_contest_factory as contest_factory, -) -from test_utils.managers.contest.conftest import ( - raffle_contest_in_db as contest_in_db, -) class TestRaffleContest: diff --git a/tests/managers/thl/test_harmonized_uqa.py b/tests/managers/thl/test_harmonized_uqa.py index 3b6df48..81ac080 100644 --- a/tests/managers/thl/test_harmonized_uqa.py +++ b/tests/managers/thl/test_harmonized_uqa.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime import pytest diff --git a/tests/managers/thl/test_ipinfo.py b/tests/managers/thl/test_ipinfo.py index 847b00c..61d4d19 100644 --- a/tests/managers/thl/test_ipinfo.py +++ b/tests/managers/thl/test_ipinfo.py @@ -1,9 +1,9 @@ import faker from generalresearch.managers.thl.ipinfo import ( + GeoIpInfoManager, IPGeonameManager, IPInformationManager, - GeoIpInfoManager, ) from generalresearch.models.thl.ipinfo import IPGeoname, IPInformation diff --git a/tests/managers/thl/test_ledger/test_lm_accounts.py b/tests/managers/thl/test_ledger/test_lm_accounts.py index faef5fb..11b2835 100644 --- a/tests/managers/thl/test_ledger/test_lm_accounts.py +++ b/tests/managers/thl/test_ledger/test_lm_accounts.py @@ -1,4 +1,3 @@ -from collections.abc import Callable from itertools import product as iproduct from random import randint from typing import TYPE_CHECKING @@ -22,7 +21,6 @@ from generalresearch.models.thl.ledger import ( if TYPE_CHECKING: from pydantic import PositiveInt - from generalresearch.config import GRLSettings from generalresearch.currency import LedgerCurrency from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager from generalresearch.models.custom_types import AccountType, Direction, UUIDStr @@ -32,10 +30,6 @@ if TYPE_CHECKING: LedgerAccount, LedgerTransaction, ) - from generalresearch.models.thl.product import Product - from generalresearch.models.thl.session import Session - from generalresearch.models.thl.user import User - from generalresearch.models.thl.wallet import PayoutType @pytest.mark.parametrize( 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 07c3712..020b74a 100644 --- a/tests/managers/thl/test_ledger/test_lm_tx_locks.py +++ b/tests/managers/thl/test_ledger/test_lm_tx_locks.py @@ -1,6 +1,6 @@ import logging from collections.abc import Callable -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal import pytest @@ -23,7 +23,6 @@ from generalresearch.models.thl.session import ( WallAdjustedStatus, ) from generalresearch.models.thl.user import User -from test_utils.models.conftest import product_user_wallet_no, session, user_factory logger = logging.getLogger("LedgerManager") @@ -56,7 +55,7 @@ class TestLedgerLocks: s1 = session_factory( user=user, wall_count=3, - wall_req_cpis=[Decimal("1.23"), Decimal("3.21"), Decimal("4")], + wall_req_cpis=[Decimal("1.23"), Decimal("3.21"), Decimal(4)], wall_statuses=[Status.COMPLETE, Status.COMPLETE, Status.COMPLETE], ) @@ -187,7 +186,7 @@ class TestLedgerLocks: user_id=user.user_id, source=Source.DYNATA, req_survey_id="xxx", - req_cpi=Decimal("4"), + req_cpi=Decimal(4), session_id=1, status=Status.COMPLETE, status_code_1=StatusCode1.COMPLETE, @@ -292,13 +291,7 @@ class TestLedgerLocks: status, status_code_1 = session.determine_session_status() thl_net, commission_amount, bp_pay, user_pay = session.determine_payments() session.update( - **{ - "status": status, - "status_code_1": status_code_1, - "finished": session.started + timedelta(minutes=10), - "payout": bp_pay, - "user_payout": user_pay, - } + status=status, status_code_1=status_code_1, finished=session.started + timedelta(minutes=10), payout=bp_pay, user_payout=user_pay ) print(thl_net, commission_amount, bp_pay, user_pay) @@ -336,13 +329,7 @@ class TestLedgerLocks: status, status_code_1 = session.determine_session_status() thl_net, commission_amount, bp_pay, user_pay = session.determine_payments() session.update( - **{ - "status": status, - "status_code_1": status_code_1, - "finished": session.started + timedelta(minutes=10), - "payout": bp_pay, - "user_payout": user_pay, - } + status=status, status_code_1=status_code_1, finished=session.started + timedelta(minutes=10), payout=bp_pay, user_payout=user_pay ) print(thl_net, commission_amount, bp_pay, user_pay) diff --git a/tests/managers/thl/test_ledger/test_thl_lm_accounts.py b/tests/managers/thl/test_ledger/test_thl_lm_accounts.py index 01d5fe1..8d7d828 100644 --- a/tests/managers/thl/test_ledger/test_thl_lm_accounts.py +++ b/tests/managers/thl/test_ledger/test_thl_lm_accounts.py @@ -8,9 +8,9 @@ class TestThlLedgerManagerAccounts: def test_get_account_or_create_user_wallet(self, user, thl_lm, lm): from generalresearch.currency import LedgerCurrency from generalresearch.models.thl.ledger import ( - LedgerAccount, - Direction, AccountType, + Direction, + LedgerAccount, ) account = thl_lm.get_account_or_create_user_wallet(user=user) @@ -31,9 +31,9 @@ class TestThlLedgerManagerAccounts: def test_get_account_or_create_bp_wallet(self, product, thl_lm, lm): from generalresearch.currency import LedgerCurrency from generalresearch.models.thl.ledger import ( - LedgerAccount, - Direction, AccountType, + Direction, + LedgerAccount, ) account = thl_lm.get_account_or_create_bp_wallet(product=product) @@ -54,8 +54,8 @@ class TestThlLedgerManagerAccounts: def test_get_account_or_create_bp_commission(self, product, thl_lm, lm): from generalresearch.currency import LedgerCurrency from generalresearch.models.thl.ledger import ( - Direction, AccountType, + Direction, ) account = thl_lm.get_account_or_create_bp_commission(product=product) @@ -76,8 +76,8 @@ class TestThlLedgerManagerAccounts: def test_get_account_or_create_bp_expense(self, product, expense, thl_lm, lm): from generalresearch.currency import LedgerCurrency from generalresearch.models.thl.ledger import ( - Direction, AccountType, + Direction, ) account = thl_lm.get_account_or_create_bp_expense( @@ -98,8 +98,8 @@ class TestThlLedgerManagerAccounts: def test_get_or_create_bp_pending_payout_account(self, product, thl_lm, lm): from generalresearch.currency import LedgerCurrency from generalresearch.models.thl.ledger import ( - Direction, AccountType, + Direction, ) account = thl_lm.get_or_create_bp_pending_payout_account(product=product) @@ -132,8 +132,8 @@ class TestThlLedgerManagerAccounts: self, account_cash, account_revenue_task_complete, thl_lm, lm ): from generalresearch.models.thl.ledger import ( - LedgerAccount, AccountType, + LedgerAccount, ) res = thl_lm.get_account_task_complete_revenue() @@ -155,8 +155,8 @@ class TestThlLedgerManagerAccounts: def test_get_account_cash(self, account_cash, thl_lm, lm): from generalresearch.models.thl.ledger import ( - LedgerAccount, AccountType, + LedgerAccount, ) res = thl_lm.get_account_cash() @@ -167,10 +167,10 @@ class TestThlLedgerManagerAccounts: assert res.display_name == "Operating Cash Account" def test_get_accounts(self, setup_accounts, product, user_factory, thl_lm, lm, lam): - from generalresearch.models.thl.user import User from generalresearch.managers.thl.ledger_manager.exceptions import ( LedgerAccountDoesntExistError, ) + from generalresearch.models.thl.user import User user1: User = user_factory(product=product) user2: User = user_factory(product=product) @@ -276,9 +276,9 @@ class TestLedgerAccountManager: LedgerAccountDoesntExistError, ) from generalresearch.models.thl.ledger import ( - LedgerAccount, - Direction, AccountType, + Direction, + LedgerAccount, ) u = uuid4().hex @@ -311,8 +311,8 @@ class TestLedgerAccountManager: LedgerAccountDoesntExistError, ) from generalresearch.models.thl.ledger import ( - LedgerAccount, AccountType, + LedgerAccount, ) with pytest.raises(LedgerAccountDoesntExistError): @@ -326,10 +326,10 @@ class TestLedgerAccountManager: assert user.product.uuid == account.reference_uuid def test_get_many(self, product_factory, thl_lm, lm, lam, currency): - from generalresearch.models.thl.product import Product from generalresearch.managers.thl.ledger_manager.exceptions import ( LedgerAccountDoesntExistError, ) + from generalresearch.models.thl.product import Product p1: Product = product_factory() p2: Product = product_factory() @@ -388,9 +388,9 @@ class TestLedgerAccountManager: def test_create_account(self, thl_lm, lm, lam): from generalresearch.models.thl.ledger import ( - LedgerAccount, - Direction, AccountType, + Direction, + LedgerAccount, ) u = uuid4().hex diff --git a/tests/managers/thl/test_ledger/test_thl_lm_bp_payout.py b/tests/managers/thl/test_ledger/test_thl_lm_bp_payout.py index 1fb9c01..cfb8f8f 100644 --- a/tests/managers/thl/test_ledger/test_thl_lm_bp_payout.py +++ b/tests/managers/thl/test_ledger/test_thl_lm_bp_payout.py @@ -1,5 +1,5 @@ import logging -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal from random import randint from uuid import uuid4 @@ -78,13 +78,7 @@ class TestThlLedgerManagerBPPayout: status, status_code_1 = session.determine_session_status() thl_net, commission_amount, bp_pay, user_pay = session.determine_payments() session.update( - **{ - "status": status, - "status_code_1": status_code_1, - "finished": now + timedelta(minutes=10), - "payout": bp_pay, - "user_payout": user_pay, - } + status=status, status_code_1=status_code_1, finished=now + timedelta(minutes=10), payout=bp_pay, user_payout=user_pay ) thl_lm.create_tx_bp_payment(session=session, created=wall1.started) diff --git a/tests/managers/thl/test_ledger/test_thl_lm_tx.py b/tests/managers/thl/test_ledger/test_thl_lm_tx.py index be988a1..1130621 100644 --- a/tests/managers/thl/test_ledger/test_thl_lm_tx.py +++ b/tests/managers/thl/test_ledger/test_thl_lm_tx.py @@ -1,5 +1,5 @@ import logging -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal from random import randint from uuid import uuid4 @@ -114,7 +114,7 @@ class TestThlLedgerTxManager: user_wallet_config=UserWalletConfig(amt=True, enabled=True), ) user = user_factory(product=product) - s1 = session_factory(user=user, wall_req_cpi=Decimal("1")) + s1 = session_factory(user=user, wall_req_cpi=Decimal(1)) status, status_code_1 = s1.determine_session_status() assert status == Status.COMPLETE @@ -207,9 +207,7 @@ class TestThlLedgerTxManager: # Update the finished timestamp, but nothing else. This means that # there is no financial changes needed session.update( - **{ - "finished": datetime.now(tz=UTC) + timedelta(minutes=10), - } + finished=datetime.now(tz=UTC) + timedelta(minutes=10) ) assert session.finished with caplog.at_level(logging.INFO): @@ -829,13 +827,7 @@ class TestThlLedgerTxManagerFlows: status, status_code_1 = s1.determine_session_status() thl_net, commission_amount, bp_pay, user_pay = s1.determine_payments() s1.update( - **{ - "status": status, - "status_code_1": status_code_1, - "finished": s1.started + timedelta(minutes=10), - "payout": bp_pay, - "user_payout": user_pay, - } + status=status, status_code_1=status_code_1, finished=s1.started + timedelta(minutes=10), payout=bp_pay, user_payout=user_pay ) print(thl_net, commission_amount, bp_pay, user_pay) thl_lm.create_tx_bp_payment(session=s1, created=w1.started) @@ -892,13 +884,7 @@ class TestThlLedgerTxManagerFlows: status, status_code_1 = session.determine_session_status() thl_net, commission_amount, bp_pay, user_pay = session.determine_payments() session.update( - **{ - "status": status, - "status_code_1": status_code_1, - "finished": session.started + timedelta(minutes=10), - "payout": bp_pay, - "user_payout": user_pay, - } + status=status, status_code_1=status_code_1, finished=session.started + timedelta(minutes=10), payout=bp_pay, user_payout=user_pay ) print(thl_net, commission_amount, bp_pay, user_pay) @@ -930,13 +916,7 @@ class TestThlLedgerTxManagerFlows: status, status_code_1 = session.determine_session_status() # thl_net, commission_amount, bp_pay, user_pay = session.determine_payments() session.update( - **{ - "status": status, - "status_code_1": status_code_1, - "finished": session.started + timedelta(minutes=10), - "payout": Decimal("1.53"), - "user_payout": Decimal("1.53"), - } + status=status, status_code_1=status_code_1, finished=session.started + timedelta(minutes=10), payout=Decimal("1.53"), user_payout=Decimal("1.53") ) thl_lm.create_tx_bp_payment(session=session, created=wall1.started) @@ -971,13 +951,7 @@ class TestThlLedgerTxManagerFlows: status, status_code_1 = session.determine_session_status() # thl_net, commission_amount, bp_pay, user_pay = session.determine_payments() session.update( - **{ - "status": status, - "status_code_1": status_code_1, - "finished": session.started + timedelta(minutes=10), - "payout": Decimal("0.39"), - "user_payout": Decimal("0.26"), - } + status=status, status_code_1=status_code_1, finished=session.started + timedelta(minutes=10), payout=Decimal("0.39"), user_payout=Decimal("0.26") ) # with pytest.logs(logger, level=logging.WARNING) as cm: # tx = thl_lm.create_transaction_bp_payment(session, created=wall1.started) @@ -1310,13 +1284,7 @@ class TestThlLedgerManagerAdj: status, status_code_1 = session.determine_session_status() thl_net, commission_amount, bp_pay, user_pay = session.determine_payments() session.update( - **{ - "status": status, - "status_code_1": status_code_1, - "finished": utc_hour_ago + timedelta(minutes=10), - "payout": bp_pay, - "user_payout": user_pay, - } + status=status, status_code_1=status_code_1, finished=utc_hour_ago + timedelta(minutes=10), payout=bp_pay, user_payout=user_pay ) thl_lm.create_tx_bp_payment(session, created=wall1.started) @@ -1484,13 +1452,7 @@ class TestThlLedgerManagerAdj: assert user_pay == Decimal("1.52") session.update( - **{ - "status": status, - "status_code_1": status_code_1, - "finished": now + timedelta(minutes=10), - "payout": bp_pay, - "user_payout": user_pay, - } + status=status, status_code_1=status_code_1, finished=now + timedelta(minutes=10), payout=bp_pay, user_payout=user_pay ) tx = thl_lm.create_tx_bp_adjustment(session=session, created=wall1.started) @@ -1664,13 +1626,7 @@ class TestThlLedgerManagerAdj: status, status_code_1 = session.determine_session_status() thl_net, commission_amount, bp_pay, user_pay = session.determine_payments() session.update( - **{ - "status": status, - "status_code_1": status_code_1, - "finished": utc_hour_ago + timedelta(minutes=10), - "payout": bp_pay, - "user_payout": user_pay, - } + status=status, status_code_1=status_code_1, finished=utc_hour_ago + timedelta(minutes=10), payout=bp_pay, user_payout=user_pay ) thl_lm.create_tx_bp_payment(session, created=wall1.started) diff --git a/tests/managers/thl/test_ledger/test_thl_lm_tx__user_payouts.py b/tests/managers/thl/test_ledger/test_thl_lm_tx__user_payouts.py index 9253ff0..cd6ea79 100644 --- a/tests/managers/thl/test_ledger/test_thl_lm_tx__user_payouts.py +++ b/tests/managers/thl/test_ledger/test_thl_lm_tx__user_payouts.py @@ -1,5 +1,5 @@ import logging -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal from uuid import uuid4 @@ -12,7 +12,6 @@ from generalresearch.managers.thl.ledger_manager.exceptions import ( from generalresearch.models.thl.payout import UserPayoutEvent from generalresearch.models.thl.user import User from generalresearch.models.thl.wallet import PayoutType -from test_utils.managers.ledger.conftest import create_main_accounts class TestLedgerManagerAMT: diff --git a/tests/managers/thl/test_ledger/test_thl_pem.py b/tests/managers/thl/test_ledger/test_thl_pem.py index 5fb9e7d..7a4f0c9 100644 --- a/tests/managers/thl/test_ledger/test_thl_pem.py +++ b/tests/managers/thl/test_ledger/test_thl_pem.py @@ -1,6 +1,5 @@ -import uuid from random import randint -from uuid import uuid4, UUID +from uuid import UUID, uuid4 import pytest @@ -33,7 +32,6 @@ class TestThlPayoutEventManager: thl_lm, brokerage_product_payout_event_manager, ): - from generalresearch.models.thl.payout import UserPayoutEvent N_PRODUCTS = randint(3, 10) N_PAYOUT_EVENTS = randint(3, 10) @@ -73,7 +71,6 @@ class TestThlPayoutEventManager: brokerage_product_payout_event_manager, thl_lm, ): - from generalresearch.models.thl.payout import UserPayoutEvent N_PRODUCTS = randint(3, 10) N_PAYOUT_EVENTS = randint(3, 10) @@ -119,7 +116,6 @@ class TestThlPayoutEventManager: description can't be None """ from generalresearch.models.thl.payout import ( - UserPayoutEvent, PayoutType, ) @@ -174,7 +170,6 @@ class TestThlPayoutEventManager: brokerage_product_payout_event_manager, lm, ): - from generalresearch.models.thl.payout import UserPayoutEvent delete_ledger_db() create_main_accounts() diff --git a/tests/managers/thl/test_ledger/test_user_txs.py b/tests/managers/thl/test_ledger/test_user_txs.py index b4b0437..f83641e 100644 --- a/tests/managers/thl/test_ledger/test_user_txs.py +++ b/tests/managers/thl/test_ledger/test_user_txs.py @@ -1,5 +1,5 @@ from collections.abc import Callable -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime from decimal import Decimal from typing import TYPE_CHECKING from uuid import uuid4 @@ -8,7 +8,6 @@ from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerMana from generalresearch.managers.thl.user_compensate import user_compensate from generalresearch.models.thl.definitions import ( Status, - WallAdjustedStatus, ) from generalresearch.models.thl.ledger import ( TransactionType, diff --git a/tests/managers/thl/test_ledger/test_wallet.py b/tests/managers/thl/test_ledger/test_wallet.py index a0abd7c..bb49cd8 100644 --- a/tests/managers/thl/test_ledger/test_wallet.py +++ b/tests/managers/thl/test_ledger/test_wallet.py @@ -4,10 +4,10 @@ from uuid import uuid4 import pytest from generalresearch.models.thl.product import ( - UserWalletConfig, PayoutConfig, PayoutTransformation, PayoutTransformationPercentArgs, + UserWalletConfig, ) from generalresearch.models.thl.user import User diff --git a/tests/managers/thl/test_product.py b/tests/managers/thl/test_product.py index 78d5dde..f8dd44d 100644 --- a/tests/managers/thl/test_product.py +++ b/tests/managers/thl/test_product.py @@ -5,15 +5,14 @@ import pytest from generalresearch.models import Source from generalresearch.models.thl.product import ( Product, + ProfilingConfig, SourceConfig, - UserCreateConfig, SourcesConfig, - UserHealthConfig, - ProfilingConfig, - SupplyPolicy, SupplyConfig, + SupplyPolicy, + UserCreateConfig, + UserHealthConfig, ) -from test_utils.models.conftest import product_factory class TestProductManagerGetMethods: diff --git a/tests/managers/thl/test_product_prod.py b/tests/managers/thl/test_product_prod.py index 7b4f677..f5aa78d 100644 --- a/tests/managers/thl/test_product_prod.py +++ b/tests/managers/thl/test_product_prod.py @@ -4,7 +4,6 @@ from uuid import uuid4 import pytest from generalresearch.models.thl.product import Product -from test_utils.models.conftest import product_factory logger = logging.getLogger() @@ -79,4 +78,3 @@ class TestProductManagerGetAll: 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 - pass diff --git a/tests/managers/thl/test_profiling/test_user_upk.py b/tests/managers/thl/test_profiling/test_user_upk.py index 491e2b1..8b995b1 100644 --- a/tests/managers/thl/test_profiling/test_user_upk.py +++ b/tests/managers/thl/test_profiling/test_user_upk.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from generalresearch.managers.thl.profiling.user_upk import UserUpkManager diff --git a/tests/managers/thl/test_session_manager.py b/tests/managers/thl/test_session_manager.py index 6bedc2b..adcbe25 100644 --- a/tests/managers/thl/test_session_manager.py +++ b/tests/managers/thl/test_session_manager.py @@ -7,11 +7,10 @@ from faker import Faker from generalresearch.models import DeviceType from generalresearch.models.legacy.bucket import Bucket from generalresearch.models.thl.definitions import ( + SessionStatusCode2, Status, StatusCode1, - SessionStatusCode2, ) -from test_utils.models.conftest import user fake = Faker() @@ -21,8 +20,8 @@ class TestSessionManager: bucket = Bucket( loi_min=timedelta(seconds=60), loi_max=timedelta(seconds=120), - user_payout_min=Decimal("1"), - user_payout_max=Decimal("2"), + user_payout_min=Decimal(1), + user_payout_max=Decimal(2), ) s1 = session_manager.create( diff --git a/tests/managers/thl/test_survey.py b/tests/managers/thl/test_survey.py index 4b4a579..37f0b66 100644 --- a/tests/managers/thl/test_survey.py +++ b/tests/managers/thl/test_survey.py @@ -1,5 +1,5 @@ import uuid -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from decimal import Decimal import pytest diff --git a/tests/managers/thl/test_survey_penalty.py b/tests/managers/thl/test_survey_penalty.py index 4c7dc08..2a3cdc2 100644 --- a/tests/managers/thl/test_survey_penalty.py +++ b/tests/managers/thl/test_survey_penalty.py @@ -1,7 +1,6 @@ import uuid import pytest -from cachetools.keys import _HashedTuple from generalresearch.models import Source from generalresearch.models.thl.survey.penalty import ( diff --git a/tests/managers/thl/test_task_adjustment.py b/tests/managers/thl/test_task_adjustment.py index 71e3535..43337b6 100644 --- a/tests/managers/thl/test_task_adjustment.py +++ b/tests/managers/thl/test_task_adjustment.py @@ -1,5 +1,5 @@ import logging -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal from random import randint diff --git a/tests/managers/thl/test_task_status.py b/tests/managers/thl/test_task_status.py index 468fd5e..93a624d 100644 --- a/tests/managers/thl/test_task_status.py +++ b/tests/managers/thl/test_task_status.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal import pytest diff --git a/tests/managers/thl/test_user_manager/test_base.py b/tests/managers/thl/test_user_manager/test_base.py index 2704490..7d83c11 100644 --- a/tests/managers/thl/test_user_manager/test_base.py +++ b/tests/managers/thl/test_user_manager/test_base.py @@ -1,5 +1,5 @@ import logging -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from random import randint from uuid import uuid4 @@ -156,7 +156,7 @@ class TestCreateUserManager: # make sure thl_user row is created res_thl_user = thl_web_rw.execute_sql_query( - query=f""" + query=""" SELECT * FROM thl_user AS u WHERE u.id = %s diff --git a/tests/managers/thl/test_user_manager/test_mysql.py b/tests/managers/thl/test_user_manager/test_mysql.py index 0313bbf..d414a13 100644 --- a/tests/managers/thl/test_user_manager/test_mysql.py +++ b/tests/managers/thl/test_user_manager/test_mysql.py @@ -1,4 +1,3 @@ -from test_utils.models.conftest import user, user_manager class TestUserManagerMysqlNew: 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 a4b3d57..7c9e012 100644 --- a/tests/managers/thl/test_user_manager/test_user_fetch.py +++ b/tests/managers/thl/test_user_manager/test_user_fetch.py @@ -3,7 +3,6 @@ from uuid import uuid4 import pytest from generalresearch.models.thl.user import User -from test_utils.models.conftest import product, user_manager, user_factory class TestUserManagerFetch: 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 91dc16a..19b3d9f 100644 --- a/tests/managers/thl/test_user_manager/test_user_metadata.py +++ b/tests/managers/thl/test_user_manager/test_user_metadata.py @@ -3,7 +3,6 @@ from uuid import uuid4 import pytest from generalresearch.models.thl.user_profile import UserMetadata -from test_utils.models.conftest import user, user_manager, user_factory class TestUserMetadataManager: diff --git a/tests/managers/thl/test_user_streak.py b/tests/managers/thl/test_user_streak.py index ef25e2b..be0729c 100644 --- a/tests/managers/thl/test_user_streak.py +++ b/tests/managers/thl/test_user_streak.py @@ -1,5 +1,5 @@ import copy -from datetime import UTC, date, datetime, timedelta, timezone +from datetime import UTC, date, datetime, timedelta from decimal import Decimal from zoneinfo import ZoneInfo diff --git a/tests/managers/thl/test_userhealth.py b/tests/managers/thl/test_userhealth.py index e2ea4a8..be8f9c8 100644 --- a/tests/managers/thl/test_userhealth.py +++ b/tests/managers/thl/test_userhealth.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from uuid import uuid4 import faker diff --git a/tests/managers/thl/test_wall_manager.py b/tests/managers/thl/test_wall_manager.py index 46d199f..5abc648 100644 --- a/tests/managers/thl/test_wall_manager.py +++ b/tests/managers/thl/test_wall_manager.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal from uuid import uuid4 @@ -10,7 +10,6 @@ from generalresearch.models.thl.session import ( Status, StatusCode1, ) -from test_utils.models.conftest import session, user class TestWallManager: @@ -92,7 +91,7 @@ class TestWallManager: source=Source.DYNATA, buyer_id="123", req_survey_id="456", - req_cpi=Decimal("1"), + req_cpi=Decimal(1), ) assert w is not None @@ -110,7 +109,7 @@ class TestWallManager: source=Source.DYNATA, buyer_id="123", req_survey_id="456", - req_cpi=Decimal("1"), + req_cpi=Decimal(1), ) wall_manager.report( wall=w1, @@ -151,7 +150,7 @@ class TestWallManager: source=Source.DYNATA, buyer_id="123", req_survey_id="456", - req_cpi=Decimal("1"), + req_cpi=Decimal(1), ) finish_ts = utc_hour_ago + timedelta(minutes=10) @@ -190,7 +189,7 @@ class TestWallManager: source=Source.DYNATA, buyer_id="123", req_survey_id="456", - req_cpi=Decimal("1"), + req_cpi=Decimal(1), ) res = wall_manager.filter_wall_attempts(user_id=user.user_id) assert len(res) == 1 @@ -202,7 +201,7 @@ class TestWallManager: source=Source.DYNATA, buyer_id="123", req_survey_id="555", - req_cpi=Decimal("1"), + req_cpi=Decimal(1), ) res = wall_manager.filter_wall_attempts(user_id=user.user_id) assert len(res) == 2 diff --git a/tests/models/admin/test_report_request.py b/tests/models/admin/test_report_request.py index 4626ab4..5b2ff0d 100644 --- a/tests/models/admin/test_report_request.py +++ b/tests/models/admin/test_report_request.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime import pandas as pd import pytest diff --git a/tests/models/custom_types/test_aware_datetime.py b/tests/models/custom_types/test_aware_datetime.py index 043fba0..7c45710 100644 --- a/tests/models/custom_types/test_aware_datetime.py +++ b/tests/models/custom_types/test_aware_datetime.py @@ -1,7 +1,7 @@ from __future__ import annotations import logging -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime import pytest import pytz diff --git a/tests/models/custom_types/test_dsn.py b/tests/models/custom_types/test_dsn.py index 050976e..eb27526 100644 --- a/tests/models/custom_types/test_dsn.py +++ b/tests/models/custom_types/test_dsn.py @@ -1,4 +1,3 @@ -from typing import Optional from uuid import uuid4 import pytest diff --git a/tests/models/dynata/test_eligbility.py b/tests/models/dynata/test_eligbility.py index 16cad26..27de5b3 100644 --- a/tests/models/dynata/test_eligbility.py +++ b/tests/models/dynata/test_eligbility.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime class TestEligibility: diff --git a/tests/models/gr/test_authentication.py b/tests/models/gr/test_authentication.py index 51595a7..ac39e64 100644 --- a/tests/models/gr/test_authentication.py +++ b/tests/models/gr/test_authentication.py @@ -2,7 +2,7 @@ import binascii import json import os from collections.abc import Callable -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from random import randint from uuid import uuid4 diff --git a/tests/models/gr/test_business.py b/tests/models/gr/test_business.py index 716ec75..948acb3 100644 --- a/tests/models/gr/test_business.py +++ b/tests/models/gr/test_business.py @@ -1,7 +1,6 @@ import os -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal -from typing import Optional from uuid import uuid4 import pandas as pd diff --git a/tests/models/innovate/test_question.py b/tests/models/innovate/test_question.py index 330f919..b0c2964 100644 --- a/tests/models/innovate/test_question.py +++ b/tests/models/innovate/test_question.py @@ -1,15 +1,15 @@ from generalresearch.models import Source from generalresearch.models.innovate.question import ( InnovateQuestion, - InnovateQuestionType, InnovateQuestionOption, + InnovateQuestionType, ) from generalresearch.models.thl.profiling.upk_question import ( - UpkQuestionSelectorTE, UpkQuestion, + UpkQuestionChoice, UpkQuestionSelectorMC, + UpkQuestionSelectorTE, UpkQuestionType, - UpkQuestionChoice, ) diff --git a/tests/models/legacy/test_user_question_answer_in.py b/tests/models/legacy/test_user_question_answer_in.py index 224334a..1904798 100644 --- a/tests/models/legacy/test_user_question_answer_in.py +++ b/tests/models/legacy/test_user_question_answer_in.py @@ -263,12 +263,12 @@ class TestUserQuestionAnswerIn: UserQuestionAnswerIn, ) - for qid in { + for qid in ( "2fbedb2b9f7647b09ff5e52fa119cc5e", "4030c52371b04e80b64e058d9c5b82e9", "a91cb1dea814480dba12d9b7b48696dd", "1d1e2e8380ac474b87fb4e4c569b48df", - }: + ): # This is the UserAgent question which only allows a single answer with pytest.raises(ValueError) as cm: UserQuestionAnswerIn.model_validate( diff --git a/tests/models/morning/test.py b/tests/models/morning/test.py index 222cb93..7474766 100644 --- a/tests/models/morning/test.py +++ b/tests/models/morning/test.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from generalresearch.models.morning.question import MorningQuestion diff --git a/tests/models/network/test_mtr.py b/tests/models/network/test_mtr.py index 2965300..840a773 100644 --- a/tests/models/network/test_mtr.py +++ b/tests/models/network/test_mtr.py @@ -1,7 +1,7 @@ -from generalresearch.models.network.mtr.execute import execute_mtr import faker -from generalresearch.models.network.tool_run import ToolName, ToolClass +from generalresearch.models.network.mtr.execute import execute_mtr +from generalresearch.models.network.tool_run import ToolClass, ToolName fake = faker.Faker() diff --git a/tests/models/network/test_nmap_parser.py b/tests/models/network/test_nmap_parser.py index abc83c9..7822380 100644 --- a/tests/models/network/test_nmap_parser.py +++ b/tests/models/network/test_nmap_parser.py @@ -4,6 +4,7 @@ import pytest from generalresearch.models.network.nmap.parser import parse_nmap_xml + @pytest.fixture def nmap_raw_output_2(request) -> str: fp = os.path.join(request.config.rootpath, "data/nmaprun2.xml") diff --git a/tests/models/prodege/test_survey_participation.py b/tests/models/prodege/test_survey_participation.py index 3b35d0c..e1ba9ab 100644 --- a/tests/models/prodege/test_survey_participation.py +++ b/tests/models/prodege/test_survey_participation.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta class TestProdegeParticipation: diff --git a/tests/models/spectrum/test_question.py b/tests/models/spectrum/test_question.py index 4f92961..57d260d 100644 --- a/tests/models/spectrum/test_question.py +++ b/tests/models/spectrum/test_question.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from generalresearch.models import Source from generalresearch.models.spectrum.question import ( diff --git a/tests/models/spectrum/test_survey.py b/tests/models/spectrum/test_survey.py index 5e095a3..7365c7e 100644 --- a/tests/models/spectrum/test_survey.py +++ b/tests/models/spectrum/test_survey.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from decimal import Decimal diff --git a/tests/models/spectrum/test_survey_manager.py b/tests/models/spectrum/test_survey_manager.py index 11970bf..ce26c44 100644 --- a/tests/models/spectrum/test_survey_manager.py +++ b/tests/models/spectrum/test_survey_manager.py @@ -1,6 +1,6 @@ import copy import logging -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from decimal import Decimal from pymysql import IntegrityError diff --git a/tests/models/test_finance.py b/tests/models/test_finance.py index 6dcd441..3d3ff3a 100644 --- a/tests/models/test_finance.py +++ b/tests/models/test_finance.py @@ -1,5 +1,5 @@ from collections.abc import Callable -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from itertools import product as iter_product from random import randint from uuid import uuid4 diff --git a/tests/models/thl/question/test_question_info.py b/tests/models/thl/question/test_question_info.py index 945ee7a..b619fc3 100644 --- a/tests/models/thl/question/test_question_info.py +++ b/tests/models/thl/question/test_question_info.py @@ -1,6 +1,6 @@ from generalresearch.models.thl.profiling.upk_property import ( - UpkProperty, ProfilingInfo, + UpkProperty, ) diff --git a/tests/models/thl/test_adjustments.py b/tests/models/thl/test_adjustments.py index c2c035d..30e9bce 100644 --- a/tests/models/thl/test_adjustments.py +++ b/tests/models/thl/test_adjustments.py @@ -1,5 +1,5 @@ from collections.abc import Callable -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal import pytest @@ -426,7 +426,7 @@ class TestAdjustments: s1 = session_factory( user=user, wall_count=1, - wall_req_cpi=Decimal("1"), + wall_req_cpi=Decimal(1), final_status=Status.COMPLETE, started=utc_hour_ago, ) @@ -525,7 +525,7 @@ class TestAdjustments: s1 = session_factory( user=user, wall_count=1, - wall_req_cpi=Decimal("1"), + wall_req_cpi=Decimal(1), final_status=Status.COMPLETE, started=utc_hour_ago, ) @@ -534,13 +534,7 @@ class TestAdjustments: status, status_code_1 = s1.determine_session_status() thl_net, commission_amount, bp_pay, user_pay = s1.determine_payments() s1.update( - **{ - "status": status, - "status_code_1": status_code_1, - "finished": utc_hour_ago + timedelta(minutes=10), - "payout": bp_pay, - "user_payout": user_pay, - } + status=status, status_code_1=status_code_1, finished=utc_hour_ago + timedelta(minutes=10), payout=bp_pay, user_payout=user_pay ) w1.update( adjusted_status=WallAdjustedStatus.ADJUSTED_TO_FAIL, @@ -592,13 +586,7 @@ class TestAdjustments: thl_net = Decimal(sum(w.cpi for w in s1.wall_events if w.is_visible_complete())) payout = user.product.determine_bp_payment(thl_net=thl_net) s1.update( - **{ - "status": status, - "status_code_1": status_code_1, - "finished": utc_hour_ago + timedelta(minutes=25), - "payout": payout, - "user_payout": None, - } + status=status, status_code_1=status_code_1, finished=utc_hour_ago + timedelta(minutes=25), payout=payout, user_payout=None ) # Test. Adjust first fail to complete. Now we have 2 completes. @@ -646,13 +634,7 @@ class TestAdjustments: thl_net = Decimal(sum(w.cpi for w in s1.wall_events if w.is_visible_complete())) payout = user.product.determine_bp_payment(thl_net) s1.update( - **{ - "status": status, - "status_code_1": status_code_1, - "finished": utc_hour_ago + timedelta(minutes=25), - "payout": payout, - "user_payout": None, - } + status=status, status_code_1=status_code_1, finished=utc_hour_ago + timedelta(minutes=25), payout=payout, user_payout=None ) # Test. Adjust complete to fail. Now we have 2 fails. diff --git a/tests/models/thl/test_contest/test_leaderboard_contest.py b/tests/models/thl/test_contest/test_leaderboard_contest.py index 3efcf2f..52a4bec 100644 --- a/tests/models/thl/test_contest/test_leaderboard_contest.py +++ b/tests/models/thl/test_contest/test_leaderboard_contest.py @@ -1,4 +1,4 @@ -from datetime import UTC, timezone +from datetime import UTC from uuid import uuid4 import pytest diff --git a/tests/models/thl/test_ledger.py b/tests/models/thl/test_ledger.py index 5edcc9d..7066180 100644 --- a/tests/models/thl/test_ledger.py +++ b/tests/models/thl/test_ledger.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from uuid import uuid4 import pytest diff --git a/tests/models/thl/test_payout.py b/tests/models/thl/test_payout.py index 7068a41..d0b0acc 100644 --- a/tests/models/thl/test_payout.py +++ b/tests/models/thl/test_payout.py @@ -5,14 +5,13 @@ from pydantic import ValidationError from generalresearch.currency import USDCent from generalresearch.models.gr import Team +from generalresearch.models.gr.business import Business, BusinessAddress, BusinessType from generalresearch.models.thl.payout import ( - BusinessPayoutEvent, BrokerageProductPayoutEvent, + BusinessPayoutEvent, ) from generalresearch.models.thl.wallet import PayoutType -from generalresearch.models.gr.business import Business, BusinessAddress, BusinessType - class TestBusinessPayoutEvent: diff --git a/tests/models/thl/test_product.py b/tests/models/thl/test_product.py index 78bc10a..bc95c2d 100644 --- a/tests/models/thl/test_product.py +++ b/tests/models/thl/test_product.py @@ -3,7 +3,7 @@ from __future__ import annotations import os import shutil from collections.abc import Callable -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal from uuid import uuid4 diff --git a/tests/models/thl/test_upkquestion.py b/tests/models/thl/test_upkquestion.py index d32875c..99d7871 100644 --- a/tests/models/thl/test_upkquestion.py +++ b/tests/models/thl/test_upkquestion.py @@ -201,18 +201,10 @@ class TestUpkQuestion: ) q = MorningQuestion( - **{ - "id": "gender", - "country_iso": "us", - "language_iso": "eng", - "name": "Gender", - "text": "What is your gender?", - "type": "s", - "options": [ + id="gender", country_iso="us", language_iso="eng", name="Gender", text="What is your gender?", type="s", options=[ {"id": "1", "text": "yes", "order": 1}, {"id": "2", "text": "no", "order": 2}, - ], - } + ] ) q.to_upk_question() q = MorningQuestion( diff --git a/tests/models/thl/test_user.py b/tests/models/thl/test_user.py index a4f331a..aafce68 100644 --- a/tests/models/thl/test_user.py +++ b/tests/models/thl/test_user.py @@ -197,7 +197,7 @@ class TestUserProductUserID: assert "Input should be a valid string" in str(cm.value) with pytest.raises(ValueError) as cm: - User(user_id=self.user_id, product_user_id=Decimal("0")) + User(user_id=self.user_id, product_user_id=Decimal(0)) assert "1 validation error for User" in str(cm.value) assert "Input should be a valid string" in str(cm.value) @@ -236,7 +236,7 @@ class TestUserProductUserID: def test_invalid_chars_slash(self): from generalresearch.models.thl.user import User - product_user_id = f"{self.randomword(50)}\{self.randomword(50)}" + product_user_id = rf"{self.randomword(50)}\{self.randomword(50)}" with pytest.raises(expected_exception=ValueError) as cm: User(user_id=self.user_id, product_user_id=product_user_id) assert "1 validation error for User" in str(cm.value) @@ -310,7 +310,7 @@ class TestUserUUID: assert "Input should be a valid string" in str(cm.value) with pytest.raises(ValueError) as cm: - User(user_id=self.user_id, uuid=Decimal("0")) + User(user_id=self.user_id, uuid=Decimal(0)) assert "1 validation error for User" in str(cm.value) assert "Input should be a valid string" in str(cm.value) diff --git a/tests/models/thl/test_user_iphistory.py b/tests/models/thl/test_user_iphistory.py index 0f050b0..d6ade9d 100644 --- a/tests/models/thl/test_user_iphistory.py +++ b/tests/models/thl/test_user_iphistory.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from generalresearch.models.thl.user_iphistory import ( UserIPHistory, diff --git a/tests/models/thl/test_user_streak.py b/tests/models/thl/test_user_streak.py index 0cacd3e..26c5e25 100644 --- a/tests/models/thl/test_user_streak.py +++ b/tests/models/thl/test_user_streak.py @@ -1,8 +1,8 @@ from datetime import datetime, timedelta +from zoneinfo import ZoneInfo import pytest from pydantic import ValidationError -from zoneinfo import ZoneInfo from generalresearch.models.thl.user_streak import ( StreakFulfillment, diff --git a/tests/models/thl/test_wall.py b/tests/models/thl/test_wall.py index 9e9483b..88914ac 100644 --- a/tests/models/thl/test_wall.py +++ b/tests/models/thl/test_wall.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal from uuid import uuid4 diff --git a/tests/models/thl/test_wall_session.py b/tests/models/thl/test_wall_session.py index 10f3cba..b39ad31 100644 --- a/tests/models/thl/test_wall_session.py +++ b/tests/models/thl/test_wall_session.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal import pytest diff --git a/tests/sql_helper.py b/tests/sql_helper.py index c4cc2ca..8ab7bdb 100644 --- a/tests/sql_helper.py +++ b/tests/sql_helper.py @@ -19,7 +19,7 @@ class TestSqlHelper: def test_scheme(self): from generalresearch.sql_helper import SqlHelper - dsn = MySQLDsn(f"mysql://root@localhost/test") + dsn = MySQLDsn("mysql://root@localhost/test") instance = SqlHelper(dsn=dsn) assert instance.is_mysql() @@ -30,7 +30,7 @@ class TestSqlHelper: # self.assertTrue(instance.is_postgresql()) with pytest.raises(ValidationError): - SqlHelper(dsn=MariaDBDsn(f"maria://root@localhost/test")) + SqlHelper(dsn=MariaDBDsn("maria://root@localhost/test")) def test_row_decode(self): from generalresearch.sql_helper import decode_uuids -- cgit v1.2.3 From 3b4059135be47f7752a08e4277a85f9e57ceaa9d Mon Sep 17 00:00:00 2001 From: Max Nanis Date: Tue, 25 Aug 2026 11:09:52 -0700 Subject: Ruff typing from this morning. WIP --- test_utils/models/ledger/conftest.py | 117 +-- .../incite/collections/test_df_collection_base.py | 2 +- .../collections/test_df_collection_item_base.py | 2 +- .../collections/test_df_collection_item_thl_web.py | 61 +- .../mergers/foundations/test_enriched_session.py | 18 +- .../foundations/test_enriched_task_adjust.py | 12 +- .../mergers/foundations/test_enriched_wall.py | 36 +- tests/incite/mergers/test_pop_ledger.py | 34 +- tests/incite/mergers/test_ym_survey_merge.py | 12 +- tests/managers/gr/test_business.py | 12 +- tests/managers/gr/test_team.py | 4 +- tests/managers/leaderboard.py | 2 +- tests/managers/test_events.py | 8 +- tests/managers/test_lucid.py | 4 +- .../managers/thl/test_contest/test_leaderboard.py | 49 +- tests/managers/thl/test_contest/test_milestone.py | 79 +- tests/managers/thl/test_contest/test_raffle.py | 146 ++- tests/managers/thl/test_ipinfo.py | 11 +- tests/managers/thl/test_ledger/test_lm_accounts.py | 23 +- tests/managers/thl/test_ledger/test_lm_tx.py | 142 ++- .../managers/thl/test_ledger/test_lm_tx_entries.py | 24 +- tests/managers/thl/test_ledger/test_lm_tx_locks.py | 208 ++-- .../thl/test_ledger/test_lm_tx_metadata.py | 40 +- .../thl/test_ledger/test_thl_lm_accounts.py | 310 +++--- .../thl/test_ledger/test_thl_lm_bp_payout.py | 234 +++-- tests/managers/thl/test_ledger/test_thl_lm_tx.py | 1004 +++++++++++--------- .../test_ledger/test_thl_lm_tx__user_payouts.py | 376 ++++---- tests/managers/thl/test_ledger/test_thl_pem.py | 124 ++- tests/managers/thl/test_ledger/test_user_txs.py | 72 +- tests/managers/thl/test_ledger/test_wallet.py | 36 +- tests/managers/thl/test_maxmind.py | 4 +- tests/managers/thl/test_payout.py | 414 ++++---- tests/managers/thl/test_product.py | 14 +- tests/managers/thl/test_product_prod.py | 12 +- tests/managers/thl/test_session_manager.py | 21 +- tests/managers/thl/test_task_adjustment.py | 6 +- tests/managers/thl/test_task_status.py | 50 +- tests/managers/thl/test_user_manager/test_base.py | 11 +- tests/managers/thl/test_user_manager/test_redis.py | 4 +- .../thl/test_user_manager/test_user_fetch.py | 4 +- .../thl/test_user_manager/test_user_metadata.py | 16 +- tests/managers/thl/test_userhealth.py | 36 +- tests/models/gr/test_authentication.py | 50 +- tests/models/gr/test_business.py | 122 +-- tests/models/gr/test_team.py | 14 +- .../models/legacy/test_user_question_answer_in.py | 16 +- tests/models/test_finance.py | 28 +- tests/models/thl/test_adjustments.py | 2 +- .../thl/test_contest/test_leaderboard_contest.py | 2 +- .../models/thl/test_contest/test_raffle_contest.py | 2 +- tests/models/thl/test_payout.py | 2 +- tests/models/thl/test_product.py | 78 +- tests/models/thl/test_user.py | 6 +- 53 files changed, 2375 insertions(+), 1741 deletions(-) (limited to 'tests/incite/collections') diff --git a/test_utils/models/ledger/conftest.py b/test_utils/models/ledger/conftest.py index b428468..1c1027c 100644 --- a/test_utils/models/ledger/conftest.py +++ b/test_utils/models/ledger/conftest.py @@ -502,74 +502,77 @@ def setup_accounts( lm: LedgerManager, user: User, currency: LedgerCurrency, -) -> None: +) -> Callable[..., None]: from generalresearch.models.thl.ledger import ( AccountType, Direction, LedgerAccount, ) - # BP's wallet and a revenue from their commissions account. - p1 = product_factory() + def _inner(): + # BP's wallet and a revenue from their commissions account. + p1 = product_factory() - account = LedgerAccount( - display_name=f"Revenue from {p1.name} commission", - qualified_name=f"{currency.value}:revenue:bp_commission:{p1.uuid}", - normal_balance=Direction.CREDIT, - account_type=AccountType.REVENUE, - reference_type="bp", - reference_uuid=p1.uuid, - currency=currency, - ) - lm.get_account_or_create(account=account) + account = LedgerAccount( + display_name=f"Revenue from {p1.name} commission", + qualified_name=f"{currency.value}:revenue:bp_commission:{p1.uuid}", + normal_balance=Direction.CREDIT, + account_type=AccountType.REVENUE, + reference_type="bp", + reference_uuid=p1.uuid, + currency=currency, + ) + lm.get_account_or_create(account=account) - account = LedgerAccount.model_validate( - { - "display_name": f"{p1.name} Wallet", - "qualified_name": f"{currency.value}:bp_wallet:{p1.uuid}", - "normal_balance": Direction.CREDIT, - "account_type": AccountType.BP_WALLET, - "reference_type": "bp", - "reference_uuid": p1.uuid, - "currency": currency, - } - ) - lm.get_account_or_create(account=account) + account = LedgerAccount.model_validate( + { + "display_name": f"{p1.name} Wallet", + "qualified_name": f"{currency.value}:bp_wallet:{p1.uuid}", + "normal_balance": Direction.CREDIT, + "account_type": AccountType.BP_WALLET, + "reference_type": "bp", + "reference_uuid": p1.uuid, + "currency": currency, + } + ) + lm.get_account_or_create(account=account) - # BP's wallet, user's wallet, and a revenue from their commissions account. - p2 = product_factory() - account = LedgerAccount( - display_name=f"Revenue from {p2.name} commission", - qualified_name=f"{currency.value}:revenue:bp_commission:{p2.uuid}", - normal_balance=Direction.CREDIT, - account_type=AccountType.REVENUE, - reference_type="bp", - reference_uuid=p2.uuid, - currency=currency, - ) - lm.get_account_or_create(account) + # BP's wallet, user's wallet, and a revenue from their commissions account. + p2 = product_factory() + account = LedgerAccount( + display_name=f"Revenue from {p2.name} commission", + qualified_name=f"{currency.value}:revenue:bp_commission:{p2.uuid}", + normal_balance=Direction.CREDIT, + account_type=AccountType.REVENUE, + reference_type="bp", + reference_uuid=p2.uuid, + currency=currency, + ) + lm.get_account_or_create(account) - account = LedgerAccount( - display_name=f"{p2.name} Wallet", - qualified_name=f"{currency.value}:bp_wallet:{p2.uuid}", - normal_balance=Direction.CREDIT, - account_type=AccountType.BP_WALLET, - reference_type="bp", - reference_uuid=p2.uuid, - currency=currency, - ) - lm.get_account_or_create(account) + account = LedgerAccount( + display_name=f"{p2.name} Wallet", + qualified_name=f"{currency.value}:bp_wallet:{p2.uuid}", + normal_balance=Direction.CREDIT, + account_type=AccountType.BP_WALLET, + reference_type="bp", + reference_uuid=p2.uuid, + currency=currency, + ) + lm.get_account_or_create(account) - account = LedgerAccount( - display_name=f"{user.uuid} Wallet", - qualified_name=f"{currency.value}:user_wallet:{user.uuid}", - normal_balance=Direction.CREDIT, - account_type=AccountType.USER_WALLET, - reference_type="user", - reference_uuid=user.uuid, - currency="test", - ) - lm.get_account_or_create(account=account) + account = LedgerAccount( + display_name=f"{user.uuid} Wallet", + qualified_name=f"{currency.value}:user_wallet:{user.uuid}", + normal_balance=Direction.CREDIT, + account_type=AccountType.USER_WALLET, + reference_type="user", + reference_uuid=user.uuid, + currency="test", + ) + lm.get_account_or_create(account=account) + + return _inner @pytest.fixture diff --git a/tests/incite/collections/test_df_collection_base.py b/tests/incite/collections/test_df_collection_base.py index 176bf4b..b9f0181 100644 --- a/tests/incite/collections/test_df_collection_base.py +++ b/tests/incite/collections/test_df_collection_base.py @@ -89,7 +89,7 @@ class TestDFCollectionBaseMethods: @pytest.mark.skip def test_initial_load(self, mnt_filepath: GRLDatasets, thl_web_rr): instance = DFCollection( - pg_config=thl_web_rr, + pg_config=thl_web_rr: PostgresConfig, data_type=DFCollectionType.USER, start=datetime(year=2022, month=1, day=1, minute=0, tzinfo=UTC), finished=datetime(year=2022, month=1, day=1, minute=5, tzinfo=UTC), diff --git a/tests/incite/collections/test_df_collection_item_base.py b/tests/incite/collections/test_df_collection_item_base.py index 0218f30..9a2ecf3 100644 --- a/tests/incite/collections/test_df_collection_item_base.py +++ b/tests/incite/collections/test_df_collection_item_base.py @@ -66,7 +66,7 @@ class TestDFCollectionItemMethods: start=datetime(year=1800, month=1, day=1, tzinfo=UTC), finished=datetime(year=1900, month=1, day=1, tzinfo=UTC), archive_path=mnt_filepath.archive_path(enum_type=df_coll_type), - pg_config=thl_web_rr, + pg_config=thl_web_rr: PostgresConfig, ) # Has RR, assume unittest server is online diff --git a/tests/incite/collections/test_df_collection_item_thl_web.py b/tests/incite/collections/test_df_collection_item_thl_web.py index 3d70e56..8038d3b 100644 --- a/tests/incite/collections/test_df_collection_item_thl_web.py +++ b/tests/incite/collections/test_df_collection_item_thl_web.py @@ -143,7 +143,7 @@ class TestDFCollectionItemMethod: offset: str, duration: timedelta, df_collection_data_type, - delete_df_collection, + delete_df_collection: Callable[..., None], ): delete_df_collection(coll=df_collection) @@ -173,7 +173,7 @@ class TestDFCollectionItemMethod: duration: timedelta, thl_web_rw: PostgresConfig, df_collection_data_type, - delete_df_collection, + delete_df_collection: Callable[..., None], ): # for i in collection.items: # assert i.update_partial_archive() @@ -186,15 +186,15 @@ class TestDFCollectionItemMethod: df_collection, offset: str, duration: str, - create_main_accounts, + create_main_accounts: Callable[..., None], thl_web_rw: PostgresConfig, thl_lm, df_collection_data_type, user_factory: Callable[..., User], - product: Product, + product: product: Product, client_no_amm, incite_item_factory, - delete_df_collection, + delete_df_collection: Callable[..., None], mnt_filepath: GRLDatasets, ): assert 1 + 1 == 2 @@ -205,7 +205,7 @@ class TestDFCollectionItemMethod: offset: str, duration: timedelta, df_collection, - delete_df_collection, + delete_df_collection: Callable[..., None], ): delete_df_collection(coll=df_collection) @@ -229,12 +229,12 @@ class TestDFCollectionItemMethod: df_collection, offset: str, duration: timedelta, - create_main_accounts, + create_main_accounts: Callable[..., None], thl_web_rw: PostgresConfig, user_factory: Callable[..., User], - product: Product, + product: product: Product, incite_item_factory, - delete_df_collection, + delete_df_collection: Callable[..., None], ): if df_collection.data_type in unsupported_mock_types: @@ -275,9 +275,9 @@ class TestDFCollectionItemMethod: offset: str, duration: timedelta, user_factory: Callable[..., User], - product: Product, + product: product: Product, incite_item_factory, - delete_df_collection, + delete_df_collection: Callable[..., None], ): if df_collection.data_type in unsupported_mock_types: @@ -318,17 +318,17 @@ class TestDFCollectionItemMethod: self, df_collection, user: User, - create_main_accounts, + create_main_accounts: Callable[..., None], offset: str, duration: timedelta, thl_web_rw: PostgresConfig, thl_lm, df_collection_data_type, user_factory: Callable[..., User], - product: Product, + product: product: Product, client_no_amm, incite_item_factory, - delete_df_collection, + delete_df_collection: Callable[..., None], mnt_filepath, ): @@ -376,10 +376,10 @@ class TestDFCollectionItemMethod: duration: timedelta, df_collection_data_type, user_factory: Callable[..., User], - product: Product, + product: product: Product, client_no_amm, incite_item_factory, - delete_df_collection, + delete_df_collection: Callable[..., None], mnt_filepath, ): @@ -410,13 +410,13 @@ class TestDFCollectionItemMethod: df_collection_data_type, df_collection, user_factory: Callable[..., User], - product: Product, + product: product: Product, offset: str, duration: timedelta, client_no_amm, user: User, incite_item_factory, - delete_df_collection, + delete_df_collection: Callable[..., None], mnt_filepath, ): """We already have a test for the "non-private" version of this, @@ -757,9 +757,9 @@ class TestDFCollectionItemFunctionalTest: df_collection, user: User, user_factory: Callable[..., User], - product: Product, + product: product: Product, incite_item_factory, - delete_df_collection, + delete_df_collection: Callable[..., None], mnt_filepath: GRLDatasets, ): @@ -805,10 +805,10 @@ class TestDFCollectionItemFunctionalTest: duration: timedelta, client_no_amm, user_factory: Callable[..., User], - product: Product, + product: product: Product, df_collection_data_type, incite_item_factory, - delete_df_collection, + delete_df_collection: Callable[..., None], mnt_filepath: GRLDatasets, ): """A functional test to write some Parquet files for the @@ -823,7 +823,6 @@ class TestDFCollectionItemFunctionalTest: import pyarrow.parquet as pq - if df_collection.data_type in unsupported_mock_types: return delete_df_collection(coll=df_collection) @@ -850,12 +849,12 @@ class TestDFCollectionItemFunctionalTest: client_no_amm, df_collection, user_factory: Callable[..., User], - product: Product, + product: product: Product, offset: str, duration: timedelta, df_collection_data_type, incite_item_factory, - delete_df_collection, + delete_df_collection: Callable[..., None], mnt_filepath: GRLDatasets, ): @@ -886,7 +885,7 @@ class TestDFCollectionItemFunctionalTest: @pytest.mark.skip def test_get_items( - self, df_collection, product: Product, offset: str, duration: timedelta + self, df_collection, product: product: Product, offset: str, duration: timedelta ): with pytest.warns(expected_warning=ResourceWarning) as cm: df_collection.get_items_last365() @@ -903,9 +902,9 @@ class TestDFCollectionItemFunctionalTest: df_collection_data_type, df_collection, incite_item_factory, - delete_df_collection, + delete_df_collection: Callable[..., None], user_factory: Callable[..., User], - product: Product, + product: product: Product, offset: str, duration: timedelta, mnt_filepath: GRLDatasets, @@ -944,7 +943,7 @@ class TestDFCollectionItemFunctionalTest: df_collection_data_type, df_collection, incite_item_factory, - delete_df_collection, + delete_df_collection: Callable[..., None], user: User, offset: str, duration: timedelta, @@ -972,9 +971,9 @@ class TestDFCollectionItemFunctionalTest: df_collection_data_type, df_collection, incite_item_factory, - delete_df_collection, + delete_df_collection: Callable[..., None], user_factory: Callable[..., User], - product: Product, + product: product: Product, offset: str, duration: timedelta, mnt_filepath, diff --git a/tests/incite/mergers/foundations/test_enriched_session.py b/tests/incite/mergers/foundations/test_enriched_session.py index ba11725..8254d81 100644 --- a/tests/incite/mergers/foundations/test_enriched_session.py +++ b/tests/incite/mergers/foundations/test_enriched_session.py @@ -26,20 +26,20 @@ class TestEnrichedSession: def test_base( self, client_no_amm, - product, - user_factory, + product: Product, + user_factory: Callable[..., User], wall_collection, session_collection, enriched_session_merge, thl_web_rr: PostgresConfig, - delete_df_collection, + delete_df_collection: Callable[..., None], incite_item_factory, ): from generalresearch.models.thl.user import User delete_df_collection(coll=session_collection) - u1: User = user_factory(product=product, created=session_collection.start) + u1: User = user_factory(product=product: Product, created=session_collection.start) for item in session_collection.items: incite_item_factory(item=item, user=u1) @@ -52,7 +52,7 @@ class TestEnrichedSession: client=client_no_amm, wall_coll=wall_collection, session_coll=session_collection, - pg_config=thl_web_rr, + pg_config=thl_web_rr: PostgresConfig, ) # -- @@ -92,11 +92,11 @@ class TestEnrichedSessionAdmin: session_collection, thl_web_rr: PostgresConfig, session_report_request, - user_factory, + user_factory: Callable[..., User], start, session_factory, - product_factory, - delete_df_collection, + product_factory: Callable[..., Product], + delete_df_collection: Callable[..., None], ): delete_df_collection(coll=wall_collection) delete_df_collection(coll=session_collection) @@ -120,7 +120,7 @@ class TestEnrichedSessionAdmin: client=client_no_amm, session_coll=session_collection, wall_coll=wall_collection, - pg_config=thl_web_rr, + pg_config=thl_web_rr: PostgresConfig, ) df = enriched_session_merge.to_admin_response( diff --git a/tests/incite/mergers/foundations/test_enriched_task_adjust.py b/tests/incite/mergers/foundations/test_enriched_task_adjust.py index 8c3a647..a33a55a 100644 --- a/tests/incite/mergers/foundations/test_enriched_task_adjust.py +++ b/tests/incite/mergers/foundations/test_enriched_task_adjust.py @@ -21,16 +21,16 @@ class TestEnrichedTaskAdjust: def test_base( self, client_no_amm, - user_factory, - product, + user_factory: Callable[..., User], + product: Product, task_adj_collection, wall_collection, session_collection, enriched_wall_merge, enriched_task_adjust_merge, incite_item_factory, - delete_df_collection, - thl_web_rr, + delete_df_collection: Callable[..., None], + thl_web_rr: PostgresConfig, ): from generalresearch.models.thl.user import User @@ -48,14 +48,14 @@ class TestEnrichedTaskAdjust: client=client_no_amm, session_coll=session_collection, wall_coll=wall_collection, - pg_config=thl_web_rr, + pg_config=thl_web_rr: PostgresConfig, ) enriched_task_adjust_merge.build( client=client_no_amm, task_adjust_coll=task_adj_collection, enriched_wall=enriched_wall_merge, - pg_config=thl_web_rr, + pg_config=thl_web_rr: PostgresConfig, ) # -- diff --git a/tests/incite/mergers/foundations/test_enriched_wall.py b/tests/incite/mergers/foundations/test_enriched_wall.py index 0e28bce..a0ca4dd 100644 --- a/tests/incite/mergers/foundations/test_enriched_wall.py +++ b/tests/incite/mergers/foundations/test_enriched_wall.py @@ -21,13 +21,13 @@ class TestEnrichedWall: def test_base( self, client_no_amm, - product, - user_factory, + product: Product, + user_factory: Callable[..., User], wall_collection, - thl_web_rr, + thl_web_rr: PostgresConfig, session_collection, enriched_wall_merge, - delete_df_collection, + delete_df_collection: Callable[..., None], incite_item_factory, ): from generalresearch.models.thl.user import User @@ -35,7 +35,7 @@ class TestEnrichedWall: # -- Build & Setup delete_df_collection(coll=session_collection) delete_df_collection(coll=wall_collection) - u1: User = user_factory(product=product, created=session_collection.start) + u1: User = user_factory(product=product: Product, created=session_collection.start) for item in session_collection.items: incite_item_factory(item=item, user=u1) @@ -48,7 +48,7 @@ class TestEnrichedWall: client=client_no_amm, wall_coll=wall_collection, session_coll=session_collection, - pg_config=thl_web_rr, + pg_config=thl_web_rr: PostgresConfig, ) # -- @@ -64,18 +64,18 @@ class TestEnrichedWall: def test_base_item( self, client_no_amm, - product, - user_factory, + product: Product, + user_factory: Callable[..., User], wall_collection, session_collection, enriched_wall_merge, - delete_df_collection, - thl_web_rr, + delete_df_collection: Callable[..., None], + thl_web_rr: PostgresConfig, incite_item_factory, ): # -- Build & Setup delete_df_collection(coll=session_collection) - u = user_factory(product=product, created=session_collection.start) + u = user_factory(product=product: Product, created=session_collection.start) for item in session_collection.items: incite_item_factory(item=item, user=u) @@ -87,7 +87,7 @@ class TestEnrichedWall: client=client_no_amm, wall_coll=wall_collection, session_coll=session_collection, - pg_config=thl_web_rr, + pg_config=thl_web_rr: PostgresConfig, ) # -- @@ -106,7 +106,7 @@ class TestEnrichedWall: client=client_no_amm, wall_coll=wall_collection, session_coll=session_collection, - pg_config=thl_web_rr, + pg_config=thl_web_rr: PostgresConfig, ) modified_time2 = path.stat().st_mtime @@ -172,12 +172,12 @@ class TestEnrichedWallToAdmin: client_no_amm, wall_collection, session_collection, - thl_web_rr, + thl_web_rr: PostgresConfig, user, session_factory, - delete_df_collection, - product_factory, - user_factory, + delete_df_collection: Callable[..., None], + product_factory: Callable[..., Product], + user_factory: Callable[..., User], start, ): delete_df_collection(coll=wall_collection) @@ -203,7 +203,7 @@ class TestEnrichedWallToAdmin: client=client_no_amm, wall_coll=wall_collection, session_coll=session_collection, - pg_config=thl_web_rr, + pg_config=thl_web_rr: PostgresConfig, ) df = enriched_wall_merge.to_admin_response( diff --git a/tests/incite/mergers/test_pop_ledger.py b/tests/incite/mergers/test_pop_ledger.py index dc01179..d054eb6 100644 --- a/tests/incite/mergers/test_pop_ledger.py +++ b/tests/incite/mergers/test_pop_ledger.py @@ -33,17 +33,17 @@ class TestMergePOPLedger: client_no_amm, ledger_collection, pop_ledger_merge, - product, - user_factory, - create_main_accounts, + product: Product, + user_factory: Callable[..., User], + create_main_accounts: Callable[..., None], thl_lm, - delete_df_collection, + delete_df_collection: Callable[..., None], incite_item_factory, - delete_ledger_db, + delete_ledger_db: Callable[..., None], ): from generalresearch.models.thl.ledger import LedgerAccount - u = user_factory(product=product, created=ledger_collection.start) + u = user_factory(product=product: Product, created=ledger_collection.start) # -- Build & Setup delete_ledger_db() @@ -127,26 +127,26 @@ class TestMergePOPLedger: ledger_collection, pop_ledger_merge, mnt_filepath, - product, - user_factory, - create_main_accounts, + product: Product, + user_factory: Callable[..., User], + create_main_accounts: Callable[..., None], offset, duration, start, thl_lm, incite_item_factory, - delete_df_collection, - delete_ledger_db, + delete_df_collection: Callable[..., None], + delete_ledger_db: Callable[..., None], session_collection, ): from generalresearch.models.thl.finance import ProductBalances from generalresearch.models.thl.ledger import LedgerAccount from generalresearch.models.thl.product import Product - u = user_factory(product=product, created=session_collection.start) + u = user_factory(product=product: Product, created=session_collection.start) assert ledger_collection.finished is not None - assert isinstance(u.product, Product) + assert isinstance(u.product: Product, Product) delete_ledger_db() create_main_accounts(), delete_df_collection(coll=ledger_collection) @@ -228,14 +228,14 @@ class TestMergePOPLedger: ledger_collection, pop_ledger_merge, mnt_filepath, - user_factory, - product, - create_main_accounts, + user_factory: Callable[..., User], + product: Product, + create_main_accounts: Callable[..., None], offset, duration, start, thl_lm, - delete_df_collection, + delete_df_collection: Callable[..., None], incite_item_factory, ): from generalresearch.models.thl.user import User diff --git a/tests/incite/mergers/test_ym_survey_merge.py b/tests/incite/mergers/test_ym_survey_merge.py index 850df8a..a0b8b87 100644 --- a/tests/incite/mergers/test_ym_survey_merge.py +++ b/tests/incite/mergers/test_ym_survey_merge.py @@ -28,20 +28,20 @@ class TestYMSurveyMerge: def test_base( self, client_no_amm, - user_factory, - product, + user_factory: Callable[..., User], + product: Product, ym_survey_wall_merge, wall_collection, session_collection, enriched_session_merge, - delete_df_collection, + delete_df_collection: Callable[..., None], incite_item_factory, - thl_web_rr, + thl_web_rr: PostgresConfig, ): from generalresearch.models.thl.user import User delete_df_collection(coll=session_collection) - user: User = user_factory(product=product, created=session_collection.start) + user: User = user_factory(product=product: Product, created=session_collection.start) # -- Build & Setup assert ym_survey_wall_merge.start is None @@ -61,7 +61,7 @@ class TestYMSurveyMerge: client=client_no_amm, session_coll=session_collection, wall_coll=wall_collection, - pg_config=thl_web_rr, + pg_config=thl_web_rr: PostgresConfig, ) assert enriched_session_merge.progress.has_archive.eq(True).all() diff --git a/tests/managers/gr/test_business.py b/tests/managers/gr/test_business.py index 74a5450..3490403 100644 --- a/tests/managers/gr/test_business.py +++ b/tests/managers/gr/test_business.py @@ -8,7 +8,7 @@ class TestBusinessBankAccountManager: def test_init(self, business_bank_account_manager, gr_db): assert business_bank_account_manager.pg_config == gr_db - def test_create(self, business, business_bank_account_manager): + def test_create(self, business: Business, business_bank_account_manager): from generalresearch.models.gr.business import ( BusinessBankAccount, TransferMethod, @@ -33,7 +33,7 @@ class TestBusinessBankAccountManager: class TestBusinessAddressManager: - def test_create(self, business, business_address_manager): + def test_create(self, business: Business, business_address_manager): from generalresearch.models.gr.business import BusinessAddress res = business_address_manager.create(uuid=uuid4().hex, business_id=business.id) @@ -81,7 +81,7 @@ class TestBusinessManager: res = business_manager.get_by_user_id(user_id=gr_user.id) assert len(res) == 0 - # Create a Business, but don't add it to anything + # Create a business: Business, but don't add it to anything b1 = business_manager.create_dummy() res = business_manager.get_by_user_id(user_id=gr_user.id) assert len(res) == 0 @@ -113,11 +113,11 @@ class TestBusinessManager: def test_get_uuids_by_user_id(self): pass - def test_get_by_uuid(self, business, business_manager): + def test_get_by_uuid(self, business: Business, business_manager): instance = business_manager.get_by_uuid(business_uuid=business.uuid) assert business.id == instance.id - def test_get_by_id(self, business, business_manager): + def test_get_by_id(self, business: Business, business_manager): instance = business_manager.get_by_id(business_id=business.id) assert business.uuid == instance.uuid @@ -131,7 +131,7 @@ class TestBusinessManager: # business = BusinessManager.create( # uuid=b_uuid, # name=f"test-{b_uuid[:6]}") - # assert isinstance(business, Business) + # assert isinstance(business: Business, Business) # # # Try to make it again # with pytest.raises(expected_exception=psycopg.errors.UniqueViolation): diff --git a/tests/managers/gr/test_team.py b/tests/managers/gr/test_team.py index 0918ab8..5e5c565 100644 --- a/tests/managers/gr/test_team.py +++ b/tests/managers/gr/test_team.py @@ -89,10 +89,10 @@ class TestTeamManager: gr_user_token, gr_user, membership, - product_factory, + product_factory: Callable[..., Product], membership_factory, team, - thl_web_rr, + thl_web_rr: PostgresConfig, gr_redis_config, gr_db, ): diff --git a/tests/managers/leaderboard.py b/tests/managers/leaderboard.py index 7773030..3d1818b 100644 --- a/tests/managers/leaderboard.py +++ b/tests/managers/leaderboard.py @@ -19,7 +19,7 @@ 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 diff --git a/tests/managers/test_events.py b/tests/managers/test_events.py index cc9a1bf..a6d3a6b 100644 --- a/tests/managers/test_events.py +++ b/tests/managers/test_events.py @@ -35,8 +35,8 @@ def user_factory(product_id): @pytest.fixture(scope="function") -def event_subscriber(thl_redis_config, product_id): - return EventSubscriber(redis_config=thl_redis_config, product_id=product_id) +def event_subscriber(thl_redis_config: RedisConfig, product_id): + return EventSubscriber(redis_config=thl_redis_config: RedisConfig, product_id=product_id) def create_dummy( @@ -185,7 +185,7 @@ class TestSessionStats: "session_fail_avg_loi_last_24h": None, } - def test_run(self, event_manager, product_id, user_factory, utc_now, utc_hour_ago): + def test_run(self, event_manager, product_id, user_factory: Callable[..., User], utc_now, utc_hour_ago): event_manager.clear_global_session_stats() user: User = user_factory() @@ -448,7 +448,7 @@ class TestChannelsSubscriptions: event_manager, event_subscriber, product_id, - user_factory, + user_factory: Callable[..., User], utc_hour_ago, utc_now, ): diff --git a/tests/managers/test_lucid.py b/tests/managers/test_lucid.py index 1a1bae7..654b58d 100644 --- a/tests/managers/test_lucid.py +++ b/tests/managers/test_lucid.py @@ -10,7 +10,7 @@ class TestLucidProfiling: @pytest.mark.skip def test_get_library(self, thl_web_rr): pks = [(qid, "us", "eng") for qid in qids] - qs = get_profiling_library(thl_web_rr, pks=pks) + qs = get_profiling_library(thl_web_rr: PostgresConfig, pks=pks) assert len(qids) == len(qs) # just making sure this doesn't raise errors @@ -19,5 +19,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, country_iso="mx", language_iso="spa") + qs = get_profiling_library(thl_web_rr: PostgresConfig, country_iso="mx", language_iso="spa") assert len(qs) > 100 diff --git a/tests/managers/thl/test_contest/test_leaderboard.py b/tests/managers/thl/test_contest/test_leaderboard.py index 7adea9c..07d8d74 100644 --- a/tests/managers/thl/test_contest/test_leaderboard.py +++ b/tests/managers/thl/test_contest/test_leaderboard.py @@ -1,7 +1,13 @@ +from __future__ import annotations + from datetime import UTC, datetime, timedelta from zoneinfo import ZoneInfo from generalresearch.currency import USDCent +from generalresearch.managers.thl.contest_manager import ContestManager +from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager +from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager +from generalresearch.managers.thl.user_manager.user_manager import UserManager from generalresearch.models.thl.contest.definitions import ( ContestEndReason, ContestStatus, @@ -12,6 +18,7 @@ from generalresearch.models.thl.contest.leaderboard import ( ) from generalresearch.models.thl.product import Product from generalresearch.models.thl.user import User +from generalresearch.redis_helper import RedisConfig class TestLeaderboardContestCRUD: @@ -20,8 +27,8 @@ class TestLeaderboardContestCRUD: self, contest_create: LeaderboardContestCreate, product_user_wallet_yes: Product, - thl_lm, - contest_manager, + thl_ledger_manager: ThlLedgerManager, + contest_manager: ContestManager, ): c = contest_manager.create( product_id=product_user_wallet_yes.uuid, contest_create=contest_create @@ -41,10 +48,10 @@ class TestLeaderboardContestCRUD: self, user_with_wallet: User, contest_in_db: LeaderboardContest, - thl_lm, - contest_manager, - user_manager, - thl_redis, + thl_ledger_manager: ThlLedgerManager, + contest_manager: ContestManager, + user_manager: UserManager, + thl_redis: RedisConfig, ): contest = contest_in_db user = user_with_wallet @@ -74,10 +81,10 @@ class TestLeaderboardContestCRUD: self, user_with_wallet: User, contest_in_db: LeaderboardContest, - thl_lm, - contest_manager, - user_manager, - thl_redis, + thl_ledger_manager: ThlLedgerManager, + contest_manager: ContestManager, + user_manager: UserManager, + thl_redis: RedisConfig, ): # The contest should be over. We need to trigger it. contest = contest_in_db @@ -96,11 +103,13 @@ class TestLeaderboardContestCRUD: ) assert c.user_rank == 1 - bp_wallet = thl_lm.get_account_or_create_bp_wallet_by_uuid(user.product_id) - bp_wallet_balance = thl_lm.get_account_balance(account=bp_wallet) + bp_wallet = thl_ledger_manager.get_account_or_create_bp_wallet_by_uuid( + user.product_id + ) + bp_wallet_balance = thl_ledger_manager.get_account_balance(account=bp_wallet) assert bp_wallet_balance == 0 - user_wallet = thl_lm.get_account_or_create_user_wallet(user=user) - user_balance = thl_lm.get_account_balance(user_wallet) + user_wallet = thl_ledger_manager.get_account_or_create_user_wallet(user=user) + user_balance = thl_ledger_manager.get_account_balance(user_wallet) assert user_balance == 0 decision, reason = contest.should_end() @@ -125,10 +134,12 @@ class TestLeaderboardContestCRUD: assert w.prize.cash_amount == USDCent(15_00) # The prize is $15.00, so the user should get $15, paid by the bp - assert thl_lm.get_account_balance(account=user_wallet) == 15_00 + assert thl_ledger_manager.get_account_balance(account=user_wallet) == 15_00 # contest wallet is 0, and the BP gets 20c - contest_wallet = thl_lm.get_account_or_create_contest_wallet_by_uuid( - contest_uuid=c.uuid + contest_wallet = ( + thl_ledger_manager.get_account_or_create_contest_wallet_by_uuid( + contest_uuid=c.uuid + ) ) - assert thl_lm.get_account_balance(account=contest_wallet) == 0 - assert thl_lm.get_account_balance(account=bp_wallet) == -15_00 + assert thl_ledger_manager.get_account_balance(account=contest_wallet) == 0 + assert thl_ledger_manager.get_account_balance(account=bp_wallet) == -15_00 diff --git a/tests/managers/thl/test_contest/test_milestone.py b/tests/managers/thl/test_contest/test_milestone.py index ed0bbb5..a2d575b 100644 --- a/tests/managers/thl/test_contest/test_milestone.py +++ b/tests/managers/thl/test_contest/test_milestone.py @@ -1,5 +1,10 @@ +from __future__ import annotations + +from collections.abc import Callable from datetime import UTC, datetime +from generalresearch.managers.thl.contest_manager import ContestManager +from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager from generalresearch.models.thl.contest.definitions import ( ContestEndReason, ContestStatus, @@ -16,7 +21,12 @@ from generalresearch.models.thl.user import User class TestMilestoneContest: - def test_should_end(self, contest: MilestoneContest, thl_lm, contest_manager): + def test_should_end( + self, + contest: MilestoneContest, + thl_ledger_manager: ThlLedgerManager, + contest_manager: ContestManager, + ): # contest is active and has no entries should, msg = contest.should_end() assert not should, msg @@ -42,8 +52,8 @@ class TestMilestoneContestCRUD: self, contest_create: MilestoneContestCreate, product_user_wallet_yes: Product, - thl_lm, - contest_manager, + thl_ledger_manager: ThlLedgerManager, + contest_manager: ContestManager, ): c = contest_manager.create( product_id=product_user_wallet_yes.uuid, contest_create=contest_create @@ -63,8 +73,8 @@ class TestMilestoneContestCRUD: self, user_with_wallet: User, contest_in_db: MilestoneContest, - thl_lm, - contest_manager, + thl_ledger_manager: ThlLedgerManager, + contest_manager: ContestManager, ): # Users CANNOT directly enter a milestone contest through the api, # but we'll call this manager method when a trigger is hit. @@ -75,7 +85,7 @@ class TestMilestoneContestCRUD: contest_uuid=contest.uuid, user=user, country_iso="us", - ledger_manager=thl_lm, + ledger_manager=thl_ledger_manager, incr=1, ) @@ -90,17 +100,19 @@ class TestMilestoneContestCRUD: assert c.user_amount == 1 # Contest wallet should have 0 bc there is no ledger - contest_wallet = thl_lm.get_account_or_create_contest_wallet_by_uuid( - contest_uuid=contest.uuid + contest_wallet = ( + thl_ledger_manager.get_account_or_create_contest_wallet_by_uuid( + contest_uuid=contest.uuid + ) ) - assert thl_lm.get_account_balance(contest_wallet) == 0 + assert thl_ledger_manager.get_account_balance(contest_wallet) == 0 # Enter again! contest_manager.enter_milestone_contest( contest_uuid=contest.uuid, user=user, country_iso="us", - ledger_manager=thl_lm, + ledger_manager=thl_ledger_manager, incr=1, ) c: MilestoneUserView = contest_manager.get_milestone_user_view( @@ -117,20 +129,20 @@ class TestMilestoneContestCRUD: self, user_with_wallet: User, contest_in_db: MilestoneContest, - thl_lm, - contest_manager, + thl_ledger_manager: ThlLedgerManager, + contest_manager: ContestManager, ): # User enters contest, which brings the USER'S total amount above the limit, # and the user reaches the milestone contest = contest_in_db user = user_with_wallet - user_wallet = thl_lm.get_account_or_create_user_wallet(user=user) - user_balance = thl_lm.get_account_balance(account=user_wallet) - bp_wallet = thl_lm.get_account_or_create_bp_wallet_by_uuid( + user_wallet = thl_ledger_manager.get_account_or_create_user_wallet(user=user) + user_balance = thl_ledger_manager.get_account_balance(account=user_wallet) + bp_wallet = thl_ledger_manager.get_account_or_create_bp_wallet_by_uuid( product_uuid=user.product_id ) - bp_wallet_balance = thl_lm.get_account_balance(account=bp_wallet) + bp_wallet_balance = thl_ledger_manager.get_account_balance(account=bp_wallet) c: MilestoneUserView = contest_manager.get_milestone_user_view( contest_uuid=contest.uuid, user=user_with_wallet @@ -145,7 +157,7 @@ class TestMilestoneContestCRUD: contest_uuid=contest.uuid, user=user, country_iso="us", - ledger_manager=thl_lm, + ledger_manager=thl_ledger_manager, incr=1, ) @@ -165,9 +177,12 @@ class TestMilestoneContestCRUD: assert c.win_count == 1 # The prize was awarded! User should have won $1.00 - assert thl_lm.get_account_balance(user_wallet) - user_balance == 100 + assert thl_ledger_manager.get_account_balance(user_wallet) - user_balance == 100 # Which was paid from the BP's balance - assert thl_lm.get_account_balance(bp_wallet) - bp_wallet_balance == -100 + assert ( + thl_ledger_manager.get_account_balance(bp_wallet) - bp_wallet_balance + == -100 + ) # winnings = cm.get_winnings_by_user(user=user) # assert len(winnings) == 1 @@ -176,11 +191,11 @@ class TestMilestoneContestCRUD: def test_enter_ends( self, - user_factory, + user_factory: Callable[..., User], product_user_wallet_yes: Product, contest_in_db: MilestoneContest, - thl_lm, - contest_manager, + thl_ledger_manager: ThlLedgerManager, + contest_manager: ContestManager, ): # Multiple users reach the milestone. Contest ends after 5 wins. users = [user_factory(product=product_user_wallet_yes) for _ in range(5)] @@ -191,7 +206,7 @@ class TestMilestoneContestCRUD: contest_uuid=contest.uuid, user=u, country_iso="us", - ledger_manager=thl_lm, + ledger_manager=thl_ledger_manager, incr=3, ) @@ -203,15 +218,15 @@ class TestMilestoneContestCRUD: self, user_with_wallet: User, contest_in_db: MilestoneContest, - thl_lm, - contest_manager, + thl_ledger_manager: ThlLedgerManager, + contest_manager: ContestManager, ): # Pretend user just got a complete cnt = contest_manager.hit_milestone_triggers( country_iso="us", user=user_with_wallet, event=ContestEntryTrigger.TASK_COMPLETE, - ledger_manager=thl_lm, + ledger_manager=thl_ledger_manager, ) assert cnt == 1 @@ -224,7 +239,11 @@ class TestMilestoneContestCRUD: class TestMilestoneContestUserViews: def test_list_user_eligible_country( - self, user_with_wallet: User, contest_factory, thl_lm, contest_manager + self, + user_with_wallet: User, + contest_factory: Callable[..., Contest], + thl_ledger_manager: ThlLedgerManager, + contest_manager: ContestManager, ): # No contests exists cs = contest_manager.get_many_by_user_eligible( @@ -257,7 +276,11 @@ class TestMilestoneContestUserViews: assert len(cs) == 2 def test_list_user_eligible( - self, user_with_money: User, contest_factory, thl_lm, contest_manager + self, + user_with_money: User, + contest_factory: Callable[..., Contest], + thl_ledger_manager: ThlLedgerManager, + contest_manager: ContestManager, ): # User reaches milestone after 1 complete c = contest_factory(target_amount=1) diff --git a/tests/managers/thl/test_contest/test_raffle.py b/tests/managers/thl/test_contest/test_raffle.py index 736a5e9..b435576 100644 --- a/tests/managers/thl/test_contest/test_raffle.py +++ b/tests/managers/thl/test_contest/test_raffle.py @@ -1,3 +1,6 @@ +from __future__ import annotations + +from collections.abc import Callable from datetime import UTC, datetime import pytest @@ -5,9 +8,11 @@ from pydantic import ValidationError from pytest import approx from generalresearch.currency import USDCent +from generalresearch.managers.thl.contest_manager import ContestManager from generalresearch.managers.thl.ledger_manager.exceptions import ( LedgerTransactionConditionFailedError, ) +from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager from generalresearch.models.thl.contest import ( ContestEndCondition, ContestEntryRule, @@ -32,7 +37,12 @@ from generalresearch.models.thl.user import User class TestRaffleContest: - def test_should_end(self, contest: RaffleContest, thl_lm, contest_manager): + def test_should_end( + self, + contest: RaffleContest, + thl_ledger_manager: ThlLedgerManager, + contest_manager: ContestManager, + ): # contest is active and has no entries should, msg = contest.should_end() assert not should, msg @@ -57,8 +67,8 @@ class TestRaffleContestCRUD: self, contest_create: RaffleContestCreate, product_user_wallet_yes: Product, - thl_lm, - contest_manager, + thl_ledger_manager: ThlLedgerManager, + contest_manager: ContestManager, ): c = contest_manager.create( product_id=product_user_wallet_yes.uuid, contest_create=contest_create @@ -78,8 +88,8 @@ class TestRaffleContestCRUD: self, user_with_money: User, contest_in_db: RaffleContest, - thl_lm, - contest_manager, + thl_ledger_manager: ThlLedgerManager, + contest_manager: ContestManager, ): # Raffle ends at $1.00. User enters for $0.60 print(user_with_money.product_id) @@ -87,8 +97,10 @@ class TestRaffleContestCRUD: print(contest_in_db.uuid) contest = contest_in_db - user_wallet = thl_lm.get_account_or_create_user_wallet(user=user_with_money) - user_balance = thl_lm.get_account_balance(account=user_wallet) + user_wallet = thl_ledger_manager.get_account_or_create_user_wallet( + user=user_with_money + ) + user_balance = thl_ledger_manager.get_account_balance(account=user_wallet) entry = ContestEntry( entry_type=ContestEntryType.CASH, user=user_with_money, amount=USDCent(60) @@ -97,7 +109,7 @@ class TestRaffleContestCRUD: contest_uuid=contest.uuid, entry=entry, country_iso="us", - ledger_manager=thl_lm, + ledger_manager=thl_ledger_manager, ) c: RaffleContest = contest_manager.get(contest_uuid=contest.uuid) assert c.current_amount == USDCent(60) @@ -112,30 +124,35 @@ class TestRaffleContestCRUD: assert c.projected_win_probability == approx(60 / 100, rel=0.01) # Contest wallet should have $0.60 - contest_wallet = thl_lm.get_account_or_create_contest_wallet_by_uuid( - contest_uuid=contest.uuid + contest_wallet = ( + thl_ledger_manager.get_account_or_create_contest_wallet_by_uuid( + contest_uuid=contest.uuid + ) ) - assert thl_lm.get_account_balance(account=contest_wallet) == 60 + assert thl_ledger_manager.get_account_balance(account=contest_wallet) == 60 # User spent 60c - assert user_balance - thl_lm.get_account_balance(account=user_wallet) == 60 + assert ( + user_balance - thl_ledger_manager.get_account_balance(account=user_wallet) + == 60 + ) @pytest.mark.parametrize("user_with_money", [{"min_balance": 120}], indirect=True) def test_enter_ends( self, user_with_money: User, contest_in_db: RaffleContest, - thl_lm, - contest_manager, + thl_ledger_manager: ThlLedgerManager, + contest_manager: ContestManager, ): # User enters contest, which brings the total amount above the limit, # and the contest should end, with a winner selected contest = contest_in_db - bp_wallet = thl_lm.get_account_or_create_bp_wallet_by_uuid( + bp_wallet = thl_ledger_manager.get_account_or_create_bp_wallet_by_uuid( user_with_money.product_id ) # I bribed the user, so the balance is not 0 - bp_wallet_balance = thl_lm.get_account_balance(account=bp_wallet) + bp_wallet_balance = thl_ledger_manager.get_account_balance(account=bp_wallet) for _ in range(2): entry = ContestEntry( @@ -147,7 +164,7 @@ class TestRaffleContestCRUD: contest_uuid=contest.uuid, entry=entry, country_iso="us", - ledger_manager=thl_lm, + ledger_manager=thl_ledger_manager, ) c: RaffleContest = contest_manager.get(contest_uuid=contest.uuid) assert c.status == ContestStatus.COMPLETED @@ -167,21 +184,29 @@ class TestRaffleContestCRUD: assert win.product_user_id == user_with_money.product_user_id # Contest wallet should have gotten zeroed out - contest_wallet = thl_lm.get_account_or_create_contest_wallet_by_uuid( - contest_uuid=contest.uuid + contest_wallet = ( + thl_ledger_manager.get_account_or_create_contest_wallet_by_uuid( + contest_uuid=contest.uuid + ) ) - assert thl_lm.get_account_balance(contest_wallet) == 0 + assert thl_ledger_manager.get_account_balance(contest_wallet) == 0 # Expense wallet gets the $1.00 expense - expense_wallet = thl_lm.get_account_or_create_bp_expense_by_uuid( + expense_wallet = thl_ledger_manager.get_account_or_create_bp_expense_by_uuid( product_uuid=user_with_money.product_id, expense_name="Prize" ) - assert thl_lm.get_account_balance(expense_wallet) == -100 + assert thl_ledger_manager.get_account_balance(expense_wallet) == -100 # And the BP gets 20c - assert thl_lm.get_account_balance(bp_wallet) - bp_wallet_balance == 20 + assert ( + thl_ledger_manager.get_account_balance(bp_wallet) - bp_wallet_balance == 20 + ) @pytest.mark.parametrize("user_with_money", [{"min_balance": 120}], indirect=True) def test_enter_ends_cash_prize( - self, user_with_money: User, contest_factory, thl_lm, contest_manager + self, + user_with_money: User, + contest_factory: Callable[..., Contest], + thl_ledger_manager: ThlLedgerManager, + contest_manager: ContestManager, ): # Same as test_enter_ends, but the prize is cash. Just # testing the ledger methods @@ -197,12 +222,14 @@ class TestRaffleContestCRUD: ) assert c.prizes[0].kind == ContestPrizeKind.CASH - user_wallet = thl_lm.get_account_or_create_user_wallet(user=user_with_money) - user_balance = thl_lm.get_account_balance(user_wallet) - bp_wallet = thl_lm.get_account_or_create_bp_wallet_by_uuid( + user_wallet = thl_ledger_manager.get_account_or_create_user_wallet( + user=user_with_money + ) + user_balance = thl_ledger_manager.get_account_balance(user_wallet) + bp_wallet = thl_ledger_manager.get_account_or_create_bp_wallet_by_uuid( user_with_money.product_id ) - bp_wallet_balance = thl_lm.get_account_balance(bp_wallet) + bp_wallet_balance = thl_ledger_manager.get_account_balance(bp_wallet) ## Enter Contest entry = ContestEntry( @@ -212,26 +239,33 @@ class TestRaffleContestCRUD: contest_uuid=c.uuid, entry=entry, country_iso="us", - ledger_manager=thl_lm, + ledger_manager=thl_ledger_manager, ) # The prize is $1.00, so the user spent $1.20 entering, won, then got $1.00 back assert ( - thl_lm.get_account_balance(account=user_wallet) == user_balance + 100 - 120 + thl_ledger_manager.get_account_balance(account=user_wallet) + == user_balance + 100 - 120 ) # contest wallet is 0, and the BP gets 20c - contest_wallet = thl_lm.get_account_or_create_contest_wallet_by_uuid( - contest_uuid=c.uuid + contest_wallet = ( + thl_ledger_manager.get_account_or_create_contest_wallet_by_uuid( + contest_uuid=c.uuid + ) + ) + assert thl_ledger_manager.get_account_balance(account=contest_wallet) == 0 + assert ( + thl_ledger_manager.get_account_balance(account=bp_wallet) + - bp_wallet_balance + == 20 ) - assert thl_lm.get_account_balance(account=contest_wallet) == 0 - assert thl_lm.get_account_balance(account=bp_wallet) - bp_wallet_balance == 20 def test_enter_failure( self, user_with_wallet: User, contest_in_db: RaffleContest, - thl_lm, - contest_manager, + thl_ledger_manager: ThlLedgerManager, + contest_manager: ContestManager, ): c = contest_in_db user = user_with_wallet @@ -252,7 +286,7 @@ class TestRaffleContestCRUD: contest_uuid=c.uuid, entry=entry, country_iso="us", - ledger_manager=thl_lm, + ledger_manager=thl_ledger_manager, ) assert e.value.args[0] == "insufficient balance" @@ -263,13 +297,17 @@ class TestRaffleContestCRUD: contest_uuid=c.uuid, entry=entry, country_iso="us", - ledger_manager=thl_lm, + ledger_manager=thl_ledger_manager, ) assert "incompatible entry type" in str(e.value) @pytest.mark.parametrize("user_with_money", [{"min_balance": 100}], indirect=True) def test_enter_not_eligible( - self, user_with_money: User, contest_factory, thl_lm, contest_manager + self, + user_with_money: User, + contest_factory: Callable[..., Contest], + thl_ledger_manager: ThlLedgerManager, + contest_manager: ContestManager, ): # Max entry amount per user $0.10. Contest still ends at $1.00 c = contest_factory( @@ -335,7 +373,11 @@ class TestRaffleContestCRUD: class TestRaffleContestUserViews: def test_list_user_eligible_country( - self, user_with_wallet: User, contest_factory, thl_lm, contest_manager + self, + user_with_wallet: User, + contest_factory: Callable[..., Contest], + thl_ledger_manager: ThlLedgerManager, + contest_manager: ContestManager, ): # No contests exists cs = contest_manager.get_many_by_user_eligible( @@ -368,7 +410,11 @@ class TestRaffleContestUserViews: assert len(cs) == 2 def test_list_user_eligible( - self, user_with_money: User, contest_factory, thl_lm, contest_manager + self, + user_with_money: User, + contest_factory: Callable[..., Contest], + thl_ledger_manager: ThlLedgerManager, + contest_manager: ContestManager, ): c = contest_factory( end_condition=ContestEndCondition(target_entry_amount=USDCent(10)), @@ -390,7 +436,7 @@ class TestRaffleContestUserViews: contest_uuid=c.uuid, entry=entry, country_iso="us", - ledger_manager=thl_lm, + ledger_manager=thl_ledger_manager, ) # User isn't eligible anymore @@ -414,7 +460,11 @@ class TestRaffleContestUserViews: assert len(contest_manager.get_winnings_by_user(user_with_money)) == 0 def test_list_user_winnings( - self, user_with_money: User, contest_factory, thl_lm, contest_manager + self, + user_with_money: User, + contest_factory: Callable[..., Contest], + thl_ledger_manager: ThlLedgerManager, + contest_manager: ContestManager, ): c = contest_factory( end_condition=ContestEndCondition(target_entry_amount=USDCent(100)), @@ -428,7 +478,7 @@ class TestRaffleContestUserViews: contest_uuid=c.uuid, entry=entry, country_iso="us", - ledger_manager=thl_lm, + ledger_manager=thl_ledger_manager, ) # Contest ends after 100 entry, user enters 100 entry, user wins! ws = contest_manager.get_winnings_by_user(user_with_money) @@ -450,7 +500,11 @@ class TestRaffleContestCRUDCount: # This is a COUNT contest. No cash moves. Not really fleshed out what we'd do with this. @pytest.mark.skip def test_enter( - self, user_with_wallet: User, contest_factory, thl_lm, contest_manager + self, + user_with_wallet: User, + contest_factory: Callable[..., Contest], + thl_ledger_manager: ThlLedgerManager, + contest_manager: ContestManager, ): c = contest_factory(entry_type=ContestEntryType.COUNT) entry = ContestEntry( @@ -462,5 +516,5 @@ class TestRaffleContestCRUDCount: contest_uuid=c.uuid, entry=entry, country_iso="us", - ledger_manager=thl_lm, + ledger_manager=thl_ledger_manager, ) diff --git a/tests/managers/thl/test_ipinfo.py b/tests/managers/thl/test_ipinfo.py index 61d4d19..c89312b 100644 --- a/tests/managers/thl/test_ipinfo.py +++ b/tests/managers/thl/test_ipinfo.py @@ -12,7 +12,7 @@ fake = faker.Faker() class TestIPGeonameManager: - def test_init(self, thl_web_rr, 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 +31,7 @@ class TestIPGeonameManager: class TestIPInformationManager: - def test_init(self, thl_web_rr, 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) @@ -57,9 +57,12 @@ class TestIPInformationManager: class TestGeoIpInfoManager: def test_init( - self, thl_web_rr, thl_redis_config, geoipinfo_manager: GeoIpInfoManager + self, + thl_web_rr: PostgresConfig, + thl_redis_config: RedisConfig, + geoipinfo_manager: GeoIpInfoManager, ): - instance = GeoIpInfoManager(pg_config=thl_web_rr, redis_config=thl_redis_config) + instance = GeoIpInfoManager(pg_config=thl_web_rr: PostgresConfig, redis_config=thl_redis_config) assert isinstance(instance, GeoIpInfoManager) assert isinstance(geoipinfo_manager, GeoIpInfoManager) diff --git a/tests/managers/thl/test_ledger/test_lm_accounts.py b/tests/managers/thl/test_ledger/test_lm_accounts.py index 11b2835..540bea8 100644 --- a/tests/managers/thl/test_ledger/test_lm_accounts.py +++ b/tests/managers/thl/test_ledger/test_lm_accounts.py @@ -1,9 +1,11 @@ +from __future__ import annotations + from itertools import product as iproduct from random import randint -from typing import TYPE_CHECKING from uuid import uuid4 import pytest +from pydantic import PositiveInt from generalresearch.currency import LedgerCurrency from generalresearch.managers.base import Permission @@ -11,26 +13,13 @@ from generalresearch.managers.thl.ledger_manager.exceptions import ( LedgerAccountDoesntExistError, ) from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager +from generalresearch.models.custom_types import AccountType, Direction, UUIDStr from generalresearch.models.thl.ledger import ( - AccountType, - Direction, LedgerAccount, LedgerEntry, + LedgerTransaction, ) -if TYPE_CHECKING: - from pydantic import PositiveInt - - from generalresearch.currency import LedgerCurrency - from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager - from generalresearch.models.custom_types import AccountType, Direction, UUIDStr - from generalresearch.models.thl import Direction - from generalresearch.models.thl.ledger import ( - AccountType, - LedgerAccount, - LedgerTransaction, - ) - @pytest.mark.parametrize( argnames="currency, kind, acct_id", @@ -55,7 +44,7 @@ class TestLedgerAccountManagerNoResults: we either get the expected None result or it raises the correct exception """ - qn = ":".join([currency, kind, acct_id]) + qn = f"{currency}:{kind}:{acct_id}" # (1) .get_account is just a wrapper for .get_account_many_ but # call it either way diff --git a/tests/managers/thl/test_ledger/test_lm_tx.py b/tests/managers/thl/test_ledger/test_lm_tx.py index 37b7ba3..13495a7 100644 --- a/tests/managers/thl/test_ledger/test_lm_tx.py +++ b/tests/managers/thl/test_ledger/test_lm_tx.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from decimal import Decimal from random import randint from uuid import uuid4 @@ -5,9 +7,12 @@ from uuid import uuid4 import pytest from generalresearch.currency import LedgerCurrency -from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager +from generalresearch.managers.thl.ledger_manager.ledger import ( + LedgerManager, +) from generalresearch.models.thl.ledger import ( Direction, + LedgerAccount, LedgerEntry, LedgerTransaction, ) @@ -15,7 +20,7 @@ from generalresearch.models.thl.ledger import ( class TestLedgerManagerCreateTx: - def test_create_account_error_permission(self, lm): + def test_create_account_error_permission(self, ledger_manager: LedgerManager): """Confirm that the Permission values that are set on the Ledger Manger allow the Creation action to occur. """ @@ -23,11 +28,11 @@ class TestLedgerManagerCreateTx: # (1) With no Permissions defined test_lm = LedgerManager( - pg_config=lm.pg_config, + pg_config=ledger_manager.pg_config, permissions=[], - redis_config=lm.redis_config, - cache_prefix=lm.cache_prefix, - testing=lm.testing, + redis_config=ledger_manager.redis_config, + cache_prefix=ledger_manager.cache_prefix, + testing=ledger_manager.testing, ) with pytest.raises(expected_exception=AssertionError) as excinfo: @@ -37,9 +42,14 @@ class TestLedgerManagerCreateTx: == "LedgerTransactionManager has insufficient Permissions" ) - def test_create_assertions(self, ledger_account_debit, ledger_account_credit, lm): + def test_create_assertions( + self, + ledger_account_debit: LedgerAccount, + ledger_account_credit: LedgerAccount, + ledger_manager: LedgerManager, + ): with pytest.raises(expected_exception=ValueError) as excinfo: - lm.create_tx( + ledger_manager.create_tx( entries=[ { "direction": Direction.CREDIT, @@ -53,7 +63,12 @@ class TestLedgerManagerCreateTx: in str(excinfo.value) ) - def test_create(self, ledger_account_credit, ledger_account_debit, lm): + def test_create( + self, + ledger_account_credit: LedgerAccount, + ledger_account_debit: LedgerAccount, + ledger_manager: LedgerManager, + ): amount = int(Decimal("1.00") * 100) entries = [ @@ -70,15 +85,20 @@ class TestLedgerManagerCreateTx: ] # Create a Transaction and validate the operation was successful - tx = lm.create_tx(entries=entries) + tx = ledger_manager.create_tx(entries=entries) assert isinstance(tx, LedgerTransaction) - res = lm.get_tx_by_id(transaction_id=tx.id) + res = ledger_manager.get_tx_by_id(transaction_id=tx.id) assert isinstance(res, LedgerTransaction) assert len(res.entries) == 2 assert tx.id == res.id - def test_create_and_reverse(self, ledger_account_credit, ledger_account_debit, lm): + def test_create_and_reverse( + self, + ledger_account_credit: LedgerAccount, + ledger_account_debit: LedgerAccount, + ledger_manager: LedgerManager, + ): amount = int(Decimal("1.00") * 100) entries = [ @@ -94,13 +114,13 @@ class TestLedgerManagerCreateTx: ), ] - tx = lm.create_tx(entries=entries) - res = lm.get_tx_by_id(transaction_id=tx.id) - assert res.id == tx.id + tx = ledger_manager.create_tx(entries=entries) + res = ledger_manager.get_tx_by_id(transaction_id=tx.id) + assert ledger_manager.id == tx.id - assert lm.get_account_balance(account=ledger_account_credit) == 100 - assert lm.get_account_balance(account=ledger_account_debit) == 100 - assert lm.check_ledger_balanced() is True + assert ledger_manager.get_account_balance(account=ledger_account_credit) == 100 + assert ledger_manager.get_account_balance(account=ledger_account_debit) == 100 + assert ledger_manager.check_ledger_balanced() is True # Reverse it entries = [ @@ -116,13 +136,13 @@ class TestLedgerManagerCreateTx: ), ] - tx = lm.create_tx(entries=entries) - res = lm.get_tx_by_id(transaction_id=tx.id) + tx = ledger_manager.create_tx(entries=entries) + res = ledger_manager.get_tx_by_id(transaction_id=tx.id) assert res.id == tx.id - assert lm.get_account_balance(ledger_account_credit) == 0 - assert lm.get_account_balance(ledger_account_debit) == 0 - assert lm.check_ledger_balanced() + assert ledger_manager.get_account_balance(ledger_account_credit) == 0 + assert ledger_manager.get_account_balance(ledger_account_debit) == 0 + assert ledger_manager.check_ledger_balanced() # subtract again entries = [ @@ -137,52 +157,60 @@ class TestLedgerManagerCreateTx: amount=amount, ), ] - tx = lm.create_tx(entries=entries) - res = lm.get_tx_by_id(transaction_id=tx.id) + tx = ledger_manager.create_tx(entries=entries) + res = ledger_manager.get_tx_by_id(transaction_id=tx.id) assert res.id == tx.id - assert lm.get_account_balance(ledger_account_credit) == -100 - assert lm.get_account_balance(ledger_account_debit) == -100 - assert lm.check_ledger_balanced() + assert ledger_manager.get_account_balance(ledger_account_credit) == -100 + assert ledger_manager.get_account_balance(ledger_account_debit) == -100 + assert ledger_manager.check_ledger_balanced() class TestLedgerManagerGetTx: # @pytest.mark.parametrize("currency", [LedgerCurrency.TEST], indirect=True) - def test_get_tx_by_id(self, ledger_tx, lm): + def test_get_tx_by_id( + self, ledger_tx: LedgerTransaction, ledger_manager: LedgerManager + ): with pytest.raises(expected_exception=AssertionError): - lm.get_tx_by_id(transaction_id=ledger_tx) + ledger_manager.get_tx_by_id(transaction_id=ledger_tx) - res = lm.get_tx_by_id(transaction_id=ledger_tx.id) + res = ledger_manager.get_tx_by_id(transaction_id=ledger_tx.id) assert res.id == ledger_tx.id # @pytest.mark.parametrize("currency", [LedgerCurrency.TEST], indirect=True) - def test_get_tx_by_ids(self, ledger_tx, lm): - res = lm.get_tx_by_id(transaction_id=ledger_tx.id) + def test_get_tx_by_ids( + self, ledger_tx: LedgerTransaction, ledger_manager: LedgerManager + ): + res = ledger_manager.get_tx_by_id(transaction_id=ledger_tx.id) assert res.id == ledger_tx.id @pytest.mark.parametrize( "tag", [f"{LedgerCurrency.TEST}:{uuid4().hex}"], indirect=True ) - def test_get_tx_ids_by_tag(self, ledger_tx, tag, lm): + def test_get_tx_ids_by_tag( + self, ledger_tx: LedgerTransaction, tag: str, ledger_manager: LedgerManager + ): # (1) search for a random tag - res = lm.get_tx_ids_by_tag(tag="aaa:bbb") + res = ledger_manager.get_tx_ids_by_tag(tag="aaa:bbb") assert isinstance(res, set) assert len(res) == 0 # (2) search for the tag that was used during ledger_transaction creation - res = lm.get_tx_ids_by_tag(tag=tag) + res = ledger_manager.get_tx_ids_by_tag(tag=tag) assert isinstance(res, set) assert len(res) == 1 - def test_get_tx_by_tag(self, ledger_tx, tag, lm): + def test_get_tx_by_tag( + self, ledger_tx: LedgerTransaction, tag: str, ledger_manager: LedgerManager + ): # (1) search for a random tag - res = lm.get_tx_by_tag(tag="aaa:bbb") + res = ledger_manager.get_tx_by_tag(tag="aaa:bbb") assert isinstance(res, list) assert len(res) == 0 # (2) search for the tag that was used during ledger_transaction creation - res = lm.get_tx_by_tag(tag=tag) + res = ledger_manager.get_tx_by_tag(tag=tag) assert isinstance(res, list) assert len(res) == 1 @@ -190,42 +218,60 @@ class TestLedgerManagerGetTx: assert ledger_tx.id == res[0].id def test_get_tx_filtered_by_account( - self, ledger_tx, ledger_account, ledger_account_debit, ledger_account_credit, lm + self, + ledger_tx: LedgerTransaction, + ledger_account: LedgerAccount, + ledger_account_debit: LedgerAccount, + ledger_account_credit: LedgerAccount, + ledger_manager: LedgerManager, ): # (1) Do basic assertion checks first with pytest.raises(expected_exception=AssertionError) as excinfo: - lm.get_tx_filtered_by_account(account_uuid=ledger_account) + ledger_manager.get_tx_filtered_by_account(account_uuid=ledger_account) assert str(excinfo.value) == "account_uuid must be a str" # (2) This search doesn't return anything because this ledger account # wasn't actually used in the entries for the ledger_transaction - res = lm.get_tx_filtered_by_account(account_uuid=ledger_account.uuid) + res = ledger_manager.get_tx_filtered_by_account( + account_uuid=ledger_account.uuid + ) assert len(res) == 0 # (3) Either the credit or the debit example ledger_accounts wll work # to find this transaction because they're both used in the entries - res = lm.get_tx_filtered_by_account(account_uuid=ledger_account_debit.uuid) + res = ledger_manager.get_tx_filtered_by_account( + account_uuid=ledger_account_debit.uuid + ) assert len(res) == 1 assert res[0].id == ledger_tx.id - res = lm.get_tx_filtered_by_account(account_uuid=ledger_account_credit.uuid) + res = ledger_manager.get_tx_filtered_by_account( + account_uuid=ledger_account_credit.uuid + ) assert len(res) == 1 assert ledger_tx.id == res[0].id - res2 = lm.get_tx_by_id(transaction_id=ledger_tx.id) + res2 = ledger_manager.get_tx_by_id(transaction_id=ledger_tx.id) assert res2.model_dump_json() == res[0].model_dump_json() - def test_filter_metadata(self, ledger_tx, tx_metadata, lm): + def test_filter_metadata( + self, + ledger_tx: LedgerTransaction, + tx_metadata: dict[str, str] | None, + ledger_manager: LedgerManager, + ): key, value = next(iter(tx_metadata.items())) # (1) Confirm a random key,value pair returns nothing - res = lm.get_tx_filtered_by_metadata( + res = ledger_manager.get_tx_filtered_by_metadata( metadata_key=f"key-{uuid4().hex[:10]}", metadata_value=uuid4().hex[:12] ) assert len(res) == 0 # (2) confirm a key,value pair return the correct results - res = lm.get_tx_filtered_by_metadata(metadata_key=key, metadata_value=value) + res = ledger_manager.get_tx_filtered_by_metadata( + metadata_key=key, metadata_value=value + ) assert len(res) == 1 # assert 0 == THL_lm.get_filtered_account_balance(account2, "thl_wall", "ccc") diff --git a/tests/managers/thl/test_ledger/test_lm_tx_entries.py b/tests/managers/thl/test_ledger/test_lm_tx_entries.py index 5bf1c48..9925b87 100644 --- a/tests/managers/thl/test_ledger/test_lm_tx_entries.py +++ b/tests/managers/thl/test_ledger/test_lm_tx_entries.py @@ -1,25 +1,35 @@ -from generalresearch.models.thl.ledger import LedgerEntry +from __future__ import annotations + +from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager +from generalresearch.models.thl.ledger import ( + LedgerEntry, + LedgerTransaction, +) class TestLedgerEntryManager: - def test_get_tx_entries_by_tx(self, ledger_tx, lm): + def test_get_tx_entries_by_tx( + self, ledger_tx: LedgerTransaction, ledger_manager: LedgerManager + ): # First confirm the Ledger TX exists with 2 Entries - res = lm.get_tx_by_id(transaction_id=ledger_tx.id) + res = ledger_manager.get_tx_by_id(transaction_id=ledger_tx.id) assert len(res.entries) == 2 - tx_entries = lm.get_tx_entries_by_tx(transaction=ledger_tx) + tx_entries = ledger_manager.get_tx_entries_by_tx(transaction=ledger_tx) assert len(tx_entries) == 2 assert res.entries == tx_entries assert isinstance(tx_entries[0], LedgerEntry) - def test_get_tx_entries_by_txs(self, ledger_tx, lm): + def test_get_tx_entries_by_txs( + self, ledger_tx: LedgerTransaction, ledger_manager: LedgerManager + ): # First confirm the Ledger TX exists with 2 Entries - res = lm.get_tx_by_id(transaction_id=ledger_tx.id) + res = ledger_manager.get_tx_by_id(transaction_id=ledger_tx.id) assert len(res.entries) == 2 - tx_entries = lm.get_tx_entries_by_txs(transactions=[ledger_tx]) + tx_entries = ledger_manager.get_tx_entries_by_txs(transactions=[ledger_tx]) assert len(tx_entries) == 2 assert res.entries == tx_entries 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 020b74a..9158e15 100644 --- a/tests/managers/thl/test_ledger/test_lm_tx_locks.py +++ b/tests/managers/thl/test_ledger/test_lm_tx_locks.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import logging from collections.abc import Callable from datetime import UTC, datetime, timedelta @@ -5,6 +7,7 @@ from decimal import Decimal import pytest +from generalresearch.currency import LedgerCurrency from generalresearch.managers.thl.ledger_manager.conditions import ( generate_condition_mp_payment, ) @@ -13,8 +16,11 @@ from generalresearch.managers.thl.ledger_manager.exceptions import ( LedgerTransactionCreateLockError, LedgerTransactionFlagAlreadyExistsError, ) +from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager +from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager from generalresearch.models import Source from generalresearch.models.thl.ledger import LedgerTransaction +from generalresearch.models.thl.product import Product from generalresearch.models.thl.session import ( Session, Status, @@ -31,17 +37,17 @@ class TestLedgerLocks: def test_a( self, - user_factory, - session_factory, - product_user_wallet_no, - create_main_accounts, + user_factory: Callable[..., User], + session_factory: Callable[..., Session], + product_user_wallet_no: Product, + create_main_accounts: Callable[..., None], caplog, - thl_lm, - lm, - utc_hour_ago, - currency, - wall_factory, - delete_ledger_db, + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + utc_hour_ago: datetime, + currency: LedgerCurrency, + wall_factory: Callable[..., Wall], + delete_ledger_db: Callable[..., None], ): """ TODO: This whole test is confusing a I don't really understand. @@ -61,12 +67,16 @@ class TestLedgerLocks: # A User does a Wall Completion in Session=1 w1 = s1.wall_events[0] - tx = thl_lm.create_tx_task_complete(wall=w1, user=user, created=w1.started) + tx = thl_ledger_manager.create_tx_task_complete( + wall=w1, user=user, created=w1.started + ) assert isinstance(tx, LedgerTransaction) # A User does another Wall Completion in Session=1 w2 = s1.wall_events[1] - tx = thl_lm.create_tx_task_complete(wall=w2, user=user, created=w2.started) + tx = thl_ledger_manager.create_tx_task_complete( + wall=w2, user=user, created=w2.started + ) assert isinstance(tx, LedgerTransaction) # That first Wall Complete was "adjusted" to instead be marked @@ -76,7 +86,7 @@ class TestLedgerLocks: adjusted_cpi=0, adjusted_timestamp=utc_hour_ago + timedelta(hours=1), ) - tx = thl_lm.create_tx_task_adjustment(wall=w1, user=user) + tx = thl_ledger_manager.create_tx_task_adjustment(wall=w1, user=user) assert isinstance(tx, LedgerTransaction) # A User does another! Wall Completion in Session=1; however, we @@ -85,55 +95,57 @@ class TestLedgerLocks: # Make sure we clear any flags/locks first lock_key = f"{currency.value}:thl_wall:{w3.uuid}" - lock_name = f"{lm.cache_prefix}:transaction_lock:{lock_key}" - flag_name = f"{lm.cache_prefix}:transaction_flag:{lock_key}" - lm.redis_client.delete(lock_name) - lm.redis_client.delete(flag_name) + lock_name = f"{ledger_manager.cache_prefix}:transaction_lock:{lock_key}" + flag_name = f"{ledger_manager.cache_prefix}:transaction_flag:{lock_key}" + ledger_manager.redis_client.delete(lock_name) + ledger_manager.redis_client.delete(flag_name) # Despite the f1 = generate_condition_mp_payment(wall=w1) f2 = generate_condition_mp_payment(wall=w2) f3 = generate_condition_mp_payment(wall=w3) - assert f1(lm=lm) is False - assert f2(lm=lm) is False - assert f3(lm=lm) is True + assert f1(ledger_manager) is False + assert f2(lm=ledger_manager) is False + assert f3(lm=ledger_manager) is True condition = f3 - create_tx_func = lambda: thl_lm.create_tx_task_complete_(wall=w3, user=user) + create_tx_func = lambda: thl_ledger_manager.create_tx_task_complete_( + wall=w3, user=user + ) assert isinstance(create_tx_func, Callable) - assert f3(lm) is True + assert f3(ledger_manager) is True - lm.redis_client.delete(flag_name) - lm.redis_client.delete(lock_name) + ledger_manager.redis_client.delete(flag_name) + ledger_manager.redis_client.delete(lock_name) - tx = thl_lm.create_tx_protected( + tx = thl_ledger_manager.create_tx_protected( lock_key=lock_key, condition=condition, create_tx_func=create_tx_func ) - assert f3(lm) is False + assert f3(ledger_manager) is False # purposely hold the lock open tx = None - lm.redis_client.set(lock_name, "1") + ledger_manager.redis_client.set(lock_name, "1") with caplog.at_level(logging.ERROR): with pytest.raises(expected_exception=LedgerTransactionCreateLockError): - tx = thl_lm.create_tx_protected( + 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 - lm.redis_client.delete(lock_name) + ledger_manager.redis_client.delete(lock_name) def test_locking( self, - user_factory, - product_user_wallet_no, - create_main_accounts, - delete_ledger_db, + user_factory: Callable[..., User], + product_user_wallet_no: Product, + create_main_accounts: Callable[..., None], + delete_ledger_db: Callable[..., None], caplog, - thl_lm, - lm, + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, ): delete_ledger_db() create_main_accounts() @@ -154,7 +166,9 @@ class TestLedgerLocks: started=now, finished=now + timedelta(seconds=1), ) - thl_lm.create_tx_task_complete(wall=wall1, user=user, created=wall1.started) + thl_ledger_manager.create_tx_task_complete( + wall=wall1, user=user, created=wall1.started + ) # A User does a Wall complete on Session.id=1 and the transaction is # logged to the ledger @@ -169,7 +183,9 @@ class TestLedgerLocks: started=now, finished=now + timedelta(seconds=1), ) - thl_lm.create_tx_task_complete(wall=wall2, user=user, created=wall2.started) + thl_ledger_manager.create_tx_task_complete( + wall=wall2, user=user, created=wall2.started + ) # An hour later, the first wall complete is adjusted to a Failure and # it's tracked in the ledger @@ -178,7 +194,7 @@ class TestLedgerLocks: adjusted_cpi=0, adjusted_timestamp=now + timedelta(hours=1), ) - thl_lm.create_tx_task_adjustment(wall=wall1, user=user) + thl_ledger_manager.create_tx_task_adjustment(wall=wall1, user=user) # A User does a Wall complete on Session.id=1 and the transaction # IS NOT logged to the ledger @@ -195,52 +211,52 @@ class TestLedgerLocks: uuid="867a282d8b4d40d2a2093d75b802b629", ) - revenue_account = thl_lm.get_account_task_complete_revenue() - assert 0 == thl_lm.get_account_filtered_balance( + revenue_account = thl_ledger_manager.get_account_task_complete_revenue() + assert 0 == thl_ledger_manager.get_account_filtered_balance( account=revenue_account, metadata_key="thl_wall", metadata_value=wall3.uuid, ) # Make sure we clear any flags/locks first lock_key = f"test:thl_wall:{wall3.uuid}" - lock_name = f"{lm.cache_prefix}:transaction_lock:{lock_key}" - flag_name = f"{lm.cache_prefix}:transaction_flag:{lock_key}" - lm.redis_client.delete(lock_name) - lm.redis_client.delete(flag_name) + lock_name = f"{ledger_manager.cache_prefix}:transaction_lock:{lock_key}" + flag_name = f"{ledger_manager.cache_prefix}:transaction_flag:{lock_key}" + ledger_manager.redis_client.delete(lock_name) + ledger_manager.redis_client.delete(flag_name) # Purposely hold the lock open - lm.redis_client.set(name=lock_name, value="1") + ledger_manager.redis_client.set(name=lock_name, value="1") with caplog.at_level(logging.DEBUG): with pytest.raises(expected_exception=LedgerTransactionCreateLockError): - tx = thl_lm.create_tx_task_complete( + 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 - lm.redis_client.delete(lock_name) + ledger_manager.redis_client.delete(lock_name) # Set the redis flag to indicate it has been run - lm.redis_client.set(flag_name, "1") + ledger_manager.redis_client.set(flag_name, "1") # with self.assertLogs(logger=logger, level=logging.DEBUG) as cm2: with pytest.raises(expected_exception=LedgerTransactionFlagAlreadyExistsError): - tx = thl_lm.create_tx_task_complete( + tx = thl_ledger_manager.create_tx_task_complete( wall=wall3, user=user, created=wall3.started ) # self.assertIn("entered_lock: True, flag_set: True", cm2.output[0]) # Unset the flag - lm.redis_client.delete(flag_name) + ledger_manager.redis_client.delete(flag_name) - assert 0 == lm.get_account_filtered_balance( + assert 0 == ledger_manager.get_account_filtered_balance( account=revenue_account, metadata_key="thl_wall", metadata_value=wall3.uuid, ) # Now actually run it - tx = thl_lm.create_tx_task_complete( + tx = thl_ledger_manager.create_tx_task_complete( wall=wall3, user=user, created=wall3.started ) assert tx is not None @@ -249,29 +265,34 @@ class TestLedgerLocks: # Confirm the Exception inheritance works tx = None with pytest.raises(expected_exception=LedgerTransactionCreateError): - tx = thl_lm.create_tx_task_complete( + tx = thl_ledger_manager.create_tx_task_complete( wall=wall3, user=user, created=wall3.started ) assert tx is None # clear the redis flag, it should query the db - assert lm.redis_client.get(flag_name) is not None - lm.redis_client.delete(flag_name) - assert lm.redis_client.get(flag_name) is None + assert ledger_manager.redis_client.get(flag_name) is not None + ledger_manager.redis_client.delete(flag_name) + assert ledger_manager.redis_client.get(flag_name) is None with pytest.raises(expected_exception=LedgerTransactionCreateError): - tx = thl_lm.create_tx_task_complete( + tx = thl_ledger_manager.create_tx_task_complete( wall=wall3, user=user, created=wall3.started ) - assert 400 == thl_lm.get_account_filtered_balance( + assert 400 == thl_ledger_manager.get_account_filtered_balance( account=revenue_account, metadata_key="thl_wall", metadata_value=wall3.uuid, ) def test_bp_payment_without_locks( - self, user_factory, product_user_wallet_no, create_main_accounts, thl_lm, lm + self, + user_factory: Callable[..., User], + product_user_wallet_no: Product, + create_main_accounts: Callable[..., None], + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, ): user: User = user_factory(product=product_user_wallet_no) wall1 = Wall( @@ -286,29 +307,42 @@ class TestLedgerLocks: finished=datetime.now(UTC) + timedelta(seconds=1), ) - thl_lm.create_tx_task_complete(wall=wall1, user=user, created=wall1.started) + thl_ledger_manager.create_tx_task_complete( + wall=wall1, user=user, created=wall1.started + ) session = Session(started=wall1.started, user=user, wall_events=[wall1]) status, status_code_1 = session.determine_session_status() thl_net, commission_amount, bp_pay, user_pay = session.determine_payments() session.update( - status=status, status_code_1=status_code_1, finished=session.started + timedelta(minutes=10), payout=bp_pay, user_payout=user_pay + status=status, + status_code_1=status_code_1, + finished=session.started + timedelta(minutes=10), + payout=bp_pay, + user_payout=user_pay, ) print(thl_net, commission_amount, bp_pay, user_pay) # Run it 3 times without any checks, and it gets made three times! - thl_lm.create_tx_bp_payment(session=session, created=wall1.started) - thl_lm.create_tx_bp_payment_(session=session, created=wall1.started) - thl_lm.create_tx_bp_payment_(session=session, created=wall1.started) + thl_ledger_manager.create_tx_bp_payment(session=session, created=wall1.started) + thl_ledger_manager.create_tx_bp_payment_(session=session, created=wall1.started) + thl_ledger_manager.create_tx_bp_payment_(session=session, created=wall1.started) - bp_wallet = thl_lm.get_account_or_create_bp_wallet(product=user.product) - assert 48 * 3 == lm.get_account_balance(account=bp_wallet) - assert 48 * 3 == thl_lm.get_account_filtered_balance( + bp_wallet = thl_ledger_manager.get_account_or_create_bp_wallet( + product=user.product + ) + assert 48 * 3 == ledger_manager.get_account_balance(account=bp_wallet) + assert 48 * 3 == thl_ledger_manager.get_account_filtered_balance( account=bp_wallet, metadata_key="thl_session", metadata_value=session.uuid ) - assert lm.check_ledger_balanced() + assert ledger_manager.check_ledger_balanced() def test_bp_payment_with_locks( - self, user_factory, product_user_wallet_no, create_main_accounts, thl_lm, lm + self, + user_factory: Callable[..., User], + product_user_wallet_no: Product, + create_main_accounts: Callable[..., None], + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, ): user: User = user_factory(product=product_user_wallet_no) @@ -324,35 +358,45 @@ class TestLedgerLocks: finished=datetime.now(UTC) + timedelta(seconds=1), ) - thl_lm.create_tx_task_complete(wall1, user, created=wall1.started) + thl_ledger_manager.create_tx_task_complete(wall1, user, created=wall1.started) session = Session(started=wall1.started, user=user, wall_events=[wall1]) status, status_code_1 = session.determine_session_status() thl_net, commission_amount, bp_pay, user_pay = session.determine_payments() session.update( - status=status, status_code_1=status_code_1, finished=session.started + timedelta(minutes=10), payout=bp_pay, user_payout=user_pay + status=status, + status_code_1=status_code_1, + finished=session.started + timedelta(minutes=10), + payout=bp_pay, + user_payout=user_pay, ) print(thl_net, commission_amount, bp_pay, user_pay) # Make sure we clear any flags/locks first lock_key = f"test:thl_wall:{wall1.uuid}" - lock_name = f"{lm.cache_prefix}:transaction_lock:{lock_key}" - flag_name = f"{lm.cache_prefix}:transaction_flag:{lock_key}" - lm.redis_client.delete(lock_name) - lm.redis_client.delete(flag_name) + lock_name = f"{ledger_manager.cache_prefix}:transaction_lock:{lock_key}" + flag_name = f"{ledger_manager.cache_prefix}:transaction_flag:{lock_key}" + ledger_manager.redis_client.delete(lock_name) + ledger_manager.redis_client.delete(flag_name) # Run it 3 times with check, and it gets made once! - thl_lm.create_tx_bp_payment(session=session, created=wall1.started) + thl_ledger_manager.create_tx_bp_payment(session=session, created=wall1.started) with pytest.raises(expected_exception=LedgerTransactionCreateError): - thl_lm.create_tx_bp_payment(session=session, created=wall1.started) + thl_ledger_manager.create_tx_bp_payment( + session=session, created=wall1.started + ) with pytest.raises(expected_exception=LedgerTransactionCreateError): - thl_lm.create_tx_bp_payment(session=session, created=wall1.started) + thl_ledger_manager.create_tx_bp_payment( + session=session, created=wall1.started + ) - bp_wallet = thl_lm.get_account_or_create_bp_wallet(product=user.product) - assert 48 == thl_lm.get_account_balance(bp_wallet) - assert 48 == thl_lm.get_account_filtered_balance( + bp_wallet = thl_ledger_manager.get_account_or_create_bp_wallet( + product=user.product + ) + assert 48 == thl_ledger_manager.get_account_balance(bp_wallet) + assert 48 == thl_ledger_manager.get_account_filtered_balance( account=bp_wallet, metadata_key="thl_session", metadata_value=session.uuid, ) - assert lm.check_ledger_balanced() + assert ledger_manager.check_ledger_balanced() diff --git a/tests/managers/thl/test_ledger/test_lm_tx_metadata.py b/tests/managers/thl/test_ledger/test_lm_tx_metadata.py index 5d12633..f63efa4 100644 --- a/tests/managers/thl/test_ledger/test_lm_tx_metadata.py +++ b/tests/managers/thl/test_ledger/test_lm_tx_metadata.py @@ -1,34 +1,52 @@ +from __future__ import annotations + +from generalresearch.managers.thl.ledger_manager.ledger import ( + LedgerManager, + LedgerTransaction, +) + + class TestLedgerMetadataManager: - def test_get_tx_metadata_by_txs(self, ledger_tx, lm): + def test_get_tx_metadata_by_txs( + self, ledger_tx: LedgerTransaction, ledger_manager: LedgerManager + ): # First confirm the Ledger TX exists with 2 Entries - res = lm.get_tx_by_id(transaction_id=ledger_tx.id) + res = ledger_manager.get_tx_by_id(transaction_id=ledger_tx.id) assert isinstance(res.metadata, dict) - tx_metadatas = lm.get_tx_metadata_by_txs(transactions=[ledger_tx]) + tx_metadatas = ledger_manager.get_tx_metadata_by_txs(transactions=[ledger_tx]) assert isinstance(tx_metadatas, dict) assert isinstance(tx_metadatas[ledger_tx.id], dict) assert res.metadata == tx_metadatas[ledger_tx.id] - def test_get_tx_metadata_ids_by_tx(self, ledger_tx, lm): + def test_get_tx_metadata_ids_by_tx( + self, ledger_tx: LedgerTransaction, ledger_manager: LedgerManager + ): # First confirm the Ledger TX exists with 2 Entries - res = lm.get_tx_by_id(transaction_id=ledger_tx.id) + res = ledger_manager.get_tx_by_id(transaction_id=ledger_tx.id) tx_metadata_cnt = len(res.metadata.keys()) - tx_metadata_ids = lm.get_tx_metadata_ids_by_tx(transaction=ledger_tx) + tx_metadata_ids = ledger_manager.get_tx_metadata_ids_by_tx( + transaction=ledger_tx + ) assert isinstance(tx_metadata_ids, set) - assert isinstance(list(tx_metadata_ids)[0], int) + assert isinstance(next(iter(tx_metadata_ids)), int) assert tx_metadata_cnt == len(tx_metadata_ids) - def test_get_tx_metadata_ids_by_txs(self, ledger_tx, lm): + def test_get_tx_metadata_ids_by_txs( + self, ledger_tx: LedgerTransaction, ledger_manager: LedgerManager + ): # First confirm the Ledger TX exists with 2 Entries - res = lm.get_tx_by_id(transaction_id=ledger_tx.id) + res = ledger_manager.get_tx_by_id(transaction_id=ledger_tx.id) tx_metadata_cnt = len(res.metadata.keys()) - tx_metadata_ids = lm.get_tx_metadata_ids_by_txs(transactions=[ledger_tx]) + tx_metadata_ids = ledger_manager.get_tx_metadata_ids_by_txs( + transactions=[ledger_tx] + ) assert isinstance(tx_metadata_ids, set) - assert isinstance(list(tx_metadata_ids)[0], int) + assert isinstance(next(iter(tx_metadata_ids)), int) assert tx_metadata_cnt == len(tx_metadata_ids) diff --git a/tests/managers/thl/test_ledger/test_thl_lm_accounts.py b/tests/managers/thl/test_ledger/test_thl_lm_accounts.py index 8d7d828..dce9116 100644 --- a/tests/managers/thl/test_ledger/test_thl_lm_accounts.py +++ b/tests/managers/thl/test_ledger/test_thl_lm_accounts.py @@ -1,19 +1,38 @@ +from __future__ import annotations + +from collections.abc import Callable from uuid import uuid4 import pytest +from generalresearch.currency import LedgerCurrency +from generalresearch.managers.thl.ledger_manager.exceptions import ( + LedgerAccountDoesntExistError, +) +from generalresearch.managers.thl.ledger_manager.ledger import ( + LedgerAccountManager, + LedgerManager, +) +from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager +from generalresearch.models.thl.ledger import ( + AccountType, + Direction, + LedgerAccount, +) +from generalresearch.models.thl.product import Product +from generalresearch.models.thl.user import User + class TestThlLedgerManagerAccounts: - def test_get_account_or_create_user_wallet(self, user, thl_lm, lm): - from generalresearch.currency import LedgerCurrency - from generalresearch.models.thl.ledger import ( - AccountType, - Direction, - LedgerAccount, - ) + def test_get_account_or_create_user_wallet( + self, + user: User, + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + ): - account = thl_lm.get_account_or_create_user_wallet(user=user) + account = thl_ledger_manager.get_account_or_create_user_wallet(user=user) assert isinstance(account, LedgerAccount) assert user.uuid in account.qualified_name @@ -25,18 +44,20 @@ class TestThlLedgerManagerAccounts: assert account.currency == LedgerCurrency.TEST # Actually query for it to confirm - res = lm.get_account(qualified_name=account.qualified_name, raise_on_error=True) + res = ledger_manager.get_account( + qualified_name=account.qualified_name, raise_on_error=True + ) + assert isinstance(res, LedgerAccount) assert res.model_dump_json() == account.model_dump_json() - def test_get_account_or_create_bp_wallet(self, product, thl_lm, lm): - from generalresearch.currency import LedgerCurrency - from generalresearch.models.thl.ledger import ( - AccountType, - Direction, - LedgerAccount, - ) + def test_get_account_or_create_bp_wallet( + self, + product: Product, + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + ): - account = thl_lm.get_account_or_create_bp_wallet(product=product) + account = thl_ledger_manager.get_account_or_create_bp_wallet(product=product) assert isinstance(account, LedgerAccount) assert product.uuid in account.qualified_name @@ -48,17 +69,22 @@ class TestThlLedgerManagerAccounts: assert account.currency == LedgerCurrency.TEST # Actually query for it to confirm - res = lm.get_account(qualified_name=account.qualified_name, raise_on_error=True) + res = ledger_manager.get_account( + qualified_name=account.qualified_name, raise_on_error=True + ) + assert isinstance(res, LedgerAccount) assert res.model_dump_json() == account.model_dump_json() - def test_get_account_or_create_bp_commission(self, product, thl_lm, lm): - from generalresearch.currency import LedgerCurrency - from generalresearch.models.thl.ledger import ( - AccountType, - Direction, - ) + def test_get_account_or_create_bp_commission( + self, + product: Product, + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + ): - account = thl_lm.get_account_or_create_bp_commission(product=product) + account = thl_ledger_manager.get_account_or_create_bp_commission( + product=product + ) assert product.uuid in account.qualified_name assert account.display_name == f"Revenue from commission {product.uuid}" @@ -69,18 +95,21 @@ class TestThlLedgerManagerAccounts: assert account.currency == LedgerCurrency.TEST # Actually query for it to confirm - res = lm.get_account(qualified_name=account.qualified_name, raise_on_error=True) + res = ledger_manager.get_account( + qualified_name=account.qualified_name, raise_on_error=True + ) + assert isinstance(res, LedgerAccount) assert res.model_dump_json() == account.model_dump_json() @pytest.mark.parametrize("expense", ["tango", "paypal", "gift", "tremendous"]) - def test_get_account_or_create_bp_expense(self, product, expense, thl_lm, lm): - from generalresearch.currency import LedgerCurrency - from generalresearch.models.thl.ledger import ( - AccountType, - Direction, - ) - - account = thl_lm.get_account_or_create_bp_expense( + def test_get_account_or_create_bp_expense( + self, + product: Product, + expense, + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + ): + account = thl_ledger_manager.get_account_or_create_bp_expense( product=product, expense_name=expense ) assert product.uuid in account.qualified_name @@ -92,17 +121,22 @@ class TestThlLedgerManagerAccounts: assert account.currency == LedgerCurrency.TEST # Actually query for it to confirm - res = lm.get_account(qualified_name=account.qualified_name, raise_on_error=True) + res = ledger_manager.get_account( + qualified_name=account.qualified_name, raise_on_error=True + ) + assert isinstance(res, LedgerAccount) assert res.model_dump_json() == account.model_dump_json() - def test_get_or_create_bp_pending_payout_account(self, product, thl_lm, lm): - from generalresearch.currency import LedgerCurrency - from generalresearch.models.thl.ledger import ( - AccountType, - Direction, - ) + def test_get_or_create_bp_pending_payout_account( + self, + product: Product, + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + ): - account = thl_lm.get_or_create_bp_pending_payout_account(product=product) + account = thl_ledger_manager.get_or_create_bp_pending_payout_account( + product=product + ) assert product.uuid in account.qualified_name assert account.display_name == f"BP Wallet Pending {product.uuid}" @@ -113,11 +147,17 @@ class TestThlLedgerManagerAccounts: assert account.currency == LedgerCurrency.TEST # Actually query for it to confirm - res = lm.get_account(qualified_name=account.qualified_name, raise_on_error=True) + res = ledger_manager.get_account( + qualified_name=account.qualified_name, raise_on_error=True + ) + assert isinstance(res, LedgerAccount) assert res.model_dump_json() == account.model_dump_json() def test_get_account_task_complete_revenue_raises( - self, delete_ledger_db, thl_lm, lm + self, + delete_ledger_db: Callable[..., None], + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, ): from generalresearch.managers.thl.ledger_manager.exceptions import ( LedgerAccountDoesntExistError, @@ -126,63 +166,75 @@ class TestThlLedgerManagerAccounts: delete_ledger_db() with pytest.raises(expected_exception=LedgerAccountDoesntExistError): - thl_lm.get_account_task_complete_revenue() + thl_ledger_manager.get_account_task_complete_revenue() def test_get_account_task_complete_revenue( - self, account_cash, account_revenue_task_complete, thl_lm, lm + self, + thl_ledger_manager: ThlLedgerManager, ): from generalresearch.models.thl.ledger import ( AccountType, LedgerAccount, ) - res = thl_lm.get_account_task_complete_revenue() + res = thl_ledger_manager.get_account_task_complete_revenue() assert isinstance(res, LedgerAccount) assert res.reference_type is None assert res.reference_uuid is None assert res.account_type == AccountType.REVENUE assert res.display_name == "Cash flow task complete" - def test_get_account_cash_raises(self, delete_ledger_db, thl_lm, lm): - from generalresearch.managers.thl.ledger_manager.exceptions import ( - LedgerAccountDoesntExistError, - ) + def test_get_account_cash_raises( + self, + delete_ledger_db: Callable[..., None], + thl_ledger_manager: ThlLedgerManager, + ): delete_ledger_db() with pytest.raises(expected_exception=LedgerAccountDoesntExistError): - thl_lm.get_account_cash() + thl_ledger_manager.get_account_cash() - def test_get_account_cash(self, account_cash, thl_lm, lm): + def test_get_account_cash( + self, + thl_ledger_manager: ThlLedgerManager, + ): from generalresearch.models.thl.ledger import ( AccountType, LedgerAccount, ) - res = thl_lm.get_account_cash() + res = thl_ledger_manager.get_account_cash() assert isinstance(res, LedgerAccount) assert res.reference_type is None assert res.reference_uuid is None assert res.account_type == AccountType.CASH assert res.display_name == "Operating Cash Account" - def test_get_accounts(self, setup_accounts, product, user_factory, thl_lm, lm, lam): - from generalresearch.managers.thl.ledger_manager.exceptions import ( - LedgerAccountDoesntExistError, - ) - from generalresearch.models.thl.user import User + def test_get_accounts( + self, + setup_accounts: Callable[..., None], + product: Product, + user_factory: Callable[..., User], + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + ledger_account_manager: LedgerAccountManager, + ): + setup_accounts() - user1: User = user_factory(product=product) - user2: User = user_factory(product=product) + _: User = user_factory(product=product) + _: User = user_factory(product=product) - account1 = thl_lm.get_account_or_create_bp_wallet(product=product) + account1 = thl_ledger_manager.get_account_or_create_bp_wallet(product=product) # (1) known account and confirm it comes back - res = lm.get_account(qualified_name=account1.qualified_name) + res = ledger_manager.get_account(qualified_name=account1.qualified_name) assert account1.model_dump_json() == res.model_dump_json() # (2) known accounts and confirm they both come back - res = lam.get_accounts(qualified_names=[account1.qualified_name]) + res = ledger_account_manager.get_accounts( + qualified_names=[account1.qualified_name] + ) assert isinstance(res, list) assert len(res) == 1 assert account1 in res @@ -190,28 +242,34 @@ class TestThlLedgerManagerAccounts: # Get 2 known and 1 made up qualified names, and confirm it raises # an error with pytest.raises(LedgerAccountDoesntExistError): - lam.get_accounts( + ledger_account_manager.get_accounts( qualified_names=[ account1.qualified_name, f"test:bp_wall:{uuid4().hex}", ] ) - def test_get_accounts_if_exists(self, product_factory, currency, thl_lm, lm): - from generalresearch.models.thl.product import Product + def test_get_accounts_if_exists( + self, + product_factory: Callable[..., Product], + currency: LedgerCurrency, + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + ): p1: Product = product_factory() p2: Product = product_factory() - account1 = thl_lm.get_account_or_create_bp_wallet(product=p1) - account2 = thl_lm.get_account_or_create_bp_wallet(product=p2) + account1 = thl_ledger_manager.get_account_or_create_bp_wallet(product=p1) + account2 = thl_ledger_manager.get_account_or_create_bp_wallet(product=p2) # (1) known account and confirm it comes back - res = lm.get_account(qualified_name=account1.qualified_name) + res = ledger_manager.get_account(qualified_name=account1.qualified_name) + assert isinstance(res, LedgerAccount) assert account1.model_dump_json() == res.model_dump_json() # (2) known accounts and confirm they both come back - res = lm.get_accounts( + res = ledger_manager.get_accounts( qualified_names=[account1.qualified_name, account2.qualified_name] ) assert isinstance(res, list) @@ -221,7 +279,7 @@ class TestThlLedgerManagerAccounts: # Get 2 known and 1 made up qualified names, and confirm only 2 # come back - lm.get_accounts_if_exists( + ledger_manager.get_accounts_if_exists( qualified_names=[ account1.qualified_name, account2.qualified_name, @@ -233,53 +291,49 @@ class TestThlLedgerManagerAccounts: assert len(res) == 2 # Confirm an empty array comes back for all unknown qualified names - res = lm.get_accounts_if_exists( + res = ledger_manager.get_accounts_if_exists( qualified_names=[ - f"{lm.currency.value}:bp_wall:{uuid4().hex}" for i in range(5) + f"{ledger_manager.currency.value}:bp_wall:{uuid4().hex}" + for _ in range(5) ] ) assert isinstance(res, list) assert len(res) == 0 - def test_get_accounts_for_products(self, product_factory, thl_lm, lm): - from generalresearch.managers.thl.ledger_manager.exceptions import ( - LedgerAccountDoesntExistError, - ) - from generalresearch.models.thl.ledger import ( - LedgerAccount, - ) - + def test_get_accounts_for_products( + self, + product_factory: Callable[..., Product], + thl_ledger_manager: ThlLedgerManager, + ): # Create 5 Products product_uuids = [] - for i in range(5): + for _ in range(5): _p = product_factory() product_uuids.append(_p.uuid) # Confirm that this fails.. because none of those accounts have been # created yet with pytest.raises(expected_exception=LedgerAccountDoesntExistError): - thl_lm.get_accounts_bp_wallet_for_products(product_uuids=product_uuids) + thl_ledger_manager.get_accounts_bp_wallet_for_products( + product_uuids=product_uuids + ) # Create the bp_wallet accounts and then try again for p_uuid in product_uuids: - thl_lm.get_account_or_create_bp_wallet_by_uuid(product_uuid=p_uuid) + thl_ledger_manager.get_account_or_create_bp_wallet_by_uuid( + product_uuid=p_uuid + ) - res = thl_lm.get_accounts_bp_wallet_for_products(product_uuids=product_uuids) + res = thl_ledger_manager.get_accounts_bp_wallet_for_products( + product_uuids=product_uuids + ) assert len(res) == len(product_uuids) assert all([isinstance(i, LedgerAccount) for i in res]) class TestLedgerAccountManager: - def test_get_or_create(self, thl_lm, lm, lam): - from generalresearch.managers.thl.ledger_manager.exceptions import ( - LedgerAccountDoesntExistError, - ) - from generalresearch.models.thl.ledger import ( - AccountType, - Direction, - LedgerAccount, - ) + def test_get_or_create(self, ledger_account_manager: LedgerAccountManager): u = uuid4().hex name = f"test-{u[:8]}" @@ -306,39 +360,42 @@ class TestLedgerAccountManager: assert isinstance(instance, LedgerAccount) assert instance.reference_uuid == u - def test_get(self, user, thl_lm, lm, lam): - from generalresearch.managers.thl.ledger_manager.exceptions import ( - LedgerAccountDoesntExistError, - ) - from generalresearch.models.thl.ledger import ( - AccountType, - LedgerAccount, - ) + def test_get( + self, + user: User, + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + ledger_account_manager: LedgerAccountManager, + ): with pytest.raises(LedgerAccountDoesntExistError): - lam.get_account(qualified_name=f"test:bp_wallet:{user.product.id}") + ledger_account_manager.get_account( + qualified_name=f"test:bp_wallet:{user.product.id}" + ) - thl_lm.get_account_or_create_bp_wallet(product=user.product) - account = lam.get_account(qualified_name=f"test:bp_wallet:{user.product.id}") + thl_ledger_manager.get_account_or_create_bp_wallet(product=user.product) + account = ledger_account_manager.get_account( + qualified_name=f"test:bp_wallet:{user.product.id}" + ) assert isinstance(account, LedgerAccount) assert AccountType.BP_WALLET == account.account_type assert user.product.uuid == account.reference_uuid - def test_get_many(self, product_factory, thl_lm, lm, lam, currency): - from generalresearch.managers.thl.ledger_manager.exceptions import ( - LedgerAccountDoesntExistError, - ) - from generalresearch.models.thl.product import Product - + def test_get_many( + self, + product_factory: Callable[..., Product], + thl_ledger_manager: ThlLedgerManager, + ledger_account_manager: LedgerAccountManager, + ): p1: Product = product_factory() p2: Product = product_factory() - account1 = thl_lm.get_account_or_create_bp_wallet(product=p1) - account2 = thl_lm.get_account_or_create_bp_wallet(product=p2) + account1 = thl_ledger_manager.get_account_or_create_bp_wallet(product=p1) + account2 = thl_ledger_manager.get_account_or_create_bp_wallet(product=p2) # Get 1 known account and confirm it comes back - res = lam.get_account_many( + res = ledger_account_manager.get_account_many( qualified_names=[account1.qualified_name, account2.qualified_name] ) assert isinstance(res, list) @@ -346,7 +403,7 @@ class TestLedgerAccountManager: assert account1 in res # Get 2 known accounts and confirm they both come back - res = lam.get_account_many( + res = ledger_account_manager.get_account_many( qualified_names=[account1.qualified_name, account2.qualified_name] ) assert isinstance(res, list) @@ -356,7 +413,7 @@ class TestLedgerAccountManager: # Get 2 known and 1 made up qualified names, and confirm only 2 come # back. Don't raise on error, so we can confirm the array is "short" - res = lam.get_account_many( + res = ledger_account_manager.get_account_many( qualified_names=[ account1.qualified_name, account2.qualified_name, @@ -369,7 +426,7 @@ class TestLedgerAccountManager: # Same as above, but confirm the raise works on checking res length with pytest.raises(LedgerAccountDoesntExistError): - lam.get_account_many( + ledger_account_manager.get_account_many( qualified_names=[ account1.qualified_name, account2.qualified_name, @@ -379,19 +436,14 @@ class TestLedgerAccountManager: ) # Confirm an empty array comes back for all unknown qualified names - res = lam.get_account_many( - qualified_names=[f"test:bp_wall:{uuid4().hex}" for i in range(5)], + res = ledger_account_manager.get_account_many( + qualified_names=[f"test:bp_wall:{uuid4().hex}" for _ in range(5)], raise_on_error=False, ) assert isinstance(res, list) assert len(res) == 0 - def test_create_account(self, thl_lm, lm, lam): - from generalresearch.models.thl.ledger import ( - AccountType, - Direction, - LedgerAccount, - ) + def test_create_account(self, ledger_account_manager: LedgerAccountManager): u = uuid4().hex name = f"test-{u[:8]}" @@ -406,6 +458,6 @@ class TestLedgerAccountManager: reference_uuid=u, ) - lam.create_account(account=account) - assert lam.get_account(f"test:bp_wallet:{u}") == account - assert lam.get_account_or_create(account) == account + ledger_account_manager.create_account(account=account) + assert ledger_account_manager.get_account(f"test:bp_wallet:{u}") == account + assert ledger_account_manager.get_account_or_create(account) == account diff --git a/tests/managers/thl/test_ledger/test_thl_lm_bp_payout.py b/tests/managers/thl/test_ledger/test_thl_lm_bp_payout.py index cfb8f8f..e4a25a3 100644 --- a/tests/managers/thl/test_ledger/test_thl_lm_bp_payout.py +++ b/tests/managers/thl/test_ledger/test_thl_lm_bp_payout.py @@ -1,4 +1,7 @@ +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 @@ -9,7 +12,7 @@ import redis from pydantic import RedisDsn from redis.lock import Lock -from generalresearch.currency import USDCent +from generalresearch.currency import LedgerCurrency, USDCent from generalresearch.managers.base import Permission from generalresearch.managers.thl.ledger_manager.exceptions import ( LedgerTransactionConditionFailedError, @@ -19,9 +22,13 @@ from generalresearch.managers.thl.ledger_manager.exceptions import ( ) from generalresearch.managers.thl.ledger_manager.ledger import LedgerTransaction from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager +from generalresearch.managers.thl.payout import ( + BrokerageProductPayoutEventManager, +) from generalresearch.models import Source from generalresearch.models.thl.definitions import PayoutStatus from generalresearch.models.thl.ledger import Direction, TransactionType +from generalresearch.models.thl.product import Product from generalresearch.models.thl.session import ( Session, Status, @@ -30,6 +37,7 @@ from generalresearch.models.thl.session import ( ) from generalresearch.models.thl.user import User from generalresearch.models.thl.wallet import PayoutType +from generalresearch.pg_helper import PostgresConfig from generalresearch.redis_helper import RedisConfig @@ -45,12 +53,12 @@ class TestThlLedgerManagerBPPayout: def test_create_tx_with_bp_payment( self, - user_factory, - product_user_wallet_no, - create_main_accounts, + user_factory: Callable[..., User], + product_user_wallet_no: Product, + create_main_accounts: Callable[..., None], caplog, - thl_lm, - delete_ledger_db, + thl_ledger_manager: ThlLedgerManager, + delete_ledger_db: Callable[..., None], ): delete_ledger_db() create_main_accounts() @@ -69,25 +77,29 @@ class TestThlLedgerManagerBPPayout: started=now, finished=now + timedelta(seconds=1), ) - tx = thl_lm.create_tx_task_complete( + tx = thl_ledger_manager.create_tx_task_complete( wall=wall1, user=user, created=wall1.started ) assert isinstance(tx, LedgerTransaction) session = Session(started=wall1.started, user=user, wall_events=[wall1]) status, status_code_1 = session.determine_session_status() - thl_net, commission_amount, bp_pay, user_pay = session.determine_payments() + _, _, bp_pay, user_pay = session.determine_payments() session.update( - status=status, status_code_1=status_code_1, finished=now + timedelta(minutes=10), payout=bp_pay, user_payout=user_pay + status=status, + status_code_1=status_code_1, + finished=now + timedelta(minutes=10), + payout=bp_pay, + user_payout=user_pay, ) - thl_lm.create_tx_bp_payment(session=session, created=wall1.started) + thl_ledger_manager.create_tx_bp_payment(session=session, created=wall1.started) lock_key = f"test:bp_payout:{user.product.id}" - flag_name = f"{thl_lm.cache_prefix}:transaction_flag:{lock_key}" - thl_lm.redis_client.delete(flag_name) + flag_name = f"{thl_ledger_manager.cache_prefix}:transaction_flag:{lock_key}" + thl_ledger_manager.redis_client.delete(flag_name) payoutevent_uuid = uuid4().hex - thl_lm.create_tx_bp_payout( + thl_ledger_manager.create_tx_bp_payout( product=user.product, amount=USDCent(200), created=now, @@ -95,7 +107,7 @@ class TestThlLedgerManagerBPPayout: ) payoutevent_uuid = uuid4().hex - thl_lm.create_tx_bp_payout( + thl_ledger_manager.create_tx_bp_payout( product=user.product, amount=USDCent(200), created=now + timedelta(minutes=2), @@ -103,13 +115,15 @@ class TestThlLedgerManagerBPPayout: payoutevent_uuid=payoutevent_uuid, ) - cash = thl_lm.get_account_cash() - bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(user.product) - assert 170 == thl_lm.get_account_balance(bp_wallet_account) - assert 200 == thl_lm.get_account_balance(cash) + cash = thl_ledger_manager.get_account_cash() + bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet( + user.product + ) + assert 170 == thl_ledger_manager.get_account_balance(bp_wallet_account) + assert 200 == thl_ledger_manager.get_account_balance(cash) with pytest.raises(expected_exception=LedgerTransactionFlagAlreadyExistsError): - thl_lm.create_tx_bp_payout( + thl_ledger_manager.create_tx_bp_payout( user.product, amount=USDCent(200), created=now + timedelta(minutes=2), @@ -121,7 +135,7 @@ class TestThlLedgerManagerBPPayout: payoutevent_uuid = uuid4().hex with caplog.at_level(logging.INFO): with pytest.raises(LedgerTransactionConditionFailedError): - thl_lm.create_tx_bp_payout( + thl_ledger_manager.create_tx_bp_payout( user.product, amount=USDCent(10_000), created=now + timedelta(minutes=2), @@ -131,7 +145,7 @@ class TestThlLedgerManagerBPPayout: ) assert "failed condition check balance:" in caplog.text - thl_lm.create_tx_bp_payout( + thl_ledger_manager.create_tx_bp_payout( product=user.product, amount=USDCent(10_00), created=now + timedelta(minutes=2), @@ -139,16 +153,22 @@ class TestThlLedgerManagerBPPayout: skip_wallet_balance_check=True, payoutevent_uuid=payoutevent_uuid, ) - assert 170 - 1000 == thl_lm.get_account_balance(bp_wallet_account) + assert 170 - 1000 == thl_ledger_manager.get_account_balance(bp_wallet_account) - def test_create_tx(self, product, caplog, thl_lm, currency): + def test_create_tx( + self, + product: Product, + caplog, + thl_ledger_manager: ThlLedgerManager, + currency: LedgerCurrency, + ): rand_amount: USDCent = USDCent(randint(100, 1_000)) payoutevent_uuid = uuid4().hex # Create a BP Payout for a Product without any activity. By issuing, # the skip_* checks, we should be able to force it to work, and will # then ultimately result in a negative balance - tx = thl_lm.create_tx_bp_payout( + tx = thl_ledger_manager.create_tx_bp_payout( product=product, amount=rand_amount, payoutevent_uuid=payoutevent_uuid, @@ -171,15 +191,15 @@ class TestThlLedgerManagerBPPayout: # Check the Product's balance, it should be negative the amount that was # paid out. That's because the Product earned nothing.. and then was # sent something. - balance = thl_lm.get_account_balance( - account=thl_lm.get_account_or_create_bp_wallet(product=product) + balance = thl_ledger_manager.get_account_balance( + account=thl_ledger_manager.get_account_or_create_bp_wallet(product=product) ) assert balance == int(rand_amount) * -1 # Test some basic assertions with caplog.at_level(logging.INFO): with pytest.raises(expected_exception=Exception): - thl_lm.create_tx_bp_payout( + thl_ledger_manager.create_tx_bp_payout( product=product, amount=rand_amount, payoutevent_uuid=uuid4().hex, @@ -190,12 +210,17 @@ class TestThlLedgerManagerBPPayout: ) assert "failed condition check >1 tx per day" in caplog.text - def test_create_tx_redis_failure(self, product, thl_web_rw, thl_lm): + def test_create_tx_redis_failure( + self, + product: Product, + thl_web_rw: PostgresConfig, + thl_ledger_manager: ThlLedgerManager, + ): rand_amount: USDCent = USDCent(randint(100, 1_000)) payoutevent_uuid = uuid4().hex now = datetime.now(tz=UTC) - thl_lm.create_tx_plug_bp_wallet( + thl_ledger_manager.create_tx_plug_bp_wallet( product, rand_amount, now, direction=Direction.CREDIT ) @@ -216,7 +241,7 @@ class TestThlLedgerManagerBPPayout: ) with pytest.raises(expected_exception=Exception) as e: - tx = thl_lm_redis_0.create_tx_bp_payout( + thl_lm_redis_0.create_tx_bp_payout( product=product, amount=rand_amount, payoutevent_uuid=payoutevent_uuid, @@ -224,21 +249,27 @@ class TestThlLedgerManagerBPPayout: ) assert e.type is redis.exceptions.TimeoutError # No txs were created - bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(product=product) - txs = thl_lm.get_tx_filtered_by_account(account_uuid=bp_wallet_account.uuid) + bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet( + product=product + ) + txs = thl_ledger_manager.get_tx_filtered_by_account( + account_uuid=bp_wallet_account.uuid + ) txs = [tx for tx in txs if tx.metadata["tx_type"] != "plug"] assert len(txs) == 0 - def test_create_tx_multiple_per_day(self, product, thl_lm): + def test_create_tx_multiple_per_day( + self, product: Product, thl_ledger_manager: ThlLedgerManager + ): rand_amount: USDCent = USDCent(randint(100, 1_000)) payoutevent_uuid = uuid4().hex now = datetime.now(tz=UTC) - thl_lm.create_tx_plug_bp_wallet( + thl_ledger_manager.create_tx_plug_bp_wallet( product, rand_amount * USDCent(2), now, direction=Direction.CREDIT ) - tx = thl_lm.create_tx_bp_payout( + thl_ledger_manager.create_tx_bp_payout( product=product, amount=rand_amount, payoutevent_uuid=payoutevent_uuid, @@ -248,7 +279,7 @@ class TestThlLedgerManagerBPPayout: # Try to create another # Will fail b/c it has the same payout event uuid with pytest.raises(expected_exception=Exception) as e: - tx = thl_lm.create_tx_bp_payout( + thl_ledger_manager.create_tx_bp_payout( product=product, amount=rand_amount, payoutevent_uuid=payoutevent_uuid, @@ -260,7 +291,7 @@ class TestThlLedgerManagerBPPayout: # Will fail due to multiple per day payoutevent_uuid2 = uuid4().hex with pytest.raises(expected_exception=Exception) as e: - tx = thl_lm.create_tx_bp_payout( + tx = thl_ledger_manager.create_tx_bp_payout( product=product, amount=rand_amount, payoutevent_uuid=payoutevent_uuid2, @@ -270,7 +301,7 @@ class TestThlLedgerManagerBPPayout: assert str(e.value) == ">1 tx per day" # Make it run by skipping one per day check - tx = thl_lm.create_tx_bp_payout( + thl_ledger_manager.create_tx_bp_payout( product=product, amount=rand_amount, payoutevent_uuid=payoutevent_uuid2, @@ -278,13 +309,17 @@ class TestThlLedgerManagerBPPayout: skip_one_per_day_check=True, ) - def test_create_tx_redis_lock_release_error(self, product, thl_lm): + def test_create_tx_redis_lock_release_error( + self, product: Product, thl_ledger_manager: ThlLedgerManager + ): rand_amount: USDCent = USDCent(randint(100, 1_000)) payoutevent_uuid = uuid4().hex now = datetime.now(tz=UTC) - bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(product=product) + bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet( + product=product + ) - thl_lm.create_tx_plug_bp_wallet( + thl_ledger_manager.create_tx_plug_bp_wallet( product, rand_amount * USDCent(2), now, direction=Direction.CREDIT ) @@ -294,7 +329,7 @@ class TestThlLedgerManagerBPPayout: # Create TX will fail on lock enter, no tx will actually get created with pytest.raises(expected_exception=Exception) as e: - tx = thl_lm.create_tx_bp_payout( + thl_ledger_manager.create_tx_bp_payout( product=product, amount=rand_amount, payoutevent_uuid=payoutevent_uuid, @@ -302,7 +337,9 @@ class TestThlLedgerManagerBPPayout: ) assert e.type is LedgerTransactionCreateError assert str(e.value) == "Redis error: Simulated timeout during acquire" - txs = thl_lm.get_tx_filtered_by_account(account_uuid=bp_wallet_account.uuid) + txs = thl_ledger_manager.get_tx_filtered_by_account( + account_uuid=bp_wallet_account.uuid + ) txs = [tx for tx in txs if tx.metadata["tx_type"] != "plug"] assert len(txs) == 0 @@ -311,7 +348,7 @@ class TestThlLedgerManagerBPPayout: # Create TX will fail on lock exit, after the tx was created! with pytest.raises(expected_exception=Exception) as e: - tx = thl_lm.create_tx_bp_payout( + tx = thl_ledger_manager.create_tx_bp_payout( product=product, amount=rand_amount, payoutevent_uuid=payoutevent_uuid, @@ -321,7 +358,9 @@ class TestThlLedgerManagerBPPayout: assert str(e.value) == "Redis error: Simulated timeout during release" # Transaction was still created! - txs = thl_lm.get_tx_filtered_by_account(account_uuid=bp_wallet_account.uuid) + txs = thl_ledger_manager.get_tx_filtered_by_account( + account_uuid=bp_wallet_account.uuid + ) txs = [tx for tx in txs if tx.metadata["tx_type"] != "plug"] assert len(txs) == 1 Lock.release = original_release @@ -329,34 +368,45 @@ class TestThlLedgerManagerBPPayout: class TestPayoutEventManagerBPPayout: - def test_create(self, product, thl_lm, brokerage_product_payout_event_manager): + def test_create( + self, + product: Product, + thl_ledger_manager: ThlLedgerManager, + brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager, + ): rand_amount: USDCent = USDCent(randint(100, 1_000)) now = datetime.now(tz=UTC) - bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(product=product) - assert thl_lm.get_account_balance(bp_wallet_account) == 0 - thl_lm.create_tx_plug_bp_wallet( + bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet( + product=product + ) + assert thl_ledger_manager.get_account_balance(bp_wallet_account) == 0 + thl_ledger_manager.create_tx_plug_bp_wallet( product, rand_amount, now, direction=Direction.CREDIT ) - assert thl_lm.get_account_balance(bp_wallet_account) == rand_amount + assert thl_ledger_manager.get_account_balance(bp_wallet_account) == rand_amount brokerage_product_payout_event_manager.set_account_lookup_table(thl_lm=thl_lm) pe = brokerage_product_payout_event_manager.create_bp_payout_event( - thl_ledger_manager=thl_lm, + thl_ledger_manager=thl_ledger_manager, product=product, created=now, amount=rand_amount, payout_type=PayoutType.ACH, ) assert brokerage_product_payout_event_manager.check_for_ledger_tx( - thl_ledger_manager=thl_lm, + thl_ledger_manager=thl_ledger_manager, product_id=product.id, amount=rand_amount, payout_event=pe, ) - assert thl_lm.get_account_balance(bp_wallet_account) == 0 + assert thl_ledger_manager.get_account_balance(bp_wallet_account) == 0 def test_create_with_redis_error( - self, product, caplog, thl_lm, brokerage_product_payout_event_manager + self, + product: Product, + caplog, + thl_ledger_manager: ThlLedgerManager, + brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager, ): caplog.set_level("WARNING") original_acquire = Lock.acquire @@ -364,19 +414,23 @@ class TestPayoutEventManagerBPPayout: rand_amount: USDCent = USDCent(randint(100, 1_000)) now = datetime.now(tz=UTC) - bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(product=product) - assert thl_lm.get_account_balance(bp_wallet_account) == 0 - thl_lm.create_tx_plug_bp_wallet( - product, rand_amount, now, direction=Direction.CREDIT + bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet( + product=product + ) + assert thl_ledger_manager.get_account_balance(bp_wallet_account) == 0 + thl_ledger_manager.create_tx_plug_bp_wallet( + product=product, amount=rand_amount, now=now, direction=Direction.CREDIT + ) + assert thl_ledger_manager.get_account_balance(bp_wallet_account) == rand_amount + brokerage_product_payout_event_manager.set_account_lookup_table( + thl_lm=thl_ledger_manager ) - assert thl_lm.get_account_balance(bp_wallet_account) == rand_amount - brokerage_product_payout_event_manager.set_account_lookup_table(thl_lm=thl_lm) # Will fail on lock enter, no tx will actually get created Lock.acquire = broken_acquire with pytest.raises(expected_exception=Exception) as e: pe = brokerage_product_payout_event_manager.create_bp_payout_event( - thl_ledger_manager=thl_lm, + thl_ledger_manager=thl_ledger_manager, product=product, created=now, amount=rand_amount, @@ -389,13 +443,15 @@ class TestPayoutEventManagerBPPayout: for m in caplog.messages ) - txs = thl_lm.get_tx_filtered_by_account(account_uuid=bp_wallet_account.uuid) + txs = thl_ledger_manager.get_tx_filtered_by_account( + account_uuid=bp_wallet_account.uuid + ) txs = [tx for tx in txs if tx.metadata["tx_type"] != "plug"] # One payout event is created, status is failed, and no ledger txs exist assert len(txs) == 0 pes = ( brokerage_product_payout_event_manager.get_bp_bp_payout_events_for_products( - thl_ledger_manager=thl_lm, product_uuids=[product.id] + thl_ledger_manager=thl_ledger_manager, product_uuids=[product.id] ) ) assert len(pes) == 1 @@ -407,17 +463,21 @@ class TestPayoutEventManagerBPPayout: # Try to fix the failed payout, by trying ledger tx again brokerage_product_payout_event_manager.retry_create_bp_payout_event_tx( - product=product, thl_ledger_manager=thl_lm, payout_event_uuid=pe.uuid + product=product, + thl_ledger_manager=thl_ledger_manager, + payout_event_uuid=pe.uuid, + ) + txs = thl_ledger_manager.get_tx_filtered_by_account( + account_uuid=bp_wallet_account.uuid ) - txs = thl_lm.get_tx_filtered_by_account(account_uuid=bp_wallet_account.uuid) txs = [tx for tx in txs if tx.metadata["tx_type"] != "plug"] assert len(txs) == 1 - assert thl_lm.get_account_balance(bp_wallet_account) == 0 + assert thl_ledger_manager.get_account_balance(bp_wallet_account) == 0 # And then try to run it again, it'll fail because a payout event with the same info exists with pytest.raises(expected_exception=Exception) as e: pe = brokerage_product_payout_event_manager.create_bp_payout_event( - thl_ledger_manager=thl_lm, + thl_ledger_manager=thl_ledger_manager, product=product, created=now, amount=rand_amount, @@ -432,7 +492,7 @@ class TestPayoutEventManagerBPPayout: now = datetime.now(tz=UTC) with pytest.raises(LedgerTransactionConditionFailedError) as e: pe = brokerage_product_payout_event_manager.create_bp_payout_event( - thl_ledger_manager=thl_lm, + thl_ledger_manager=thl_ledger_manager, product=product, created=now, amount=rand_amount, @@ -446,7 +506,7 @@ class TestPayoutEventManagerBPPayout: # And if we really want to, we can make it again now = datetime.now(tz=UTC) pe = brokerage_product_payout_event_manager.create_bp_payout_event( - thl_ledger_manager=thl_lm, + thl_ledger_manager=thl_ledger_manager, product=product, created=now, amount=rand_amount, @@ -455,17 +515,25 @@ class TestPayoutEventManagerBPPayout: skip_wallet_balance_check=True, ) - txs = thl_lm.get_tx_filtered_by_account(account_uuid=bp_wallet_account.uuid) + txs = thl_ledger_manager.get_tx_filtered_by_account( + account_uuid=bp_wallet_account.uuid + ) txs = [tx for tx in txs if tx.metadata["tx_type"] != "plug"] assert len(txs) == 2 # since they were paid twice - assert thl_lm.get_account_balance(bp_wallet_account) == 0 - rand_amount + assert ( + thl_ledger_manager.get_account_balance(bp_wallet_account) == 0 - rand_amount + ) Lock.release = original_release Lock.acquire = original_acquire def test_create_with_redis_error_release( - self, product, caplog, thl_lm, brokerage_product_payout_event_manager + self, + product: Product, + caplog, + thl_ledger_manager: ThlLedgerManager, + brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager, ): caplog.set_level("WARNING") @@ -473,20 +541,24 @@ class TestPayoutEventManagerBPPayout: rand_amount: USDCent = USDCent(randint(100, 1_000)) now = datetime.now(tz=UTC) - bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(product=product) - brokerage_product_payout_event_manager.set_account_lookup_table(thl_lm=thl_lm) + bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet( + product=product + ) + brokerage_product_payout_event_manager.set_account_lookup_table( + thl_lm=thl_ledger_manager + ) - assert thl_lm.get_account_balance(bp_wallet_account) == 0 - thl_lm.create_tx_plug_bp_wallet( + assert thl_ledger_manager.get_account_balance(bp_wallet_account) == 0 + thl_ledger_manager.create_tx_plug_bp_wallet( product, rand_amount, now, direction=Direction.CREDIT ) - assert thl_lm.get_account_balance(bp_wallet_account) == rand_amount + assert thl_ledger_manager.get_account_balance(bp_wallet_account) == rand_amount # Will fail on lock exit, after the tx was created! # But it'll see that the tx was created and so everything will be fine Lock.release = broken_release pe = brokerage_product_payout_event_manager.create_bp_payout_event( - thl_ledger_manager=thl_lm, + thl_ledger_manager=thl_ledger_manager, product=product, created=now, amount=rand_amount, @@ -497,12 +569,14 @@ class TestPayoutEventManagerBPPayout: for m in caplog.messages ) - txs = thl_lm.get_tx_filtered_by_account(account_uuid=bp_wallet_account.uuid) + txs = thl_ledger_manager.get_tx_filtered_by_account( + account_uuid=bp_wallet_account.uuid + ) txs = [tx for tx in txs if tx.metadata["tx_type"] != "plug"] assert len(txs) == 1 pes = ( brokerage_product_payout_event_manager.get_bp_bp_payout_events_for_products( - thl_ledger_manager=thl_lm, product_uuids=[product.uuid] + thl_ledger_manager=thl_ledger_manager, product_uuids=[product.uuid] ) ) assert len(pes) == 1 diff --git a/tests/managers/thl/test_ledger/test_thl_lm_tx.py b/tests/managers/thl/test_ledger/test_thl_lm_tx.py index 6fb0a0f..89adb0b 100644 --- a/tests/managers/thl/test_ledger/test_thl_lm_tx.py +++ b/tests/managers/thl/test_ledger/test_thl_lm_tx.py @@ -1,4 +1,7 @@ +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 @@ -6,19 +9,30 @@ from uuid import uuid4 import pytest -from generalresearch.currency import USDCent +from generalresearch.currency import LedgerCurrency, USDCent from generalresearch.managers.thl.ledger_manager.ledger import ( + LedgerManager, LedgerTransaction, ) +from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager +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 ( WALL_ALLOWED_STATUS_STATUS_CODE, ) -from generalresearch.models.thl.ledger import Direction, TransactionType +from generalresearch.models.thl.ledger import ( + AccountType, + Direction, + LedgerAccount, + TransactionType, +) from generalresearch.models.thl.payout import UserPayoutEvent from generalresearch.models.thl.product import ( PayoutConfig, PayoutTransformation, + Product, UserWalletConfig, ) from generalresearch.models.thl.session import ( @@ -38,45 +52,50 @@ class TestThlLedgerTxManager: def test_create_tx_task_complete( self, - wall, - user, - account_revenue_task_complete, - create_main_accounts, - thl_lm, - lm, + wall: Wall, + user: User, + account_revenue_task_complete: LedgerAccount, + create_main_accounts: Callable[..., None], + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, ): create_main_accounts() - tx = thl_lm.create_tx_task_complete(wall=wall, user=user) + tx = thl_ledger_manager.create_tx_task_complete(wall=wall, user=user) assert isinstance(tx, LedgerTransaction) - res = lm.get_tx_by_id(transaction_id=tx.id) + res = ledger_manager.get_tx_by_id(transaction_id=tx.id) assert res.created == tx.created def test_create_tx_task_complete_( - self, wall, user, account_revenue_task_complete, thl_lm, lm + self, + wall: Wall, + user: User, + account_revenue_task_complete: LedgerAccount, + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, ): - tx = thl_lm.create_tx_task_complete_(wall=wall, user=user) + tx = thl_ledger_manager.create_tx_task_complete_(wall=wall, user=user) assert isinstance(tx, LedgerTransaction) - res = lm.get_tx_by_id(transaction_id=tx.id) + res = ledger_manager.get_tx_by_id(transaction_id=tx.id) assert res.created == tx.created def test_create_tx_bp_payment( self, - session_factory, - user, - create_main_accounts, - delete_ledger_db, - thl_lm, - lm, - session_manager, + session_factory: Callable[..., Session], + user: User, + create_main_accounts: Callable[..., None], + delete_ledger_db: Callable[..., None], + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + session_manager: SessionManager, ): delete_ledger_db() create_main_accounts() s1 = session_factory(user=user) - status, status_code_1 = s1.determine_session_status() - thl_net, commission_amount, bp_pay, user_pay = s1.determine_payments() + _, status_code_1 = s1.determine_session_status() + _, _, bp_pay, user_pay = s1.determine_payments() session_manager.finish_with_status( session=s1, status=Status.COMPLETE, @@ -86,22 +105,22 @@ class TestThlLedgerTxManager: user_payout=user_pay, ) - tx = thl_lm.create_tx_bp_payment(session=s1) + tx = thl_ledger_manager.create_tx_bp_payment(session=s1) assert isinstance(tx, LedgerTransaction) - res = lm.get_tx_by_id(transaction_id=tx.id) + res = ledger_manager.get_tx_by_id(transaction_id=tx.id) assert res.created == tx.created def test_create_tx_bp_payment_amt( self, - session_factory, - user_factory, - product_manager, - create_main_accounts, - delete_ledger_db, - thl_lm, - lm, - session_manager, + session_factory: Callable[..., Session], + user_factory: Callable[..., User], + product_manager: ProductManager, + create_main_accounts: Callable[..., None], + delete_ledger_db: Callable[..., None], + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + session_manager: SessionManager, ): delete_ledger_db() create_main_accounts() @@ -119,7 +138,7 @@ class TestThlLedgerTxManager: status, status_code_1 = s1.determine_session_status() assert status == Status.COMPLETE thl_net, commission_amount, bp_pay, user_pay = s1.determine_payments( - thl_ledger_manager=thl_lm + thl_ledger_manager=thl_ledger_manager ) print(thl_net, commission_amount, bp_pay, user_pay) session_manager.finish_with_status( @@ -131,25 +150,25 @@ class TestThlLedgerTxManager: user_payout=user_pay, ) - tx = thl_lm.create_tx_bp_payment(session=s1) + tx = thl_ledger_manager.create_tx_bp_payment(session=s1) assert isinstance(tx, LedgerTransaction) - res = lm.get_tx_by_id(transaction_id=tx.id) + res = ledger_manager.get_tx_by_id(transaction_id=tx.id) assert res.created == tx.created def test_create_tx_bp_payment_( self, - session_factory, - user, - create_main_accounts, - thl_lm, - lm, - session_manager, - utc_hour_ago, + session_factory: Callable[..., Session], + user: User, + create_main_accounts: Callable[..., None], + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + session_manager: SessionManager, + utc_hour_ago: datetime, ): s1 = session_factory(user=user) status, status_code_1 = s1.determine_session_status() - thl_net, commission_amount, bp_pay, user_pay = s1.determine_payments() + _, _, bp_pay, user_pay = s1.determine_payments() session_manager.finish_with_status( session=s1, status=status, @@ -160,14 +179,20 @@ class TestThlLedgerTxManager: ) s1.determine_payments() - tx = thl_lm.create_tx_bp_payment_(session=s1) + tx = thl_ledger_manager.create_tx_bp_payment_(session=s1) assert isinstance(tx, LedgerTransaction) - res = lm.get_tx_by_id(transaction_id=tx.id) + res = ledger_manager.get_tx_by_id(transaction_id=tx.id) assert res.created == tx.created def test_create_tx_task_adjustment( - self, wall_factory, session, user, create_main_accounts, thl_lm, lm + self, + wall_factory: Callable[..., Wall], + session: Session, + user: User, + create_main_accounts: Callable[..., None], + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, ): """Create Wall event Complete, and Create a Tx Task Adjustment @@ -179,16 +204,23 @@ class TestThlLedgerTxManager: wall_status = Status.COMPLETE wall: Wall = wall_factory(session=session, wall_status=wall_status) - tx = thl_lm.create_tx_task_adjustment(wall=wall, user=user) + tx = thl_ledger_manager.create_tx_task_adjustment(wall=wall, user=user) assert isinstance(tx, LedgerTransaction) - res = lm.get_tx_by_id(transaction_id=tx.id) + res = ledger_manager.get_tx_by_id(transaction_id=tx.id) assert res.entries[0].amount == int(wall.cpi * 100) assert res.entries[1].amount == int(wall.cpi * 100) assert wall.source.name in res.ext_description assert res.created == tx.created - def test_create_tx_bp_adjustment(self, session, user, caplog, thl_lm, lm): + def test_create_tx_bp_adjustment( + self, + session: Session, + user: User, + caplog, + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + ): status, status_code_1 = session.determine_session_status() thl_net, commission_amount, bp_pay, user_pay = session.determine_payments() @@ -209,18 +241,24 @@ class TestThlLedgerTxManager: session.update(finished=datetime.now(tz=UTC) + timedelta(minutes=10)) assert session.finished with caplog.at_level(logging.INFO): - tx = thl_lm.create_tx_bp_adjustment(session=session) + tx = thl_ledger_manager.create_tx_bp_adjustment(session=session) assert tx is None assert "No transactions needed." in caplog.text - def test_create_tx_bp_payout(self, product, caplog, thl_lm, currency): + def test_create_tx_bp_payout( + self, + product: Product, + caplog, + thl_ledger_manager: ThlLedgerManager, + currency: LedgerCurrency, + ): rand_amount: USDCent = USDCent(randint(100, 1_000)) payoutevent_uuid = uuid4().hex # Create a BP Payout for a Product without any activity. By issuing, # the skip_* checks, we should be able to force it to work, and will # then ultimately result in a negative balance - tx = thl_lm.create_tx_bp_payout( + tx = thl_ledger_manager.create_tx_bp_payout( product=product, amount=rand_amount, payoutevent_uuid=payoutevent_uuid, @@ -235,7 +273,7 @@ class TestThlLedgerTxManager: assert tx.ext_description == "BP Payout" assert ( tx.tag - == f"{thl_lm.currency.value}:{TransactionType.BP_PAYOUT.value}:{payoutevent_uuid}" + == f"{thl_ledger_manager.currency.value}:{TransactionType.BP_PAYOUT.value}:{payoutevent_uuid}" ) assert tx.entries[0].amount == rand_amount assert tx.entries[1].amount == rand_amount @@ -243,15 +281,15 @@ class TestThlLedgerTxManager: # Check the Product's balance, it should be negative the amount that was # paid out. That's because the Product earned nothing.. and then was # sent something. - balance = thl_lm.get_account_balance( - account=thl_lm.get_account_or_create_bp_wallet(product=product) + balance = thl_ledger_manager.get_account_balance( + account=thl_ledger_manager.get_account_or_create_bp_wallet(product=product) ) assert balance == int(rand_amount) * -1 # Test some basic assertions with caplog.at_level(logging.INFO): with pytest.raises(expected_exception=Exception): - thl_lm.create_tx_bp_payout( + thl_ledger_manager.create_tx_bp_payout( product=product, amount=rand_amount, payoutevent_uuid=uuid4().hex, @@ -262,7 +300,13 @@ class TestThlLedgerTxManager: ) assert "failed condition check >1 tx per day" in caplog.text - def test_create_tx_bp_payout_(self, product, thl_lm, lm, currency): + def test_create_tx_bp_payout_( + self, + product: Product, + thl_lm: ThlLedgerManager, + ledger_manager: LedgerManager, + currency: LedgerCurrency, + ): rand_amount: USDCent = USDCent(randint(100, 1_000)) payoutevent_uuid = uuid4().hex @@ -285,14 +329,19 @@ class TestThlLedgerTxManager: assert tx.entries[1].amount == rand_amount def test_create_tx_plug_bp_wallet( - self, product, create_main_accounts, thl_lm, lm, currency + self, + product: Product, + create_main_accounts: Callable[..., None], + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + currency: LedgerCurrency, ): """A BP Wallet "plug" is a way to makeup discrepancies and simply add or remove money """ rand_amount: USDCent = USDCent(randint(100, 1_000)) - tx = thl_lm.create_tx_plug_bp_wallet( + tx = thl_ledger_manager.create_tx_plug_bp_wallet( product=product, amount=rand_amount, created=datetime.now(tz=UTC), @@ -304,13 +353,18 @@ class TestThlLedgerTxManager: # We issued the BP money they didn't earn, so now they have a # negative balance - balance = thl_lm.get_account_balance( - account=thl_lm.get_account_or_create_bp_wallet(product=product) + balance = thl_ledger_manager.get_account_balance( + account=thl_ledger_manager.get_account_or_create_bp_wallet(product=product) ) assert balance == int(rand_amount) * -1 def test_create_tx_plug_bp_wallet_( - self, product, create_main_accounts, thl_lm, lm, currency + self, + product: Product, + create_main_accounts: Callable[..., None], + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + currency: LedgerCurrency, ): """A BP Wallet "plug" is a way to fix discrepancies and simply add or remove money. @@ -320,7 +374,7 @@ class TestThlLedgerTxManager: """ rand_amount: USDCent = USDCent(randint(100, 1_000)) - tx = thl_lm.create_tx_plug_bp_wallet_( + tx = thl_ledger_manager.create_tx_plug_bp_wallet_( product=product, amount=rand_amount, created=datetime.now(tz=UTC), @@ -331,32 +385,32 @@ class TestThlLedgerTxManager: # We issued the BP money they didn't earn, so now they have a # negative balance - balance = thl_lm.get_account_balance( - account=thl_lm.get_account_or_create_bp_wallet(product=product) + balance = thl_ledger_manager.get_account_balance( + account=thl_ledger_manager.get_account_or_create_bp_wallet(product=product) ) assert balance == int(rand_amount) * -1 # Issue a positive one now, and confirm the balance goes positive - thl_lm.create_tx_plug_bp_wallet_( + thl_ledger_manager.create_tx_plug_bp_wallet_( product=product, amount=rand_amount + rand_amount, created=datetime.now(tz=UTC), direction=Direction.CREDIT, ) - balance = thl_lm.get_account_balance( - account=thl_lm.get_account_or_create_bp_wallet(product=product) + balance = thl_ledger_manager.get_account_balance( + account=thl_ledger_manager.get_account_or_create_bp_wallet(product=product) ) assert balance == int(rand_amount) def test_create_tx_user_payout_request( self, - user, - product_user_wallet_yes, - user_factory, - delete_df_collection, - thl_lm, - lm, - currency, + user: User, + product_user_wallet_yes: Product, + user_factory: Callable[..., User], + delete_df_collection: Callable[..., None], + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + currency: LedgerCurrency, ): pe = UserPayoutEvent( uuid=uuid4().hex, @@ -369,7 +423,7 @@ class TestThlLedgerTxManager: # The default user fixture uses a product that doesn't have wallet # mode enabled with pytest.raises(expected_exception=AssertionError): - thl_lm.create_tx_user_payout_request( + thl_ledger_manager.create_tx_user_payout_request( user=user, payout_event=pe, skip_flag_check=True, @@ -380,12 +434,12 @@ class TestThlLedgerTxManager: u2 = user_factory(product=product_user_wallet_yes) # User's pre-balance is 0 because no activity has occurred yet - pre_balance = lm.get_account_balance( - account=thl_lm.get_account_or_create_user_wallet(user=u2) + pre_balance = ledger_manager.get_account_balance( + account=thl_ledger_manager.get_account_or_create_user_wallet(user=u2) ) assert pre_balance == 0 - tx = thl_lm.create_tx_user_payout_request( + tx = thl_ledger_manager.create_tx_user_payout_request( user=u2, payout_event=pe, skip_flag_check=True, @@ -406,19 +460,19 @@ class TestThlLedgerTxManager: # Post balance is -$5.00 because it comes out of the wallet before # it's Approved or Completed - post_balance = lm.get_account_balance( - account=thl_lm.get_account_or_create_user_wallet(user=u2) + post_balance = ledger_manager.get_account_balance( + account=thl_ledger_manager.get_account_or_create_user_wallet(user=u2) ) assert post_balance == -500 def test_create_tx_user_payout_request_( self, - user, - product_user_wallet_yes, - user_factory, - delete_ledger_db, - thl_lm, - lm, + user: User, + product_user_wallet_yes: Product, + user_factory: Callable[..., User], + delete_ledger_db: Callable[..., None], + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, ): delete_ledger_db() @@ -431,36 +485,36 @@ class TestThlLedgerTxManager: ) rand_description = uuid4().hex - tx = thl_lm.create_tx_user_payout_request_( + tx = thl_ledger_manager.create_tx_user_payout_request_( user=user, payout_event=pe, description=rand_description ) assert tx.ext_description == rand_description - post_balance = lm.get_account_balance( - account=thl_lm.get_account_or_create_user_wallet(user=user) + post_balance = ledger_manager.get_account_balance( + account=thl_ledger_manager.get_account_or_create_user_wallet(user=user) ) assert post_balance == -500 def test_create_tx_user_payout_complete( self, - user_factory, - product_user_wallet_yes, - create_main_accounts, - delete_ledger_db, - thl_lm, - lm, - currency, + user_factory: Callable[..., User], + product_user_wallet_yes: Product, + create_main_accounts: Callable[..., None], + delete_ledger_db: Callable[..., None], + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + currency: LedgerCurrency, ): delete_ledger_db() create_main_accounts() user: User = user_factory(product=product_user_wallet_yes) - user_account = thl_lm.get_account_or_create_user_wallet(user=user) + user_account = thl_ledger_manager.get_account_or_create_user_wallet(user=user) rand_amount = randint(100, 1_000) # Ensure the user starts out with nothing... - assert lm.get_account_balance(account=user_account) == 0 + assert ledger_manager.get_account_balance(account=user_account) == 0 pe = UserPayoutEvent( uuid=uuid4().hex, @@ -472,7 +526,7 @@ class TestThlLedgerTxManager: # Confirm it's not possible unless a request occurred happen with pytest.raises(expected_exception=ValueError): - thl_lm.create_tx_user_payout_complete( + thl_ledger_manager.create_tx_user_payout_complete( user=user, payout_event=pe, fee_amount=None, @@ -480,17 +534,19 @@ class TestThlLedgerTxManager: ) # (1) Make a request first - thl_lm.create_tx_user_payout_request( + thl_ledger_manager.create_tx_user_payout_request( user=user, payout_event=pe, skip_flag_check=True, skip_wallet_balance_check=True, ) # Assert the balance came out of their user wallet - assert lm.get_account_balance(account=user_account) == rand_amount * -1 + assert ( + ledger_manager.get_account_balance(account=user_account) == rand_amount * -1 + ) # (2) Complete the request - tx = thl_lm.create_tx_user_payout_complete( + tx = thl_ledger_manager.create_tx_user_payout_complete( user=user, payout_event=pe, fee_amount=Decimal(0), @@ -503,18 +559,20 @@ class TestThlLedgerTxManager: # The amount that comes out of the user wallet doesn't change after # it's approved becuase it's already been withdrawn - assert lm.get_account_balance(account=user_account) == rand_amount * -1 + assert ( + ledger_manager.get_account_balance(account=user_account) == rand_amount * -1 + ) def test_create_tx_user_payout_complete_( self, - user_factory, - product_user_wallet_yes, - create_main_accounts, - thl_lm, - lm, + user_factory: Callable[..., User], + product_user_wallet_yes: Product, + create_main_accounts: Callable[..., None], + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, ): user: User = user_factory(product=product_user_wallet_yes) - user_account = thl_lm.get_account_or_create_user_wallet(user=user) + user_account = thl_ledger_manager.get_account_or_create_user_wallet(user=user) rand_amount = randint(100, 1_000) pe = UserPayoutEvent( @@ -526,7 +584,7 @@ class TestThlLedgerTxManager: ) # (1) Make a request first - thl_lm.create_tx_user_payout_request( + thl_ledger_manager.create_tx_user_payout_request( user=user, payout_event=pe, skip_flag_check=True, @@ -536,12 +594,14 @@ class TestThlLedgerTxManager: # (2) Complete the request rand_desc = uuid4().hex - bp_expense_account = thl_lm.get_account_or_create_bp_expense( + bp_expense_account = thl_ledger_manager.get_account_or_create_bp_expense( product=user.product, expense_name="paypal" ) - bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(product=user.product) + bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet( + product=user.product + ) - tx = thl_lm.create_tx_user_payout_complete_( + tx = thl_ledger_manager.create_tx_user_payout_complete_( user=user, payout_event=pe, fee_amount=Decimal("0.00"), @@ -550,19 +610,21 @@ class TestThlLedgerTxManager: description=rand_desc, ) assert tx.ext_description == rand_desc - assert lm.get_account_balance(account=user_account) == rand_amount * -1 + assert ( + ledger_manager.get_account_balance(account=user_account) == rand_amount * -1 + ) def test_create_tx_user_payout_cancelled( self, - user_factory, - product_user_wallet_yes, - create_main_accounts, - thl_lm, - lm, - currency, + user_factory: Callable[..., User], + product_user_wallet_yes: Product, + create_main_accounts: Callable[..., None], + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + currency: LedgerCurrency, ): user: User = user_factory(product=product_user_wallet_yes) - user_account = thl_lm.get_account_or_create_user_wallet(user=user) + user_account = thl_ledger_manager.get_account_or_create_user_wallet(user=user) rand_amount = randint(100, 1_000) pe = UserPayoutEvent( @@ -574,17 +636,19 @@ class TestThlLedgerTxManager: ) # (1) Make a request first - thl_lm.create_tx_user_payout_request( + thl_ledger_manager.create_tx_user_payout_request( user=user, payout_event=pe, skip_flag_check=True, skip_wallet_balance_check=True, ) # Assert the balance came out of their user wallet - assert lm.get_account_balance(account=user_account) == rand_amount * -1 + assert ( + ledger_manager.get_account_balance(account=user_account) == rand_amount * -1 + ) # (2) Cancel the request - tx = thl_lm.create_tx_user_payout_cancelled( + tx = thl_ledger_manager.create_tx_user_payout_cancelled( user=user, payout_event=pe, skip_flag_check=False, @@ -595,19 +659,19 @@ class TestThlLedgerTxManager: assert isinstance(tx, LedgerTransaction) # Assert the balance comes back to 0 after it was cancelled - assert lm.get_account_balance(account=user_account) == 0 + assert ledger_manager.get_account_balance(account=user_account) == 0 def test_create_tx_user_payout_cancelled_( self, - user_factory, - product_user_wallet_yes, - create_main_accounts, - thl_lm, - lm, - currency, + user_factory: Callable[..., User], + product_user_wallet_yes: Product, + create_main_accounts: Callable[..., None], + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + currency: LedgerCurrency, ): user: User = user_factory(product=product_user_wallet_yes) - user_account = thl_lm.get_account_or_create_user_wallet(user=user) + user_account = thl_ledger_manager.get_account_or_create_user_wallet(user=user) rand_amount = randint(100, 1_000) pe = UserPayoutEvent( @@ -619,43 +683,45 @@ class TestThlLedgerTxManager: ) # (1) Make a request first - thl_lm.create_tx_user_payout_request( + thl_ledger_manager.create_tx_user_payout_request( user=user, payout_event=pe, skip_flag_check=True, skip_wallet_balance_check=True, ) # Assert the balance came out of their user wallet - assert lm.get_account_balance(account=user_account) == rand_amount * -1 + assert ( + ledger_manager.get_account_balance(account=user_account) == rand_amount * -1 + ) # (2) Cancel the request rand_desc = uuid4().hex - tx = thl_lm.create_tx_user_payout_cancelled_( + tx = thl_ledger_manager.create_tx_user_payout_cancelled_( user=user, payout_event=pe, description=rand_desc ) assert isinstance(tx, LedgerTransaction) assert tx.ext_description == rand_desc - assert lm.get_account_balance(account=user_account) == 0 + assert ledger_manager.get_account_balance(account=user_account) == 0 def test_create_tx_user_bonus( self, - user_factory, - product_user_wallet_yes, - create_main_accounts, - thl_lm, - lm, - currency, + user_factory: Callable[..., User], + product_user_wallet_yes: Product, + create_main_accounts: Callable[..., None], + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + currency: LedgerCurrency, ): user: User = user_factory(product=product_user_wallet_yes) - user_account = thl_lm.get_account_or_create_user_wallet(user=user) + user_account = thl_ledger_manager.get_account_or_create_user_wallet(user=user) rand_amount = randint(100, 1_000) rand_ref_uuid = uuid4().hex rand_desc = uuid4().hex # Assert the balance came out of their user wallet - assert lm.get_account_balance(account=user_account) == 0 + assert ledger_manager.get_account_balance(account=user_account) == 0 - tx = thl_lm.create_tx_user_bonus( + tx = thl_ledger_manager.create_tx_user_bonus( user=user, amount=Decimal(rand_amount / 100), ref_uuid=rand_ref_uuid, @@ -663,44 +729,48 @@ class TestThlLedgerTxManager: skip_flag_check=True, ) assert tx.ext_description == rand_desc - assert tx.tag == f"{thl_lm.currency.value}:user_bonus:{rand_ref_uuid}" + assert ( + tx.tag == f"{thl_ledger_manager.currency.value}:user_bonus:{rand_ref_uuid}" + ) assert tx.entries[0].amount == rand_amount assert tx.entries[1].amount == rand_amount # Assert the balance came out of their user wallet - assert lm.get_account_balance(account=user_account) == rand_amount + assert ledger_manager.get_account_balance(account=user_account) == rand_amount def test_create_tx_user_bonus_( self, - user_factory, - product_user_wallet_yes, - create_main_accounts, - thl_lm, - lm, - currency, + user_factory: Callable[..., User], + product_user_wallet_yes: Product, + create_main_accounts: Callable[..., None], + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + currency: LedgerCurrency, ): user: User = user_factory(product=product_user_wallet_yes) - user_account = thl_lm.get_account_or_create_user_wallet(user=user) + user_account = thl_ledger_manager.get_account_or_create_user_wallet(user=user) rand_amount = randint(100, 1_000) rand_ref_uuid = uuid4().hex rand_desc = uuid4().hex # Assert the balance came out of their user wallet - assert lm.get_account_balance(account=user_account) == 0 + assert ledger_manager.get_account_balance(account=user_account) == 0 - tx = thl_lm.create_tx_user_bonus_( + tx = thl_ledger_manager.create_tx_user_bonus_( user=user, amount=Decimal(rand_amount / 100), ref_uuid=rand_ref_uuid, description=rand_desc, ) assert tx.ext_description == rand_desc - assert tx.tag == f"{thl_lm.currency.value}:user_bonus:{rand_ref_uuid}" + assert ( + tx.tag == f"{thl_ledger_manager.currency.value}:user_bonus:{rand_ref_uuid}" + ) assert tx.entries[0].amount == rand_amount assert tx.entries[1].amount == rand_amount # Assert the balance came out of their user wallet - assert lm.get_account_balance(account=user_account) == rand_amount + assert ledger_manager.get_account_balance(account=user_account) == rand_amount class TestThlLedgerTxManagerFlows: @@ -709,7 +779,13 @@ class TestThlLedgerTxManagerFlows: """ def test_create_tx_task_complete( - self, user, create_main_accounts, thl_lm, lm, currency, delete_ledger_db + self, + user: User, + create_main_accounts: Callable[..., None], + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + currency: LedgerCurrency, + delete_ledger_db: Callable[..., None], ): delete_ledger_db() create_main_accounts() @@ -725,7 +801,9 @@ class TestThlLedgerTxManagerFlows: started=datetime.now(UTC), finished=datetime.now(UTC) + timedelta(seconds=1), ) - thl_lm.create_tx_task_complete(wall=wall1, user=user, created=wall1.started) + thl_ledger_manager.create_tx_task_complete( + wall=wall1, user=user, created=wall1.started + ) wall2 = Wall( user_id=1, @@ -738,38 +816,40 @@ class TestThlLedgerTxManagerFlows: started=datetime.now(UTC), finished=datetime.now(UTC) + timedelta(seconds=1), ) - thl_lm.create_tx_task_complete(wall=wall2, user=user, created=wall2.started) + thl_ledger_manager.create_tx_task_complete( + wall=wall2, user=user, created=wall2.started + ) - cash = thl_lm.get_account_cash() - revenue = thl_lm.get_account_task_complete_revenue() + cash = thl_ledger_manager.get_account_cash() + revenue = thl_ledger_manager.get_account_task_complete_revenue() - assert lm.get_account_balance(cash) == 123 + 321 - assert lm.get_account_balance(revenue) == 123 + 321 - assert lm.check_ledger_balanced() + assert ledger_manager.get_account_balance(cash) == 123 + 321 + assert ledger_manager.get_account_balance(revenue) == 123 + 321 + assert ledger_manager.check_ledger_balanced() assert ( - lm.get_account_filtered_balance( + ledger_manager.get_account_filtered_balance( account=revenue, metadata_key="source", metadata_value="d" ) == 123 ) assert ( - lm.get_account_filtered_balance( + ledger_manager.get_account_filtered_balance( account=revenue, metadata_key="source", metadata_value="f" ) == 321 ) assert ( - lm.get_account_filtered_balance( + ledger_manager.get_account_filtered_balance( account=revenue, metadata_key="source", metadata_value="x" ) == 0 ) assert ( - thl_lm.get_account_filtered_balance( + thl_ledger_manager.get_account_filtered_balance( account=revenue, metadata_key="thl_wall", metadata_value=wall1.uuid, @@ -778,7 +858,12 @@ class TestThlLedgerTxManagerFlows: ) def test_create_transaction_task_complete_1_cent( - self, user, create_main_accounts, thl_lm, lm, currency + self, + user: User, + create_main_accounts: Callable[..., None], + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + currency: LedgerCurrency, ): wall1 = Wall( user_id=1, @@ -791,7 +876,7 @@ class TestThlLedgerTxManagerFlows: started=datetime.now(UTC), finished=datetime.now(UTC) + timedelta(seconds=1), ) - tx = thl_lm.create_tx_task_complete( + tx = thl_ledger_manager.create_tx_task_complete( wall=wall1, user=user, created=wall1.started ) @@ -799,14 +884,14 @@ class TestThlLedgerTxManagerFlows: def test_create_transaction_bp_payment( self, - user, - create_main_accounts, - thl_lm, - lm, - currency, - delete_ledger_db, - session_factory, - utc_hour_ago, + user: User, + create_main_accounts: Callable[..., None], + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + currency: LedgerCurrency, + delete_ledger_db: Callable[..., None], + session_factory: Callable[..., Session], + utc_hour_ago: datetime, ): delete_ledger_db() create_main_accounts() @@ -819,7 +904,9 @@ class TestThlLedgerTxManagerFlows: ) w1: Wall = s1.wall_events[0] - tx = thl_lm.create_tx_task_complete(wall=w1, user=user, created=w1.started) + tx = thl_ledger_manager.create_tx_task_complete( + wall=w1, user=user, created=w1.started + ) assert isinstance(tx, LedgerTransaction) status, status_code_1 = s1.determine_session_status() @@ -832,35 +919,39 @@ class TestThlLedgerTxManagerFlows: user_payout=user_pay, ) print(thl_net, commission_amount, bp_pay, user_pay) - thl_lm.create_tx_bp_payment(session=s1, created=w1.started) + thl_ledger_manager.create_tx_bp_payment(session=s1, created=w1.started) - revenue = thl_lm.get_account_task_complete_revenue() - bp_wallet = thl_lm.get_account_or_create_bp_wallet(product=user.product) - bp_commission = thl_lm.get_account_or_create_bp_commission(product=user.product) + revenue = thl_ledger_manager.get_account_task_complete_revenue() + bp_wallet = thl_ledger_manager.get_account_or_create_bp_wallet( + product=user.product + ) + bp_commission = thl_ledger_manager.get_account_or_create_bp_commission( + product=user.product + ) - assert 0 == lm.get_account_balance(account=revenue) - assert 50 == lm.get_account_filtered_balance( + assert 0 == ledger_manager.get_account_balance(account=revenue) + assert 50 == ledger_manager.get_account_filtered_balance( account=revenue, metadata_key="source", metadata_value=Source.TESTING, ) - assert 48 == lm.get_account_balance(account=bp_wallet) - assert 48 == lm.get_account_filtered_balance( + assert 48 == ledger_manager.get_account_balance(account=bp_wallet) + assert 48 == ledger_manager.get_account_filtered_balance( account=bp_wallet, metadata_key="thl_session", metadata_value=s1.uuid, ) - assert 2 == thl_lm.get_account_balance(account=bp_commission) - assert thl_lm.check_ledger_balanced() + assert 2 == thl_ledger_manager.get_account_balance(account=bp_commission) + assert thl_ledger_manager.check_ledger_balanced() def test_create_transaction_bp_payment_round( self, - user_factory, - product_user_wallet_no, - create_main_accounts, - thl_lm, - lm, - currency, + user_factory: Callable[..., User], + product_user_wallet_no: Product, + create_main_accounts: Callable[..., None], + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + currency: LedgerCurrency, ): product_user_wallet_no.commission_pct = Decimal("0.085") user: User = user_factory(product=product_user_wallet_no) @@ -877,7 +968,7 @@ class TestThlLedgerTxManagerFlows: finished=datetime.now(UTC) + timedelta(seconds=1), ) - tx = thl_lm.create_tx_task_complete( + tx = thl_ledger_manager.create_tx_task_complete( wall=wall1, user=user, created=wall1.started ) assert isinstance(tx, LedgerTransaction) @@ -894,11 +985,19 @@ class TestThlLedgerTxManagerFlows: ) print(thl_net, commission_amount, bp_pay, user_pay) - tx = thl_lm.create_tx_bp_payment(session=session, created=wall1.started) + tx = thl_ledger_manager.create_tx_bp_payment( + session=session, created=wall1.started + ) assert isinstance(tx, LedgerTransaction) def test_create_transaction_bp_payment_round2( - self, delete_ledger_db, user, create_main_accounts, thl_lm, lm, currency + self, + delete_ledger_db: Callable[..., None], + user: User, + create_main_accounts: Callable[..., None], + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + currency: LedgerCurrency, ): delete_ledger_db() create_main_accounts() @@ -917,7 +1016,9 @@ class TestThlLedgerTxManagerFlows: finished=datetime.now(UTC) + timedelta(seconds=1), ) - thl_lm.create_tx_task_complete(wall=wall1, user=user, created=wall1.started) + thl_ledger_manager.create_tx_task_complete( + wall=wall1, user=user, created=wall1.started + ) session = Session(started=wall1.started, user=user, wall_events=[wall1]) status, status_code_1 = session.determine_session_status() # thl_net, commission_amount, bp_pay, user_pay = session.determine_payments() @@ -929,16 +1030,16 @@ class TestThlLedgerTxManagerFlows: user_payout=Decimal("1.53"), ) - thl_lm.create_tx_bp_payment(session=session, created=wall1.started) + thl_ledger_manager.create_tx_bp_payment(session=session, created=wall1.started) def test_create_transaction_bp_payment_round3( self, - user_factory, - product_user_wallet_yes, - create_main_accounts, - thl_lm, - lm, - currency, + user_factory: Callable[..., User], + product_user_wallet_yes: Product, + create_main_accounts: Callable[..., None], + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + currency: LedgerCurrency, ): # e.g. session ___ fails b/c we rounded incorrectly # before, and now we are off by a penny... @@ -955,7 +1056,9 @@ class TestThlLedgerTxManagerFlows: started=datetime.now(UTC), finished=datetime.now(UTC) + timedelta(seconds=1), ) - thl_lm.create_tx_task_complete(wall=wall1, user=user, created=wall1.started) + thl_ledger_manager.create_tx_task_complete( + wall=wall1, user=user, created=wall1.started + ) session = Session(started=wall1.started, user=user, wall_events=[wall1]) status, status_code_1 = session.determine_session_status() @@ -973,22 +1076,23 @@ class TestThlLedgerTxManagerFlows: def test_create_transaction_bp_payment_user_wallet( self, - user_factory, - product_user_wallet_yes, - create_main_accounts, - delete_ledger_db, - thl_lm, - session_manager, - wall_manager, - lm, - session_factory, - currency, - utc_hour_ago, + user_factory: Callable[..., User], + product_user_wallet_yes: Product, + create_main_accounts: Callable[..., None], + delete_ledger_db: Callable[..., None], + thl_ledger_manager: ThlLedgerManager, + session_manager: SessionManager, + wall_manager: WallManager, + ledger_manager: LedgerManager, + session_factory: Callable[..., Session], + currency: LedgerCurrency, + utc_hour_ago: datetime, ): delete_ledger_db() create_main_accounts() user: User = user_factory(product=product_user_wallet_yes) + assert isinstance(user.product, Product) assert user.product.user_wallet_enabled s1: Session = session_factory( @@ -1000,10 +1104,12 @@ class TestThlLedgerTxManagerFlows: ) w1: Wall = s1.wall_events[0] - thl_lm.create_tx_task_complete(wall=w1, user=user, created=w1.started) + thl_ledger_manager.create_tx_task_complete( + wall=w1, user=user, created=w1.started + ) status, status_code_1 = s1.determine_session_status() - thl_net, commission_amount, bp_pay, user_pay = s1.determine_payments() + _, _, bp_pay, user_pay = s1.determine_payments() session_manager.finish_with_status( session=s1, status=status, @@ -1012,52 +1118,56 @@ class TestThlLedgerTxManagerFlows: payout=bp_pay, user_payout=user_pay, ) - thl_lm.create_tx_bp_payment(session=s1, created=w1.started) + thl_ledger_manager.create_tx_bp_payment(session=s1, created=w1.started) - revenue = thl_lm.get_account_task_complete_revenue() - bp_wallet = thl_lm.get_account_or_create_bp_wallet(product=user.product) - bp_commission = thl_lm.get_account_or_create_bp_commission(product=user.product) - user_wallet = thl_lm.get_account_or_create_user_wallet(user=user) + revenue = thl_ledger_manager.get_account_task_complete_revenue() + bp_wallet = thl_ledger_manager.get_account_or_create_bp_wallet( + product=user.product + ) + bp_commission = thl_ledger_manager.get_account_or_create_bp_commission( + product=user.product + ) + user_wallet = thl_ledger_manager.get_account_or_create_user_wallet(user=user) - assert 0 == thl_lm.get_account_balance(account=revenue) - assert 50 == thl_lm.get_account_filtered_balance( + assert 0 == thl_ledger_manager.get_account_balance(account=revenue) + assert 50 == thl_ledger_manager.get_account_filtered_balance( account=revenue, metadata_key="source", metadata_value=Source.TESTING, ) - assert 48 - 19 == thl_lm.get_account_balance(account=bp_wallet) - assert 48 - 19 == thl_lm.get_account_filtered_balance( + assert 48 - 19 == thl_ledger_manager.get_account_balance(account=bp_wallet) + assert 48 - 19 == thl_ledger_manager.get_account_filtered_balance( account=bp_wallet, metadata_key="thl_session", metadata_value=s1.uuid, ) - assert 2 == thl_lm.get_account_balance(bp_commission) - assert 19 == thl_lm.get_account_balance(user_wallet) - assert 19 == thl_lm.get_account_filtered_balance( + assert 2 == thl_ledger_manager.get_account_balance(bp_commission) + assert 19 == thl_ledger_manager.get_account_balance(user_wallet) + assert 19 == thl_ledger_manager.get_account_filtered_balance( account=user_wallet, metadata_key="thl_session", metadata_value=s1.uuid, ) - assert 0 == thl_lm.get_account_filtered_balance( + assert 0 == thl_ledger_manager.get_account_filtered_balance( account=user_wallet, metadata_key="thl_session", metadata_value="x" ) - assert thl_lm.check_ledger_balanced() + assert thl_ledger_manager.check_ledger_balanced() class TestThlLedgerManagerAdj: def test_create_tx_task_adjustment( self, - user_factory, - product_user_wallet_no, - create_main_accounts, - delete_ledger_db, - thl_lm, - lm, - utc_hour_ago, - currency, + user_factory: Callable[..., User], + product_user_wallet_no: Product, + create_main_accounts: Callable[..., None], + delete_ledger_db: Callable[..., None], + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + utc_hour_ago: datetime, + currency: LedgerCurrency, ): delete_ledger_db() create_main_accounts() @@ -1076,7 +1186,7 @@ class TestThlLedgerManagerAdj: finished=utc_hour_ago + timedelta(seconds=1), ) - thl_lm.create_tx_task_complete(wall1, user, created=wall1.started) + thl_ledger_manager.create_tx_task_complete(wall1, user, created=wall1.started) wall2 = Wall( user_id=1, @@ -1089,7 +1199,7 @@ class TestThlLedgerManagerAdj: started=utc_hour_ago, finished=utc_hour_ago + timedelta(seconds=1), ) - thl_lm.create_tx_task_complete(wall2, user, created=wall2.started) + thl_ledger_manager.create_tx_task_complete(wall2, user, created=wall2.started) wall1.update( adjusted_status=WallAdjustedStatus.ADJUSTED_TO_FAIL, @@ -1097,24 +1207,26 @@ class TestThlLedgerManagerAdj: adjusted_timestamp=utc_hour_ago + timedelta(hours=1), ) print(wall1.get_cpi_after_adjustment()) - thl_lm.create_tx_task_adjustment(wall1, user) + thl_ledger_manager.create_tx_task_adjustment(wall1, user) - cash = thl_lm.get_account_cash() - revenue = thl_lm.get_account_task_complete_revenue() + cash = thl_ledger_manager.get_account_cash() + revenue = thl_ledger_manager.get_account_task_complete_revenue() - assert 123 + 321 - 123 == thl_lm.get_account_balance(account=cash) - assert 123 + 321 - 123 == thl_lm.get_account_balance(account=revenue) - assert thl_lm.check_ledger_balanced() - assert 0 == thl_lm.get_account_filtered_balance( + assert 123 + 321 - 123 == thl_ledger_manager.get_account_balance(account=cash) + assert 123 + 321 - 123 == thl_ledger_manager.get_account_balance( + account=revenue + ) + assert thl_ledger_manager.check_ledger_balanced() + assert 0 == thl_ledger_manager.get_account_filtered_balance( revenue, metadata_key="source", metadata_value="d" ) - assert 321 == thl_lm.get_account_filtered_balance( + assert 321 == thl_ledger_manager.get_account_filtered_balance( revenue, metadata_key="source", metadata_value="f" ) - assert 0 == thl_lm.get_account_filtered_balance( + assert 0 == thl_ledger_manager.get_account_filtered_balance( revenue, metadata_key="source", metadata_value="x" ) - assert 123 - 123 == thl_lm.get_account_filtered_balance( + assert 123 - 123 == thl_ledger_manager.get_account_filtered_balance( account=revenue, metadata_key="thl_wall", metadata_value=wall1.uuid ) @@ -1125,43 +1237,43 @@ class TestThlLedgerManagerAdj: adjusted_timestamp=utc_hour_ago + timedelta(minutes=45), ) print(wall1.get_cpi_after_adjustment()) - thl_lm.create_tx_task_adjustment(wall1, user) + thl_ledger_manager.create_tx_task_adjustment(wall1, user) # and then run it again to make sure it does nothing - thl_lm.create_tx_task_adjustment(wall1, user) + thl_ledger_manager.create_tx_task_adjustment(wall1, user) - cash = thl_lm.get_account_cash() - revenue = thl_lm.get_account_task_complete_revenue() + cash = thl_ledger_manager.get_account_cash() + revenue = thl_ledger_manager.get_account_task_complete_revenue() - assert 123 + 321 - 123 + 123 == thl_lm.get_account_balance(cash) - assert 123 + 321 - 123 + 123 == thl_lm.get_account_balance(revenue) - assert thl_lm.check_ledger_balanced() - assert 123 == thl_lm.get_account_filtered_balance( + assert 123 + 321 - 123 + 123 == thl_ledger_manager.get_account_balance(cash) + assert 123 + 321 - 123 + 123 == thl_ledger_manager.get_account_balance(revenue) + assert thl_ledger_manager.check_ledger_balanced() + assert 123 == thl_ledger_manager.get_account_filtered_balance( account=revenue, metadata_key="source", metadata_value="d" ) - assert 321 == thl_lm.get_account_filtered_balance( + assert 321 == thl_ledger_manager.get_account_filtered_balance( account=revenue, metadata_key="source", metadata_value="f" ) - assert 0 == thl_lm.get_account_filtered_balance( + assert 0 == thl_ledger_manager.get_account_filtered_balance( account=revenue, metadata_key="source", metadata_value="x" ) - assert 123 - 123 + 123 == thl_lm.get_account_filtered_balance( + assert 123 - 123 + 123 == thl_ledger_manager.get_account_filtered_balance( account=revenue, metadata_key="thl_wall", metadata_value=wall1.uuid ) def test_create_tx_bp_adjustment( self, - user, - product_user_wallet_no, - create_main_accounts, + user: User, + product_user_wallet_no: Product, + create_main_accounts: Callable[..., None], caplog, - thl_lm, - lm, - currency, - session_manager, - wall_manager, - session_factory, - utc_hour_ago, - delete_ledger_db, + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + currency: LedgerCurrency, + session_manager: SessionManager, + wall_manager: WallManager, + session_factory: Callable[..., Session], + utc_hour_ago: datetime, + delete_ledger_db: Callable[..., None], ): delete_ledger_db() create_main_accounts() @@ -1177,8 +1289,12 @@ class TestThlLedgerManagerAdj: w1: Wall = s1.wall_events[0] w2: Wall = s1.wall_events[1] - thl_lm.create_tx_task_complete(wall=w1, user=user, created=w1.started) - thl_lm.create_tx_task_complete(wall=w2, user=user, created=w2.started) + thl_ledger_manager.create_tx_task_complete( + wall=w1, user=user, created=w1.started + ) + thl_ledger_manager.create_tx_task_complete( + wall=w2, user=user, created=w2.started + ) status, status_code_1 = s1.determine_session_status() _, _, bp_pay, user_pay = s1.determine_payments() @@ -1190,21 +1306,25 @@ class TestThlLedgerManagerAdj: payout=bp_pay, user_payout=user_pay, ) - thl_lm.create_tx_bp_payment(session=s1, created=w1.started) - revenue = thl_lm.get_account_task_complete_revenue() - bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(product=user.product) - bp_commission_account = thl_lm.get_account_or_create_bp_commission( + thl_ledger_manager.create_tx_bp_payment(session=s1, created=w1.started) + revenue = thl_ledger_manager.get_account_task_complete_revenue() + bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet( product=user.product ) - assert 380 == thl_lm.get_account_balance(account=bp_wallet_account) - assert 0 == thl_lm.get_account_balance(account=revenue) - assert 20 == thl_lm.get_account_balance(account=bp_commission_account) - thl_lm.check_ledger_balanced() + bp_commission_account = thl_ledger_manager.get_account_or_create_bp_commission( + product=user.product + ) + assert 380 == thl_ledger_manager.get_account_balance(account=bp_wallet_account) + assert 0 == thl_ledger_manager.get_account_balance(account=revenue) + assert 20 == thl_ledger_manager.get_account_balance( + account=bp_commission_account + ) + thl_ledger_manager.check_ledger_balanced() # This should do nothing (since we haven't adjusted any wall events) s1.adjust_status() with caplog.at_level(logging.INFO): - thl_lm.create_tx_bp_adjustment(session=s1) + thl_ledger_manager.create_tx_bp_adjustment(session=s1) assert ( "create_transaction_bp_adjustment. No transactions needed." in caplog.text @@ -1222,22 +1342,22 @@ class TestThlLedgerManagerAdj: adjusted_cpi=Decimal(0), adjusted_timestamp=utc_hour_ago + timedelta(hours=1), ) - thl_lm.create_tx_task_adjustment(wall=w1, user=user) + thl_ledger_manager.create_tx_task_adjustment(wall=w1, user=user) # -$1.00 b/c the MP took the $1 back, but we haven't yet taken the BP payment back - assert -100 == thl_lm.get_account_balance(revenue) + assert -100 == thl_ledger_manager.get_account_balance(revenue) s1.adjust_status() - thl_lm.create_tx_bp_adjustment(session=s1) + thl_ledger_manager.create_tx_bp_adjustment(session=s1) with caplog.at_level(logging.INFO): - thl_lm.create_tx_bp_adjustment(session=s1) + thl_ledger_manager.create_tx_bp_adjustment(session=s1) assert ( "create_transaction_bp_adjustment. No transactions needed." in caplog.text ) - assert 380 - 95 == thl_lm.get_account_balance(bp_wallet_account) - assert 0 == thl_lm.get_account_balance(revenue) - assert 20 - 5 == thl_lm.get_account_balance(bp_commission_account) - assert thl_lm.check_ledger_balanced() + assert 380 - 95 == thl_ledger_manager.get_account_balance(bp_wallet_account) + assert 0 == thl_ledger_manager.get_account_balance(revenue) + assert 20 - 5 == thl_ledger_manager.get_account_balance(bp_commission_account) + assert thl_ledger_manager.check_ledger_balanced() # unrecon the $1 survey wall_manager.adjust_status( @@ -1246,29 +1366,29 @@ class TestThlLedgerManagerAdj: adjusted_cpi=None, adjusted_timestamp=utc_hour_ago + timedelta(minutes=45), ) - thl_lm.create_tx_task_adjustment( + thl_ledger_manager.create_tx_task_adjustment( wall=w1, user=user, created=utc_hour_ago + timedelta(minutes=45), ) _, _, _ = s1.determine_new_status_and_payouts() s1.adjust_status() - thl_lm.create_tx_bp_adjustment(session=s1) - assert 380 == thl_lm.get_account_balance(bp_wallet_account) - assert 0 == thl_lm.get_account_balance(revenue) - assert 20, thl_lm.get_account_balance(bp_commission_account) - assert thl_lm.check_ledger_balanced() + thl_ledger_manager.create_tx_bp_adjustment(session=s1) + assert 380 == thl_ledger_manager.get_account_balance(bp_wallet_account) + assert 0 == thl_ledger_manager.get_account_balance(revenue) + assert 20, thl_ledger_manager.get_account_balance(bp_commission_account) + assert thl_ledger_manager.check_ledger_balanced() def test_create_tx_bp_adjustment_small( self, user_factory: Callable[..., User], - product_user_wallet_no, - create_main_accounts, - delete_ledger_db, - thl_ledger_manager, - ledger_manager, + product_user_wallet_no: Product, + create_main_accounts: Callable[..., None], + delete_ledger_db: Callable[..., None], + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, utc_hour_ago: datetime, - currency, + currency: LedgerCurrency, ): delete_ledger_db() create_main_accounts() @@ -1289,7 +1409,7 @@ class TestThlLedgerManagerAdj: finished=utc_hour_ago + timedelta(seconds=1), ) - tx = thl_lm.create_tx_task_complete( + tx = thl_ledger_manager.create_tx_task_complete( wall=wall1, user=user, created=wall1.started ) assert isinstance(tx, LedgerTransaction) @@ -1304,31 +1424,31 @@ class TestThlLedgerManagerAdj: payout=bp_pay, user_payout=user_pay, ) - thl_lm.create_tx_bp_payment(session, created=wall1.started) + thl_ledger_manager.create_tx_bp_payment(session, created=wall1.started) wall1.update( adjusted_status=WallAdjustedStatus.ADJUSTED_TO_FAIL, adjusted_cpi=0, adjusted_timestamp=utc_hour_ago + timedelta(hours=1), ) - thl_lm.create_tx_task_adjustment(wall1, user) + thl_ledger_manager.create_tx_task_adjustment(wall1, user) session.adjust_status() - thl_lm.create_tx_bp_adjustment(session) + thl_ledger_manager.create_tx_bp_adjustment(session) def test_create_tx_bp_adjustment_abandon( self, - user_factory, - product_user_wallet_no, - delete_ledger_db, - session_factory, - create_main_accounts, + user_factory: Callable[..., User], + product_user_wallet_no: Product, + delete_ledger_db: Callable[..., None], + session_factory: Callable[..., Session], + create_main_accounts: Callable[..., None], caplog, - thl_lm, - lm, - currency, - utc_hour_ago, - session_manager, - wall_manager, + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + currency: LedgerCurrency, + utc_hour_ago: datetime, + session_manager: SessionManager, + wall_manager: WallManager, ): delete_ledger_db() create_main_accounts() @@ -1345,9 +1465,9 @@ class TestThlLedgerManagerAdj: adjusted_cpi=w1.cpi, adjusted_timestamp=utc_hour_ago + timedelta(hours=1), ) - thl_lm.create_tx_task_adjustment(wall=w1, user=user) + thl_ledger_manager.create_tx_task_adjustment(wall=w1, user=user) s1.adjust_status() - thl_lm.create_tx_bp_adjustment(session=s1) + thl_ledger_manager.create_tx_bp_adjustment(session=s1) # And then adjust it back (it was abandon before, but now it should be # fail (?) or back to abandon?) wall_manager.adjust_status( @@ -1356,24 +1476,26 @@ class TestThlLedgerManagerAdj: adjusted_cpi=None, adjusted_timestamp=utc_hour_ago + timedelta(hours=1), ) - thl_lm.create_tx_task_adjustment(wall=w1, user=user) + thl_ledger_manager.create_tx_task_adjustment(wall=w1, user=user) s1.adjust_status() - thl_lm.create_tx_bp_adjustment(session=s1) + thl_ledger_manager.create_tx_bp_adjustment(session=s1) - revenue = thl_lm.get_account_task_complete_revenue() - bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(product=user.product) - bp_commission_account = thl_lm.get_account_or_create_bp_commission( + revenue = thl_ledger_manager.get_account_task_complete_revenue() + bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet( + product=user.product + ) + bp_commission_account = thl_ledger_manager.get_account_or_create_bp_commission( product=user.product ) - assert 0 == thl_lm.get_account_balance(bp_wallet_account) - assert 0 == thl_lm.get_account_balance(revenue) - assert 0 == thl_lm.get_account_balance(bp_commission_account) - assert thl_lm.check_ledger_balanced() + assert 0 == thl_ledger_manager.get_account_balance(bp_wallet_account) + assert 0 == thl_ledger_manager.get_account_balance(revenue) + assert 0 == thl_ledger_manager.get_account_balance(bp_commission_account) + assert thl_ledger_manager.check_ledger_balanced() # This should do nothing s1.adjust_status() with caplog.at_level(logging.INFO): - thl_lm.create_tx_bp_adjustment(session=s1) + thl_ledger_manager.create_tx_bp_adjustment(session=s1) assert "No transactions needed" in caplog.text # Now back to complete again @@ -1384,19 +1506,19 @@ class TestThlLedgerManagerAdj: adjusted_timestamp=utc_hour_ago + timedelta(hours=1), ) s1.adjust_status() - thl_lm.create_tx_bp_adjustment(session=s1) - assert 95 == thl_lm.get_account_balance(bp_wallet_account) + thl_ledger_manager.create_tx_bp_adjustment(session=s1) + assert 95 == thl_ledger_manager.get_account_balance(bp_wallet_account) def test_create_tx_bp_adjustment_user_wallet( self, - user_factory, - product_user_wallet_yes, - create_main_accounts, - delete_ledger_db, + user_factory: Callable[..., User], + product_user_wallet_yes: Product, + create_main_accounts: Callable[..., None], + delete_ledger_db: Callable[..., None], caplog, - thl_lm, - lm, - currency, + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + currency: LedgerCurrency, ): delete_ledger_db() create_main_accounts() @@ -1432,7 +1554,7 @@ class TestThlLedgerManagerAdj: started=now_w1, finished=now_w1 + timedelta(minutes=1), ) - tx = thl_lm.create_tx_task_complete( + tx = thl_ledger_manager.create_tx_task_complete( wall=wall1, user=user, created=wall1.started ) assert isinstance(tx, LedgerTransaction) @@ -1449,7 +1571,7 @@ class TestThlLedgerManagerAdj: started=now_w2, finished=now_w2 + timedelta(minutes=1), ) - tx = thl_lm.create_tx_task_complete( + tx = thl_ledger_manager.create_tx_task_complete( wall=wall2, user=user, created=wall2.started ) assert isinstance(tx, LedgerTransaction) @@ -1477,25 +1599,31 @@ class TestThlLedgerManagerAdj: user_payout=user_pay, ) - tx = thl_lm.create_tx_bp_adjustment(session=session, created=wall1.started) + tx = thl_ledger_manager.create_tx_bp_adjustment( + session=session, created=wall1.started + ) assert isinstance(tx, LedgerTransaction) - bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(product=user.product) - assert 228 == thl_lm.get_account_balance(account=bp_wallet_account) + bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet( + product=user.product + ) + assert 228 == thl_ledger_manager.get_account_balance(account=bp_wallet_account) - user_account = thl_lm.get_account_or_create_user_wallet(user=user) - assert 152 == thl_lm.get_account_balance(account=user_account) + user_account = thl_ledger_manager.get_account_or_create_user_wallet(user=user) + assert 152 == thl_ledger_manager.get_account_balance(account=user_account) - revenue = thl_lm.get_account_task_complete_revenue() - assert 0 == thl_lm.get_account_balance(account=revenue) + revenue = thl_ledger_manager.get_account_task_complete_revenue() + assert 0 == thl_ledger_manager.get_account_balance(account=revenue) - bp_commission_account = thl_lm.get_account_or_create_bp_commission( + bp_commission_account = thl_ledger_manager.get_account_or_create_bp_commission( product=user.product ) - assert 20 == thl_lm.get_account_balance(account=bp_commission_account) + assert 20 == thl_ledger_manager.get_account_balance( + account=bp_commission_account + ) # the total (4.00) = 2.28 + 1.52 + .20 - assert thl_lm.check_ledger_balanced() + assert thl_ledger_manager.check_ledger_balanced() # This should do nothing (since we haven't adjusted any wall events) session.adjust_status() @@ -1505,7 +1633,7 @@ class TestThlLedgerManagerAdj: session.get_user_payout_after_adjustment(), ) with caplog.at_level(logging.INFO): - thl_lm.create_tx_bp_adjustment(session) + thl_ledger_manager.create_tx_bp_adjustment(session) assert ( "create_transaction_bp_adjustment. No transactions needed." in caplog.text ) @@ -1516,16 +1644,16 @@ class TestThlLedgerManagerAdj: adjusted_cpi=0, adjusted_timestamp=now + timedelta(hours=1), ) - thl_lm.create_tx_task_adjustment(wall1, user) + thl_ledger_manager.create_tx_task_adjustment(wall1, user) # -$1.00 b/c the MP took the $1 back, but we haven't yet taken the BP payment back - assert -100 == thl_lm.get_account_balance(revenue) + assert -100 == thl_ledger_manager.get_account_balance(revenue) session.adjust_status() print( session.get_status_after_adjustment(), session.get_payout_after_adjustment(), session.get_user_payout_after_adjustment(), ) - thl_lm.create_tx_bp_adjustment(session) + thl_ledger_manager.create_tx_bp_adjustment(session) # running this twice b/c it should do nothing the 2nd time print( @@ -1534,16 +1662,16 @@ class TestThlLedgerManagerAdj: session.get_user_payout_after_adjustment(), ) with caplog.at_level(logging.INFO): - thl_lm.create_tx_bp_adjustment(session) + thl_ledger_manager.create_tx_bp_adjustment(session) assert ( "create_transaction_bp_adjustment. No transactions needed." in caplog.text ) - assert 228 - 57 == thl_lm.get_account_balance(bp_wallet_account) - assert 152 - 38 == thl_lm.get_account_balance(user_account) - assert 0 == thl_lm.get_account_balance(revenue) - assert 20 - 5 == thl_lm.get_account_balance(bp_commission_account) - assert thl_lm.check_ledger_balanced() + assert 228 - 57 == thl_ledger_manager.get_account_balance(bp_wallet_account) + assert 152 - 38 == thl_ledger_manager.get_account_balance(user_account) + assert 0 == thl_ledger_manager.get_account_balance(revenue) + assert 20 - 5 == thl_ledger_manager.get_account_balance(bp_commission_account) + assert thl_ledger_manager.check_ledger_balanced() # unrecon the $1 survey wall1.update( @@ -1551,7 +1679,7 @@ class TestThlLedgerManagerAdj: adjusted_cpi=None, adjusted_timestamp=now + timedelta(hours=2), ) - tx = thl_lm.create_tx_task_adjustment(wall=wall1, user=user) + tx = thl_ledger_manager.create_tx_task_adjustment(wall=wall1, user=user) assert isinstance(tx, LedgerTransaction) new_status, new_payout, new_user_payout = ( @@ -1564,13 +1692,17 @@ class TestThlLedgerManagerAdj: session.get_payout_after_adjustment(), session.get_user_payout_after_adjustment(), ) - thl_lm.create_tx_bp_adjustment(session) + thl_ledger_manager.create_tx_bp_adjustment(session) - assert 228 - 57 + 57 == thl_lm.get_account_balance(bp_wallet_account) - assert 152 - 38 + 38 == thl_lm.get_account_balance(user_account) - assert 0 == thl_lm.get_account_balance(revenue) - assert 20 - 5 + 5 == thl_lm.get_account_balance(bp_commission_account) - assert thl_lm.check_ledger_balanced() + assert 228 - 57 + 57 == thl_ledger_manager.get_account_balance( + bp_wallet_account + ) + assert 152 - 38 + 38 == thl_ledger_manager.get_account_balance(user_account) + assert 0 == thl_ledger_manager.get_account_balance(revenue) + assert 20 - 5 + 5 == thl_ledger_manager.get_account_balance( + bp_commission_account + ) + assert thl_ledger_manager.check_ledger_balanced() # make the $2 failure into a complete also wall3.update( @@ -1578,7 +1710,7 @@ class TestThlLedgerManagerAdj: adjusted_cpi=wall3.cpi, adjusted_timestamp=now + timedelta(hours=2), ) - thl_lm.create_tx_task_adjustment(wall3, user) + thl_ledger_manager.create_tx_task_adjustment(wall3, user) new_status, new_payout, new_user_payout = ( session.determine_new_status_and_payouts() ) @@ -1589,24 +1721,30 @@ class TestThlLedgerManagerAdj: session.get_payout_after_adjustment(), session.get_user_payout_after_adjustment(), ) - thl_lm.create_tx_bp_adjustment(session) - assert 228 - 57 + 57 + 114 == thl_lm.get_account_balance(bp_wallet_account) - assert 152 - 38 + 38 + 76 == thl_lm.get_account_balance(user_account) - assert 0 == thl_lm.get_account_balance(revenue) - assert 20 - 5 + 5 + 10 == thl_lm.get_account_balance(bp_commission_account) - assert thl_lm.check_ledger_balanced() + thl_ledger_manager.create_tx_bp_adjustment(session) + assert 228 - 57 + 57 + 114 == thl_ledger_manager.get_account_balance( + bp_wallet_account + ) + assert 152 - 38 + 38 + 76 == thl_ledger_manager.get_account_balance( + user_account + ) + assert 0 == thl_ledger_manager.get_account_balance(revenue) + assert 20 - 5 + 5 + 10 == thl_ledger_manager.get_account_balance( + bp_commission_account + ) + assert thl_ledger_manager.check_ledger_balanced() def test_create_transaction_bp_adjustment_cpi_adjustment( self, - user_factory, - product_user_wallet_no, - create_main_accounts, - delete_ledger_db, + user_factory: Callable[..., User], + product_user_wallet_no: Product, + create_main_accounts: Callable[..., None], + delete_ledger_db: Callable[..., None], caplog, - thl_lm, - lm, - utc_hour_ago, - currency, + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + utc_hour_ago: datetime, + currency: LedgerCurrency, ): delete_ledger_db() create_main_accounts() @@ -1623,7 +1761,7 @@ class TestThlLedgerManagerAdj: started=utc_hour_ago, finished=utc_hour_ago + timedelta(seconds=1), ) - tx = thl_lm.create_tx_task_complete( + tx = thl_ledger_manager.create_tx_task_complete( wall=wall1, user=user, created=wall1.started ) assert isinstance(tx, LedgerTransaction) @@ -1639,7 +1777,7 @@ class TestThlLedgerManagerAdj: started=utc_hour_ago, finished=utc_hour_ago + timedelta(seconds=1), ) - tx = thl_lm.create_tx_task_complete( + tx = thl_ledger_manager.create_tx_task_complete( wall=wall2, user=user, created=wall2.started ) assert isinstance(tx, LedgerTransaction) @@ -1654,15 +1792,19 @@ class TestThlLedgerManagerAdj: payout=bp_pay, user_payout=user_pay, ) - thl_lm.create_tx_bp_payment(session, created=wall1.started) + thl_ledger_manager.create_tx_bp_payment(session, created=wall1.started) - revenue = thl_lm.get_account_task_complete_revenue() - bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(user.product) - bp_commission_account = thl_lm.get_account_or_create_bp_commission(user.product) - assert 380 == thl_lm.get_account_balance(bp_wallet_account) - assert 0 == thl_lm.get_account_balance(revenue) - assert 20 == thl_lm.get_account_balance(bp_commission_account) - assert thl_lm.check_ledger_balanced() + revenue = ththl_ledger_managerl_lm.get_account_task_complete_revenue() + bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet( + user.product + ) + bp_commission_account = thl_ledger_manager.get_account_or_create_bp_commission( + user.product + ) + assert 380 == thl_ledger_manager.get_account_balance(bp_wallet_account) + assert 0 == thl_ledger_manager.get_account_balance(revenue) + assert 20 == thl_ledger_manager.get_account_balance(bp_commission_account) + assert thl_ledger_manager.check_ledger_balanced() # cpi adjustment $1 -> $.60. wall1.update( @@ -1670,17 +1812,17 @@ class TestThlLedgerManagerAdj: adjusted_cpi=Decimal("0.60"), adjusted_timestamp=utc_hour_ago + timedelta(minutes=30), ) - thl_lm.create_tx_task_adjustment(wall1, user) + thl_ledger_manager.create_tx_task_adjustment(wall1, user) # -$0.40 b/c the MP took $0.40 back, but we haven't yet taken the BP payment back - assert -40 == thl_lm.get_account_balance(revenue) + assert -40 == thl_ledger_manager.get_account_balance(revenue) session.adjust_status() print( session.get_status_after_adjustment(), session.get_payout_after_adjustment(), session.get_user_payout_after_adjustment(), ) - thl_lm.create_tx_bp_adjustment(session) + thl_ledger_manager.create_tx_bp_adjustment(session) # running this twice b/c it should do nothing the 2nd time print( @@ -1689,14 +1831,14 @@ class TestThlLedgerManagerAdj: session.get_user_payout_after_adjustment(), ) with caplog.at_level(logging.INFO): - thl_lm.create_tx_bp_adjustment(session) + thl_ledger_manager.create_tx_bp_adjustment(session) assert "create_transaction_bp_adjustment." in caplog.text assert "No transactions needed." in caplog.text - assert 380 - 38 == thl_lm.get_account_balance(bp_wallet_account) - assert 0 == thl_lm.get_account_balance(revenue) - assert 20 - 2 == thl_lm.get_account_balance(bp_commission_account) - assert thl_lm.check_ledger_balanced() + assert 380 - 38 == thl_ledger_manager.get_account_balance(bp_wallet_account) + assert 0 == thl_ledger_manager.get_account_balance(revenue) + assert 20 - 2 == thl_ledger_manager.get_account_balance(bp_commission_account) + assert thl_ledger_manager.check_ledger_balanced() # adjust it to failure wall1.update( @@ -1704,13 +1846,17 @@ class TestThlLedgerManagerAdj: adjusted_cpi=0, adjusted_timestamp=utc_hour_ago + timedelta(minutes=45), ) - thl_lm.create_tx_task_adjustment(wall1, user) + thl_ledger_manager.create_tx_task_adjustment(wall1, user) session.adjust_status() - thl_lm.create_tx_bp_adjustment(session) - assert 300 - (300 * 0.05) == thl_lm.get_account_balance(bp_wallet_account) - assert 0 == thl_lm.get_account_balance(revenue) - assert 300 * 0.05 == thl_lm.get_account_balance(bp_commission_account) - assert thl_lm.check_ledger_balanced() + thl_ledger_manager.create_tx_bp_adjustment(session) + assert 300 - (300 * 0.05) == thl_ledger_manager.get_account_balance( + bp_wallet_account + ) + assert 0 == thl_ledger_manager.get_account_balance(revenue) + assert 300 * 0.05 == thl_ledger_manager.get_account_balance( + bp_commission_account + ) + assert thl_ledger_manager.check_ledger_balanced() # and then back to cpi adj again, but this time for more than the orig amount wall1.update( @@ -1718,13 +1864,17 @@ class TestThlLedgerManagerAdj: adjusted_cpi=Decimal("2.00"), adjusted_timestamp=utc_hour_ago + timedelta(minutes=45), ) - thl_lm.create_tx_task_adjustment(wall1, user) + thl_ledger_manager.create_tx_task_adjustment(wall1, user) session.adjust_status() - thl_lm.create_tx_bp_adjustment(session) - assert 500 - (500 * 0.05) == thl_lm.get_account_balance(bp_wallet_account) - assert 0 == thl_lm.get_account_balance(revenue) - assert 500 * 0.05 == thl_lm.get_account_balance(bp_commission_account) - assert thl_lm.check_ledger_balanced() + thl_ledger_manager.create_tx_bp_adjustment(session) + assert 500 - (500 * 0.05) == thl_ledger_manager.get_account_balance( + bp_wallet_account + ) + assert 0 == thl_ledger_manager.get_account_balance(revenue) + assert 500 * 0.05 == thl_ledger_manager.get_account_balance( + bp_commission_account + ) + assert thl_ledger_manager.check_ledger_balanced() # And adjust again wall1.update( @@ -1732,12 +1882,14 @@ class TestThlLedgerManagerAdj: adjusted_cpi=Decimal("3.00"), adjusted_timestamp=utc_hour_ago + timedelta(minutes=45), ) - thl_lm.create_tx_task_adjustment(wall=wall1, user=user) + thl_ledger_manager.create_tx_task_adjustment(wall=wall1, user=user) session.adjust_status() - thl_lm.create_tx_bp_adjustment(session=session) - assert 600 - (600 * 0.05) == thl_lm.get_account_balance( + thl_ledger_manager.create_tx_bp_adjustment(session=session) + assert 600 - (600 * 0.05) == thl_ledger_manager.get_account_balance( account=bp_wallet_account ) - assert 0 == thl_lm.get_account_balance(account=revenue) - assert 600 * 0.05 == thl_lm.get_account_balance(account=bp_commission_account) - assert thl_lm.check_ledger_balanced() + assert 0 == thl_ledger_manager.get_account_balance(account=revenue) + assert 600 * 0.05 == thl_ledger_manager.get_account_balance( + account=bp_commission_account + ) + assert thl_ledger_manager.check_ledger_balanced() diff --git a/tests/managers/thl/test_ledger/test_thl_lm_tx__user_payouts.py b/tests/managers/thl/test_ledger/test_thl_lm_tx__user_payouts.py index cd6ea79..5fb6935 100644 --- a/tests/managers/thl/test_ledger/test_thl_lm_tx__user_payouts.py +++ b/tests/managers/thl/test_ledger/test_thl_lm_tx__user_payouts.py @@ -1,5 +1,7 @@ +from __future__ import annotations + import logging -from datetime import UTC, datetime, timedelta +from collections.abc import Callable from decimal import Decimal from uuid import uuid4 @@ -9,7 +11,10 @@ from generalresearch.managers.thl.ledger_manager.exceptions import ( LedgerTransactionConditionFailedError, LedgerTransactionFlagAlreadyExistsError, ) +from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager +from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager from generalresearch.models.thl.payout import UserPayoutEvent +from generalresearch.models.thl.product import Product from generalresearch.models.thl.user import User from generalresearch.models.thl.wallet import PayoutType @@ -18,12 +23,12 @@ class TestLedgerManagerAMT: def test_create_transaction_amt_ass_request( self, - user_factory, - product_amt_true, - create_main_accounts, - thl_lm, - lm, - delete_ledger_db, + user_factory: Callable[..., User], + product_amt_true: Product, + create_main_accounts: Callable[..., None], + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + delete_ledger_db: Callable[..., None], ): delete_ledger_db() create_main_accounts() @@ -39,16 +44,16 @@ class TestLedgerManagerAMT: ) flag_key = f"test:user_payout:{pe.uuid}:request" flag_name = f"ledger-manager:transaction_flag:{flag_key}" - lm.redis_client.delete(flag_name) + ledger_manager.redis_client.delete(flag_name) # User has $0 in their wallet. They are allowed amt_assignment payouts until -$1.00 - thl_lm.create_tx_user_payout_request(user=user, payout_event=pe) + thl_ledger_manager.create_tx_user_payout_request(user=user, payout_event=pe) with pytest.raises(expected_exception=LedgerTransactionFlagAlreadyExistsError): - thl_lm.create_tx_user_payout_request( + thl_ledger_manager.create_tx_user_payout_request( user=user, payout_event=pe, skip_flag_check=False ) with pytest.raises(expected_exception=LedgerTransactionConditionFailedError): - thl_lm.create_tx_user_payout_request( + thl_ledger_manager.create_tx_user_payout_request( user=user, payout_event=pe, skip_flag_check=True ) pe2 = UserPayoutEvent( @@ -61,36 +66,40 @@ class TestLedgerManagerAMT: flag_key = f"test:user_payout:{pe2.uuid}:request" flag_name = f"ledger-manager:transaction_flag:{flag_key}" - lm.redis_client.delete(flag_name) + ledger_manager.redis_client.delete(flag_name) # 96 cents would put them over the -$1.00 limit with pytest.raises(expected_exception=LedgerTransactionConditionFailedError): - thl_lm.create_tx_user_payout_request(user, payout_event=pe2) + thl_ledger_manager.create_tx_user_payout_request(user, payout_event=pe2) # But they could do 0.95 cents pe2.amount = 95 - thl_lm.create_tx_user_payout_request( + thl_ledger_manager.create_tx_user_payout_request( user, payout_event=pe2, skip_flag_check=True ) - cash = thl_lm.get_account_cash() - bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(user.product) - bp_pending_account = thl_lm.get_or_create_bp_pending_payout_account( + cash = thl_ledger_manager.get_account_cash() + bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet( + user.product + ) + bp_pending_account = thl_ledger_manager.get_or_create_bp_pending_payout_account( product=user.product ) - user_wallet_account = thl_lm.get_account_or_create_user_wallet(user=user) + user_wallet_account = thl_ledger_manager.get_account_or_create_user_wallet( + user=user + ) - assert 0 == lm.get_account_balance(account=bp_wallet_account) - assert 0 == lm.get_account_balance(account=cash) - assert 100 == lm.get_account_balance(account=bp_pending_account) - assert -100 == lm.get_account_balance(account=user_wallet_account) - assert thl_lm.check_ledger_balanced() - assert -5 == thl_lm.get_account_filtered_balance( + assert 0 == ledger_manager.get_account_balance(account=bp_wallet_account) + assert 0 == ledger_manager.get_account_balance(account=cash) + assert 100 == ledger_manager.get_account_balance(account=bp_pending_account) + assert -100 == ledger_manager.get_account_balance(account=user_wallet_account) + assert thl_ledger_manager.check_ledger_balanced() + assert -5 == thl_ledger_manager.get_account_filtered_balance( account=user_wallet_account, metadata_key="payoutevent", metadata_value=pe.uuid, ) - assert -95 == thl_lm.get_account_filtered_balance( + assert -95 == thl_ledger_manager.get_account_filtered_balance( account=user_wallet_account, metadata_key="payoutevent", metadata_value=pe2.uuid, @@ -98,12 +107,12 @@ class TestLedgerManagerAMT: def test_create_transaction_amt_ass_complete( self, - user_factory, - product_amt_true, - create_main_accounts, - thl_lm, - lm, - delete_ledger_db, + user_factory: Callable[..., User], + product_amt_true: Product, + create_main_accounts: Callable[..., None], + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + delete_ledger_db: Callable[..., None], ): delete_ledger_db() create_main_accounts() @@ -117,40 +126,42 @@ class TestLedgerManagerAMT: debit_account_uuid=uuid4().hex, ) flag = f"ledger-manager:transaction_flag:test:user_payout:{pe.uuid}:request" - lm.redis_client.delete(flag) + ledger_manager.redis_client.delete(flag) flag = f"ledger-manager:transaction_flag:test:user_payout:{pe.uuid}:complete" - lm.redis_client.delete(flag) + ledger_manager.redis_client.delete(flag) # User has $0 in their wallet. They are allowed amt_assignment payouts until -$1.00 - thl_lm.create_tx_user_payout_request(user, payout_event=pe) - thl_lm.create_tx_user_payout_complete(user, payout_event=pe) + thl_ledger_manager.create_tx_user_payout_request(user, payout_event=pe) + thl_ledger_manager.create_tx_user_payout_complete(user, payout_event=pe) - cash = thl_lm.get_account_cash() - bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(user.product) - bp_pending_account = thl_lm.get_or_create_bp_pending_payout_account( + cash = thl_ledger_manager.get_account_cash() + bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet( + user.product + ) + bp_pending_account = thl_ledger_manager.get_or_create_bp_pending_payout_account( user.product ) - bp_amt_expense_account = thl_lm.get_account_or_create_bp_expense( + bp_amt_expense_account = thl_ledger_manager.get_account_or_create_bp_expense( user.product, expense_name="amt" ) - user_wallet_account = thl_lm.get_account_or_create_user_wallet(user) + user_wallet_account = thl_ledger_manager.get_account_or_create_user_wallet(user) # BP wallet pays the 1cent fee - assert -1 == thl_lm.get_account_balance(bp_wallet_account) - assert -5 == thl_lm.get_account_balance(cash) - assert -1 == thl_lm.get_account_balance(bp_amt_expense_account) - assert 0 == thl_lm.get_account_balance(bp_pending_account) - assert -5 == lm.get_account_balance(user_wallet_account) - assert thl_lm.check_ledger_balanced() + assert -1 == thl_ledger_manager.get_account_balance(bp_wallet_account) + assert -5 == thl_ledger_manager.get_account_balance(cash) + assert -1 == thl_ledger_manager.get_account_balance(bp_amt_expense_account) + assert 0 == thl_ledger_manager.get_account_balance(bp_pending_account) + assert -5 == ledger_manager.get_account_balance(user_wallet_account) + assert thl_ledger_manager.check_ledger_balanced() def test_create_transaction_amt_bonus( self, - user_factory, - product_amt_true, - create_main_accounts, - thl_lm, - lm, - delete_ledger_db, + user_factory: Callable[..., User], + product_amt_true: Product, + create_main_accounts: Callable[..., None], + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + delete_ledger_db: Callable[..., None], ): delete_ledger_db() create_main_accounts() @@ -165,15 +176,15 @@ class TestLedgerManagerAMT: debit_account_uuid=uuid4().hex, ) flag = f"ledger-manager:transaction_flag:test:user_payout:{pe.uuid}:request" - lm.redis_client.delete(flag) + ledger_manager.redis_client.delete(flag) flag = f"ledger-manager:transaction_flag:test:user_payout:{pe.uuid}:complete" - lm.redis_client.delete(flag) + ledger_manager.redis_client.delete(flag) with pytest.raises(expected_exception=LedgerTransactionConditionFailedError): # User has $0 in their wallet. No amt bonus allowed - thl_lm.create_tx_user_payout_request(user, payout_event=pe) + thl_ledger_manager.create_tx_user_payout_request(user, payout_event=pe) - thl_lm.create_tx_user_bonus( + thl_ledger_manager.create_tx_user_bonus( user, amount=Decimal(5), ref_uuid="e703830dec124f17abed2d697d8d7701", @@ -181,68 +192,68 @@ class TestLedgerManagerAMT: skip_flag_check=True, ) pe.amount = 101 - thl_lm.create_tx_user_payout_request( + thl_ledger_manager.create_tx_user_payout_request( user, payout_event=pe, skip_flag_check=False ) - thl_lm.create_tx_user_payout_complete( + thl_ledger_manager.create_tx_user_payout_complete( user, payout_event=pe, skip_flag_check=False ) with pytest.raises(expected_exception=LedgerTransactionFlagAlreadyExistsError): # duplicate, even if amount changed pe.amount = 200 - thl_lm.create_tx_user_payout_complete( + thl_ledger_manager.create_tx_user_payout_complete( user, payout_event=pe, skip_flag_check=False ) with pytest.raises(expected_exception=LedgerTransactionConditionFailedError): # duplicate - thl_lm.create_tx_user_payout_complete( + thl_ledger_manager.create_tx_user_payout_complete( user, payout_event=pe, skip_flag_check=True ) pe.uuid = "533364150de4451198e5774e221a2acb" pe.amount = 9900 with pytest.raises(expected_exception=ValueError): # Trying to complete payout with no pending tx - thl_lm.create_tx_user_payout_complete( + thl_ledger_manager.create_tx_user_payout_complete( user, payout_event=pe, skip_flag_check=True ) with pytest.raises(expected_exception=LedgerTransactionConditionFailedError): # trying to payout $99 with only a $5 balance - thl_lm.create_tx_user_payout_request( + thl_ledger_manager.create_tx_user_payout_request( user, payout_event=pe, skip_flag_check=True ) - cash = thl_lm.get_account_cash() - bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(user.product) - bp_pending_account = thl_lm.get_or_create_bp_pending_payout_account( + cash = thl_ledger_manager.get_account_cash() + bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet( user.product ) - bp_amt_expense_account = thl_lm.get_account_or_create_bp_expense( + bp_pending_account = thl_ledger_manager.get_or_create_bp_pending_payout_account( + user.product + ) + bp_amt_expense_account = thl_ledger_manager.get_account_or_create_bp_expense( user.product, expense_name="amt" ) - user_wallet_account = thl_lm.get_account_or_create_user_wallet(user) - assert -500 + round(-101 * 0.20) == thl_lm.get_account_balance( + user_wallet_account = thl_ledger_manager.get_account_or_create_user_wallet(user) + assert -500 + round(-101 * 0.20) == thl_ledger_manager.get_account_balance( bp_wallet_account ) - assert -101 == lm.get_account_balance(cash) - assert -20 == lm.get_account_balance(bp_amt_expense_account) - assert 0 == lm.get_account_balance(bp_pending_account) - assert 500 - 101 == lm.get_account_balance(user_wallet_account) - assert lm.check_ledger_balanced() is True + assert -101 == ledger_manager.get_account_balance(cash) + assert -20 == ledger_manager.get_account_balance(bp_amt_expense_account) + assert 0 == ledger_manager.get_account_balance(bp_pending_account) + assert 500 - 101 == ledger_manager.get_account_balance(user_wallet_account) + assert ledger_manager.check_ledger_balanced() is True def test_create_transaction_amt_bonus_cancel( self, - user_factory, - product_amt_true, - create_main_accounts, + user_factory: Callable[..., User], + product_amt_true: Product, + create_main_accounts: Callable[..., None], caplog, - thl_lm, - lm, - delete_ledger_db, + thl_ledger_manager: ThlLedgerManager, + delete_ledger_db: Callable[..., None], ): delete_ledger_db() create_main_accounts() - now = datetime.now(UTC) - timedelta(hours=1) user: User = user_factory(product=product_amt_true) pe = UserPayoutEvent( @@ -253,41 +264,47 @@ class TestLedgerManagerAMT: debit_account_uuid=uuid4().hex, ) - thl_lm.create_tx_user_bonus( + thl_ledger_manager.create_tx_user_bonus( user, amount=Decimal(5), ref_uuid="c44f4da2db1d421ebc6a5e5241ca4ce6", description="Bribe", skip_flag_check=True, ) - thl_lm.create_tx_user_payout_request( + thl_ledger_manager.create_tx_user_payout_request( user, payout_event=pe, skip_flag_check=True ) - thl_lm.create_tx_user_payout_cancelled( + thl_ledger_manager.create_tx_user_payout_cancelled( user, payout_event=pe, skip_flag_check=True ) with pytest.raises(expected_exception=LedgerTransactionConditionFailedError): with caplog.at_level(logging.WARNING): - thl_lm.create_tx_user_payout_complete( + thl_ledger_manager.create_tx_user_payout_complete( user, payout_event=pe, skip_flag_check=True ) assert "trying to complete payout that was already cancelled" in caplog.text - cash = thl_lm.get_account_cash() - bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(user.product) - bp_pending_account = thl_lm.get_or_create_bp_pending_payout_account( + cash = thl_ledger_manager.get_account_cash() + bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet( user.product ) - bp_amt_expense_account = thl_lm.get_account_or_create_bp_expense( + bp_pending_account = thl_ledger_manager.get_or_create_bp_pending_payout_account( + user.product + ) + bp_amt_expense_account = thl_ledger_manager.get_account_or_create_bp_expense( user.product, expense_name="amt" ) - user_wallet_account = thl_lm.get_account_or_create_user_wallet(user) - assert -500 == thl_lm.get_account_balance(account=bp_wallet_account) - assert 0 == thl_lm.get_account_balance(account=cash) - assert 0 == thl_lm.get_account_balance(account=bp_amt_expense_account) - assert 0 == thl_lm.get_account_balance(account=bp_pending_account) - assert 500 == thl_lm.get_account_balance(account=user_wallet_account) - assert thl_lm.check_ledger_balanced() + user_wallet_account = thl_ledger_manager.get_account_or_create_user_wallet(user) + assert -500 == thl_ledger_manager.get_account_balance(account=bp_wallet_account) + assert 0 == thl_ledger_manager.get_account_balance(account=cash) + assert 0 == thl_ledger_manager.get_account_balance( + account=bp_amt_expense_account + ) + assert 0 == thl_ledger_manager.get_account_balance(account=bp_pending_account) + assert 500 == thl_ledger_manager.get_account_balance( + account=user_wallet_account + ) + assert thl_ledger_manager.check_ledger_balanced() pe2 = UserPayoutEvent( uuid=uuid4().hex, @@ -296,15 +313,15 @@ class TestLedgerManagerAMT: cashout_method_uuid=uuid4().hex, debit_account_uuid=uuid4().hex, ) - thl_lm.create_tx_user_payout_request( + thl_ledger_manager.create_tx_user_payout_request( user, payout_event=pe2, skip_flag_check=True ) - thl_lm.create_tx_user_payout_complete( + thl_ledger_manager.create_tx_user_payout_complete( user, payout_event=pe2, skip_flag_check=True ) with pytest.raises(expected_exception=LedgerTransactionConditionFailedError): with caplog.at_level(logging.WARNING): - thl_lm.create_tx_user_payout_cancelled( + thl_ledger_manager.create_tx_user_payout_cancelled( user, payout_event=pe2, skip_flag_check=True ) assert "trying to cancel payout that was already completed" in caplog.text @@ -314,12 +331,12 @@ class TestLedgerManagerTango: def test_create_transaction_tango_request( self, - user_factory, - product_amt_true, - create_main_accounts, - thl_lm, - lm, - delete_ledger_db, + user_factory: Callable[..., User], + product_amt_true: Product, + create_main_accounts: Callable[..., None], + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + delete_ledger_db: Callable[..., None], ): delete_ledger_db() create_main_accounts() @@ -336,64 +353,65 @@ class TestLedgerManagerTango: ) flag_key = f"test:user_payout:{pe.uuid}:request" flag_name = f"ledger-manager:transaction_flag:{flag_key}" - lm.redis_client.delete(flag_name) - thl_lm.create_tx_user_bonus( + ledger_manager.redis_client.delete(flag_name) + thl_ledger_manager.create_tx_user_bonus( user, amount=Decimal(6), ref_uuid="e703830dec124f17abed2d697d8d7701", description="Bribe", skip_flag_check=True, ) - thl_lm.create_tx_user_payout_request( + thl_ledger_manager.create_tx_user_payout_request( user, payout_event=pe, skip_flag_check=True ) - cash = thl_lm.get_account_cash() - bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(user.product) - bp_pending_account = thl_lm.get_or_create_bp_pending_payout_account( + cash = thl_ledger_manager.get_account_cash() + bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet( + user.product + ) + bp_pending_account = thl_ledger_manager.get_or_create_bp_pending_payout_account( user.product ) - bp_tango_expense_account = thl_lm.get_account_or_create_bp_expense( + bp_tango_expense_account = thl_ledger_manager.get_account_or_create_bp_expense( user.product, expense_name="tango" ) - user_wallet_account = thl_lm.get_account_or_create_user_wallet(user) - assert -600 == thl_lm.get_account_balance(bp_wallet_account) - assert 0 == thl_lm.get_account_balance(cash) - assert 0 == thl_lm.get_account_balance(bp_tango_expense_account) - assert 500 == thl_lm.get_account_balance(bp_pending_account) - assert 600 - 500 == thl_lm.get_account_balance(user_wallet_account) - assert thl_lm.check_ledger_balanced() + user_wallet_account = thl_ledger_manager.get_account_or_create_user_wallet(user) + assert -600 == thl_ledger_manager.get_account_balance(bp_wallet_account) + assert 0 == thl_ledger_manager.get_account_balance(cash) + assert 0 == thl_ledger_manager.get_account_balance(bp_tango_expense_account) + assert 500 == thl_ledger_manager.get_account_balance(bp_pending_account) + assert 600 - 500 == thl_ledger_manager.get_account_balance(user_wallet_account) + assert thl_ledger_manager.check_ledger_balanced() - thl_lm.create_tx_user_payout_complete( + thl_ledger_manager.create_tx_user_payout_complete( user, payout_event=pe, skip_flag_check=True ) - assert -600 - round(500 * 0.035) == thl_lm.get_account_balance( + assert -600 - round(500 * 0.035) == thl_ledger_manager.get_account_balance( bp_wallet_account ) - assert -500, thl_lm.get_account_balance(cash) - assert round(-500 * 0.035) == thl_lm.get_account_balance( + assert -500, thl_ledger_manager.get_account_balance(cash) + assert round(-500 * 0.035) == thl_ledger_manager.get_account_balance( bp_tango_expense_account ) - assert 0 == lm.get_account_balance(bp_pending_account) - assert 100 == lm.get_account_balance(user_wallet_account) - assert lm.check_ledger_balanced() + assert 0 == ledger_manager.get_account_balance(bp_pending_account) + assert 100 == ledger_manager.get_account_balance(user_wallet_account) + assert ledger_manager.check_ledger_balanced() class TestLedgerManagerPaypal: def test_create_transaction_paypal_request( self, - user_factory, - product_amt_true, - create_main_accounts, - thl_lm, - lm, - delete_ledger_db, + user_factory: Callable[..., User], + product_amt_true: Product, + create_main_accounts: Callable[..., None], + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + delete_ledger_db: Callable[..., None], ): delete_ledger_db() create_main_accounts() - now = datetime.now(tz=UTC) - timedelta(hours=1) user: User = user_factory(product=product_amt_true) # debit_account_uuid nothing checks they match the ledger ... todo? @@ -406,8 +424,8 @@ class TestLedgerManagerPaypal: ) flag_key = f"test:user_payout:{pe.uuid}:request" flag_name = f"ledger-manager:transaction_flag:{flag_key}" - lm.redis_client.delete(flag_name) - thl_lm.create_tx_user_bonus( + ledger_manager.redis_client.delete(flag_name) + thl_ledger_manager.create_tx_user_bonus( user=user, amount=Decimal(6), ref_uuid="e703830dec124f17abed2d697d8d7701", @@ -415,79 +433,91 @@ class TestLedgerManagerPaypal: skip_flag_check=True, ) - thl_lm.create_tx_user_payout_request( + thl_ledger_manager.create_tx_user_payout_request( user, payout_event=pe, skip_flag_check=True ) - cash = thl_lm.get_account_cash() - bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(user.product) - bp_pending_account = thl_lm.get_or_create_bp_pending_payout_account( + cash = thl_ledger_manager.get_account_cash() + bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet( + user.product + ) + bp_pending_account = thl_ledger_manager.get_or_create_bp_pending_payout_account( product=user.product ) - bp_paypal_expense_account = thl_lm.get_account_or_create_bp_expense( + bp_paypal_expense_account = thl_ledger_manager.get_account_or_create_bp_expense( product=user.product, expense_name="paypal" ) - user_wallet_account = thl_lm.get_account_or_create_user_wallet(user=user) - assert -600 == lm.get_account_balance(account=bp_wallet_account) - assert 0 == lm.get_account_balance(account=cash) - assert 0 == lm.get_account_balance(account=bp_paypal_expense_account) - assert 500 == lm.get_account_balance(account=bp_pending_account) - assert 600 - 500 == lm.get_account_balance(account=user_wallet_account) - assert thl_lm.check_ledger_balanced() + user_wallet_account = thl_ledger_manager.get_account_or_create_user_wallet( + user=user + ) + assert -600 == ledger_manager.get_account_balance(account=bp_wallet_account) + assert 0 == ledger_manager.get_account_balance(account=cash) + assert 0 == ledger_manager.get_account_balance( + account=bp_paypal_expense_account + ) + assert 500 == ledger_manager.get_account_balance(account=bp_pending_account) + assert 600 - 500 == ledger_manager.get_account_balance( + account=user_wallet_account + ) + assert thl_ledger_manager.check_ledger_balanced() - thl_lm.create_tx_user_payout_complete( + thl_ledger_manager.create_tx_user_payout_complete( user=user, payout_event=pe, skip_flag_check=True, fee_amount=Decimal("0.50") ) - assert -600 - 50 == thl_lm.get_account_balance(bp_wallet_account) - assert -500 == thl_lm.get_account_balance(cash) - assert -50 == thl_lm.get_account_balance(bp_paypal_expense_account) - assert 0 == thl_lm.get_account_balance(bp_pending_account) - assert 100 == thl_lm.get_account_balance(user_wallet_account) - assert thl_lm.check_ledger_balanced() + assert -600 - 50 == thl_ledger_manager.get_account_balance(bp_wallet_account) + assert -500 == thl_ledger_manager.get_account_balance(cash) + assert -50 == thl_ledger_manager.get_account_balance(bp_paypal_expense_account) + assert 0 == thl_ledger_manager.get_account_balance(bp_pending_account) + assert 100 == thl_ledger_manager.get_account_balance(user_wallet_account) + assert thl_ledger_manager.check_ledger_balanced() class TestLedgerManagerBonus: def test_create_transaction_bonus( self, - user_factory, - product_user_wallet_yes, - create_main_accounts, - thl_lm, - lm, - delete_ledger_db, + user_factory: Callable[..., User], + product_user_wallet_yes: Product, + create_main_accounts: Callable[..., None], + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, + delete_ledger_db: Callable[..., None], ): delete_ledger_db() create_main_accounts() user: User = user_factory(product=product_user_wallet_yes) - thl_lm.create_tx_user_bonus( + thl_ledger_manager.create_tx_user_bonus( user=user, amount=Decimal(5), ref_uuid="8d0aaf612462448a9ebdd57fab0fc660", description="Bribe", skip_flag_check=True, ) - cash = thl_lm.get_account_cash() - bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(user.product) - bp_pending_account = thl_lm.get_or_create_bp_pending_payout_account( + cash = thl_ledger_manager.get_account_cash() + bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet( + user.product + ) + bp_pending_account = thl_ledger_manager.get_or_create_bp_pending_payout_account( product=user.product ) - bp_amt_expense_account = thl_lm.get_account_or_create_bp_expense( + bp_amt_expense_account = thl_ledger_manager.get_account_or_create_bp_expense( user.product, expense_name="amt" ) - user_wallet_account = thl_lm.get_account_or_create_user_wallet(user=user) + user_wallet_account = thl_ledger_manager.get_account_or_create_user_wallet( + user=user + ) - assert -500 == lm.get_account_balance(account=bp_wallet_account) - assert 0 == lm.get_account_balance(account=cash) - assert 0 == lm.get_account_balance(account=bp_amt_expense_account) - assert 0 == lm.get_account_balance(account=bp_pending_account) - assert 500 == lm.get_account_balance(account=user_wallet_account) - assert thl_lm.check_ledger_balanced() + assert -500 == ledger_manager.get_account_balance(account=bp_wallet_account) + assert 0 == ledger_manager.get_account_balance(account=cash) + assert 0 == ledger_manager.get_account_balance(account=bp_amt_expense_account) + assert 0 == ledger_manager.get_account_balance(account=bp_pending_account) + assert 500 == ledger_manager.get_account_balance(account=user_wallet_account) + assert thl_ledger_manager.check_ledger_balanced() with pytest.raises(expected_exception=LedgerTransactionFlagAlreadyExistsError): - thl_lm.create_tx_user_bonus( + thl_ledger_manager.create_tx_user_bonus( user=user, amount=Decimal(5), ref_uuid="8d0aaf612462448a9ebdd57fab0fc660", @@ -495,7 +525,7 @@ class TestLedgerManagerBonus: skip_flag_check=False, ) with pytest.raises(expected_exception=LedgerTransactionConditionFailedError): - thl_lm.create_tx_user_bonus( + thl_ledger_manager.create_tx_user_bonus( user=user, amount=Decimal(5), ref_uuid="8d0aaf612462448a9ebdd57fab0fc660", diff --git a/tests/managers/thl/test_ledger/test_thl_pem.py b/tests/managers/thl/test_ledger/test_thl_pem.py index 7a4f0c9..29341cf 100644 --- a/tests/managers/thl/test_ledger/test_thl_pem.py +++ b/tests/managers/thl/test_ledger/test_thl_pem.py @@ -1,11 +1,24 @@ +from __future__ import annotations + +from collections.abc import Callable from random import randint from uuid import UUID, uuid4 import pytest from generalresearch.currency import USDCent +from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager +from generalresearch.managers.thl.ledger_manager.thl_ledger import ( + ThlLedgerManager, +) +from generalresearch.managers.thl.payout import ( + BrokerageProductPayoutEventManager, + UserPayoutEventManager, +) from generalresearch.models.thl.definitions import PayoutStatus -from generalresearch.models.thl.payout import BrokerageProductPayoutEvent +from generalresearch.models.thl.payout import ( + BrokerageProductPayoutEvent, +) from generalresearch.models.thl.product import Product from generalresearch.models.thl.wallet.cashout_method import ( CashoutRequestInfo, @@ -14,7 +27,9 @@ from generalresearch.models.thl.wallet.cashout_method import ( class TestThlPayoutEventManager: - def test_get_by_uuid(self, brokerage_product_payout_event_manager, thl_lm): + def test_get_by_uuid( + self, brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager + ): """This validates that the method raises an exception if it fails. There are plenty of other tests that use this method so it seems silly to duplicate it here again @@ -26,11 +41,11 @@ class TestThlPayoutEventManager: def test_filter_by( self, - product_factory, - usd_cent, - bp_payout_event_factory, - thl_lm, - brokerage_product_payout_event_manager, + product_factory: Callable[..., Product], + usd_cent: USDCent, + bp_payout_event_factory: Callable[..., BrokerageProductPayoutEvent], + thl_ledger_manager: ThlLedgerManager, + brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager, ): N_PRODUCTS = randint(3, 10) @@ -38,22 +53,22 @@ class TestThlPayoutEventManager: amounts = [] products = [] - for x_idx in range(N_PRODUCTS): + for _ in range(N_PRODUCTS): product: Product = product_factory() - thl_lm.get_account_or_create_bp_wallet(product=product) + thl_ledger_manager.get_account_or_create_bp_wallet(product=product) products.append(product) brokerage_product_payout_event_manager.set_account_lookup_table( - thl_lm=thl_lm + thl_lm=thl_ledger_manager ) - for y_idx in range(N_PAYOUT_EVENTS): + for _ in range(N_PAYOUT_EVENTS): pe = bp_payout_event_factory(product=product, usd_cent=usd_cent) amounts.append(int(usd_cent)) assert isinstance(pe, BrokerageProductPayoutEvent) # We just added Payout Events for Products, now go ahead and # query for them - accounts = thl_lm.get_accounts_bp_wallet_for_products( + accounts = thl_ledger_manager.get_accounts_bp_wallet_for_products( product_uuids=[i.uuid for i in products] ) res = brokerage_product_payout_event_manager.filter_by( @@ -65,11 +80,11 @@ class TestThlPayoutEventManager: def test_get_bp_payout_events_for_product( self, - product_factory, - usd_cent, - bp_payout_event_factory, - brokerage_product_payout_event_manager, - thl_lm, + product_factory: Callable[..., Product], + usd_cent: USDCent, + bp_payout_event_factory: Callable[..., BrokerageProductPayoutEvent], + brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager, + thl_ledger_manager: ThlLedgerManager, ): N_PRODUCTS = randint(3, 10) @@ -77,23 +92,23 @@ class TestThlPayoutEventManager: amounts = [] products = [] - for x_idx in range(N_PRODUCTS): + for _ in range(N_PRODUCTS): product: Product = product_factory() products.append(product) - thl_lm.get_account_or_create_bp_wallet(product=product) + thl_ledger_manager.get_account_or_create_bp_wallet(product=product) brokerage_product_payout_event_manager.set_account_lookup_table( - thl_lm=thl_lm + thl_lm=thl_ledger_manager ) - for y_idx in range(N_PAYOUT_EVENTS): + for _ in range(N_PAYOUT_EVENTS): pe = bp_payout_event_factory(product=product, usd_cent=usd_cent) amounts.append(usd_cent) assert isinstance(pe, BrokerageProductPayoutEvent) - # We just added 5 Payouts for a specific Product, now go + # We just added 5 Payouts for a specific product: Product, now go # ahead and query for them res = brokerage_product_payout_event_manager.get_bp_bp_payout_events_for_products( - thl_ledger_manager=thl_lm, product_uuids=[product.id] + thl_ledger_manager=thl_ledger_manager, product_uuids=[product.id] ) assert len(res) == N_PAYOUT_EVENTS @@ -102,7 +117,8 @@ class TestThlPayoutEventManager: # ahead and query for them res = ( brokerage_product_payout_event_manager.get_bp_bp_payout_events_for_products( - thl_ledger_manager=thl_lm, product_uuids=[i.uuid for i in products] + thl_ledger_manager=thl_ledger_manager, + product_uuids=[i.uuid for i in products], ) ) @@ -110,7 +126,7 @@ class TestThlPayoutEventManager: assert sum([i.amount for i in res]) == sum(amounts) @pytest.mark.skip - def test_get_payout_detail(self, user_payout_event_manager): + def test_get_payout_detail(self, user_payout_event_manager: UserPayoutEventManager): """This fails because the description coming back is None, but then it tries to return a PayoutEvent which validates that the description can't be None @@ -141,7 +157,7 @@ class TestThlPayoutEventManager: # def test_filter_by(self): # raise NotImplementedError - def test_create(self, user_payout_event_manager): + def test_create(self, user_payout_event_manager: UserPayoutEventManager): from generalresearch.models.thl.payout import UserPayoutEvent # Confirm the creation method returns back an instance. @@ -163,26 +179,30 @@ class TestThlPayoutEventManager: def test_create_bp_payout( self, - product, - delete_ledger_db, - create_main_accounts, - thl_lm, - brokerage_product_payout_event_manager, - lm, + product: Product, + delete_ledger_db: Callable[..., None], + create_main_accounts: Callable[..., None], + thl_ledger_manager: ThlLedgerManager, + brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager, + ledger_manager: LedgerManager, ): delete_ledger_db() create_main_accounts() - account_bp_wallet = thl_lm.get_account_or_create_bp_wallet(product=product) - brokerage_product_payout_event_manager.set_account_lookup_table(thl_lm=thl_lm) + account_bp_wallet = thl_ledger_manager.get_account_or_create_bp_wallet( + product=product + ) + brokerage_product_payout_event_manager.set_account_lookup_table( + thl_lm=thl_ledger_manager + ) rand_amount = randint(a=99, b=999) # Save a Brokerage Product Payout, so we have something in the # Payout Event table and the respective ledger TX and Entry rows for it pe = brokerage_product_payout_event_manager.create_bp_payout_event( - thl_ledger_manager=thl_lm, + thl_ledger_manager=thl_ledger_manager, product=product, amount=USDCent(rand_amount), skip_wallet_balance_check=True, @@ -191,15 +211,17 @@ class TestThlPayoutEventManager: assert isinstance(pe, BrokerageProductPayoutEvent) # Now try to query for it! - res = thl_lm.get_tx_bp_payouts(account_uuids=[account_bp_wallet.uuid]) + res = thl_ledger_manager.get_tx_bp_payouts( + account_uuids=[account_bp_wallet.uuid] + ) assert len(res) == 1 - res = thl_lm.get_tx_bp_payouts(account_uuids=[uuid4().hex]) + res = thl_ledger_manager.get_tx_bp_payouts(account_uuids=[uuid4().hex]) assert len(res) == 0 # Confirm it added to the users balance. The amount is negative because - # money was sent to the Brokerage Product, but they didn't have + # money was sent to the Brokerage product: Product, but they didn't have # any activity that earned them money - bal = lm.get_account_balance(account=account_bp_wallet) + bal = ledger_manager.get_account_balance(account=account_bp_wallet) assert rand_amount == bal * -1 @@ -207,13 +229,13 @@ class TestBPPayoutEvent: def test_get_bp_bp_payout_events_for_products( self, - product_factory, - bp_payout_event_factory, - usd_cent, - delete_ledger_db, - create_main_accounts, - brokerage_product_payout_event_manager, - thl_lm, + product_factory: Callable[..., Product], + bp_payout_event_factory: Callable[..., BrokerageProductPayoutEvent], + usd_cent: USDCent, + delete_ledger_db: Callable[..., None], + create_main_accounts: Callable[..., None], + brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager, + thl_ledger_manager: ThlLedgerManager, ): delete_ledger_db() create_main_accounts() @@ -222,10 +244,12 @@ class TestBPPayoutEvent: amounts = [] product: Product = product_factory() - thl_lm.get_account_or_create_bp_wallet(product=product) - brokerage_product_payout_event_manager.set_account_lookup_table(thl_lm=thl_lm) + thl_ledger_manager.get_account_or_create_bp_wallet(product=product) + brokerage_product_payout_event_manager.set_account_lookup_table( + thl_lm=thl_ledger_manager + ) - for y_idx in range(N_PAYOUT_EVENTS): + for _ in range(N_PAYOUT_EVENTS): bp_payout_event_factory(product=product, usd_cent=usd_cent) amounts.append(usd_cent) @@ -233,7 +257,7 @@ class TestBPPayoutEvent: # array of BPPayoutEvents bp_bp_res = ( brokerage_product_payout_event_manager.get_bp_bp_payout_events_for_products( - thl_ledger_manager=thl_lm, product_uuids=[product.uuid] + thl_ledger_manager=thl_ledger_manager, product_uuids=[product.uuid] ) ) assert isinstance(bp_bp_res, list) diff --git a/tests/managers/thl/test_ledger/test_user_txs.py b/tests/managers/thl/test_ledger/test_user_txs.py index a6bfa79..56dc485 100644 --- a/tests/managers/thl/test_ledger/test_user_txs.py +++ b/tests/managers/thl/test_ledger/test_user_txs.py @@ -1,11 +1,14 @@ +from __future__ import annotations + from collections.abc import Callable from datetime import UTC, datetime from decimal import Decimal -from typing import TYPE_CHECKING from uuid import uuid4 from generalresearch.config import GRLBaseSettings +from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager +from generalresearch.managers.thl.payout import UserPayoutEventManager from generalresearch.managers.thl.user_compensate import user_compensate from generalresearch.models.thl.definitions import ( Status, @@ -15,41 +18,38 @@ from generalresearch.models.thl.ledger import ( UserLedgerTransactionTypesSummary, UserLedgerTransactionTypeSummary, ) - -if TYPE_CHECKING: - from generalresearch.config import GRLSettings - from generalresearch.models.thl.product import Product - from generalresearch.models.thl.session import Session - from generalresearch.models.thl.user import User - from generalresearch.models.thl.wallet import PayoutType +from generalresearch.models.thl.product import Product +from generalresearch.models.thl.session import Session +from generalresearch.models.thl.user import User +from generalresearch.models.thl.wallet import PayoutType def test_user_txs( user_factory: Callable[..., User], product_amt_true: Product, create_main_accounts: Callable[..., None], - thl_lm: ThlLedgerManager, - lm, + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, delete_ledger_db: Callable[..., None], - session_with_tx_factory, - adj_to_fail_with_tx_factory, - adj_to_complete_with_tx_factory, - session_factory, - user_payout_event_manager, + session_with_tx_factory: Callable[..., Session], + adj_to_fail_with_tx_factory: Callable[..., None], + adj_to_complete_with_tx_factory: Callable[..., None], + session_factory: Callable[..., Session], + user_payout_event_manager: UserPayoutEventManager, utc_now: datetime, - settings: GRLSettings, + settings: GRLBaseSettings, ): delete_ledger_db() create_main_accounts() user: User = user_factory(product=product_amt_true) - account = thl_lm.get_account_or_create_user_wallet(user) + account = thl_ledger_manager.get_account_or_create_user_wallet(user) print(f"{account.uuid=}") s: Session = session_with_tx_factory(user=user, wall_req_cpi=Decimal("1.00")) user_compensate( - ledger_manager=thl_lm, + ledger_manager=ledger_manager, user=user, amount_int=100, ) @@ -63,7 +63,7 @@ def test_user_txs( payout_type=PayoutType.AMT_HIT, request_data={}, ) - thl_lm.create_tx_user_payout_request( + thl_ledger_manager.create_tx_user_payout_request( user=user, payout_event=pe, ) @@ -76,7 +76,7 @@ def test_user_txs( payout_type=PayoutType.AMT_BONUS, request_data={}, ) - thl_lm.create_tx_user_payout_request( + thl_ledger_manager.create_tx_user_payout_request( user=user, payout_event=pe, ) @@ -93,16 +93,16 @@ def test_user_txs( ) adj_to_complete_with_tx_factory(session=s_fail, created=utc_now) - # txs = thl_lm.get_tx_filtered_by_account(account.uuid) + # txs = thl_ledger_manager.get_tx_filtered_by_account(account.uuid) # print(len(txs), txs) - txs = thl_lm.get_user_txs(user) + txs = thl_ledger_manager.get_user_txs(user) assert len(txs.transactions) == 6 assert txs.total == 6 assert txs.page == 1 assert txs.size == 50 # print(len(txs.transactions), txs) - d = txs.model_dump_json() + # d = txs.model_dump_json() # print(d) descriptions = {x.description for x in txs.transactions} @@ -140,25 +140,26 @@ def test_user_txs_pagination( user_factory: Callable[..., User], product_amt_true: Product, create_main_accounts: Callable[..., None], - thl_lm: ThlLedgerManager, + ledger_manager: LedgerManager, + thl_ledger_manager: ThlLedgerManager, delete_ledger_db: Callable[..., None], ): delete_ledger_db() create_main_accounts() user: User = user_factory(product=product_amt_true) - account = thl_lm.get_account_or_create_user_wallet(user) + account = thl_ledger_manager.get_account_or_create_user_wallet(user) print(f"{account.uuid=}") for _ in range(12): user_compensate( - ledger_manager=thl_lm, + ledger_manager=ledger_manager, user=user, amount_int=100, skip_flag_check=True, ) - txs = thl_lm.get_user_txs(user, page=1, size=5) + txs = thl_ledger_manager.get_user_txs(user, page=1, size=5) assert len(txs.transactions) == 5 assert txs.total == 12 assert txs.page == 1 @@ -167,7 +168,7 @@ def test_user_txs_pagination( assert txs.summary.user_bonus.entry_count == 12 # Skip to the 3rd page. We made 12, so there are 2 left - txs = thl_lm.get_user_txs(user, page=3, size=5) + txs = thl_ledger_manager.get_user_txs(user, page=3, size=5) assert len(txs.transactions) == 2 assert txs.total == 12 assert txs.page == 3 @@ -175,7 +176,7 @@ def test_user_txs_pagination( assert txs.summary.user_bonus.entry_count == 12 # Should be empty, not fail - txs = thl_lm.get_user_txs(user, page=4, size=5) + txs = thl_ledger_manager.get_user_txs(user, page=4, size=5) assert len(txs.transactions) == 0 assert txs.total == 12 assert txs.page == 4 @@ -185,12 +186,12 @@ def test_user_txs_pagination( # Test filtering. We should pull back only this one now = datetime.now(tz=UTC) user_compensate( - ledger_manager=thl_lm, + ledger_manager=ledger_manager, user=user, amount_int=100, skip_flag_check=True, ) - txs = thl_lm.get_user_txs(user, page=1, size=5, time_start=now) + txs = thl_ledger_manager.get_user_txs(user, page=1, size=5, time_start=now) assert len(txs.transactions) == 1 assert txs.total == 1 assert txs.page == 1 @@ -200,7 +201,7 @@ def test_user_txs_pagination( # And filtering with 0 results now = datetime.now(tz=UTC) - txs = thl_lm.get_user_txs(user, page=1, size=5, time_start=now) + txs = thl_ledger_manager.get_user_txs(user, page=1, size=5, time_start=now) assert len(txs.transactions) == 0 assert txs.total == 0 assert txs.page == 1 @@ -213,13 +214,10 @@ def test_user_txs_pagination( def test_user_txs_rolling_balance( user_factory: Callable[..., User], product_amt_true: Product, - create_main_accounts, + create_main_accounts: Callable[..., None], thl_ledger_manager: ThlLedgerManager, - ledger_manager: LedgerManager, delete_ledger_db: Callable[..., None], - session_with_tx_factory, - adj_to_fail_with_tx_factory, - user_payout_event_manager, + user_payout_event_manager: UserPayoutEventManager, settings: GRLBaseSettings, ): """ diff --git a/tests/managers/thl/test_ledger/test_wallet.py b/tests/managers/thl/test_ledger/test_wallet.py index bb49cd8..9e886db 100644 --- a/tests/managers/thl/test_ledger/test_wallet.py +++ b/tests/managers/thl/test_ledger/test_wallet.py @@ -1,19 +1,26 @@ +from __future__ import annotations + +from collections.abc import Callable from decimal import Decimal 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 ( PayoutConfig, PayoutTransformation, PayoutTransformationPercentArgs, + Product, UserWalletConfig, ) from generalresearch.models.thl.user import User @pytest.fixture() -def schrute_product(product_manager): +def schrute_product(product_manager: ProductManager) -> Product: return product_manager.create_dummy( user_wallet_config=UserWalletConfig(enabled=True, amt=False), payout_config=PayoutConfig( @@ -27,25 +34,30 @@ def schrute_product(product_manager): class TestGetUserWalletBalance: - def test_get_user_wallet_balance_non_managed(self, user, thl_lm): + def test_get_user_wallet_balance_non_managed( + self, user: User, thl_ledger_manager: ThlLedgerManager + ): with pytest.raises( AssertionError, match="Can't get wallet balance on non-managed account.", ): - thl_lm.get_user_wallet_balance(user=user) + thl_ledger_manager.get_user_wallet_balance(user=user) def test_get_user_wallet_balance_managed_0( - self, schrute_product, user_factory, thl_lm + self, + schrute_product: Product, + user_factory: Callable[..., User], + thl_ledger_manager: ThlLedgerManager, ): assert ( schrute_product.payout_config.payout_format == "{payout:,.0f} Schrute Bucks" ) user: User = user_factory(schrute_product) - balance = thl_lm.get_user_wallet_balance(user=user) + balance = thl_ledger_manager.get_user_wallet_balance(user=user) assert balance == 0 balance_string = user.product.format_payout_format(Decimal(balance) / 100) assert balance_string == "0 Schrute Bucks" - redeemable_balance = thl_lm.get_user_redeemable_wallet_balance( + redeemable_balance = thl_ledger_manager.get_user_redeemable_wallet_balance( user=user, user_wallet_balance=balance ) assert redeemable_balance == 0 @@ -55,10 +67,14 @@ class TestGetUserWalletBalance: assert redeemable_balance_string == "0 Schrute Bucks" def test_get_user_wallet_balance_managed( - self, schrute_product, user_factory, thl_lm, session_with_tx_factory + self, + schrute_product: Product, + user_factory: Callable[..., User], + thl_ledger_manager: ThlLedgerManager, + session_with_tx_factory: Callable[..., None], ): user: User = user_factory(schrute_product) - thl_lm.create_tx_user_bonus( + thl_ledger_manager.create_tx_user_bonus( user=user, amount=Decimal(1), ref_uuid=uuid4().hex, @@ -69,10 +85,10 @@ class TestGetUserWalletBalance: # This product has a payout xform of 40% and commission of 5% # 1.23 * 0.05 = 0.06 of commission # 1.17 of payout * 0.40 = 0.47 of user pay and (1.17-0.47) 0.70 bp pay - balance = thl_lm.get_user_wallet_balance(user=user) + balance = thl_ledger_manager.get_user_wallet_balance(user=user) assert balance == 47 + 100 # plus the $1 bribe - redeemable_balance = thl_lm.get_user_redeemable_wallet_balance( + redeemable_balance = thl_ledger_manager.get_user_redeemable_wallet_balance( user=user, user_wallet_balance=balance ) assert redeemable_balance == 20 + 100 diff --git a/tests/managers/thl/test_maxmind.py b/tests/managers/thl/test_maxmind.py index 3e85cc3..e44fe49 100644 --- a/tests/managers/thl/test_maxmind.py +++ b/tests/managers/thl/test_maxmind.py @@ -70,8 +70,8 @@ IP_v6_US_SAME_64 = "2600:1700:ece0:9410:55d:faf3:c15d:aaaa" # class TestMaxmindManager: -# def test_init(self, thl_web_rr, thl_redis_config, maxmind_manager: MaxmindManager): -# instance = MaxmindManager(pg_config=thl_web_rr, redis_config=thl_redis_config) +# def test_init(self, thl_web_rr: PostgresConfig, thl_redis_config: RedisConfig, maxmind_manager: MaxmindManager): +# instance = MaxmindManager(pg_config=thl_web_rr: PostgresConfig, redis_config=thl_redis_config) # assert isinstance(instance, MaxmindManager) # assert isinstance(maxmind_manager, MaxmindManager) diff --git a/tests/managers/thl/test_payout.py b/tests/managers/thl/test_payout.py index 39bbe6b..153bee9 100644 --- a/tests/managers/thl/test_payout.py +++ b/tests/managers/thl/test_payout.py @@ -2,7 +2,6 @@ import io import logging import os from collections.abc import Callable -from dask.distributed import Client as DaskClient from datetime import UTC, datetime, timedelta from decimal import Decimal from random import choice as rand_choice @@ -11,12 +10,26 @@ from uuid import uuid4 import pandas as pd import pytest +from dask.distributed import Client as DaskClient from generalresearch.currency import USDCent +from generalresearch.incite.base import GRLDatasets +from generalresearch.incite.collections.thl_web import ( + LedgerDFCollection, +) +from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager -from generalresearch.managers.thl.payout import UserPayoutEventManager +from generalresearch.managers.thl.payout import ( + BrokerageProductPayoutEventManager, + BusinessPayoutEventManager, + PayoutEventManager, + UserPayoutEventManager, +) +from generalresearch.managers.thl.product import ProductManager +from generalresearch.models.gr.business import Business from generalresearch.models.thl.definitions import PayoutStatus +from generalresearch.models.thl.finance import BusinessBalances from generalresearch.models.thl.ledger import LedgerAccount from generalresearch.models.thl.payout import ( BrokerageProductPayoutEvent, @@ -24,10 +37,11 @@ from generalresearch.models.thl.payout import ( UserPayoutEvent, ) from generalresearch.models.thl.product import Product -from generalresearch.models.gr.business import Business +from generalresearch.models.thl.session import Session from generalresearch.models.thl.user import User from generalresearch.models.thl.wallet import PayoutType from generalresearch.pg_helper import PostgresConfig +from generalresearch.redis_helper import RedisConfig logger = logging.getLogger() @@ -66,14 +80,11 @@ class TestPayout: def test_update( self, user: User, - user_payout_event_manager, + user_payout_event_manager: UserPayoutEventManager, ledger_manager: LedgerManager, thl_ledger_manager: ThlLedgerManager, utc_now: datetime, ): - from generalresearch.models.thl.definitions import PayoutStatus - from generalresearch.models.thl.wallet import PayoutType - user_account = thl_ledger_manager.get_account_or_create_user_wallet(user=user) pe1 = user_payout_event_manager.create( @@ -113,7 +124,7 @@ class TestPayout: thl_web_rw: PostgresConfig, product: Product, thl_lm: ThlLedgerManager, - brokerage_product_payout_event_manager, + brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager, utc_now: datetime, ) -> BrokerageProductPayoutEvent: account = thl_lm.get_account_or_create_bp_wallet(product=product) @@ -144,11 +155,11 @@ class TestPayout: def test_create_bp_payout_quick_dupe( self, product: Product, - brokerage_product_payout_event_manager, + brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager, thl_lm: ThlLedgerManager, - ledger_manager, + ledger_manager: LedgerManager, utc_now: datetime, - pending_bp_pe, + pending_bp_pe: BrokerageProductPayoutEvent, ): thl_lm.get_account_or_create_bp_wallet(product=product) @@ -171,10 +182,10 @@ class TestPayout: def test_filter( self, thl_ledger_manager: ThlLedgerManager, - ledger_manager, + ledger_manager: LedgerManager, product: Product, user: User, - user_payout_event_manager, + user_payout_event_manager: UserPayoutEventManager, utc_now: datetime, ): from generalresearch.models.thl.definitions import PayoutStatus @@ -264,14 +275,14 @@ class TestBusinessPayoutEventManager: def test_base( self, - brokerage_product_payout_event_manager, - business_payout_event_manager, + brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager, + business_payout_event_manager: BusinessPayoutEventManager, delete_ledger_db: Callable[..., None], - create_main_accounts, + create_main_accounts: Callable[..., None], thl_ledger_manager: ThlLedgerManager, product_factory: Callable[..., Product], - bp_payout_factory, - business, + bp_payout_factory: Callable[..., BrokerageProductPayoutEvent], + business: Business, ): delete_ledger_db() create_main_accounts() @@ -298,6 +309,7 @@ class TestBusinessPayoutEventManager: bpem=business_payout_event_manager, ) + assert isinstance(business.payouts, list) assert len(business.payouts) == 3 assert business.payouts_total == sum([pe.amount for pe in business.payouts]) assert business.payouts[0].created > business.payouts[1].created @@ -313,20 +325,20 @@ class TestBusinessPayoutEventManager: def test_update_ext_reference_ids( self, - business_payout_event_manager, - delete_ledger_db: Callable[..., None],, + business_payout_event_manager: BusinessPayoutEventManager, + delete_ledger_db: Callable[..., None], create_main_accounts: Callable[..., None], - thl_ledger_manager, + thl_ledger_manager: ThlLedgerManager, thl_web_rr: PostgresConfig, product_factory: Callable[..., Product], - delete_df_collection, - user_factory, - ledger_collection, - session_with_tx_factory, - pop_ledger_merge, - client_no_amm, - mnt_filepath: GRLDataset, - product_manager, + delete_df_collection: Callable[..., None], + user_factory: Callable[..., User], + ledger_collection: LedgerDFCollection, + session_with_tx_factory: Callable[..., Session], + pop_ledger_merge: PopLedgerMerge, + client_no_amm: DaskClient, + mnt_filepath: GRLDatasets, + product_manager: ProductManager, start: datetime, business: Business, ): @@ -392,8 +404,10 @@ class TestBusinessPayoutEventManager: assert business_payout_event_manager.get_by_ext_ref_id(ext_ref_id=ach_id2) - def test_recoup_empty(self, business_payout_event_manager): - res = {uuid4().hex: USDCent(0) for i in range(100)} + def test_recoup_empty( + self, business_payout_event_manager: BusinessPayoutEventManager + ): + res = {uuid4().hex: USDCent(0) for _ in range(100)} df = pd.DataFrame.from_dict(res, orient="index").reset_index() df.columns = ["product_id", "available_balance"] @@ -403,10 +417,12 @@ class TestBusinessPayoutEventManager: ) assert "Total available amount is empty, cannot recoup" in str(cm) - def test_recoup_exceeds(self, business_payout_event_manager): + def test_recoup_exceeds( + self, business_payout_event_manager: BusinessPayoutEventManager + ): from random import randint - res = {uuid4().hex: USDCent(randint(a=0, b=1_000_00)) for i in range(100)} + res = {uuid4().hex: USDCent(randint(a=0, b=1_000_00)) for _ in range(100)} df = pd.DataFrame.from_dict(res, orient="index").reset_index() df.columns = ["product_id", "available_balance"] @@ -418,10 +434,10 @@ class TestBusinessPayoutEventManager: ) assert " exceeds total available " in str(cm) - def test_recoup(self, business_payout_event_manager): + def test_recoup(self, business_payout_event_manager: BusinessPayoutEventManager): from random import randint, random - res = {uuid4().hex: USDCent(randint(a=0, b=1_000_00)) for i in range(100)} + res = {uuid4().hex: USDCent(randint(a=0, b=1_000_00)) for _ in range(100)} df = pd.DataFrame.from_dict(res, orient="index").reset_index() df.columns = ["product_id", "available_balance"] @@ -437,7 +453,9 @@ class TestBusinessPayoutEventManager: assert res.deduction.sum() == random_recoup_amount assert res.remaining_balance.sum() == avail_balance - random_recoup_amount - def test_recoup_loop(self, business_payout_event_manager, request): + def test_recoup_loop( + self, business_payout_event_manager: BusinessPayoutEventManager, request + ): # TODO: Generate this file at random fp = os.path.join( request.config.rootpath, "data/pytest_recoup_proportional.csv" @@ -451,9 +469,11 @@ class TestBusinessPayoutEventManager: assert int(res.deduction.sum()) == 1416089 - def test_recoup_loop_single_profitable_account(self, business_payout_event_manager): - res = [{"product_id": uuid4().hex, "available_balance": 0} for i in range(1000)] - for x in range(100): + def test_recoup_loop_single_profitable_account( + self, business_payout_event_manager: BusinessPayoutEventManager + ): + res = [{"product_id": uuid4().hex, "available_balance": 0} for _ in range(1000)] + for _ in range(100): item = rand_choice(res) item["available_balance"] = randint(8, 12) @@ -464,14 +484,16 @@ class TestBusinessPayoutEventManager: # res = res[res["remaining_balance"] > 0] assert int(res.deduction.sum()) == 500 - def test_recoup_loop_assertions(self, business_payout_event_manager): + def test_recoup_loop_assertions( + self, business_payout_event_manager: BusinessPayoutEventManager + ): df = pd.DataFrame( [ { "product_id": uuid4().hex, "available_balance": randint(0, 999_999), } - for i in range(10_000) + for _ in range(10_000) ] ) available_balance = int(df.available_balance.sum()) @@ -501,7 +523,9 @@ class TestBusinessPayoutEventManager: assert res.remaining_balance.sum() == available_balance assert int(res.deduction.sum()) == 0 - def test_distribute_amount(self, business_payout_event_manager): + def test_distribute_amount( + self, business_payout_event_manager: BusinessPayoutEventManager + ): df = pd.read_csv( io.StringIO( @@ -517,29 +541,29 @@ class TestBusinessPayoutEventManager: def test_ach_payment_min_amount( self, - product, - mnt_filepath, - thl_lm, - client_no_amm, - thl_redis_config, - payout_event_manager, - brokerage_product_payout_event_manager, - business_payout_event_manager, - delete_ledger_db, - create_main_accounts, - delete_df_collection, - ledger_collection, - business, - user_factory, - product_factory, - session_with_tx_factory, - pop_ledger_merge, - start, - bp_payout_factory, - adj_to_fail_with_tx_factory, - thl_web_rr, - lm, - product_manager, + 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], + delete_df_collection: Callable[..., None], + ledger_collection: LedgerDFCollection, + business: Business, + user_factory: Callable[..., User], + product_factory: Callable[..., Product], + session_with_tx_factory: Callable[..., Session], + 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, ): """Test having a Business with three products. One that lost money and two that gained money. Ensure that the Business balance @@ -554,7 +578,7 @@ class TestBusinessPayoutEventManager: p1: Product = product_factory(business=business) u1: User = user_factory(product=p1) - thl_lm.get_account_or_create_bp_wallet(product=p1) + thl_ledger_manager.get_account_or_create_bp_wallet(product=p1) session_with_tx_factory( user=u1, @@ -576,7 +600,7 @@ class TestBusinessPayoutEventManager: pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) business.prebuild_balance( thl_pg_config=thl_web_rr, - lm=lm, + lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, pop_ledger=pop_ledger_merge, @@ -587,38 +611,30 @@ class TestBusinessPayoutEventManager: business=business, amount=USDCent(500), pm=product_manager, - thl_lm=thl_lm, + thl_lm=thl_ledger_manager, transaction_id=uuid4().hex, ) assert "Must issue Supplier Payouts at least $100 minimum." in str(cm) def test_create_from_ach_or_wire( self, - product, - mnt_filepath, - thl_lm, - client_no_amm, - thl_redis_config, - payout_event_manager, - brokerage_product_payout_event_manager, - business_payout_event_manager, - delete_ledger_db, - create_main_accounts, - delete_df_collection, - ledger_collection, - business, - user_factory, - product_factory, - session_with_tx_factory, - pop_ledger_merge, - start, - bp_payout_factory, - adj_to_fail_with_tx_factory, - thl_web_rr, - lm, - product_manager, - rm_ledger_collection, - rm_pop_ledger_merge, + mnt_filepath: GRLDatasets, + thl_ledger_manager: ThlLedgerManager, + client_no_amm: DaskClient, + business_payout_event_manager: BusinessPayoutEventManager, + delete_ledger_db: Callable[..., None], + create_main_accounts: Callable[..., None], + delete_df_collection: Callable[..., None], + ledger_collection: LedgerDFCollection, + business: Business, + user_factory: Callable[..., User], + product_factory: Callable[..., Product], + session_with_tx_factory: Callable[..., None], + pop_ledger_merge: PopLedgerMerge, + start: datetime, + thl_web_rr: PostgresConfig, + ledger_manager: LedgerManager, + product_manager: ProductManager, caplog, ): """Test having a Business with three products""" @@ -630,12 +646,12 @@ class TestBusinessPayoutEventManager: p1: Product = product_factory(business=business) p2: Product = product_factory(business=business) p3: Product = product_factory(business=business) - u1: User = user_factory(product=p1) + _: User = user_factory(product=p1) u2: User = user_factory(product=p2) u3: User = user_factory(product=p3) - thl_lm.get_account_or_create_bp_wallet(product=p1) - thl_lm.get_account_or_create_bp_wallet(product=p2) - thl_lm.get_account_or_create_bp_wallet(product=p3) + thl_ledger_manager.get_account_or_create_bp_wallet(product=p1) + thl_ledger_manager.get_account_or_create_bp_wallet(product=p2) + thl_ledger_manager.get_account_or_create_bp_wallet(product=p3) ach_id1 = uuid4().hex ach_id2 = uuid4().hex @@ -660,13 +676,14 @@ class TestBusinessPayoutEventManager: pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) business.prebuild_balance( thl_pg_config=thl_web_rr, - lm=lm, + lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, pop_ledger=pop_ledger_merge, ) bb = business.balance + assert isinstance(bb, BusinessBalances) assert bb.payout == 475_00 # $500 * .95% = $475 assert bb.net == 475_00 @@ -674,7 +691,7 @@ class TestBusinessPayoutEventManager: business=business, amount=USDCent(100_00), pm=product_manager, - thl_lm=thl_lm, + thl_lm=thl_ledger_manager, created=start + timedelta(days=1, hours=5), transaction_id=ach_id1, ) @@ -686,7 +703,7 @@ class TestBusinessPayoutEventManager: business=business, amount=USDCent(bb.available_balance), pm=product_manager, - thl_lm=thl_lm, + thl_lm=thl_ledger_manager, created=start + timedelta(days=2, hours=5), transaction_id=ach_id2, ) @@ -696,7 +713,7 @@ class TestBusinessPayoutEventManager: with caplog.at_level(logging.WARNING): business_payout_event_manager.resume_failed_business_payout( - ext_ref_id=ach_id1, thl_lm=thl_lm, pm=product_manager + ext_ref_id=ach_id1, thl_lm=thl_ledger_manager, pm=product_manager ) assert "Nothing to do!" in caplog.text @@ -714,31 +731,31 @@ class TestBusinessPayoutEventManager: def test_ach_payment( self, - product, - mnt_filepath, - thl_lm, - client_no_amm, - thl_redis_config, - payout_event_manager, - brokerage_product_payout_event_manager, - business_payout_event_manager, - delete_ledger_db, - create_main_accounts, - delete_df_collection, - ledger_collection, - business, - user_factory, - product_factory, - session_with_tx_factory, - pop_ledger_merge, - start, - bp_payout_factory, - adj_to_fail_with_tx_factory, - thl_web_rr, - lm, - product_manager, - rm_ledger_collection, - rm_pop_ledger_merge, + 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], + delete_df_collection: Callable[..., None], + ledger_collection: LedgerDFCollection, + business: Business, + user_factory: Callable[..., User], + product_factory: Callable[..., Product], + 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, + rm_ledger_collection: Callable[..., None], + rm_pop_ledger_merge: Callable[..., None], ): """Test having a Business with three products.. one that lost money and two that gained money. Ensure that the Business balance @@ -757,9 +774,9 @@ class TestBusinessPayoutEventManager: u1: User = user_factory(product=p1) u2: User = user_factory(product=p2) u3: User = user_factory(product=p3) - thl_lm.get_account_or_create_bp_wallet(product=p1) - thl_lm.get_account_or_create_bp_wallet(product=p2) - thl_lm.get_account_or_create_bp_wallet(product=p3) + thl_ledger_manager.get_account_or_create_bp_wallet(product=p1) + thl_ledger_manager.get_account_or_create_bp_wallet(product=p2) + thl_ledger_manager.get_account_or_create_bp_wallet(product=p3) ach_id1 = uuid4().hex @@ -796,18 +813,19 @@ class TestBusinessPayoutEventManager: started=start + timedelta(days=1, hours=3, minutes=1 + idx), ) - # Now that we paid out the business, let's confirm the updated balances + # Now that we paid out the business: Business, let's confirm the updated balances ledger_collection.initial_load(client=None, sync=True) pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) business.prebuild_balance( thl_pg_config=thl_web_rr, - lm=lm, + lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, pop_ledger=pop_ledger_merge, ) bb1 = business.balance + assert isinstance(bb1, BusinessBalances) pb1 = bb1.product_balances[0] pb2 = bb1.product_balances[1] pb3 = bb1.product_balances[2] @@ -833,9 +851,10 @@ class TestBusinessPayoutEventManager: assert business.payouts is None business.prebuild_payouts( thl_pg_config=thl_web_rr, - thl_lm=thl_lm, + thl_lm=thl_ledger_manager, bpem=business_payout_event_manager, ) + assert isinstance(business.payouts, list) assert len(business.payouts) == 1 assert business.payouts[0].ext_ref_id == ach_id1 @@ -843,7 +862,7 @@ class TestBusinessPayoutEventManager: business=business, amount=USDCent(bb1.available_balance), pm=product_manager, - thl_lm=thl_lm, + thl_lm=thl_ledger_manager, created=start + timedelta(days=1, hours=5), ) assert isinstance(bp1, BusinessPayoutEvent) @@ -851,7 +870,7 @@ class TestBusinessPayoutEventManager: assert bp1.bp_payouts[0].status == PayoutStatus.COMPLETE assert bp1.bp_payouts[1].status == PayoutStatus.COMPLETE bp1_tx = brokerage_product_payout_event_manager.check_for_ledger_tx( - thl_ledger_manager=thl_lm, + thl_ledger_manager=thl_ledger_manager, payout_event=bp1.bp_payouts[0], product_id=bp1.bp_payouts[0].product_id, amount=bp1.bp_payouts[0].amount, @@ -859,14 +878,14 @@ class TestBusinessPayoutEventManager: assert bp1_tx bp2_tx = brokerage_product_payout_event_manager.check_for_ledger_tx( - thl_ledger_manager=thl_lm, + thl_ledger_manager=thl_ledger_manager, payout_event=bp1.bp_payouts[1], product_id=bp1.bp_payouts[1].product_id, amount=bp1.bp_payouts[1].amount, ) assert bp2_tx - # Now that we paid out the business, let's confirm the updated balances + # Now that we paid out the business: Business, let's confirm the updated balances rm_ledger_collection() rm_pop_ledger_merge() ledger_collection.initial_load(client=None, sync=True) @@ -874,16 +893,17 @@ class TestBusinessPayoutEventManager: business.prebuild_balance( thl_pg_config=thl_web_rr, - lm=lm, + lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, pop_ledger=pop_ledger_merge, ) business.prebuild_payouts( thl_pg_config=thl_web_rr, - thl_lm=thl_lm, + thl_lm=thl_ledger_manager, bpem=business_payout_event_manager, ) + assert isinstance(business.payouts, list) assert len(business.payouts) == 2 assert len(business.payouts[0].bp_payouts) == 2 assert len(business.payouts[1].bp_payouts) == 1 @@ -892,6 +912,8 @@ class TestBusinessPayoutEventManager: # Okay os we have the balance before, and after the Business Payout # of bb1.available_balance worth.. + assert isinstance(bb1, BusinessBalances) + assert isinstance(bb2, BusinessBalances) assert bb1.payout == bb2.payout assert bb1.adjustment == bb2.adjustment assert bb1.net == bb2.net @@ -908,34 +930,34 @@ class TestBusinessPayoutEventManager: def test_ach_payment_partial_amount( self, - product, - mnt_filepath, + product: Product, + mnt_filepath: GRLDatasets, thl_ledger_manager: ThlLedgerManager, client_no_amm: DaskClient, - thl_redis_config, - payout_event_manager, - brokerage_product_payout_event_manager, - business_payout_event_manager, - delete_ledger_db, - create_main_accounts, - delete_df_collection, - ledger_collection, - business, - user_factory, - product_factory, - session_with_tx_factory, - pop_ledger_merge, - start, - bp_payout_factory, - adj_to_fail_with_tx_factory, - thl_web_rr, - ledger_manager, - product_manager, - rm_ledger_collection, - rm_pop_ledger_merge, + 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], + delete_df_collection: Callable[..., None], + ledger_collection: LedgerDFCollection, + business: Business, + user_factory: Callable[..., User], + product_factory: Callable[..., Product], + 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, + rm_ledger_collection: Callable[..., None], + rm_pop_ledger_merge: Callable[..., None], ): """There are valid instances when we want issue a ACH or Wire to a - Business, but not for the full Available Balance amount in their + business: Business, but not for the full Available Balance amount in their account. To test this, we'll create a Business with multiple Products, and @@ -956,9 +978,9 @@ class TestBusinessPayoutEventManager: u1: User = user_factory(product=p1) u2: User = user_factory(product=p2) u3: User = user_factory(product=p3) - thl_lm.get_account_or_create_bp_wallet(product=p1) - thl_lm.get_account_or_create_bp_wallet(product=p2) - thl_lm.get_account_or_create_bp_wallet(product=p3) + thl_ledger_manager.get_account_or_create_bp_wallet(product=p1) + thl_ledger_manager.get_account_or_create_bp_wallet(product=p2) + thl_ledger_manager.get_account_or_create_bp_wallet(product=p3) # Product 1, 2, 3: Complete, and Payout multiple times. for idx in range(5): @@ -968,9 +990,9 @@ class TestBusinessPayoutEventManager: wall_req_cpi=Decimal("50.00"), started=start + timedelta(days=1, hours=2, minutes=1 + idx), ) - payout_event_manager.set_account_lookup_table(thl_lm=thl_lm) + payout_event_manager.set_account_lookup_table(thl_lm=thl_ledger_manager) - # Now that we paid out the business, let's confirm the updated balances + # Now that we paid out the business: Business, let's confirm the updated balances ledger_collection.initial_load(client=None, sync=True) pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) business.prebuild_balance( @@ -987,6 +1009,8 @@ class TestBusinessPayoutEventManager: # Confirm the initial amounts. assert len(business.payouts) == 0 bb1 = business.balance + + assert isinstance(bb1, BusinessBalances) assert bb1.payout == 3 * 5 * 4750 assert bb1.adjustment == 0 assert bb1.payout == bb1.net @@ -999,6 +1023,7 @@ class TestBusinessPayoutEventManager: assert bb1.product_balances[x].available_balance_usd_str == "$178.13" assert business.payouts_total_str == "$0.00" + assert isinstance(business.balance, BusinessBalances) assert business.balance.payment_usd_str == "$0.00" assert business.balance.available_balance_usd_str == "$534.39" @@ -1009,13 +1034,13 @@ class TestBusinessPayoutEventManager: business=business, amount=USDCent(250_00), pm=product_manager, - thl_lm=thl_lm, + thl_lm=thl_ledger_manager, created=start + timedelta(days=1, hours=3), ) assert isinstance(bp1, BusinessPayoutEvent) assert len(bp1.bp_payouts) == 3 - # Now that we paid out the business, let's confirm the updated + # Now that we paid out the business: Business, let's confirm the updated # balances. Clear and rebuild the parquet files. rm_ledger_collection() rm_pop_ledger_merge() @@ -1034,40 +1059,40 @@ class TestBusinessPayoutEventManager: ) business.prebuild_payouts( thl_pg_config=thl_web_rr, - thl_lm=thl_lm, + thl_lm=thl_ledger_manager, bpem=business_payout_event_manager, ) + assert isinstance(business.payouts, list) assert len(business.payouts) == 1 assert len(business.payouts[0].bp_payouts) == 3 assert business.payouts_total_str == "$250.00" + assert isinstance(business.balance, BusinessBalances) assert business.balance.payment_usd_str == "$250.00" assert business.balance.available_balance_usd_str == "$346.88" def test_ach_tx_id_reference( self, - mnt_filepath, - thl_ledger_manager, - client_no_amm, - payout_event_manager, - brokerage_product_payout_event_manager, - business_payout_event_manager, - delete_ledger_db, - create_main_accounts, - delete_df_collection, - ledger_collection, + mnt_filepath: GRLDatasets, + 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], + delete_df_collection: Callable[..., None], + ledger_collection: LedgerDFCollection, business: Business, - user_factory, - product_factory, - session_with_tx_factory, - pop_ledger_merge, + user_factory: Callable[..., User], + product_factory: Callable[..., Product], + session_with_tx_factory: Callable[..., Session], + pop_ledger_merge: PopLedgerMerge, start: datetime, - bp_payout_factory, - adj_to_fail_with_tx_factory, - thl_web_rr, - lm, - product_manager, - rm_ledger_collection, - rm_pop_ledger_merge, + thl_web_rr: PostgresConfig, + ledger_manager: LedgerManager, + product_manager: ProductManager, + rm_ledger_collection: Callable[..., None], + rm_pop_ledger_merge: Callable[..., None], ): # Now let's load it up and actually test some things @@ -1103,7 +1128,7 @@ class TestBusinessPayoutEventManager: pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) business.prebuild_balance( thl_pg_config=thl_web_rr, - lm=lm, + lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, pop_ledger=pop_ledger_merge, @@ -1124,7 +1149,7 @@ class TestBusinessPayoutEventManager: pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) business.prebuild_balance( thl_pg_config=thl_web_rr, - lm=lm, + lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, pop_ledger=pop_ledger_merge, @@ -1153,10 +1178,11 @@ class TestBusinessPayoutEventManager: ) business.prebuild_balance( thl_pg_config=thl_web_rr, - lm=lm, + lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, pop_ledger=pop_ledger_merge, ) + assert isinstance(business.payouts, list) assert business.payouts[0].ext_ref_id == ach_id2 assert business.payouts[1].ext_ref_id == ach_id1 diff --git a/tests/managers/thl/test_product.py b/tests/managers/thl/test_product.py index f8dd44d..31b7b73 100644 --- a/tests/managers/thl/test_product.py +++ b/tests/managers/thl/test_product.py @@ -4,7 +4,7 @@ import pytest from generalresearch.models import Source from generalresearch.models.thl.product import ( - Product, + product: Product, ProfilingConfig, SourceConfig, SourcesConfig, @@ -178,7 +178,13 @@ class TestProductManager: ] ] - def test_get_by_uuid1(self, product_manager, team, product, product_factory): + def test_get_by_uuid1( + self, + product_manager: ProductManager, + team, + product: Product, + product_factory, + ): p1 = product_factory(team=team) instance = product_manager.get_by_uuid(product_uuid=p1.uuid) assert instance.id == p1.id @@ -191,7 +197,7 @@ 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_uuid2(self, product_manager, product_factory): + def test_get_by_uuid2(self, product_manager: ProductManager, product_factory): p2 = product_factory() instance = product_manager.get_by_uuid(p2.id) assert instance.id, p2.id @@ -203,7 +209,7 @@ 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, product_factory): + def test_get_by_uuid3(self, product_manager: ProductManager, product_factory): p3 = product_factory() instance = product_manager.get_by_uuid(p3.id) assert instance.id == p3.id diff --git a/tests/managers/thl/test_product_prod.py b/tests/managers/thl/test_product_prod.py index f5aa78d..0f622b6 100644 --- a/tests/managers/thl/test_product_prod.py +++ b/tests/managers/thl/test_product_prod.py @@ -10,7 +10,7 @@ logger = logging.getLogger() class TestProductManagerGetMethods: - def test_get_by_uuid(self, product_manager, product_factory): + def test_get_by_uuid(self, product_manager: ProductManager, product_factory): # 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 +22,7 @@ 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, product_factory): + def test_get_by_uuids(self, product_manager: ProductManager, product_factory): products = [product_factory(), product_factory(), product_factory()] cnt = len(products) res = product_manager.get_by_uuids(product_uuids=[p.id for p in products]) @@ -42,7 +42,9 @@ class TestProductManagerGetMethods: ) assert "invalid uuid passed" in str(cm.value) - def test_get_by_uuid_if_exists(self, product_factory, product_manager): + def test_get_by_uuid_if_exists( + self, product_factory: Callable[..., Product], product_manager + ): products = [product_factory(), product_factory(), product_factory()] instance = product_manager.get_by_uuid_if_exists(product_uuid=products[0].id) @@ -51,7 +53,9 @@ class TestProductManagerGetMethods: instance = product_manager.get_by_uuid_if_exists(product_uuid="abc123") assert instance is None - def test_get_by_uuids_if_exists(self, product_manager, product_factory): + def test_get_by_uuids_if_exists( + self, product_manager: ProductManager, product_factory + ): products = [product_factory(), product_factory(), product_factory()] res = product_manager.get_by_uuids_if_exists( diff --git a/tests/managers/thl/test_session_manager.py b/tests/managers/thl/test_session_manager.py index adcbe25..6c5f820 100644 --- a/tests/managers/thl/test_session_manager.py +++ b/tests/managers/thl/test_session_manager.py @@ -75,7 +75,12 @@ class TestSessionManagerFilter: assert len(res) == 2 def test_product( - self, product_factory, user_factory, session_manager, user, utc_hour_ago + self, + product_factory: Callable[..., Product], + user_factory: Callable[..., User], + session_manager, + user, + utc_hour_ago, ): from generalresearch.models.thl.session import Session from generalresearch.models.thl.user import User @@ -95,13 +100,13 @@ class TestSessionManagerFilter: def test_team( self, - product_factory, - user_factory, + product_factory: Callable[..., Product], + user_factory: Callable[..., User], team, session_manager, user, utc_hour_ago, - thl_web_rr, + thl_web_rr: PostgresConfig, ): p1 = product_factory(team=team) @@ -116,13 +121,13 @@ class TestSessionManagerFilter: def test_business( self, - product_factory, - business, - user_factory, + product_factory: Callable[..., Product], + business: Business, + user_factory: Callable[..., User], session_manager, user, utc_hour_ago, - thl_web_rr, + thl_web_rr: PostgresConfig, ): p1 = product_factory(business=business) diff --git a/tests/managers/thl/test_task_adjustment.py b/tests/managers/thl/test_task_adjustment.py index 43337b6..7b77a68 100644 --- a/tests/managers/thl/test_task_adjustment.py +++ b/tests/managers/thl/test_task_adjustment.py @@ -14,14 +14,16 @@ from generalresearch.models.thl.definitions import ( @pytest.fixture() -def session_complete(session_with_tx_factory, user): +def session_complete(session_with_tx_factory: Callable[..., None], user): return session_with_tx_factory( user=user, final_status=Status.COMPLETE, wall_req_cpi=Decimal("1.23") ) @pytest.fixture() -def session_complete_with_wallet(session_with_tx_factory, user_with_wallet): +def session_complete_with_wallet( + session_with_tx_factory: Callable[..., None], user_with_wallet +): return session_with_tx_factory( user=user_with_wallet, final_status=Status.COMPLETE, diff --git a/tests/managers/thl/test_task_status.py b/tests/managers/thl/test_task_status.py index 93a624d..44938a6 100644 --- a/tests/managers/thl/test_task_status.py +++ b/tests/managers/thl/test_task_status.py @@ -71,7 +71,7 @@ class TestTaskStatus: def test_task_status_complete_1( self, bp1, - user_factory, + user_factory: Callable[..., User], finished_session_factory, session_manager: SessionManager, ): @@ -130,7 +130,11 @@ class TestTaskStatus: assert tsr == expected_tsr def test_task_status_complete_2( - self, bp2, user_factory, finished_session_factory, session_manager + self, + bp2, + user_factory: Callable[..., User], + finished_session_factory, + session_manager, ): # User Payout xform 40% user2: User = user_factory(product=bp2) @@ -197,7 +201,11 @@ class TestTaskStatus: assert tsr == expected_tsr def test_task_status_complete_3( - self, bp3, user_factory, finished_session_factory, session_manager + self, + bp3, + user_factory: Callable[..., User], + finished_session_factory, + session_manager, ): # Wallet enabled User Payout xform 50% (the response is identical # to the user wallet disabled w same xform) @@ -232,7 +240,11 @@ class TestTaskStatus: assert tsr == expected_tsr def test_task_status_fail( - self, bp1, user_factory, finished_session_factory, session_manager + self, + bp1, + user_factory: Callable[..., User], + finished_session_factory, + session_manager, ): # User Payout xform NULL: user payout is None always user1: User = user_factory(product=bp1) @@ -268,7 +280,11 @@ class TestTaskStatus: assert tsr == expected_tsr def test_task_status_fail_xform( - self, bp2, user_factory, finished_session_factory, session_manager + self, + bp2, + user_factory: Callable[..., User], + finished_session_factory, + session_manager, ): # User Payout xform 40%: user_payout is 0 (not None) @@ -303,7 +319,11 @@ class TestTaskStatus: assert tsr == expected_tsr def test_task_status_abandon( - self, bp1, user_factory, session_factory, session_manager + self, + bp1, + user_factory: Callable[..., User], + session_factory, + session_manager, ): # User Payout xform NULL: all payout fields are None user: User = user_factory(product=bp1) @@ -337,7 +357,11 @@ class TestTaskStatus: assert tsr == expected_tsr def test_task_status_abandon_xform( - self, bp2, user_factory, session_factory, session_manager + self, + bp2, + user_factory: Callable[..., User], + session_factory, + session_manager, ): # User Payout xform 40%: all payout fields are None (same as when payout xform is null) user: User = user_factory(product=bp2) @@ -376,7 +400,7 @@ class TestTaskStatus: def test_task_status_adj_fail( self, bp1, - user_factory, + user_factory: Callable[..., User], finished_session_factory, wall_manager, session_manager, @@ -425,7 +449,7 @@ class TestTaskStatus: def test_task_status_adj_fail_xform( self, bp2, - user_factory, + user_factory: Callable[..., User], finished_session_factory, wall_manager, session_manager, @@ -477,7 +501,7 @@ class TestTaskStatus: def test_task_status_adj_complete_from_abandon( self, bp1, - user_factory, + user_factory: Callable[..., User], session_factory, wall_manager, session_manager, @@ -531,7 +555,7 @@ class TestTaskStatus: def test_task_status_adj_complete_from_abandon_xform( self, bp2, - user_factory, + user_factory: Callable[..., User], session_factory, wall_manager, session_manager, @@ -588,7 +612,7 @@ class TestTaskStatus: def test_task_status_adj_complete_from_fail( self, bp1, - user_factory, + user_factory: Callable[..., User], finished_session_factory, wall_manager, session_manager, @@ -642,7 +666,7 @@ class TestTaskStatus: def test_task_status_adj_complete_from_fail_xform( self, bp2, - user_factory, + user_factory: Callable[..., User], finished_session_factory, wall_manager, session_manager, diff --git a/tests/managers/thl/test_user_manager/test_base.py b/tests/managers/thl/test_user_manager/test_base.py index 5235a0f..6b259ff 100644 --- a/tests/managers/thl/test_user_manager/test_base.py +++ b/tests/managers/thl/test_user_manager/test_base.py @@ -17,7 +17,7 @@ from generalresearch.managers.thl.user_manager.user_manager import ( UserManager, ) from generalresearch.managers.thl.userhealth import AuditLogManager -from generalresearch.models.thl.product import Product, UserCreateConfig +from generalresearch.models.thl.product import product: Product, UserCreateConfig from generalresearch.models.thl.user import User from generalresearch.pg_helper import PostgresConfig @@ -86,7 +86,7 @@ class TestUserManager: class TestBlockUserManager: - def test_block_user(self, product: Product, user_manager: UserManager): + def test_block_user(self, product: product: Product, user_manager: UserManager): product_user_id = f"user-{uuid4().hex[:10]}" # mysql_user_manager to skip user creation limit check @@ -113,7 +113,7 @@ class TestBlockUserManager: assert user.blocked def test_block_user_whitelist( - self, product: Product, user_manager: UserManager, thl_web_rw: PostgresConfig + self, product: product: Product, user_manager: UserManager, thl_web_rw: PostgresConfig ): product_user_id = f"user-{uuid4().hex[:10]}" @@ -183,7 +183,10 @@ class TestCreateUserManager: assert u2.uuid == user.uuid def test_create_user_integrity_error( - self, product_manager, user_manager: UserManager, caplog + self, + product_manager: ProductManager, + user_manager: UserManager, + caplog, ): product: Product = product_manager.create_dummy( product_id=uuid4().hex, diff --git a/tests/managers/thl/test_user_manager/test_redis.py b/tests/managers/thl/test_user_manager/test_redis.py index a69519e..0731438 100644 --- a/tests/managers/thl/test_user_manager/test_redis.py +++ b/tests/managers/thl/test_user_manager/test_redis.py @@ -47,7 +47,7 @@ class TestUserManagerRedis: um1 = UserManager( pg_config=thl_web_rw, - pg_config_rr=thl_web_rr, + pg_config_rr=thl_web_rr: PostgresConfig, sql_permissions=[Permission.UPDATE, Permission.CREATE], redis=settings.redis, redis_timeout=settings.redis_timeout, @@ -55,7 +55,7 @@ class TestUserManagerRedis: um2 = UserManager( pg_config=thl_web_rw, - pg_config_rr=thl_web_rr, + pg_config_rr=thl_web_rr: PostgresConfig, sql_permissions=[Permission.UPDATE, Permission.CREATE], redis=settings.redis, redis_timeout=settings.redis_timeout, 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 7c9e012..5c608b3 100644 --- a/tests/managers/thl/test_user_manager/test_user_fetch.py +++ b/tests/managers/thl/test_user_manager/test_user_fetch.py @@ -7,7 +7,9 @@ from generalresearch.models.thl.user import User class TestUserManagerFetch: - def test_fetch(self, user_factory, product, user_manager): + def test_fetch( + self, user_factory: Callable[..., User], product: Product, user_manager + ): user1: User = user_factory(product=product) user2: User = user_factory(product=product) res = user_manager.fetch_by_bpuids( 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 19b3d9f..0b99afe 100644 --- a/tests/managers/thl/test_user_manager/test_user_metadata.py +++ b/tests/managers/thl/test_user_manager/test_user_metadata.py @@ -12,7 +12,9 @@ class TestUserMetadataManager: um1 = user_metadata_manager.get(user_id=user.user_id) assert um1 == UserMetadata(user_id=user.user_id) - def test_create(self, user_factory, product, user_metadata_manager): + def test_create( + self, user_factory: Callable[..., User], product: Product, user_metadata_manager + ): from generalresearch.models.thl.user import User u1: User = user_factory(product=product) @@ -26,7 +28,9 @@ class TestUserMetadataManager: um2 = user_metadata_manager.get(email_address=email_address) assert um == um2 - def test_create_no_email(self, product, user_factory, user_metadata_manager): + def test_create_no_email( + self, product: Product, user_factory: Callable[..., User], user_metadata_manager + ): from generalresearch.models.thl.user import User u1: User = user_factory(product=product) @@ -37,7 +41,9 @@ class TestUserMetadataManager: um2 = user_metadata_manager.get(user_id=u1.user_id) assert um == um2 - def test_update(self, product, user_factory, user_metadata_manager): + def test_update( + self, product: Product, user_factory: Callable[..., User], user_metadata_manager + ): from generalresearch.models.thl.user import User u: User = user_factory(product=product) @@ -57,7 +63,9 @@ class TestUserMetadataManager: email_address=email_address.replace("example1", "example2"), ) - def test_filter(self, user_factory, product, user_metadata_manager): + 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) diff --git a/tests/managers/thl/test_userhealth.py b/tests/managers/thl/test_userhealth.py index be8f9c8..98d8f25 100644 --- a/tests/managers/thl/test_userhealth.py +++ b/tests/managers/thl/test_userhealth.py @@ -19,7 +19,7 @@ fake = faker.Faker() class TestAuditLog: - def test_init(self, thl_web_rr, audit_log_manager): + def test_init(self, thl_web_rr: PostgresConfig, audit_log_manager): from generalresearch.managers.thl.userhealth import AuditLogManager alm = AuditLogManager(pg_config=thl_web_rr) @@ -55,8 +55,8 @@ class TestAuditLog: def test_filter_by_product( self, - user_factory, - product_factory, + user_factory: Callable[..., User], + product_factory: Callable[..., Product], audit_log_factory, audit_log_manager, ): @@ -82,7 +82,7 @@ class TestAuditLog: assert len(res) == 1 def test_filter_by_user_id( - self, user_factory, product, audit_log_factory, audit_log_manager + self, user_factory: Callable[..., User], product: Product, audit_log_factory, audit_log_manager ): u1 = user_factory(product=product) u2 = user_factory(product=product) @@ -108,8 +108,8 @@ class TestAuditLog: def test_filter( self, - user_factory, - product_factory, + user_factory: Callable[..., User], + product_factory: Callable[..., Product], audit_log_factory, audit_log_manager, ): @@ -142,8 +142,8 @@ class TestAuditLog: def test_filter_count( self, - user_factory, - product_factory, + user_factory: Callable[..., User], + product_factory: Callable[..., Product], audit_log_factory, audit_log_manager, ): @@ -205,8 +205,8 @@ class TestAuditLog: class TestIPRecordManager: - def test_init(self, thl_web_rr, thl_redis_config, ip_record_manager): - instance = IPRecordManager(pg_config=thl_web_rr, redis_config=thl_redis_config) + 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) assert isinstance(instance, IPRecordManager) assert isinstance(ip_record_manager, IPRecordManager) @@ -232,8 +232,8 @@ class TestIPRecordManager: ip_information_factory, ip_geoname, user, - thl_web_rr, - thl_redis_config, + thl_web_rr: PostgresConfig, + thl_redis_config: RedisConfig, ): ip = fake.ipv4_public() @@ -246,8 +246,8 @@ class TestIPRecordManager: assert fipr.information is None ipr.prefetch_ipinfo( - pg_config=thl_web_rr, - redis_config=thl_redis_config, + pg_config=thl_web_rr: PostgresConfig, + redis_config=thl_redis_config: RedisConfig, include_forwarded=True, ) assert isinstance(ipr.information, GeoIPInformation) @@ -256,8 +256,8 @@ class TestIPRecordManager: ip_information_factory(ip=fipr.ip, geoname=ip_geoname) ipr.prefetch_ipinfo( - pg_config=thl_web_rr, - redis_config=thl_redis_config, + pg_config=thl_web_rr: PostgresConfig, + redis_config=thl_redis_config: RedisConfig, include_forwarded=True, ) assert fipr.information is not None @@ -265,9 +265,9 @@ class TestIPRecordManager: @pytest.mark.usefixtures("user_iphistory_manager_clear_cache") class TestUserIpHistoryManager: - def test_init(self, thl_web_rr, thl_redis_config, user_iphistory_manager): + def test_init(self, thl_web_rr: PostgresConfig, thl_redis_config: RedisConfig, user_iphistory_manager): instance = UserIpHistoryManager( - pg_config=thl_web_rr, redis_config=thl_redis_config + pg_config=thl_web_rr: PostgresConfig, redis_config=thl_redis_config ) assert isinstance(instance, UserIpHistoryManager) assert isinstance(user_iphistory_manager, UserIpHistoryManager) diff --git a/tests/models/gr/test_authentication.py b/tests/models/gr/test_authentication.py index ac39e64..d4db112 100644 --- a/tests/models/gr/test_authentication.py +++ b/tests/models/gr/test_authentication.py @@ -44,10 +44,10 @@ class TestGRUser: gr_user_token, gr_user: GRUser, membership: Membership, - product_factory, + product_factory: Callable[..., Product], membership_factory, team: Team, - thl_web_rr, + thl_web_rr: PostgresConfig, gr_redis_config, gr_db, ): @@ -64,11 +64,11 @@ class TestGRUser: def test_products( self, gr_user: GRUser, - product_factory, + product_factory: Callable[..., Product], team: Team, membership: Membership, gr_db, - thl_web_rr, + thl_web_rr: PostgresConfig, gr_redis_config, ): from generalresearch.models.thl.product import Product @@ -87,7 +87,7 @@ class TestGRUser: gr_user.prefetch_products( pg_config=gr_db, - thl_pg_config=thl_web_rr, + thl_pg_config=thl_web_rr: PostgresConfig, redis_config=gr_redis_config, ) assert isinstance(gr_user.products, list) @@ -107,8 +107,8 @@ class TestGRUserMethods: gr_user: GRUser, gr_redis, team: Team, - business, - product_factory, + business: Business, + product_factory: Callable[..., Product], membership_factory: Callable[Membership], ): product_factory(team=team, business=business) @@ -128,7 +128,7 @@ class TestGRUserMethods: gr_user_token, gr_redis, gr_db, - thl_web_rr, + thl_web_rr: PostgresConfig, gr_redis_config, ): assert gr_redis.get(name=gr_user.cache_key) is None @@ -137,7 +137,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, redis_config=gr_redis_config + pg_config=gr_db, thl_web_rr=thl_web_rr: PostgresConfig, redis_config=gr_redis_config ) assert gr_redis.get(name=gr_user.cache_key) is not None @@ -152,11 +152,11 @@ class TestGRUserMethods: gr_redis, gr_redis_config, gr_db, - thl_web_rr, - product_factory, + thl_web_rr: PostgresConfig, + product_factory: Callable[..., Product], team, membership_factory, - thl_redis_config, + thl_redis_config: RedisConfig, ): from generalresearch.models.gr.authentication import GRUser @@ -164,7 +164,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, redis_config=gr_redis_config + pg_config=gr_db, thl_web_rr=thl_web_rr: PostgresConfig, redis_config=gr_redis_config ) res: str = gr_redis.get(name=gr_user.cache_key) @@ -176,8 +176,8 @@ class TestGRUserMethods: gru2.prefetch_products( pg_config=gr_db, - thl_pg_config=thl_web_rr, - redis_config=thl_redis_config, + thl_pg_config=thl_web_rr: PostgresConfig, + redis_config=thl_redis_config: RedisConfig, ) assert gru2.product_uuids == [p1.uuid] @@ -188,15 +188,15 @@ class TestGRUserMethods: gr_user_token, gr_redis, gr_db, - thl_web_rr, - product_factory, + thl_web_rr: PostgresConfig, + product_factory: Callable[..., Product], team, gr_redis_config, ): product_factory(team=team) gr_user.set_cache( - pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config + pg_config=gr_db, thl_web_rr=thl_web_rr: PostgresConfig, redis_config=gr_redis_config ) res = json.loads(gr_redis.get(name=f"{gr_user.cache_key}:team_uuids")) assert len(res) == 1 @@ -208,16 +208,16 @@ class TestGRUserMethods: gr_user: GRUser, gr_redis, gr_db, - thl_web_rr, - product_factory, - business, + thl_web_rr: PostgresConfig, + product_factory: Callable[..., Product], + business: Business, team, gr_redis_config, ): product_factory(team=team, business=business) gr_user.set_cache( - pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config + pg_config=gr_db, thl_web_rr=thl_web_rr: PostgresConfig, redis_config=gr_redis_config ) res = json.loads(gr_redis.get(name=f"{gr_user.cache_key}:business_uuids")) assert len(res) == 1 @@ -230,15 +230,15 @@ class TestGRUserMethods: gr_user_token, gr_redis, gr_db, - thl_web_rr, - product_factory, + thl_web_rr: PostgresConfig, + product_factory: Callable[..., Product], team, gr_redis_config, ): product_factory(team=team) gr_user.set_cache( - pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config + pg_config=gr_db, thl_web_rr=thl_web_rr: PostgresConfig, redis_config=gr_redis_config ) res = json.loads(gr_redis.get(name=f"{gr_user.cache_key}:product_uuids")) assert len(res) == 1 diff --git a/tests/models/gr/test_business.py b/tests/models/gr/test_business.py index 48a7bb0..5239ac2 100644 --- a/tests/models/gr/test_business.py +++ b/tests/models/gr/test_business.py @@ -29,7 +29,7 @@ from generalresearch.managers.thl.payout import ( PayoutEventManager, ) from generalresearch.models.gr.business import ( - Business, + business: Business, BusinessAddress, BusinessBankAccount, BusinessContact, @@ -50,7 +50,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 +68,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 +79,7 @@ class TestBusinessBankAccount: business_bank_account.prefetch_business( pg_config=gr_db, redis_config=gr_redis_config ) - assert isinstance(business_bank_account.business, Business) + assert isinstance(business_bank_account.business: Business, Business) assert business_bank_account.business.uuid == business.uuid @@ -112,13 +112,13 @@ class TestBusiness: def test_init(self, business: Business): - assert isinstance(business, Business) + assert isinstance(business: 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 +181,12 @@ class TestBusiness: pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) business.prebuild_payouts( - thl_pg_config=thl_web_rr, + thl_pg_config=thl_web_rr: PostgresConfig, thl_lm=thl_ledger_manager, bpem=business_payout_event_manager, ) business.prebuild_balance( - thl_pg_config=thl_web_rr, + thl_pg_config=thl_web_rr: PostgresConfig, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -198,7 +198,7 @@ class TestBusiness: def test_addresses( self, - business: Business, + business: business: Business, business_address: BusinessAddress, gr_db: PostgresConfig, ): @@ -213,7 +213,7 @@ class TestBusiness: def test_teams( self, - business: Business, + business: business: Business, team: Team, team_manager: TeamManager, gr_db: PostgresConfig, @@ -231,7 +231,7 @@ class TestBusiness: def test_products( self, - business: Business, + business: business: Business, product_factory: Callable[..., Product], thl_web_rr: PostgresConfig, ): @@ -254,7 +254,7 @@ class TestBusiness: business.prefetch_products(thl_pg_config=thl_web_rr) assert len(business.products) == 3 - def test_bank_accounts(self, business: Business, gr_db: PostgresConfig): + def test_bank_accounts(self, business: business: Business, gr_db: PostgresConfig): assert business.products is None # It's an empty list after prefetch @@ -264,7 +264,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 +275,7 @@ class TestBusiness: with pytest.raises(expected_exception=AssertionError) as cm: business.prebuild_balance( - thl_pg_config=thl_web_rr, + thl_pg_config=thl_web_rr: PostgresConfig, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -289,7 +289,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 +299,7 @@ class TestBusiness: with pytest.raises(expected_exception=AssertionError) as cm: business.prebuild_payouts( - thl_pg_config=thl_web_rr, + thl_pg_config=thl_web_rr: PostgresConfig, thl_lm=thl_ledger_manager, bpem=business_payout_event_manager, ) @@ -309,7 +309,7 @@ class TestBusiness: thl_ledger_manager.get_account_or_create_bp_wallet(product=p) business.prebuild_payouts( - thl_pg_config=thl_web_rr, + thl_pg_config=thl_web_rr: PostgresConfig, thl_lm=thl_ledger_manager, bpem=business_payout_event_manager, ) @@ -318,7 +318,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 +338,7 @@ class TestBusiness: ) business.prebuild_payouts( - thl_pg_config=thl_web_rr, + thl_pg_config=thl_web_rr: PostgresConfig, thl_lm=thl_ledger_manager, bpem=business_payout_event_manager, ) @@ -356,7 +356,7 @@ class TestBusiness: thl_lm=thl_ledger_manager ) business.prebuild_payouts( - thl_pg_config=thl_web_rr, + thl_pg_config=thl_web_rr: PostgresConfig, thl_lm=thl_ledger_manager, bpem=business_payout_event_manager, ) @@ -367,7 +367,7 @@ 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, @@ -406,7 +406,7 @@ class TestBusiness: ) business.prebuild_payouts( - thl_pg_config=thl_web_rr, + thl_pg_config=thl_web_rr: PostgresConfig, thl_lm=thl_ledger_manager, bpem=business_payout_event_manager, ) @@ -419,7 +419,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 +428,7 @@ class TestBusiness: ): assert business.pop_financial is None business.prebuild_pop_financial( - thl_pg_config=thl_web_rr, + thl_pg_config=thl_web_rr: PostgresConfig, thl_lm=thl_ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -438,7 +438,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 +480,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 +519,7 @@ class TestBusinessBalance: pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) business.prebuild_balance( - thl_pg_config=thl_web_rr, + thl_pg_config=thl_web_rr: PostgresConfig, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -541,7 +541,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 +579,7 @@ class TestBusinessBalance: pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) business.prebuild_balance( - thl_pg_config=thl_web_rr, + thl_pg_config=thl_web_rr: PostgresConfig, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -625,7 +625,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 +665,7 @@ class TestBusinessBalance: payout_event_manager.set_account_lookup_table(thl_lm=thl_ledger_manager) bp_payout_factory( - product=u1.product, + product=u1.product: Product, amount=USDCent(5), created=start + timedelta(days=4), skip_wallet_balance_check=True, @@ -673,7 +673,7 @@ class TestBusinessBalance: ) bp_payout_factory( - product=u2.product, + product=u2.product: Product, amount=USDCent(50), created=start + timedelta(days=4), skip_wallet_balance_check=True, @@ -684,7 +684,7 @@ class TestBusinessBalance: pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) business.prebuild_balance( - thl_pg_config=thl_web_rr, + thl_pg_config=thl_web_rr: PostgresConfig, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -699,7 +699,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 +758,7 @@ class TestBusinessBalance: payout_event_manager.set_account_lookup_table(thl_lm=thl_ledger_manager) bp_payout_factory( - product=u1.product, + product=u1.product: Product, amount=USDCent(250), created=start + timedelta(days=3), skip_wallet_balance_check=True, @@ -766,7 +766,7 @@ class TestBusinessBalance: ) bp_payout_factory( - product=u2.product, + product=u2.product: Product, amount=USDCent(50), created=start + timedelta(days=4), skip_wallet_balance_check=True, @@ -796,7 +796,7 @@ class TestBusinessBalance: assert df.shape == (20, 28) business.prebuild_balance( - thl_pg_config=thl_web_rr, + thl_pg_config=thl_web_rr: PostgresConfig, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -833,7 +833,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 +869,7 @@ class TestBusinessBalance: ) payout_event_manager.set_account_lookup_table(thl_lm=thl_ledger_manager) bp_payout_factory( - product=u1.product, + product=u1.product: Product, amount=USDCent(71), ext_ref_id=uuid4().hex, created=start + timedelta(days=1, minutes=1), @@ -898,7 +898,7 @@ class TestBusinessBalance: pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) business.prebuild_balance( - thl_pg_config=thl_web_rr, + thl_pg_config=thl_web_rr: PostgresConfig, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -946,7 +946,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, @@ -956,7 +956,7 @@ class TestBusinessBalance: start: datetime, thl_web_rr: PostgresConfig, payout_event_manager, - session_with_tx_factory, + session_with_tx_factory: Callable[..., None], delete_ledger_db: Callable[..., None], create_main_accounts: Callable[..., None], client_no_amm: DaskClient, @@ -1022,7 +1022,7 @@ class TestBusinessBalance: payout_event_manager.set_account_lookup_table(thl_lm=thl_ledger_manager) bp_payout_factory( - product=u1.product, + product=u1.product: Product, amount=USDCent(250), created=start + timedelta(days=3), skip_wallet_balance_check=True, @@ -1030,7 +1030,7 @@ class TestBusinessBalance: ) bp_payout_factory( - product=u2.product, + product=u2.product: Product, amount=USDCent(50), created=start + timedelta(days=4), skip_wallet_balance_check=True, @@ -1060,7 +1060,7 @@ class TestBusinessBalance: assert df.shape == (20, 28) business.prebuild_balance( - thl_pg_config=thl_web_rr, + thl_pg_config=thl_web_rr: PostgresConfig, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -1068,7 +1068,7 @@ class TestBusinessBalance: ) business.prebuild_balance( - thl_pg_config=thl_web_rr, + thl_pg_config=thl_web_rr: PostgresConfig, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -1078,7 +1078,7 @@ class TestBusinessBalance: day1_bal = business.balance business.prebuild_balance( - thl_pg_config=thl_web_rr, + thl_pg_config=thl_web_rr: PostgresConfig, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -1088,7 +1088,7 @@ class TestBusinessBalance: day2_bal = business.balance business.prebuild_balance( - thl_pg_config=thl_web_rr, + thl_pg_config=thl_web_rr: PostgresConfig, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -1098,7 +1098,7 @@ class TestBusinessBalance: day3_bal = business.balance business.prebuild_balance( - thl_pg_config=thl_web_rr, + thl_pg_config=thl_web_rr: PostgresConfig, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -1108,7 +1108,7 @@ class TestBusinessBalance: day4_bal = business.balance business.prebuild_balance( - thl_pg_config=thl_web_rr, + thl_pg_config=thl_web_rr: PostgresConfig, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -1118,7 +1118,7 @@ class TestBusinessBalance: day5_bal = business.balance business.prebuild_balance( - thl_pg_config=thl_web_rr, + thl_pg_config=thl_web_rr: PostgresConfig, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -1183,7 +1183,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 +1219,7 @@ class TestBusinessMethods: business.set_cache( pg_config=gr_db, - thl_web_rr=thl_web_rr, + thl_web_rr=thl_web_rr: PostgresConfig, redis_config=gr_redis_config, client=client_no_amm, ds=mnt_filepath, @@ -1245,7 +1245,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 +1282,7 @@ class TestBusinessMethods: business.set_cache( pg_config=gr_db, - thl_web_rr=thl_web_rr, + thl_web_rr=thl_web_rr: PostgresConfig, redis_config=gr_redis_config, client=client_no_amm, ds=mnt_filepath, @@ -1353,7 +1353,7 @@ class TestBusinessMethods: 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 +1380,11 @@ class TestBusinessMethods: client=client_no_amm, session_coll=session_collection, wall_coll=wall_collection, - pg_config=thl_web_rr, + pg_config=thl_web_rr: PostgresConfig, ) business.prebuild_enriched_session_parquet( - thl_pg_config=thl_web_rr, + thl_pg_config=thl_web_rr: PostgresConfig, ds=mnt_filepath, client=client_no_amm, mnt_gr_api=mnt_gr_api_dir, @@ -1409,7 +1409,7 @@ class TestBusinessMethods: 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 +1436,11 @@ class TestBusinessMethods: client=client_no_amm, session_coll=session_collection, wall_coll=wall_collection, - pg_config=thl_web_rr, + pg_config=thl_web_rr: PostgresConfig, ) business.prebuild_enriched_wall_parquet( - thl_pg_config=thl_web_rr, + thl_pg_config=thl_web_rr: PostgresConfig, 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 dc7d4b9..26300b9 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, + thl_web_rr=thl_web_rr: PostgresConfig, 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, + thl_web_rr=thl_web_rr: PostgresConfig, 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, + pg_config=thl_web_rr: PostgresConfig, ) team.prebuild_enriched_session_parquet( - thl_pg_config=thl_web_rr, + thl_pg_config=thl_web_rr: PostgresConfig, 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, + pg_config=thl_web_rr: PostgresConfig, ) team.prebuild_enriched_wall_parquet( - thl_pg_config=thl_web_rr, + thl_pg_config=thl_web_rr: PostgresConfig, ds=mnt_filepath, client=client_no_amm, mnt_gr_api=mnt_gr_api_dir, diff --git a/tests/models/legacy/test_user_question_answer_in.py b/tests/models/legacy/test_user_question_answer_in.py index ee70d81..313862c 100644 --- a/tests/models/legacy/test_user_question_answer_in.py +++ b/tests/models/legacy/test_user_question_answer_in.py @@ -15,12 +15,12 @@ class TestUserQuestionAnswers: def test_json_init( self, - product_manager, + product_manager: ProductManager, user_manager, session_manager, wall_manager, - user_factory, - product, + user_factory: Callable[..., User], + product: Product, session_factory, utc_hour_ago, ): @@ -60,7 +60,11 @@ class TestUserQuestionAnswers: assert isinstance(instance, UserQuestionAnswers) def test_simple_validation_errors( - self, product_manager, user_manager, session_manager, wall_manager + self, + product_manager: ProductManager, + user_manager, + session_manager, + wall_manager, ): from generalresearch.models.legacy.questions import ( UserQuestionAnswers, @@ -162,8 +166,8 @@ class TestUserQuestionAnswers: def test_allow_answer_failures_silent( self, user_manager, - product, - user_factory, + product: Product, + user_factory: Callable[..., User], utc_hour_ago, session_factory, ): diff --git a/tests/models/test_finance.py b/tests/models/test_finance.py index 3a313e2..f84d0b6 100644 --- a/tests/models/test_finance.py +++ b/tests/models/test_finance.py @@ -31,7 +31,7 @@ from generalresearch.models.thl.user import User from test_utils.incite.collections.conftest import ledger_collection from test_utils.incite.mergers.conftest import pop_ledger_merge from test_utils.managers.ledger.conftest import ( - session_with_tx_factory, + session_with_tx_factory: Callable[..., None], ) fake = Faker() @@ -665,7 +665,7 @@ class TestProductFinanceData: def test_base( self, - product: Product, + product: product: Product, user_factory: Callable[..., User], start: datetime, duration: timedelta, @@ -675,7 +675,7 @@ class TestProductFinanceData: # -- Build & Setup # assert ledger_collection.start is None # assert ledger_collection.offset is None - u: User = user_factory(product=product, created=ledger_collection.start) + u: User = user_factory(product=product: Product, created=ledger_collection.start) for item in ledger_collection.items: @@ -748,14 +748,14 @@ class TestPOPFinancialData: ledger_collection: LedgerDFCollection, pop_ledger_merge: PopLedgerMerge, user_factory: Callable[..., User], - product: Product, + product: product: Product, start: datetime, duration: timedelta, - create_main_accounts, + create_main_accounts: Callable[..., None], session_with_tx_factory: Callable[..., Session], thl_lm: ThlLedgerManager, - delete_df_collection, - delete_ledger_db, + delete_df_collection: Callable[..., None], + delete_ledger_db: Callable[..., None], ): # -- Build & Setup delete_ledger_db() @@ -820,7 +820,7 @@ class TestPOPFinancialData: # This does not return the AccountID, it's the Product ID assert i.product_id in [u.product_id for u in users] - # 1 Product, multiple Users + # 1 product: Product, multiple Users assert len(users) == len(accounts) # We group on days, and duration is a parameter to parametrize @@ -846,12 +846,12 @@ class TestBusinessBalanceData: ledger_collection: LedgerDFCollection, pop_ledger_merge: PopLedgerMerge, user_factory: Callable[..., User], - product: Product, - create_main_accounts, + product: product: Product, + create_main_accounts: Callable[..., None], thl_lm: ThlLedgerManager, - thl_web_rr, - delete_df_collection, - delete_ledger_db, + thl_web_rr: PostgresConfig, + delete_df_collection: Callable[..., None], + delete_ledger_db: Callable[..., None], session_with_tx_factory: Callable[..., Session], rm_ledger_collection, ): @@ -863,7 +863,7 @@ class TestBusinessBalanceData: rm_ledger_collection() for _ in range(5): - u: User = user_factory(product=product, created=ledger_collection.start) + u: User = user_factory(product=product: Product, created=ledger_collection.start) for item in ledger_collection.items: item_time = fake.date_time_between( diff --git a/tests/models/thl/test_adjustments.py b/tests/models/thl/test_adjustments.py index 96b67d7..91e5316 100644 --- a/tests/models/thl/test_adjustments.py +++ b/tests/models/thl/test_adjustments.py @@ -459,7 +459,7 @@ class TestAdjustments: assert Status.FAIL == new_status assert Decimal(0) == new_payout - assert isinstance(user.product, Product) + assert isinstance(user.product: Product, Product) assert not user.product.user_wallet_config.enabled assert new_user_payout is None diff --git a/tests/models/thl/test_contest/test_leaderboard_contest.py b/tests/models/thl/test_contest/test_leaderboard_contest.py index 52a4bec..5bab060 100644 --- a/tests/models/thl/test_contest/test_leaderboard_contest.py +++ b/tests/models/thl/test_contest/test_leaderboard_contest.py @@ -25,7 +25,7 @@ class TestLeaderboardContest(TestContest): @pytest.fixture def leaderboard_contest( - self, product: Product, thl_redis, user_manager + self, product: product: Product, thl_redis, user_manager ) -> LeaderboardContest: board_key = f"leaderboard:{product.uuid}:us:weekly:2025-05-26:complete_count" diff --git a/tests/models/thl/test_contest/test_raffle_contest.py b/tests/models/thl/test_contest/test_raffle_contest.py index d7920f0..f85ba75 100644 --- a/tests/models/thl/test_contest/test_raffle_contest.py +++ b/tests/models/thl/test_contest/test_raffle_contest.py @@ -243,7 +243,7 @@ class TestRaffleContestWinners(TestRaffleContest): assert len(winners) == 2 def test_winners_3_prizes_3_entries( - self, ended_raffle_contest, product, user_1, user_2, user_3 + self, ended_raffle_contest, product: Product, user_1, user_2, user_3 ): ended_raffle_contest.prizes = [ ContestPrize( diff --git a/tests/models/thl/test_payout.py b/tests/models/thl/test_payout.py index f1046cb..dd0065c 100644 --- a/tests/models/thl/test_payout.py +++ b/tests/models/thl/test_payout.py @@ -5,7 +5,7 @@ from pydantic import ValidationError from generalresearch.currency import USDCent from generalresearch.models.gr import Team -from generalresearch.models.gr.business import Business, BusinessAddress, BusinessType +from generalresearch.models.gr.business import business: Business, BusinessAddress, BusinessType from generalresearch.models.thl.payout import ( BrokerageProductPayoutEvent, BusinessPayoutEvent, diff --git a/tests/models/thl/test_product.py b/tests/models/thl/test_product.py index bc95c2d..b7ee654 100644 --- a/tests/models/thl/test_product.py +++ b/tests/models/thl/test_product.py @@ -28,7 +28,7 @@ from generalresearch.models.thl.product import ( PayoutConfig, PayoutTransformation, PayoutTransformationPercentArgs, - Product, + product: Product, ProfilingConfig, SourceConfig, SourcesConfig, @@ -287,7 +287,7 @@ class TestProduct: p.profiling_config = ProfilingConfig(max_questions=1) assert p.profiling_config.max_questions == 1 - def test_bp_account(self, product, thl_lm): + def test_bp_account(self, product: Product, thl_lm): assert product.bp_account is None product.prefetch_bp_account(thl_lm=thl_lm) @@ -391,7 +391,7 @@ class TestGlobalProduct: random_product = uuid4().hex random_team = uuid4().hex res = instance.sources_config.get_policies_for( - product_id=random_product, team_id=random_team + product_id=random_product: Product, team_id=random_team ) assert res == s.global_scoped_policies_dict @@ -598,7 +598,7 @@ class TestProductFinancials: def test_balance( self, - business: Business, + business: business: Business, product_factory: Callable[..., Product], user_factory: Callable[..., User], mnt_filepath: GRLDatasets, @@ -607,12 +607,12 @@ class TestProductFinancials: start: datetime, brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager, session_with_tx_factory: Callable[..., Session], - delete_ledger_db, - create_main_accounts, + delete_ledger_db: Callable[..., None], + create_main_accounts: Callable[..., None], client_no_amm: DaskClient, ledger_collection, pop_ledger_merge: PopLedgerMerge, - delete_df_collection, + delete_df_collection: Callable[..., None], ): delete_ledger_db() create_main_accounts() @@ -781,13 +781,13 @@ class TestProductBalance: def test_inconsistent( self, - product: Product, + product: product: Product, mnt_filepath: GRLDatasets, thl_lm: ThlLedgerManager, client_no_amm: DaskClient, - delete_ledger_db, - create_main_accounts, - delete_df_collection, + delete_ledger_db: Callable[..., None], + create_main_accounts: Callable[..., None], + delete_df_collection: Callable[..., None], ledger_collection, user_factory: Callable[..., User], session_with_tx_factory: Callable[..., Session], @@ -815,7 +815,7 @@ class TestProductBalance: # 2. Payout and build Parquets 2nd time payout_event_manager.set_account_lookup_table(thl_lm=thl_lm) bp_payout_factory( - product=product, + product=product: Product, amount=USDCent(71), ext_ref_id=uuid4().hex, created=start + timedelta(days=1, minutes=1), @@ -833,16 +833,16 @@ class TestProductBalance: def test_not_inconsistent( self, - product: Product, + product: product: Product, mnt_filepath: GRLDatasets, thl_lm: ThlLedgerManager, client_no_amm: DaskClient, - delete_ledger_db, - create_main_accounts, - delete_df_collection, + delete_ledger_db: Callable[..., None], + create_main_accounts: Callable[..., None], + delete_df_collection: Callable[..., None], ledger_collection, user_factory: Callable[..., User], - session_with_tx_factory, + session_with_tx_factory: Callable[..., None], pop_ledger_merge: PopLedgerMerge, start: datetime, bp_payout_factory, @@ -874,7 +874,7 @@ class TestProductBalance: # so it hasn't already been archived payout_event_manager.set_account_lookup_table(thl_lm=thl_lm) bp_payout_factory( - product=product, + product=product: Product, amount=USDCent(71), ext_ref_id=uuid4().hex, created=datetime.now(tz=UTC), @@ -904,16 +904,16 @@ class TestProductPOPFinancial: def test_base( self, - product: Product, + product: product: Product, mnt_filepath: GRLDatasets, thl_lm: ThlLedgerManager, client_no_amm: DaskClient, - delete_ledger_db, - create_main_accounts, - delete_df_collection, + delete_ledger_db: Callable[..., None], + create_main_accounts: Callable[..., None], + delete_df_collection: Callable[..., None], ledger_collection, user_factory: Callable[..., User], - session_with_tx_factory, + session_with_tx_factory: Callable[..., None], pop_ledger_merge: PopLedgerMerge, start: datetime, ): @@ -977,18 +977,18 @@ class TestProductCache: def test_basic( self, - product: Product, + product: product: Product, mnt_filepath, thl_lm, client_no_amm: DaskClient, - thl_redis_config, + thl_redis_config: RedisConfig, brokerage_product_payout_event_manager, - delete_ledger_db, - create_main_accounts, - delete_df_collection, + delete_ledger_db: Callable[..., None], + create_main_accounts: Callable[..., None], + delete_df_collection: Callable[..., None], ledger_collection, user_factory: Callable[..., User], - session_with_tx_factory, + session_with_tx_factory: Callable[..., None], pop_ledger_merge: PopLedgerMerge, start: datetime, ): @@ -1007,7 +1007,7 @@ class TestProductCache: ds=mnt_filepath, client=client_no_amm, bp_pem=brokerage_product_payout_event_manager, - redis_config=thl_redis_config, + redis_config=thl_redis_config: RedisConfig, ) from generalresearch.models.thl.product import Product @@ -1029,7 +1029,7 @@ class TestProductCache: ds=mnt_filepath, client=client_no_amm, bp_pem=brokerage_product_payout_event_manager, - redis_config=thl_redis_config, + redis_config=thl_redis_config: RedisConfig, ) # Fetch from cache and assert the instance loaded from redis @@ -1048,18 +1048,18 @@ class TestProductCache: def test_neg_balance_cache( self, - product: Product, + product: product: Product, mnt_filepath: GRLDatasets, thl_lm, client_no_amm: DaskClient, - thl_redis_config, + thl_redis_config: RedisConfig, brokerage_product_payout_event_manager, - delete_ledger_db, - create_main_accounts, - delete_df_collection, + delete_ledger_db: Callable[..., None], + create_main_accounts: Callable[..., None], + delete_df_collection: Callable[..., None], ledger_collection, user_factory: Callable[..., User], - session_with_tx_factory, + session_with_tx_factory: Callable[..., None], pop_ledger_merge: PopLedgerMerge, start: datetime, bp_payout_factory, @@ -1085,7 +1085,7 @@ class TestProductCache: # 2. Payout payout_event_manager.set_account_lookup_table(thl_lm=thl_lm) bp_payout_factory( - product=product, + product=product: Product, amount=USDCent(71), ext_ref_id=uuid4().hex, created=start + timedelta(days=1, minutes=1), @@ -1108,7 +1108,7 @@ class TestProductCache: ds=mnt_filepath, client=client_no_amm, bp_pem=brokerage_product_payout_event_manager, - redis_config=thl_redis_config, + redis_config=thl_redis_config: RedisConfig, ) # Fetch from cache and assert the instance loaded from redis diff --git a/tests/models/thl/test_user.py b/tests/models/thl/test_user.py index e0ba6f8..0b8634a 100644 --- a/tests/models/thl/test_user.py +++ b/tests/models/thl/test_user.py @@ -666,7 +666,11 @@ class TestUserMethods: assert len(user.audit_log) == 1 def test_transactions( - self, user_factory, thl_lm, session_with_tx_factory, product_user_wallet_yes + self, + user_factory: Callable[..., User], + thl_lm, + session_with_tx_factory: Callable[..., None], + product_user_wallet_yes, ): u1 = user_factory(product=product_user_wallet_yes) -- cgit v1.2.3 From aeeb7fef2594ccd34fbe96a77f6c5b392299fed7 Mon Sep 17 00:00:00 2001 From: Max Nanis Date: Thu, 27 Aug 2026 17:46:11 -0700 Subject: Ruff afternoon --- generalresearch/managers/thl/product.py | 1 - generalresearch/models/network/nmap/parser.py | 8 +- generalresearch/models/precision/question.py | 4 +- generalresearch/models/precision/survey.py | 17 +- generalresearch/models/prodege/survey.py | 3 +- generalresearch/models/spectrum/question.py | 9 +- generalresearch/models/spectrum/survey.py | 13 +- generalresearch/models/spectrum/task_collection.py | 2 +- generalresearch/models/thl/contest/contest.py | 10 +- generalresearch/models/thl/contest/leaderboard.py | 13 +- generalresearch/models/thl/contest/milestone.py | 20 +- generalresearch/models/thl/contest/raffle.py | 46 ++-- generalresearch/models/thl/finance.py | 12 +- generalresearch/models/thl/product.py | 36 +-- .../models/thl/profiling/other_option.py | 4 +- .../models/thl/profiling/upk_question.py | 13 +- .../models/thl/profiling/user_question_answer.py | 31 +-- pyproject.toml | 5 +- .../collections/test_df_collection_item_thl_web.py | 273 ++++++++++----------- .../mergers/foundations/test_enriched_session.py | 52 ++-- .../foundations/test_enriched_task_adjust.py | 38 ++- .../mergers/foundations/test_enriched_wall.py | 73 +++--- tests/incite/mergers/test_merge_collection.py | 53 +++- tests/incite/mergers/test_merge_collection_item.py | 25 +- tests/incite/mergers/test_pop_ledger.py | 109 ++++---- tests/incite/mergers/test_ym_survey_merge.py | 55 +++-- tests/incite/schemas/test_admin_responses.py | 36 +-- tests/incite/schemas/test_thl_web.py | 8 +- tests/incite/test_collection_base.py | 62 ++--- tests/incite/test_collection_base_item.py | 74 +++--- tests/incite/test_interval_idx.py | 6 +- tests/managers/gr/test_business.py | 60 +++-- tests/managers/gr/test_team.py | 58 +++-- tests/managers/network/test_label.py | 4 +- .../managers/thl/test_contest/test_leaderboard.py | 5 +- tests/managers/thl/test_contest/test_milestone.py | 5 +- tests/managers/thl/test_contest/test_raffle.py | 12 +- tests/managers/thl/test_ledger/test_lm_accounts.py | 2 +- tests/managers/thl/test_ledger/test_lm_tx.py | 4 - .../thl/test_ledger/test_thl_lm_accounts.py | 11 +- .../thl/test_ledger/test_thl_lm_bp_payout.py | 54 ++-- tests/managers/thl/test_ledger/test_thl_lm_tx.py | 26 +- tests/managers/thl/test_payout.py | 2 - tests/managers/thl/test_survey.py | 6 +- tests/managers/thl/test_user_manager/test_base.py | 2 +- tests/models/network/test_nmap.py | 2 +- tests/models/spectrum/test_survey.py | 44 ++++ tests/models/thl/test_product.py | 4 +- 48 files changed, 815 insertions(+), 597 deletions(-) (limited to 'tests/incite/collections') diff --git a/generalresearch/managers/thl/product.py b/generalresearch/managers/thl/product.py index d924e17..54fa7c8 100644 --- a/generalresearch/managers/thl/product.py +++ b/generalresearch/managers/thl/product.py @@ -33,7 +33,6 @@ if TYPE_CHECKING: ProfilingConfig, SessionConfig, SourcesConfig, - SupplyConfigs, UserCreateConfig, UserHealthConfig, UserWalletConfig, diff --git a/generalresearch/models/network/nmap/parser.py b/generalresearch/models/network/nmap/parser.py index ecaf2d1..866b4bd 100644 --- a/generalresearch/models/network/nmap/parser.py +++ b/generalresearch/models/network/nmap/parser.py @@ -48,7 +48,7 @@ class NmapXmlParser: try: root = ET.fromstring(nmap_data) - except Exception as e: + except ET.ParseError as e: emsg = f"Wrong XML structure: cannot parse data: {e}" raise NmapParserException(emsg) @@ -103,7 +103,7 @@ class NmapXmlParser: @classmethod def _parse_scaninfo(cls, scaninfo_el: ET.Element) -> NmapScanInfo: - data = dict() + data = {} data["type"] = NmapScanType(scaninfo_el.attrib["type"]) data["protocol"] = IPProtocol(scaninfo_el.attrib["protocol"]) data["num_services"] = scaninfo_el.attrib["numservices"] @@ -132,7 +132,7 @@ class NmapXmlParser: @classmethod def _parse_nmaprun(cls, nmaprun_el: ET.Element) -> dict: - nmap_data = dict() + nmap_data = {} nmaprun = dict(nmaprun_el.attrib) nmap_data["command_line"] = nmaprun["args"] nmap_data["started_at"] = datetime.fromtimestamp( @@ -148,7 +148,7 @@ class NmapXmlParser: Receives a XML tag representing a scanned host with its services. """ - data = dict() + data = {} # status_el = host_el.find("status") diff --git a/generalresearch/models/precision/question.py b/generalresearch/models/precision/question.py index a2189d5..cc90aa9 100644 --- a/generalresearch/models/precision/question.py +++ b/generalresearch/models/precision/question.py @@ -6,7 +6,7 @@ import logging from enum import StrEnum from typing import TYPE_CHECKING, Any, Literal -from pydantic import BaseModel, Field, field_validator, model_validator +from pydantic import BaseModel, Field, ValidationError, field_validator, model_validator from generalresearch.models import Source, string_utils from generalresearch.models.precision import PrecisionQuestionID @@ -112,7 +112,7 @@ class PrecisionQuestion(MarketplaceQuestion): """ try: return cls._from_api(d) - except Exception as e: + except ValidationError as e: logger.warning(f"Unable to parse question: {d}. {e}") return None diff --git a/generalresearch/models/precision/survey.py b/generalresearch/models/precision/survey.py index f515552..b27b8c4 100644 --- a/generalresearch/models/precision/survey.py +++ b/generalresearch/models/precision/survey.py @@ -99,11 +99,11 @@ class PrecisionQuota(BaseModel): self, criteria_evaluation: dict[str, bool | None] ) -> tuple[bool | None, list[str]]: # Passes back "matches" (T/F/none) and a list of unknown criterion hashes - unknowns = list() + unknowns = [] for c in self.condition_hashes: eval_value = criteria_evaluation.get(c) if eval_value is False: - return False, list() + return False, [] if eval_value is None: unknowns.append(c) if unknowns: @@ -245,11 +245,10 @@ class PrecisionSurvey(MarketplaceTask): # Fancy repr that abbreviates exclude_pids and excluded_surveys repr_args = list(self.__repr_args__()) for n, (k, v) in enumerate(repr_args): - if k in {"excluded_surveys"}: - if v and len(v) > 6: - v = sorted(v) - v = v[:3] + ["…"] + v[-3:] - repr_args[n] = (k, v) + if k in {"excluded_surveys"} and v and len(v) > 6: + v = sorted(v) + v = v[:3] + ["…"] + v[-3:] + repr_args[n] = (k, v) join_str = ", " repr_str = join_str.join( repr(v) if a is None else f"{a}={v!r}" for a, v in repr_args @@ -369,6 +368,4 @@ class PrecisionSurvey(MarketplaceTask): return False if self.group_id in att_group_ids: return False - if self.excluded_surveys & att_survey_ids: - return False - return True + return not self.excluded_surveys & att_survey_ids diff --git a/generalresearch/models/prodege/survey.py b/generalresearch/models/prodege/survey.py index 3f4c88f..7ab6df6 100644 --- a/generalresearch/models/prodege/survey.py +++ b/generalresearch/models/prodege/survey.py @@ -13,6 +13,7 @@ from pydantic import ( BaseModel, ConfigDict, Field, + ValidationError, computed_field, field_validator, model_validator, @@ -513,7 +514,7 @@ class ProdegeSurvey(MarketplaceTask): def from_api(cls, d: dict[str, Any]) -> ProdegeSurvey | None: try: return cls._from_api(d) - except Exception as e: + except ValidationError as e: logger.warning(f"Unable to parse survey: {d}. {e}") return None diff --git a/generalresearch/models/spectrum/question.py b/generalresearch/models/spectrum/question.py index c8eea4a..7add692 100644 --- a/generalresearch/models/spectrum/question.py +++ b/generalresearch/models/spectrum/question.py @@ -13,6 +13,7 @@ from pydantic import ( BaseModel, Field, PositiveInt, + ValidationError, field_validator, model_validator, ) @@ -132,7 +133,7 @@ class SpectrumQuestionType(StrEnum): @classmethod def from_api(cls, a: int): api_type_map = cls.get_api_map() - return api_type_map[a] if a in api_type_map else None + return api_type_map.get(a, None) class SpectrumQuestionClass(IntEnum): @@ -260,7 +261,7 @@ class SpectrumQuestion(MarketplaceQuestion): return None try: return cls._from_api(d, country_iso, language_iso) - except Exception as e: + except ValidationError as e: logger.warning(f"Unable to parse question: {d}. {e}") return None @@ -280,7 +281,9 @@ class SpectrumQuestion(MarketplaceQuestion): ] created = ( - datetime.utcfromtimestamp(d["crtd_on"] / 1000).replace(tzinfo=UTC) + datetime.fromtimestamp(timestamp=d["crtd_on"] / 1000, tz=UTC).replace( + tzinfo=UTC + ) if d.get("crtd_on") else None ) diff --git a/generalresearch/models/spectrum/survey.py b/generalresearch/models/spectrum/survey.py index f9a5e27..424d206 100644 --- a/generalresearch/models/spectrum/survey.py +++ b/generalresearch/models/spectrum/survey.py @@ -75,8 +75,7 @@ class SpectrumCondition(MarketplaceCondition): rs["from"] = round(rs["from"] / 12) rs["to"] = round(rs["to"] / 12) d["values"] = [ - f"{rs["from"] or "inf"}-{rs["to"] or "inf"}" - for rs in d["range_sets"] + f"{rs["from"] or "inf"}-{rs["to"] or "inf"}" for rs in d["range_sets"] ] d["value_type"] = ConditionValueType.RANGE return cls.model_validate(d) @@ -103,7 +102,7 @@ class SpectrumQuota(BaseModel): # There is no explicit status. The quota is closed if the count is 0 def __hash__(self) -> int: - return hash(tuple((tuple(self.condition_hashes), self.remaining_count))) + return hash((tuple(self.condition_hashes), self.remaining_count)) @property def is_open(self) -> bool: @@ -113,7 +112,7 @@ class SpectrumQuota(BaseModel): return self.remaining_count >= min_open_spots @classmethod - def from_api(cls, d: dict) -> Self: + def from_api(cls, d: dict[str, Any]) -> Self: d["remaining_count"] = d["quantities"]["currently_open"] return cls.model_validate(d) @@ -323,7 +322,7 @@ class SpectrumSurvey(MarketplaceTask): def from_api(cls, d: dict[str, Any]) -> SpectrumSurvey | None: try: return cls._from_api(d) - except Exception as e: + except (AssertionError, ValueError) as e: logger.warning(f"Unable to parse survey: {d}. {e}") return None @@ -336,7 +335,7 @@ class SpectrumSurvey(MarketplaceTask): else TaskCalculationType.COMPLETES ) - d["conditions"] = dict() + d["conditions"] = {} # If we haven't hit the "detail" endpoint, we won't get this d.setdefault("qualifications", []) @@ -454,7 +453,7 @@ class SpectrumSurvey(MarketplaceTask): quota_eval = { quota: quota.matches_soft(criteria_evaluation) for quota in self.quotas } - evals = set(g[0] for g in quota_eval.values()) + evals = {g[0] for g in quota_eval.values()} if any(m[0] is True and not q.is_open for q, m in quota_eval.items()): # matched a full quota return False, set() diff --git a/generalresearch/models/spectrum/task_collection.py b/generalresearch/models/spectrum/task_collection.py index 6715378..8ca5a93 100644 --- a/generalresearch/models/spectrum/task_collection.py +++ b/generalresearch/models/spectrum/task_collection.py @@ -91,7 +91,7 @@ class SpectrumTaskCollection(TaskCollection): "survey_id", ] rows = [] - d = dict() + d = {} for k in fields: d[k] = getattr(s, k) if hasattr(s, k) else None d["used_question_ids"] = list(s.used_question_ids) diff --git a/generalresearch/models/thl/contest/contest.py b/generalresearch/models/thl/contest/contest.py index 2a8853d..bd0fc04 100644 --- a/generalresearch/models/thl/contest/contest.py +++ b/generalresearch/models/thl/contest/contest.py @@ -136,10 +136,12 @@ class Contest(ContestBase): # return False def should_end(self) -> tuple[bool, ContestEndReason | None]: - if self.status == ContestStatus.ACTIVE: - if self.end_condition.ends_at: - if datetime.now(tz=UTC) >= self.end_condition.ends_at: - return True, ContestEndReason.ENDS_AT + if ( + self.status == ContestStatus.ACTIVE + and self.end_condition.ends_at + and datetime.now(tz=UTC) >= self.end_condition.ends_at + ): + return True, ContestEndReason.ENDS_AT return False, None diff --git a/generalresearch/models/thl/contest/leaderboard.py b/generalresearch/models/thl/contest/leaderboard.py index 696cdea..e923383 100644 --- a/generalresearch/models/thl/contest/leaderboard.py +++ b/generalresearch/models/thl/contest/leaderboard.py @@ -151,7 +151,8 @@ class LeaderboardContest(LeaderboardContestCreate, Contest): len(self.country_isos) == 1 ), "Can only set 1 country_iso in a leaderboard contest" assert ( - list(self.country_isos)[0] == self.leaderboard_key_parts["country_iso"] + next(iter(self.country_isos)) + == self.leaderboard_key_parts["country_iso"] ), "leaderboard_key country_iso must match the country_isos" else: self.country_isos = {self.leaderboard_key_parts["country_iso"]} @@ -192,10 +193,12 @@ class LeaderboardContest(LeaderboardContestCreate, Contest): return lbm def should_end(self) -> tuple[bool, ContestEndReason | None]: - if self.status == ContestStatus.ACTIVE: - if self.end_condition.ends_at: - if datetime.now(tz=UTC) >= self.end_condition.ends_at: - return True, ContestEndReason.ENDS_AT + if ( + self.status == ContestStatus.ACTIVE + and self.end_condition.ends_at + and datetime.now(tz=UTC) >= self.end_condition.ends_at + ): + return True, ContestEndReason.ENDS_AT return False, None diff --git a/generalresearch/models/thl/contest/milestone.py b/generalresearch/models/thl/contest/milestone.py index 8d96fcb..5fc27fa 100644 --- a/generalresearch/models/thl/contest/milestone.py +++ b/generalresearch/models/thl/contest/milestone.py @@ -132,10 +132,12 @@ class MilestoneContest(MilestoneContestCreate, Contest): if res: return res, msg - if self.status == ContestStatus.ACTIVE: - if self.end_condition.max_winners: - if self.win_count >= self.end_condition.max_winners: - return True, ContestEndReason.MAX_WINNERS + if ( + self.status == ContestStatus.ACTIVE + and self.end_condition.max_winners + and self.win_count >= self.end_condition.max_winners + ): + return True, ContestEndReason.MAX_WINNERS return False, None @@ -189,16 +191,10 @@ class MilestoneUserView(MilestoneContest, ContestUserView): ) def should_award(self): - if self.status == ContestStatus.ACTIVE: - if self.should_have_awarded(): - return True - return False + return bool(self.status == ContestStatus.ACTIVE and self.should_have_awarded()) def should_have_awarded(self): - if self.target_amount: - if self.user_amount >= self.target_amount: - return True - return False + return bool(self.target_amount and self.user_amount >= self.target_amount) def is_user_eligible(self, country_iso: str) -> tuple[bool, str]: passes, msg = super().is_user_eligible(country_iso=country_iso) diff --git a/generalresearch/models/thl/contest/raffle.py b/generalresearch/models/thl/contest/raffle.py index 08243f4..16a0a47 100644 --- a/generalresearch/models/thl/contest/raffle.py +++ b/generalresearch/models/thl/contest/raffle.py @@ -127,7 +127,7 @@ class RaffleContest(RaffleContestCreate, Contest): # If there is more than 1 prize, the winning entry is subtracted # from the user's entry count user_amount = defaultdict(int) - user_id_user = dict() + user_id_user = {} for entry in self.entries: user_amount[entry.user.user_id] += entry.amount user_id_user[entry.user.user_id] = entry.user @@ -149,10 +149,12 @@ class RaffleContest(RaffleContestCreate, Contest): res, msg = super().should_end() if res: return res, msg - if self.status == ContestStatus.ACTIVE: - if self.end_condition.target_entry_amount: - if self.current_amount >= self.end_condition.target_entry_amount: - return True, ContestEndReason.TARGET_ENTRY_AMOUNT + if ( + self.status == ContestStatus.ACTIVE + and self.end_condition.target_entry_amount + and self.current_amount >= self.end_condition.target_entry_amount + ): + return True, ContestEndReason.TARGET_ENTRY_AMOUNT return False, None @staticmethod @@ -278,17 +280,19 @@ class RaffleUserView(RaffleContest, ContestUserView): return probs def is_entry_eligible(self, entry: ContestEntry) -> tuple[bool, str]: - if self.entry_rule.max_entry_amount_per_user: - if ( - self.user_amount + entry.amount - ) > self.entry_rule.max_entry_amount_per_user: - return False, "Entry would exceed max amount per user." - - if self.entry_rule.max_daily_entries_per_user: - if ( - self.user_amount_today + entry.amount - ) > self.entry_rule.max_daily_entries_per_user: - return False, "Entry would exceed max amount per user per day." + if ( + self.entry_rule.max_entry_amount_per_user + and (self.user_amount + entry.amount) + > self.entry_rule.max_entry_amount_per_user + ): + return False, "Entry would exceed max amount per user." + + if ( + self.entry_rule.max_daily_entries_per_user + and (self.user_amount_today + entry.amount) + > self.entry_rule.max_daily_entries_per_user + ): + return False, "Entry would exceed max amount per user per day." return True, "" def is_user_eligible(self, country_iso: str) -> tuple[bool, str]: @@ -296,16 +300,18 @@ class RaffleUserView(RaffleContest, ContestUserView): if not passes: return False, msg - if self.entry_rule.max_entry_amount_per_user: + if self.entry_rule.max_entry_amount_per_user: # noqa: SIM102 # Greater or equal b/c we're asking if the user is eligible to # enter MORE, now! If it equals, nothing is wrong, just that they # are not eligible anymore. if self.user_amount >= self.entry_rule.max_entry_amount_per_user: return False, "Reached max amount per user." - if self.entry_rule.max_daily_entries_per_user: - if self.user_amount_today >= self.entry_rule.max_daily_entries_per_user: - return False, "Reached max amount today." + if ( + self.entry_rule.max_daily_entries_per_user + and self.user_amount_today >= self.entry_rule.max_daily_entries_per_user + ): + return False, "Reached max amount today." # This would indicate something is wrong, as something else should have done this e, _ = self.should_end() diff --git a/generalresearch/models/thl/finance.py b/generalresearch/models/thl/finance.py index 0856825..8c94390 100644 --- a/generalresearch/models/thl/finance.py +++ b/generalresearch/models/thl/finance.py @@ -124,10 +124,10 @@ class POPFinancial(BaseModel): Direction, ) - assert all([a.account_type == AccountType.BP_WALLET for a in accounts]) - assert all([a.normal_balance == Direction.CREDIT for a in accounts]) + assert all(a.account_type == AccountType.BP_WALLET for a in accounts) + assert all(a.normal_balance == Direction.CREDIT for a in accounts) if not is_debug(): - assert all([a.currency == "USD" for a in accounts]) + assert all(a.currency == "USD" for a in accounts) if input_data.empty: return [] @@ -850,11 +850,11 @@ class BusinessBalances(BaseModel): # Validate the input accounts assert len(accounts) > 0, "Must provide accounts" - assert all([a.account_type == AccountType.BP_WALLET for a in accounts]) - assert all([a.normal_balance == Direction.CREDIT for a in accounts]) + assert all(a.account_type == AccountType.BP_WALLET for a in accounts) + assert all(a.normal_balance == Direction.CREDIT for a in accounts) if not is_debug(): - assert all([a.currency == "USD" for a in accounts]) + assert all(a.currency == "USD" for a in accounts) # Validate the input dataframe assert input_data.index.name == "account_id" diff --git a/generalresearch/models/thl/product.py b/generalresearch/models/thl/product.py index 76a8e83..a7ecd55 100644 --- a/generalresearch/models/thl/product.py +++ b/generalresearch/models/thl/product.py @@ -430,8 +430,8 @@ class UserWalletConfig(BaseModel): @field_serializer("supported_payout_types", when_used="json") def serialize_supported_payout_types_in_order( self, supported_payout_types: set[PayoutType] - ) -> set[PayoutType]: - return set(sorted(supported_payout_types)) + ) -> list[PayoutType]: + return sorted(supported_payout_types) @field_validator("min_cashout", mode="after") @classmethod @@ -552,14 +552,14 @@ class PayoutTransformation(BaseModel): min_payout = Decimal(0) pct = Decimal(pct) - payout = Decimal(payout) + _payout = Decimal(payout) min_payout = Decimal(min_payout) max_payout = Decimal(max_payout) if max_payout else None - payout: Decimal = payout * pct - payout: Decimal = max([payout, min_payout]) - payout: Decimal = min([payout, max_payout]) if max_payout else payout - return payout + _payout: Decimal = _payout * pct + _payout: Decimal = max([_payout, min_payout]) + _payout: Decimal = min([_payout, max_payout]) if max_payout else payout + return _payout def payout_transformation_amt( self, payout: Decimal, user_wallet_balance: Decimal | None = None @@ -569,22 +569,22 @@ class PayoutTransformation(BaseModel): # (display, adjustment) so ignore the 7-cent rounding. if user_wallet_balance is None: return self.payout_transformation_percent(payout=payout, pct=Decimal(".95")) - payout = Decimal(payout) + _payout = Decimal(payout) - payout: Decimal = payout * Decimal("0.95") - new_balance = payout + user_wallet_balance + _payout: Decimal = _payout * Decimal("0.95") + new_balance = _payout + user_wallet_balance # If the new_balance is <0, we aren't paying anything, so use the # full amount if new_balance < 0: - return payout + return _payout amt = (5 * math.floor((int(new_balance * 100) - 2) / 5)) + 2 rounded_new_balance = Decimal(amt / 100).quantize(Decimal("0.00")) - payout = rounded_new_balance - user_wallet_balance - if payout < Decimal(0): + _payout = rounded_new_balance - user_wallet_balance + if _payout < Decimal(0): return Decimal(0) - return payout + return _payout class SourceConfig(BaseModel): @@ -731,8 +731,8 @@ class SupplyConfig(BaseModel): Use global config. """ d = self.global_scoped_policies_dict.copy() - d.update(self.team_scoped_policies_dict.get(team_id, dict())) - d.update(self.product_scoped_policies_dict.get(product_id, dict())) + d.update(self.team_scoped_policies_dict.get(team_id, {})) + d.update(self.product_scoped_policies_dict.get(product_id, {})) return d def get_config_for_product(self, product: Product) -> MergedSupplyConfig: @@ -751,7 +751,7 @@ class SupplyConfig(BaseModel): supply_policy=policy_dict[source], source_config=sources_dict[source], ) - for source in policy_dict.keys() + for source in policy_dict ] ) @@ -1000,10 +1000,12 @@ class Product(BaseModel, validate_assignment=True): @property def business_uuid(self) -> UUIDStr: + assert self.business_id return self.business_id @property def team_uuid(self) -> UUIDStr: + assert self.team_id return self.team_id @property diff --git a/generalresearch/models/thl/profiling/other_option.py b/generalresearch/models/thl/profiling/other_option.py index 6d789e5..2f3cac9 100644 --- a/generalresearch/models/thl/profiling/other_option.py +++ b/generalresearch/models/thl/profiling/other_option.py @@ -51,6 +51,4 @@ def option_is_catch_all(c: UpkQuestionChoice) -> bool: return True if c.text.lower() in texts_exact: return True - if any(t in c.text.lower() for t in texts_in): - return True - return False + return bool(any(t in c.text.lower() for t in texts_in)) diff --git a/generalresearch/models/thl/profiling/upk_question.py b/generalresearch/models/thl/profiling/upk_question.py index 77bba6f..3bb0733 100644 --- a/generalresearch/models/thl/profiling/upk_question.py +++ b/generalresearch/models/thl/profiling/upk_question.py @@ -475,10 +475,9 @@ class UpkQuestion(BaseModel): # Almost nothing has >1k options, besides location stuff (cities, # etc.) which should get harmonized. When presenting them, we'll # filter down options to at most 50. - if self.choices and (len(self.choices) <= 1 or len(self.choices) > 1000): - return False - - return True + return not ( + self.choices and (len(self.choices) <= 1 or len(self.choices) > 1000) + ) @property def md5sum(self): @@ -534,7 +533,7 @@ class UpkQuestion(BaseModel): ), "Multiple of the same answer submitted" if self.type == UpkQuestionType.MULTIPLE_CHOICE: assert len(answer) >= 1, "MC question with no selected answers" - choice_codes = set(x.id for x in self.choices) + choice_codes = {x.id for x in self.choices} if self.selector == UpkQuestionSelectorMC.SINGLE_ANSWER: assert ( len(answer) == 1 @@ -563,9 +562,7 @@ class UpkQuestion(BaseModel): assert len(answer) == 1, "Only one answer allowed" answer = answer[0] assert len(answer) > 0, "Must provide answer" - max_length = ( - self.configuration.max_length if self.configuration else 0 or 100000 - ) + max_length = self.configuration.max_length if self.configuration else 100000 assert len(answer) <= max_length, "Answer longer than allowed" if self.validation and self.validation.patterns: for pattern in self.validation.patterns: diff --git a/generalresearch/models/thl/profiling/user_question_answer.py b/generalresearch/models/thl/profiling/user_question_answer.py index 2db07b7..378345e 100644 --- a/generalresearch/models/thl/profiling/user_question_answer.py +++ b/generalresearch/models/thl/profiling/user_question_answer.py @@ -3,7 +3,7 @@ from __future__ import annotations import json from collections.abc import Iterator from datetime import UTC, datetime, timedelta -from typing import Any, Literal, Self +from typing import Any, Literal from pydantic import ( BaseModel, @@ -14,7 +14,6 @@ from pydantic import ( model_validator, ) -from generalresearch.grpc import timestamp_to_datetime from generalresearch.models import MAX_INT32, Source from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.thl.locales import CountryISO, LanguageISO @@ -39,19 +38,23 @@ class UserQuestionAnswer(BaseModel): calc_answers: dict[str, tuple[str, ...]] | None = Field(default=None) @field_validator("calc_answers") - def sorted_calc_answers(cls, calc_answers) -> dict[str, tuple[str, ...]] | None: + def sorted_calc_answers( + cls, calc_answers: dict[str, tuple[str, ...]] | None + ) -> dict[str, tuple[str, ...]] | None: if calc_answers is None: return None return {k: tuple(sorted(v)) for k, v in calc_answers.items()} @field_validator("calc_answers") - def validate_keys(cls, calc_answers) -> dict[str, tuple[str, ...]] | None: + def validate_keys( + cls, calc_answers: dict[str, tuple[str, ...]] | None + ) -> dict[str, tuple[str, ...]] | None: if calc_answers is None: return None assert all( - ":" in k for k in calc_answers.keys() + ":" in k for k in calc_answers ), "calc_answers expects the keys to be in format source:question_code" return calc_answers @@ -66,6 +69,7 @@ class UserQuestionAnswer(BaseModel): return d def get_mrpqs(self) -> Iterator[MarketplaceResearchProfileQuestion]: + assert self.calc_answers for k, v in self.calc_answers.items(): source, question_code = k.split(":", 1) yield MarketplaceResearchProfileQuestion( @@ -105,21 +109,6 @@ class UserQuestionAnswer(BaseModel): def is_stale(self) -> bool: return self.timestamp < datetime.now(tz=UTC) - timedelta(days=30) - @classmethod - def from_grpc(cls, msg, default_timestamp: datetime) -> Self: - """ - Handles correctly issues with grpc timestamps - :param msg: "thl.protos.generalresearch_pb2.ProfilingQuestionAnswer" - """ - assert default_timestamp.tzinfo is not None, "must use tz-aware timestamps" - timestamp = timestamp_to_datetime(msg.timestamp) - timestamp = default_timestamp if timestamp < datetime(2000, 1, 1) else timestamp - return cls( - question_id=msg.question_id, - answer=tuple(msg.answer), - timestamp=timestamp, - ) - # We can't set a redis list to [] vs None. We'll push this dummy answer into # the cache to signify the user has no answered questions. It'll get removed @@ -131,7 +120,7 @@ DUMMY_UQA = UserQuestionAnswer( country_iso="xx", language_iso="xxx", property_code="dummy", - calc_answers=dict(), + calc_answers={}, ) diff --git a/pyproject.toml b/pyproject.toml index dbdf3b9..03a1a1f 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -56,4 +56,7 @@ testpaths = ["tests"] addopts = "-v --tb=short" [tool.ruff] -target-version = "py314" \ No newline at end of file +target-version = "py314" +exclude = [ + "generalresearch/thl_django", +] \ No newline at end of file diff --git a/tests/incite/collections/test_df_collection_item_thl_web.py b/tests/incite/collections/test_df_collection_item_thl_web.py index 8038d3b..edf90f7 100644 --- a/tests/incite/collections/test_df_collection_item_thl_web.py +++ b/tests/incite/collections/test_df_collection_item_thl_web.py @@ -5,12 +5,12 @@ from datetime import UTC, datetime, timedelta from itertools import product as iter_product from os.path import join as pjoin from pathlib import Path, PurePath -from typing import TYPE_CHECKING from uuid import uuid4 import dask.dataframe as dd import pandas as pd import pytest +from dask.distributed import Client as DaskClient from distributed import Client, Scheduler, Worker # noinspection PyUnresolvedReferences @@ -21,20 +21,19 @@ from faker import Faker from pandera.pandas import DataFrameSchema from pydantic import FilePath -from generalresearch.incite.base import CollectionItemBase +from generalresearch.incite.base import CollectionItemBase, GRLDatasets from generalresearch.incite.collections import ( + DFCollection, DFCollectionItem, DFCollectionType, ) from generalresearch.incite.schemas import ARCHIVE_AFTER +from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager from generalresearch.models.thl.product import Product from generalresearch.models.thl.user import User from generalresearch.pg_helper import PostgresConfig from generalresearch.sql_helper import PostgresDsn -if TYPE_CHECKING: - from generalresearch.incite.base import GRLDatasets - fake = Faker() df_collections = [ @@ -72,7 +71,12 @@ class TestDFCollectionItemBase: ) class TestDFCollectionItemProperties: - def test_filename(self, df_collection_data_type, df_collection, offset: str): + def test_filename( + self, + df_collection_data_type: DFCollectionType, + df_collection: DFCollection, + offset: str, + ): for i in df_collection.items: assert isinstance(i.filename, str) @@ -89,37 +93,59 @@ class TestDFCollectionItemProperties: ) class TestDFCollectionItemPropertiesBase: - def test_name(self, df_collection_data_type, offset: str, df_collection): + def test_name( + self, + df_collection: DFCollection, + ): for i in df_collection.items: assert isinstance(i.name, str) - def test_finish(self, df_collection_data_type, offset: str, df_collection): + def test_finish( + self, + df_collection: DFCollection, + ): for i in df_collection.items: assert isinstance(i.finish, datetime) - def test_interval(self, df_collection_data_type, offset: str, df_collection): + def test_interval( + self, + df_collection: DFCollection, + ): for i in df_collection.items: assert isinstance(i.interval, pd.Interval) def test_partial_filename( - self, df_collection_data_type, offset: str, df_collection + self, + df_collection: DFCollection, ): for i in df_collection.items: assert isinstance(i.partial_filename, str) - def test_empty_filename(self, df_collection_data_type, offset: str, df_collection): + def test_empty_filename( + self, + df_collection: DFCollection, + ): for i in df_collection.items: assert isinstance(i.empty_filename, str) - def test_path(self, df_collection_data_type, offset: str, df_collection): + def test_path( + self, + df_collection: DFCollection, + ): for i in df_collection.items: assert isinstance(i.path, FilePath) - def test_partial_path(self, df_collection_data_type, offset: str, df_collection): + def test_partial_path( + self, + df_collection: DFCollection, + ): for i in df_collection.items: assert isinstance(i.partial_path, FilePath) - def test_empty_path(self, df_collection_data_type, offset: str, df_collection): + def test_empty_path( + self, + df_collection: DFCollection, + ): for i in df_collection.items: assert isinstance(i.empty_path, FilePath) @@ -138,11 +164,8 @@ class TestDFCollectionItemMethod: def test_has_mysql( self, - df_collection, + df_collection: DFCollection, thl_web_rr: PostgresConfig, - offset: str, - duration: timedelta, - df_collection_data_type, delete_df_collection: Callable[..., None], ): delete_df_collection(coll=df_collection) @@ -168,12 +191,6 @@ class TestDFCollectionItemMethod: @pytest.mark.skip def test_update_partial_archive( self, - df_collection, - offset: str, - duration: timedelta, - thl_web_rw: PostgresConfig, - df_collection_data_type, - delete_df_collection: Callable[..., None], ): # for i in collection.items: # assert i.update_partial_archive() @@ -183,28 +200,12 @@ class TestDFCollectionItemMethod: @pytest.mark.skip def test_create_partial_archive( self, - df_collection, - offset: str, - duration: str, - create_main_accounts: Callable[..., None], - thl_web_rw: PostgresConfig, - thl_lm, - df_collection_data_type, - user_factory: Callable[..., User], - product: product: Product, - client_no_amm, - incite_item_factory, - delete_df_collection: Callable[..., None], - mnt_filepath: GRLDatasets, ): assert 1 + 1 == 2 def test_dict( self, - df_collection_data_type, - offset: str, - duration: timedelta, - df_collection, + df_collection: DFCollection, delete_df_collection: Callable[..., None], ): delete_df_collection(coll=df_collection) @@ -225,15 +226,15 @@ class TestDFCollectionItemMethod: def test_from_mysql( self, - df_collection_data_type, - df_collection, + df_collection_data_type: DFCollectionType, + df_collection: DFCollection, offset: str, duration: timedelta, create_main_accounts: Callable[..., None], thl_web_rw: PostgresConfig, user_factory: Callable[..., User], - product: product: Product, - incite_item_factory, + product: Product, + incite_item_factory: Callable[..., None], delete_df_collection: Callable[..., None], ): @@ -253,12 +254,14 @@ class TestDFCollectionItemMethod: if df_collection.data_type == DFCollectionType.LEDGER: assert df is None else: + assert isinstance(df, pd.DataFrame) assert df.empty assert set(df.columns) == set(df_collection._schema.columns.keys()) incite_item_factory(user=u1, item=item) df = item.from_mysql() + assert isinstance(df, pd.DataFrame) assert not df.empty assert set(df.columns) == set(df_collection._schema.columns.keys()) if df_collection.data_type == DFCollectionType.LEDGER: @@ -270,13 +273,13 @@ class TestDFCollectionItemMethod: def test_from_mysql_standard( self, - df_collection_data_type, - df_collection, + df_collection_data_type: DFCollectionType, + df_collection: DFCollection, offset: str, duration: timedelta, user_factory: Callable[..., User], - product: product: Product, - incite_item_factory, + product: Product, + incite_item_factory: Callable[..., None], delete_df_collection: Callable[..., None], ): @@ -293,7 +296,7 @@ class TestDFCollectionItemMethod: # We're using parametrize, so this If statement is just to # confirm other Item Types will always raise an assertion with pytest.raises(expected_exception=AssertionError) as cm: - res = item.from_mysql_standard() + _ = item.from_mysql_standard() assert ( "Can't call from_mysql_standard for Ledger DFCollectionItem" in str(cm.value) @@ -304,32 +307,34 @@ class TestDFCollectionItemMethod: # Unlike .from_mysql_ledger(), .from_mysql_standard() will return # back and empty df with the correct columns in place df = item.from_mysql_standard() + assert isinstance(df, pd.DataFrame) assert df.empty assert set(df.columns) == set(df_collection._schema.columns.keys()) incite_item_factory(user=u1, item=item) df = item.from_mysql_standard() + assert isinstance(df, pd.DataFrame) assert not df.empty assert set(df.columns) == set(df_collection._schema.columns.keys()) assert df.shape[0] > 0 def test_from_mysql_ledger( self, - df_collection, + df_collection: DFCollection, user: User, create_main_accounts: Callable[..., None], offset: str, duration: timedelta, thl_web_rw: PostgresConfig, - thl_lm, - df_collection_data_type, + thl_ledger_manager: ThlLedgerManager, + df_collection_data_type: DFCollectionType, user_factory: Callable[..., User], - product: product: Product, - client_no_amm, - incite_item_factory, + product: Product, + client_no_amm: DaskClient, + incite_item_factory: Callable[..., None], delete_df_collection: Callable[..., None], - mnt_filepath, + mnt_filepath: GRLDatasets, ): if df_collection.data_type != DFCollectionType.LEDGER: @@ -370,17 +375,17 @@ class TestDFCollectionItemMethod: def test_to_archive( self, - df_collection, + df_collection: DFCollection, user: User, offset: str, duration: timedelta, - df_collection_data_type, + df_collection_data_type: DFCollectionType, user_factory: Callable[..., User], - product: product: Product, - client_no_amm, - incite_item_factory, + product: Product, + client_no_amm: DaskClient, + incite_item_factory: Callable[..., None], delete_df_collection: Callable[..., None], - mnt_filepath, + mnt_filepath: GRLDatasets, ): if df_collection.data_type in unsupported_mock_types: @@ -407,17 +412,17 @@ class TestDFCollectionItemMethod: def test__to_archive( self, - df_collection_data_type, - df_collection, + df_collection_data_type: DFCollectionType, + df_collection: DFCollection, user_factory: Callable[..., User], - product: product: Product, + product: Product, offset: str, duration: timedelta, - client_no_amm, + client_no_amm: DaskClient, user: User, - incite_item_factory, + incite_item_factory: Callable[..., None], delete_df_collection: Callable[..., None], - mnt_filepath, + mnt_filepath: GRLDatasets, ): """We already have a test for the "non-private" version of this, which primarily just uses the respective Client to determine if @@ -480,19 +485,19 @@ class TestDFCollectionItemMethod: @pytest.mark.skip def test_to_archive_numbered_partial( - self, df_collection_data_type, df_collection, offset: str, duration: timedelta + self, ): pass @pytest.mark.skip def test_initial_load( - self, df_collection_data_type, df_collection, offset: str, duration: timedelta + self, ): pass @pytest.mark.skip def test_clear_corrupt_archive( - self, df_collection_data_type, df_collection, offset: str, duration: timedelta + self, ): pass @@ -505,34 +510,40 @@ class TestDFCollectionItemMethodBase: @pytest.mark.skip def test_path_exists( - self, df_collection_data_type, offset: str, duration: timedelta + self, ): pass @pytest.mark.skip def test_next_numbered_path( - self, df_collection_data_type, offset: str, duration: timedelta + self, ): pass @pytest.mark.skip def test_search_highest_numbered_path( - self, df_collection_data_type, offset: str, duration: timedelta + self, + df_collection_data_type: DFCollectionType, + offset: str, + duration: timedelta, ): pass @pytest.mark.skip def test_tmp_filename( - self, df_collection_data_type, offset: str, duration: timedelta + self, ): pass @pytest.mark.skip - def test_tmp_path(self, df_collection_data_type, offset: str, duration: timedelta): + def test_tmp_path( + self, + ): pass def test_is_empty( - self, df_collection_data_type, df_collection, offset: str, duration: timedelta + self, + df_collection: DFCollection, ): """ test_has_empty was merged into this because item.has_empty is @@ -549,7 +560,8 @@ class TestDFCollectionItemMethodBase: assert item.has_empty() def test_has_partial_archive( - self, df_collection_data_type, df_collection, offset: str, duration: timedelta + self, + df_collection: DFCollection, ): for item in df_collection.items: assert not item.has_partial_archive() @@ -557,7 +569,8 @@ class TestDFCollectionItemMethodBase: assert item.has_partial_archive() def test_has_archive( - self, df_collection_data_type, df_collection, offset: str, duration: timedelta + self, + df_collection: DFCollection, ): for item in df_collection.items: # (1) Originally, nothing exists... so let's just make a file and @@ -594,7 +607,8 @@ class TestDFCollectionItemMethodBase: assert item.has_archive(include_empty=True) def test_delete_archive( - self, df_collection_data_type, df_collection, offset: str, duration: timedelta + self, + df_collection: DFCollection, ): for item in df_collection.items: item: DFCollectionItem @@ -617,7 +631,8 @@ class TestDFCollectionItemMethodBase: assert not item.partial_path.exists() def test_should_archive( - self, df_collection_data_type, df_collection, offset: str, duration: timedelta + self, + df_collection: DFCollection, ): schema: DataFrameSchema = df_collection._schema aa = schema.metadata[ARCHIVE_AFTER] @@ -635,12 +650,13 @@ class TestDFCollectionItemMethodBase: @pytest.mark.skip def test_set_empty( - self, df_collection_data_type, df_collection, offset: str, duration: timedelta + self, ): pass def test_valid_archive( - self, df_collection_data_type, df_collection, offset: str, duration: timedelta + self, + df_collection: DFCollection, ): # Originally, nothing has been saved or anything.. so confirm it # always comes back as None @@ -664,18 +680,19 @@ class TestDFCollectionItemMethodBase: @pytest.mark.skip def test_validate_df( - self, df_collection_data_type, df_collection, offset: str, duration: timedelta + self, ): pass @pytest.mark.skip def test_from_archive( - self, df_collection_data_type, df_collection, offset: str, duration: timedelta + self, ): pass def test__to_dict( - self, df_collection_data_type, df_collection, offset: str, duration: timedelta + self, + df_collection: DFCollection, ): for item in df_collection.items: @@ -694,19 +711,19 @@ class TestDFCollectionItemMethodBase: @pytest.mark.skip def test_delete_partial( - self, df_collection_data_type, df_collection, offset: str, duration: timedelta + self, ): pass @pytest.mark.skip def test_cleanup_partials( - self, df_collection_data_type, df_collection, offset: str, duration: timedelta + self, ): pass @pytest.mark.skip def test_delete_dangling_partials( - self, df_collection_data_type, df_collection, offset: str, duration: timedelta + self, ): pass @@ -726,7 +743,9 @@ async def test_client(client, s, worker): ) @gen_cluster(client=True, nthreads=[("127.0.0.1", 1)]) @pytest.mark.anyio -async def test_client_parametrize(c, s, w, df_collection_data_type, offset: str): +async def test_client_parametrize( + c, s, w, df_collection_data_type: DFCollectionType, offset: str +): """c,s,a are all required - the secondary Worker (b) is not required""" assert isinstance(c, Client), f"c is not Client, it's {type(c)}" @@ -750,17 +769,12 @@ class TestDFCollectionItemFunctionalTest: def test_to_archive_and_ddf( self, - df_collection_data_type, - offset: str, - duration: timedelta, - client_no_amm, - df_collection, - user: User, + client_no_amm: DaskClient, + df_collection: DFCollection, user_factory: Callable[..., User], - product: product: Product, - incite_item_factory, + product: Product, + incite_item_factory: Callable[..., None], delete_df_collection: Callable[..., None], - mnt_filepath: GRLDatasets, ): if df_collection.data_type in unsupported_mock_types: @@ -799,17 +813,11 @@ class TestDFCollectionItemFunctionalTest: def test_filesize_estimate( self, - df_collection, - user: User, - offset: str, - duration: timedelta, - client_no_amm, + df_collection: DFCollection, user_factory: Callable[..., User], - product: product: Product, - df_collection_data_type, - incite_item_factory, + product: Product, + incite_item_factory: Callable[..., None], delete_df_collection: Callable[..., None], - mnt_filepath: GRLDatasets, ): """A functional test to write some Parquet files for the DFCollection and then confirm that the files get written @@ -846,16 +854,12 @@ class TestDFCollectionItemFunctionalTest: def test_to_archive_client( self, - client_no_amm, - df_collection, + client_no_amm: DaskClient, + df_collection: DFCollection, user_factory: Callable[..., User], - product: product: Product, - offset: str, - duration: timedelta, - df_collection_data_type, - incite_item_factory, + product: Product, + incite_item_factory: Callable[..., None], delete_df_collection: Callable[..., None], - mnt_filepath: GRLDatasets, ): delete_df_collection(coll=df_collection) @@ -885,7 +889,8 @@ class TestDFCollectionItemFunctionalTest: @pytest.mark.skip def test_get_items( - self, df_collection, product: product: Product, offset: str, duration: timedelta + self, + df_collection: DFCollection, ): with pytest.warns(expected_warning=ResourceWarning) as cm: df_collection.get_items_last365() @@ -898,16 +903,11 @@ class TestDFCollectionItemFunctionalTest: def test_saving_protections( self, - client_no_amm, - df_collection_data_type, - df_collection, - incite_item_factory, + df_collection: DFCollection, + incite_item_factory: Callable[..., None], delete_df_collection: Callable[..., None], user_factory: Callable[..., User], - product: product: Product, - offset: str, - duration: timedelta, - mnt_filepath: GRLDatasets, + product: Product, ): """Don't allow creating an archive for data that will likely be overwritten or updated @@ -939,15 +939,8 @@ class TestDFCollectionItemFunctionalTest: def test_empty_item( self, - client_no_amm, - df_collection_data_type, - df_collection, - incite_item_factory, + df_collection: DFCollection, delete_df_collection: Callable[..., None], - user: User, - offset: str, - duration: timedelta, - mnt_filepath: GRLDatasets, ): delete_df_collection(coll=df_collection) @@ -967,16 +960,12 @@ class TestDFCollectionItemFunctionalTest: def test_file_touching( self, - client_no_amm, - df_collection_data_type, - df_collection, - incite_item_factory, + client_no_amm: DaskClient, + df_collection: DFCollection, + incite_item_factory: Callable[..., None], delete_df_collection: Callable[..., None], user_factory: Callable[..., User], - product: product: Product, - offset: str, - duration: timedelta, - mnt_filepath, + product: Product, ): delete_df_collection(coll=df_collection) diff --git a/tests/incite/mergers/foundations/test_enriched_session.py b/tests/incite/mergers/foundations/test_enriched_session.py index 8254d81..2a161e4 100644 --- a/tests/incite/mergers/foundations/test_enriched_session.py +++ b/tests/incite/mergers/foundations/test_enriched_session.py @@ -1,3 +1,6 @@ +from __future__ import annotations + +from collections.abc import Callable from datetime import UTC, datetime, timedelta from decimal import Decimal from itertools import product @@ -5,10 +8,24 @@ from itertools import product import dask.dataframe as dd import pandas as pd import pytest +from dask.distributed import Client as DaskClient +from generalresearch.incite.collections.thl_web import ( + SessionDFCollection, + WallDFCollection, +) +from generalresearch.incite.mergers.foundations.enriched_session import ( + EnrichedSessionMerge, +) from generalresearch.incite.schemas.admin_responses import ( AdminPOPSessionSchema, ) +from generalresearch.models.admin.request import ( + ReportRequest, +) +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 @@ -25,21 +42,20 @@ class TestEnrichedSession: def test_base( self, - client_no_amm, + client_no_amm: DaskClient, product: Product, user_factory: Callable[..., User], - wall_collection, - session_collection, - enriched_session_merge, + wall_collection: WallDFCollection, + session_collection: SessionDFCollection, + enriched_session_merge: EnrichedSessionMerge, thl_web_rr: PostgresConfig, delete_df_collection: Callable[..., None], - incite_item_factory, + incite_item_factory: Callable[..., None], ): - from generalresearch.models.thl.user import User delete_df_collection(coll=session_collection) - u1: User = user_factory(product=product: Product, created=session_collection.start) + u1: User = user_factory(product=product, created=session_collection.start) for item in session_collection.items: incite_item_factory(item=item, user=u1) @@ -52,7 +68,7 @@ class TestEnrichedSession: client=client_no_amm, wall_coll=wall_collection, session_coll=session_collection, - pg_config=thl_web_rr: PostgresConfig, + pg_config=thl_web_rr, ) # -- @@ -85,16 +101,16 @@ class TestEnrichedSessionAdmin: def test_to_admin_response( self, - event_report_request, - enriched_session_merge, - client_no_amm, - wall_collection, - session_collection, + event_report_request: ReportRequest, + enriched_session_merge: EnrichedSessionMerge, + client_no_amm: DaskClient, + wall_collection: WallDFCollection, + session_collection: SessionDFCollection, thl_web_rr: PostgresConfig, - session_report_request, + session_report_request: ReportRequest, user_factory: Callable[..., User], - start, - session_factory, + start: datetime, + session_factory: Callable[..., Session], product_factory: Callable[..., Product], delete_df_collection: Callable[..., None], ): @@ -107,7 +123,7 @@ class TestEnrichedSessionAdmin: for p in [p1, p2]: u = user_factory(product=p) for i in range(50): - s = session_factory( + _ = session_factory( user=u, wall_count=1, wall_req_cpi=Decimal("1.00"), @@ -120,7 +136,7 @@ class TestEnrichedSessionAdmin: client=client_no_amm, session_coll=session_collection, wall_coll=wall_collection, - pg_config=thl_web_rr: PostgresConfig, + pg_config=thl_web_rr, ) df = enriched_session_merge.to_admin_response( diff --git a/tests/incite/mergers/foundations/test_enriched_task_adjust.py b/tests/incite/mergers/foundations/test_enriched_task_adjust.py index a33a55a..0606b6f 100644 --- a/tests/incite/mergers/foundations/test_enriched_task_adjust.py +++ b/tests/incite/mergers/foundations/test_enriched_task_adjust.py @@ -1,9 +1,28 @@ +from __future__ import annotations + +from collections.abc import Callable from datetime import timedelta from itertools import product as iter_product import dask.dataframe as dd import pandas as pd import pytest +from dask.distributed import Client as DaskClient + +from generalresearch.incite.collections.thl_web import ( + SessionDFCollection, + TaskAdjustmentDFCollection, + WallDFCollection, +) +from generalresearch.incite.mergers.foundations.enriched_task_adjust import ( + EnrichedTaskAdjustMerge, +) +from generalresearch.incite.mergers.foundations.enriched_wall import ( + EnrichedWallMerge, +) +from generalresearch.models.thl.product import Product +from generalresearch.models.thl.user import User +from generalresearch.pg_helper import PostgresConfig @pytest.mark.parametrize( @@ -20,19 +39,18 @@ class TestEnrichedTaskAdjust: @pytest.mark.skip def test_base( self, - client_no_amm, + client_no_amm: DaskClient, user_factory: Callable[..., User], product: Product, - task_adj_collection, - wall_collection, - session_collection, - enriched_wall_merge, - enriched_task_adjust_merge, - incite_item_factory, + task_adj_collection: TaskAdjustmentDFCollection, + wall_collection: WallDFCollection, + session_collection: SessionDFCollection, + enriched_wall_merge: EnrichedWallMerge, + enriched_task_adjust_merge: EnrichedTaskAdjustMerge, + incite_item_factory: Callable[..., None], delete_df_collection: Callable[..., None], thl_web_rr: PostgresConfig, ): - from generalresearch.models.thl.user import User # -- Build & Setup delete_df_collection(coll=session_collection) @@ -48,14 +66,14 @@ class TestEnrichedTaskAdjust: client=client_no_amm, session_coll=session_collection, wall_coll=wall_collection, - pg_config=thl_web_rr: PostgresConfig, + pg_config=thl_web_rr, ) enriched_task_adjust_merge.build( client=client_no_amm, task_adjust_coll=task_adj_collection, enriched_wall=enriched_wall_merge, - pg_config=thl_web_rr: PostgresConfig, + pg_config=thl_web_rr, ) # -- diff --git a/tests/incite/mergers/foundations/test_enriched_wall.py b/tests/incite/mergers/foundations/test_enriched_wall.py index a0ca4dd..0cb8f60 100644 --- a/tests/incite/mergers/foundations/test_enriched_wall.py +++ b/tests/incite/mergers/foundations/test_enriched_wall.py @@ -1,3 +1,4 @@ +from collections.abc import Callable from datetime import UTC, datetime, timedelta from decimal import Decimal from itertools import product as iter_product @@ -5,11 +6,23 @@ from itertools import product as iter_product import dask.dataframe as dd import pandas as pd import pytest +from dask.distributed import Client as DaskClient + +from generalresearch.incite.collections.thl_web import ( + SessionDFCollection, + WallDFCollection, +) # noinspection PyUnresolvedReferences from generalresearch.incite.mergers.foundations.enriched_wall import ( + EnrichedWallMerge, EnrichedWallMergeItem, ) +from generalresearch.models.admin.request import ReportRequest +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 @pytest.mark.parametrize( @@ -20,22 +33,21 @@ class TestEnrichedWall: def test_base( self, - client_no_amm, + client_no_amm: DaskClient, product: Product, user_factory: Callable[..., User], - wall_collection, + wall_collection: WallDFCollection, thl_web_rr: PostgresConfig, - session_collection, - enriched_wall_merge, + session_collection: SessionDFCollection, + enriched_wall_merge: EnrichedWallMerge, delete_df_collection: Callable[..., None], - incite_item_factory, + incite_item_factory: Callable[..., None], ): - from generalresearch.models.thl.user import User # -- Build & Setup delete_df_collection(coll=session_collection) delete_df_collection(coll=wall_collection) - u1: User = user_factory(product=product: Product, created=session_collection.start) + u1: User = user_factory(product=product, created=session_collection.start) for item in session_collection.items: incite_item_factory(item=item, user=u1) @@ -48,7 +60,7 @@ class TestEnrichedWall: client=client_no_amm, wall_coll=wall_collection, session_coll=session_collection, - pg_config=thl_web_rr: PostgresConfig, + pg_config=thl_web_rr, ) # -- @@ -63,19 +75,19 @@ class TestEnrichedWall: def test_base_item( self, - client_no_amm, + client_no_amm: DaskClient, product: Product, user_factory: Callable[..., User], - wall_collection, - session_collection, - enriched_wall_merge, + wall_collection: WallDFCollection, + session_collection: SessionDFCollection, + enriched_wall_merge: EnrichedWallMerge, delete_df_collection: Callable[..., None], thl_web_rr: PostgresConfig, - incite_item_factory, + incite_item_factory: Callable[..., None], ): # -- Build & Setup delete_df_collection(coll=session_collection) - u = user_factory(product=product: Product, created=session_collection.start) + u = user_factory(product=product, created=session_collection.start) for item in session_collection.items: incite_item_factory(item=item, user=u) @@ -87,7 +99,7 @@ class TestEnrichedWall: client=client_no_amm, wall_coll=wall_collection, session_coll=session_collection, - pg_config=thl_web_rr: PostgresConfig, + pg_config=thl_web_rr, ) # -- @@ -99,14 +111,14 @@ class TestEnrichedWall: try: modified_time1 = path.stat().st_mtime - except Exception: + except OSError: modified_time1 = 0 item.build( client=client_no_amm, wall_coll=wall_collection, session_coll=session_collection, - pg_config=thl_web_rr: PostgresConfig, + pg_config=thl_web_rr, ) modified_time2 = path.stat().st_mtime @@ -150,7 +162,12 @@ class TestEnrichedWallToAdmin: def duration(self) -> timedelta | None: return timedelta(days=5) - def test_empty(self, enriched_wall_merge, client_no_amm, start): + def test_empty( + self, + enriched_wall_merge: EnrichedWallMerge, + client_no_amm: DaskClient, + start: datetime, + ): from generalresearch.models.admin.request import ReportRequest rr = ReportRequest.model_validate({"interval": "5min", "start": start}) @@ -167,18 +184,18 @@ class TestEnrichedWallToAdmin: def test_to_admin_response( self, - event_report_request, - enriched_wall_merge, - client_no_amm, - wall_collection, - session_collection, + event_report_request: ReportRequest, + enriched_wall_merge: EnrichedWallMerge, + client_no_amm: DaskClient, + wall_collection: WallDFCollection, + session_collection: SessionDFCollection, thl_web_rr: PostgresConfig, - user, - session_factory, + user: User, + session_factory: Callable[..., Session], delete_df_collection: Callable[..., None], product_factory: Callable[..., Product], user_factory: Callable[..., User], - start, + start: datetime, ): delete_df_collection(coll=wall_collection) delete_df_collection(coll=session_collection) @@ -189,7 +206,7 @@ class TestEnrichedWallToAdmin: for p in [p1, p2]: u = user_factory(product=p) for i in range(50): - s = session_factory( + _ = session_factory( user=u, wall_count=2, wall_req_cpi=Decimal("1.00"), @@ -203,7 +220,7 @@ class TestEnrichedWallToAdmin: client=client_no_amm, wall_coll=wall_collection, session_coll=session_collection, - pg_config=thl_web_rr: PostgresConfig, + pg_config=thl_web_rr, ) df = enriched_wall_merge.to_admin_response( diff --git a/tests/incite/mergers/test_merge_collection.py b/tests/incite/mergers/test_merge_collection.py index 15fa4db..cf8315f 100644 --- a/tests/incite/mergers/test_merge_collection.py +++ b/tests/incite/mergers/test_merge_collection.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from datetime import UTC, datetime, timedelta from itertools import product @@ -5,12 +7,13 @@ import pandas as pd import pytest from pandera.pandas import DataFrameSchema +from generalresearch.incite.base import GRLDatasets from generalresearch.incite.mergers import ( MergeCollection, MergeType, ) -merge_types = list(e for e in MergeType if e != MergeType.TEST) +merge_types = [e for e in MergeType if e != MergeType.TEST] @pytest.mark.parametrize( @@ -26,7 +29,11 @@ merge_types = list(e for e in MergeType if e != MergeType.TEST) ) class TestMergeCollection: - def test_init(self, mnt_filepath, merge_type, offset, duration, start): + def test_init( + self, + mnt_filepath: GRLDatasets, + merge_type: MergeType, + ): with pytest.raises(expected_exception=ValueError) as cm: MergeCollection(archive_path=mnt_filepath.data_src) assert "Must explicitly provide a merge_type" in str(cm.value) @@ -37,7 +44,14 @@ class TestMergeCollection: ) assert instance.merge_type == merge_type - def test_items(self, mnt_filepath, merge_type, offset, duration, start): + def test_items( + self, + mnt_filepath: GRLDatasets, + merge_type: MergeType, + offset: str, + duration: timedelta, + start: datetime, + ): instance = MergeCollection( merge_type=merge_type, offset=offset, @@ -48,7 +62,14 @@ class TestMergeCollection: assert len(instance.interval_range) == len(instance.items) - def test_progress(self, mnt_filepath, merge_type, offset, duration, start): + def test_progress( + self, + mnt_filepath: GRLDatasets, + merge_type: MergeType, + offset: str, + duration: timedelta, + start: datetime, + ): instance = MergeCollection( merge_type=merge_type, offset=offset, @@ -62,7 +83,11 @@ class TestMergeCollection: assert instance.progress.shape[1] == 7 assert instance.progress["group_by"].isnull().all() - def test_schema(self, mnt_filepath, merge_type, offset, duration, start): + def test_schema( + self, + mnt_filepath: GRLDatasets, + merge_type: MergeType, + ): instance = MergeCollection( merge_type=merge_type, archive_path=mnt_filepath.archive_path(enum_type=merge_type), @@ -70,7 +95,14 @@ class TestMergeCollection: assert isinstance(instance._schema, DataFrameSchema) - def test_load(self, mnt_filepath, merge_type, offset, duration, start): + def test_load( + self, + mnt_filepath: GRLDatasets, + merge_type: MergeType, + offset: str, + duration: timedelta, + start: datetime, + ): instance = MergeCollection( merge_type=merge_type, start=start, @@ -82,7 +114,14 @@ class TestMergeCollection: # Confirm that there are no archives available yet assert instance.progress.has_archive.eq(False).all() - def test_get_items(self, mnt_filepath, merge_type, offset, duration, start): + def test_get_items( + self, + mnt_filepath: GRLDatasets, + merge_type: MergeType, + offset: str, + duration: timedelta, + start: datetime, + ): instance = MergeCollection( start=start, finished=start + duration, diff --git a/tests/incite/mergers/test_merge_collection_item.py b/tests/incite/mergers/test_merge_collection_item.py index 3d0b644..5ca2f6b 100644 --- a/tests/incite/mergers/test_merge_collection_item.py +++ b/tests/incite/mergers/test_merge_collection_item.py @@ -1,10 +1,16 @@ +from __future__ import annotations + from datetime import timedelta from itertools import product from pathlib import PurePath import pytest -from generalresearch.incite.mergers import MergeCollectionItem, MergeType +from generalresearch.incite.mergers import ( + MergeCollection, + MergeCollectionItem, + MergeType, +) @pytest.mark.parametrize( @@ -19,7 +25,10 @@ from generalresearch.incite.mergers import MergeCollectionItem, MergeType ) class TestMergeCollectionItem: - def test_file_naming(self, merge_collection, offset, duration, start): + def test_file_naming( + self, + merge_collection: MergeCollection, + ): assert len(merge_collection.items) == 25 items: list[MergeCollectionItem] = merge_collection.items @@ -34,7 +43,10 @@ class TestMergeCollectionItem: assert i._collection.offset in i.filename assert i.start.strftime("%Y-%m-%d-%H-%M-%S") in i.filename - def test_archives(self, merge_collection, offset, duration, start): + def test_archives( + self, + merge_collection: MergeCollection, + ): assert len(merge_collection.items) == 25 for i in merge_collection.items: @@ -44,10 +56,13 @@ class TestMergeCollectionItem: assert not i.has_partial_archive() assert i.has_archive() == i.path_exists(generic_path=i.path) - res = set([i.should_archive() for i in merge_collection.items]) + res = {i.should_archive() for i in merge_collection.items} assert len(res) == 1 - def test_item_to_archive(self, merge_collection, offset, duration, start): + def test_item_to_archive( + self, + merge_collection: MergeCollection, + ): for item in merge_collection.items: item: MergeCollectionItem assert not item.has_archive() diff --git a/tests/incite/mergers/test_pop_ledger.py b/tests/incite/mergers/test_pop_ledger.py index d054eb6..529a641 100644 --- a/tests/incite/mergers/test_pop_ledger.py +++ b/tests/incite/mergers/test_pop_ledger.py @@ -1,12 +1,25 @@ +from __future__ import annotations + +from collections.abc import Callable from datetime import UTC, datetime, timedelta from itertools import product as iter_product import pandas as pd import pytest +from dask.distributed import Client as DaskClient +from generalresearch.incite.base import GRLDatasets +from generalresearch.incite.collections.thl_web import ( + LedgerDFCollection, + SessionDFCollection, +) +from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge from generalresearch.incite.schemas.mergers.pop_ledger import ( numerical_col_names, ) +from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager +from generalresearch.models.thl.product import Product +from generalresearch.models.thl.user import User @pytest.mark.parametrize( @@ -30,20 +43,20 @@ class TestMergePOPLedger: def test_base( self, - client_no_amm, - ledger_collection, - pop_ledger_merge, + client_no_amm: DaskClient, + ledger_collection: LedgerDFCollection, + pop_ledger_merge: PopLedgerMerge, product: Product, user_factory: Callable[..., User], create_main_accounts: Callable[..., None], - thl_lm, + thl_ledger_manager: ThlLedgerManager, delete_df_collection: Callable[..., None], - incite_item_factory, + incite_item_factory: Callable[..., None], delete_ledger_db: Callable[..., None], ): from generalresearch.models.thl.ledger import LedgerAccount - u = user_factory(product=product: Product, created=ledger_collection.start) + u = user_factory(product=product, created=ledger_collection.start) # -- Build & Setup delete_ledger_db() @@ -73,19 +86,21 @@ class TestMergePOPLedger: # -- - user_wallet_account: LedgerAccount = thl_lm.get_account_or_create_user_wallet( - user=u + user_wallet_account: LedgerAccount = ( + thl_ledger_manager.get_account_or_create_user_wallet(user=u) + ) + cash_account: LedgerAccount = thl_ledger_manager.get_account_cash() + rev_account: LedgerAccount = ( + thl_ledger_manager.get_account_task_complete_revenue() ) - cash_account: LedgerAccount = thl_lm.get_account_cash() - rev_account: LedgerAccount = thl_lm.get_account_task_complete_revenue() item_finishes = [i.finish for i in ledger_collection.items] item_finishes.sort(reverse=True) last_item_finish = item_finishes[0] # Pure SQL based lookups - cash_balance: int = thl_lm.get_account_balance(account=cash_account) - rev_balance: int = thl_lm.get_account_balance(account=rev_account) + cash_balance: int = thl_ledger_manager.get_account_balance(account=cash_account) + rev_balance: int = thl_ledger_manager.get_account_balance(account=rev_account) assert cash_balance > rev_balance # (1) Test Cash Account @@ -123,39 +138,42 @@ class TestMergePOPLedger: def test_pydantic_init( self, - client_no_amm, - ledger_collection, - pop_ledger_merge, - mnt_filepath, + client_no_amm: DaskClient, + ledger_collection: LedgerDFCollection, + pop_ledger_merge: PopLedgerMerge, + mnt_filepath: GRLDatasets, product: Product, user_factory: Callable[..., User], create_main_accounts: Callable[..., None], - offset, - duration, - start, - thl_lm, - incite_item_factory, + offset: str, + duration: timedelta, + start: datetime, + thl_ledger_manager: ThlLedgerManager, + incite_item_factory: Callable[..., None], delete_df_collection: Callable[..., None], delete_ledger_db: Callable[..., None], - session_collection, + session_collection: SessionDFCollection, ): from generalresearch.models.thl.finance import ProductBalances from generalresearch.models.thl.ledger import LedgerAccount from generalresearch.models.thl.product import Product - u = user_factory(product=product: Product, created=session_collection.start) + u = user_factory(product=product, created=session_collection.start) assert ledger_collection.finished is not None - assert isinstance(u.product: Product, Product) + assert isinstance(u.product, Product) delete_ledger_db() - create_main_accounts(), + create_main_accounts() + delete_df_collection(coll=ledger_collection) - bp_account: LedgerAccount = thl_lm.get_account_or_create_bp_wallet( + bp_account: LedgerAccount = thl_ledger_manager.get_account_or_create_bp_wallet( product=u.product ) - cash_account: LedgerAccount = thl_lm.get_account_cash() - rev_account: LedgerAccount = thl_lm.get_account_task_complete_revenue() + cash_account: LedgerAccount = thl_ledger_manager.get_account_cash() + rev_account: LedgerAccount = ( + thl_ledger_manager.get_account_task_complete_revenue() + ) for item in ledger_collection.items: incite_item_factory(item=item, user=u) @@ -185,8 +203,10 @@ class TestMergePOPLedger: assert instance.payout == instance.net == instance.bp_payment_credit assert instance.available_balance < instance.net assert instance.available_balance + instance.retainer == instance.net - assert instance.balance == thl_lm.get_account_balance(bp_account) - assert df["bp_payment.CREDIT"].sum() == thl_lm.get_account_balance(bp_account) + assert instance.balance == thl_ledger_manager.get_account_balance(bp_account) + assert df["bp_payment.CREDIT"].sum() == thl_ledger_manager.get_account_balance( + bp_account + ) # (2) Filter by the Cash Account ddf = pop_ledger_merge.ddf( @@ -199,7 +219,7 @@ class TestMergePOPLedger: ) df: pd.DataFrame = client_no_amm.compute(collections=ddf, sync=True) - cash_balance: int = thl_lm.get_account_balance(account=cash_account) + cash_balance: int = thl_ledger_manager.get_account_balance(account=cash_account) assert df["bp_payment.CREDIT"].sum() == 0 assert cash_balance > 0 assert df["mp_payment.CREDIT"].sum() == 0 @@ -216,7 +236,7 @@ class TestMergePOPLedger: ) df: pd.DataFrame = client_no_amm.compute(collections=ddf, sync=True) - rev_balance: int = thl_lm.get_account_balance(account=rev_account) + rev_balance: int = thl_ledger_manager.get_account_balance(account=rev_account) assert rev_balance == 0 assert df["bp_payment.CREDIT"].sum() == 0 assert df["mp_payment.DEBIT"].sum() == 0 @@ -224,27 +244,28 @@ class TestMergePOPLedger: def test_resample( self, - client_no_amm, - ledger_collection, - pop_ledger_merge, - mnt_filepath, + client_no_amm: DaskClient, + ledger_collection: LedgerDFCollection, + pop_ledger_merge: PopLedgerMerge, + mnt_filepath: GRLDatasets, user_factory: Callable[..., User], product: Product, create_main_accounts: Callable[..., None], - offset, - duration, - start, - thl_lm, + offset: str, + duration: timedelta, + start: datetime, + thl_ledger_manager: ThlLedgerManager, delete_df_collection: Callable[..., None], - incite_item_factory, + incite_item_factory: Callable[..., None], ): - from generalresearch.models.thl.user import User assert ledger_collection.finished is not None delete_df_collection(coll=ledger_collection) u1: User = user_factory(product=product) - bp_account = thl_lm.get_account_or_create_bp_wallet(product=u1.product) + bp_account = thl_ledger_manager.get_account_or_create_bp_wallet( + product=u1.product + ) for item in ledger_collection.items: incite_item_factory(user=u1, item=item) @@ -274,7 +295,7 @@ class TestMergePOPLedger: assert isinstance(df.index, pd.Index) assert isinstance(df.index, pd.DatetimeIndex) - bp_account_balance = thl_lm.get_account_balance(account=bp_account) + bp_account_balance = thl_ledger_manager.get_account_balance(account=bp_account) # Initial sum initial_sum = df.sum().sum() diff --git a/tests/incite/mergers/test_ym_survey_merge.py b/tests/incite/mergers/test_ym_survey_merge.py index a0b8b87..8a4897b 100644 --- a/tests/incite/mergers/test_ym_survey_merge.py +++ b/tests/incite/mergers/test_ym_survey_merge.py @@ -1,8 +1,24 @@ +from __future__ import annotations + +from collections.abc import Callable from datetime import UTC, datetime, timedelta from itertools import product import pandas as pd import pytest +from dask.distributed import Client as DaskClient + +from generalresearch.incite.collections.thl_web import ( + SessionDFCollection, + WallDFCollection, +) +from generalresearch.incite.mergers.foundations.enriched_session import ( + EnrichedSessionMerge, +) +from generalresearch.incite.mergers.ym_survey_wall import YMSurveyWallMerge +from generalresearch.models.thl.product import Product +from generalresearch.models.thl.user import User +from generalresearch.pg_helper import PostgresConfig # noinspection PyUnresolvedReferences @@ -27,21 +43,20 @@ class TestYMSurveyMerge: def test_base( self, - client_no_amm, + client_no_amm: DaskClient, user_factory: Callable[..., User], product: Product, - ym_survey_wall_merge, - wall_collection, - session_collection, - enriched_session_merge, + ym_survey_wall_merge: YMSurveyWallMerge, + wall_collection: WallDFCollection, + session_collection: SessionDFCollection, + enriched_session_merge: EnrichedSessionMerge, delete_df_collection: Callable[..., None], - incite_item_factory, + incite_item_factory: Callable[..., None], thl_web_rr: PostgresConfig, ): - from generalresearch.models.thl.user import User delete_df_collection(coll=session_collection) - user: User = user_factory(product=product: Product, created=session_collection.start) + user: User = user_factory(product=product, created=session_collection.start) # -- Build & Setup assert ym_survey_wall_merge.start is None @@ -61,15 +76,15 @@ class TestYMSurveyMerge: client=client_no_amm, session_coll=session_collection, wall_coll=wall_collection, - pg_config=thl_web_rr: PostgresConfig, + pg_config=thl_web_rr, ) assert enriched_session_merge.progress.has_archive.eq(True).all() ddf = enriched_session_merge.ddf() - df: pd.DataFrame = client_no_amm.compute(collections=ddf, sync=True) + df1: pd.DataFrame | None = client_no_amm.compute(collections=ddf, sync=True) - assert isinstance(df, pd.DataFrame) - assert not df.empty + assert isinstance(df1, pd.DataFrame) + assert not df1.empty # -- @@ -83,18 +98,18 @@ class TestYMSurveyMerge: # -- ddf = ym_survey_wall_merge.ddf() - df: pd.DataFrame = client_no_amm.compute(collections=ddf, sync=True) + df2: pd.DataFrame | None = client_no_amm.compute(collections=ddf, sync=True) - assert isinstance(df, pd.DataFrame) - assert not df.empty + assert isinstance(df2, pd.DataFrame) + assert not df2.empty # -- - assert df.product_id.nunique() == 1 - assert df.team_id.nunique() == 1 - assert df.source.nunique() > 1 + assert df2.product_id.nunique() == 1 + assert df2.team_id.nunique() == 1 + assert df2.source.nunique() > 1 - started_min_ts = df.started.min() - started_max_ts = df.started.max() + started_min_ts = df2.started.min() + started_max_ts = df2.started.max() assert type(started_min_ts) is pd.Timestamp assert type(started_max_ts) is pd.Timestamp diff --git a/tests/incite/schemas/test_admin_responses.py b/tests/incite/schemas/test_admin_responses.py index e98eecd..d2658ea 100644 --- a/tests/incite/schemas/test_admin_responses.py +++ b/tests/incite/schemas/test_admin_responses.py @@ -1,8 +1,11 @@ +from __future__ import annotations + from datetime import UTC, datetime, timedelta from random import sample import numpy as np import pandas as pd +import pandera as pa import pytest from generalresearch.incite.schemas import empty_dataframe_from_schema @@ -16,12 +19,14 @@ from generalresearch.locales import Localelator class TestAdminPOPSchema: schema_df = empty_dataframe_from_schema(AdminPOPSchema) countries = list(Localelator().get_all_countries())[:5] - dates = [datetime(year=2024, month=1, day=i, tzinfo=None) for i in range(1, 10)] + dates = [ + datetime(year=2024, month=1, day=i, tzinfo=None) for i in range(1, 10) # noqa + ] @classmethod def assign_valid_vals(cls, df: pd.DataFrame) -> pd.DataFrame: for c in df.columns: - check_attrs: dict = AdminPOPSchema.columns[c].checks[0].statistics + check_attrs = AdminPOPSchema.columns[c].checks[0].statistics df[c] = np.random.randint( check_attrs["min_value"], check_attrs["max_value"], df.shape[0] ) @@ -29,7 +34,7 @@ class TestAdminPOPSchema: return df def test_empty(self): - with pytest.raises(Exception): + with pytest.raises(pa.errors.SchemaError): AdminPOPSchema.validate(pd.DataFrame()) def test_new_empty_df(self): @@ -42,7 +47,7 @@ class TestAdminPOPSchema: def test_valid(self): # (1) Works with raw naive datetime dates = [ - datetime(year=2024, month=1, day=i, tzinfo=None).isoformat() + datetime(year=2024, month=1, day=i, tzinfo=None).isoformat() # noqa for i in range(1, 10) ] df = pd.DataFrame( @@ -57,7 +62,10 @@ class TestAdminPOPSchema: assert isinstance(df, pd.DataFrame) # (2) Works with isoformat naive datetime - dates = [datetime(year=2024, month=1, day=i, tzinfo=None) for i in range(1, 10)] + dates = [ + datetime(year=2024, month=1, day=i, tzinfo=None) # noqa + for i in range(1, 10) + ] df = pd.DataFrame( index=pd.MultiIndex.from_product( iterables=[dates, self.countries], names=["index0", "index1"] @@ -84,12 +92,12 @@ class TestAdminPOPSchema: # Initially, they're all set with a timezone timestmaps: list[pd.Timestamp] = [i for i in df.index.get_level_values(0)] - assert all([ts.tz == UTC for ts in timestmaps]) + assert all(ts.tz == UTC for ts in timestmaps) # After validation, the timezone is removed df = AdminPOPSchema.validate(df) timestmaps: list[pd.Timestamp] = [i for i in df.index.get_level_values(0)] - assert all([ts.tz is None for ts in timestmaps]) + assert all(ts.tz is None for ts in timestmaps) def test_index_tz_no_future_beyond_one_year(self): now = datetime.now(tz=UTC) @@ -123,12 +131,12 @@ class TestAdminPOPSchema: df = self.assign_valid_vals(df) vals = [i for i in df.index.get_level_values(1)] - assert all([isinstance(v, float) for v in vals]) + assert all(isinstance(v, float) for v in vals) df = AdminPOPSchema.validate(df, lazy=True) vals = [i for i in df.index.get_level_values(1)] - assert all([isinstance(v, str) for v in vals]) + assert all(isinstance(v, str) for v in vals) # --- int to str --- @@ -142,12 +150,12 @@ class TestAdminPOPSchema: df = self.assign_valid_vals(df) vals = [i for i in df.index.get_level_values(1)] - assert all([isinstance(v, int) for v in vals]) + assert all(isinstance(v, int) for v in vals) df = AdminPOPSchema.validate(df, lazy=True) vals = [i for i in df.index.get_level_values(1)] - assert all([isinstance(v, str) for v in vals]) + assert all(isinstance(v, str) for v in vals) # a = 1 assert isinstance(df, pd.DataFrame) @@ -170,7 +178,7 @@ class TestAdminPOPSchema: assert isinstance(df, pd.DataFrame) timestmaps: list[pd.Timestamp] = [i for i in df.index.get_level_values(0)] - assert all([ts.tz is None for ts in timestmaps]) + assert all(ts.tz is None for ts in timestmaps) # (2) Timezones are removed dates = [ @@ -187,12 +195,12 @@ class TestAdminPOPSchema: # Has tz before validation, and none after timestmaps: list[pd.Timestamp] = [i for i in df.index.get_level_values(0)] - assert all([ts.tz is UTC for ts in timestmaps]) + assert all(ts.tz is UTC for ts in timestmaps) df = AdminPOPSchema.validate(df, lazy=True) timestmaps: list[pd.Timestamp] = [i for i in df.index.get_level_values(0)] - assert all([ts.tz is None for ts in timestmaps]) + assert all(ts.tz is None for ts in timestmaps) def test_clipping(self): df = pd.DataFrame( diff --git a/tests/incite/schemas/test_thl_web.py b/tests/incite/schemas/test_thl_web.py index 7f4434b..9b34ce0 100644 --- a/tests/incite/schemas/test_thl_web.py +++ b/tests/incite/schemas/test_thl_web.py @@ -16,7 +16,7 @@ class TestWallSchema: df = pd.DataFrame(columns=THLWallSchema.columns.keys()) - with pytest.raises(SchemaError) as cm: + with pytest.raises(SchemaError): THLWallSchema.validate(df) def test_no_rows(self): @@ -24,7 +24,7 @@ class TestWallSchema: df = pd.DataFrame(index=["uuid"], columns=THLWallSchema.columns.keys()) - with pytest.raises(SchemaError) as cm: + with pytest.raises(SchemaError): THLWallSchema.validate(df) def test_new_empty_df(self): @@ -50,7 +50,7 @@ class TestSessionSchema: df = pd.DataFrame(columns=THLSessionSchema.columns.keys()) df.set_index("uuid", inplace=True) - with pytest.raises(SchemaError) as cm: + with pytest.raises(SchemaError): THLSessionSchema.validate(df) def test_no_rows(self): @@ -58,7 +58,7 @@ class TestSessionSchema: df = pd.DataFrame(index=["id"], columns=THLSessionSchema.columns.keys()) - with pytest.raises(SchemaError) as cm: + with pytest.raises(SchemaError): THLSessionSchema.validate(df) def test_new_empty_df(self): diff --git a/tests/incite/test_collection_base.py b/tests/incite/test_collection_base.py index 7e1577a..d6ce2b1 100644 --- a/tests/incite/test_collection_base.py +++ b/tests/incite/test_collection_base.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from datetime import UTC, datetime, timedelta, timezone from os.path import exists as pexists from os.path import join as pjoin @@ -9,7 +11,7 @@ import pandas as pd import pytest from _pytest._code.code import ExceptionInfo -from generalresearch.incite.base import CollectionBase +from generalresearch.incite.base import CollectionBase, GRLDatasets AGO_15min = (datetime.now(tz=UTC) - timedelta(minutes=15)).replace(microsecond=0) AGO_1HR = (datetime.now(tz=UTC) - timedelta(hours=1)).replace(microsecond=0) @@ -17,11 +19,11 @@ AGO_2HR = (datetime.now(tz=UTC) - timedelta(hours=2)).replace(microsecond=0) class TestCollectionBase: - def test_init(self, mnt_filepath): + def test_init(self, mnt_filepath: GRLDatasets): instance = CollectionBase(archive_path=mnt_filepath.data_src) assert instance.df.empty is True - def test_init_df(self, mnt_filepath): + def test_init_df(self, mnt_filepath: GRLDatasets): # Only an empty pd.DataFrame can ever be provided instance = CollectionBase( df=pd.DataFrame({}), archive_path=mnt_filepath.data_src @@ -43,7 +45,7 @@ class TestCollectionBase: ) assert "Do not provide a pd.DataFrame" in str(cm.value) - def test_init_start(self, mnt_filepath): + def test_init_start(self, mnt_filepath: GRLDatasets): with pytest.raises(expected_exception=ValueError) as cm: cm: ExceptionInfo CollectionBase( @@ -74,7 +76,7 @@ class TestCollectionBase: cm.value ) - def test_init_archive_path(self, mnt_filepath): + def test_init_archive_path(self, mnt_filepath: GRLDatasets): """DirectoryPath is apparently smart enough to confirm that the directory path exists. """ @@ -99,7 +101,7 @@ class TestCollectionBase: CollectionBase(archive_path=new_path) assert "Path does not point to a directory" in str(cm.value) - def test_init_offset(self, mnt_filepath): + def test_init_offset(self, mnt_filepath: GRLDatasets): with pytest.raises(expected_exception=ValueError) as cm: cm: ExceptionInfo CollectionBase(offset="1:X", archive_path=mnt_filepath.data_src) @@ -118,14 +120,14 @@ class TestCollectionBase: class TestCollectionBaseProperties: - def test_items(self, mnt_filepath): + def test_items(self, mnt_filepath: GRLDatasets): with pytest.raises(expected_exception=NotImplementedError) as cm: cm: ExceptionInfo instance = CollectionBase(archive_path=mnt_filepath.data_src) - x = instance.items + _ = instance.items assert "Must override" in str(cm.value) - def test_interval_range(self, mnt_filepath): + def test_interval_range(self, mnt_filepath: GRLDatasets): instance = CollectionBase(archive_path=mnt_filepath.data_src) # Private method requires the end parameter with pytest.raises(expected_exception=AssertionError) as cm: @@ -147,7 +149,7 @@ class TestCollectionBaseProperties: assert res.is_monotonic_increasing assert res.is_unique - def test_interval_range2(self, mnt_filepath): + def test_interval_range2(self, mnt_filepath: GRLDatasets): instance = CollectionBase(archive_path=mnt_filepath.data_src) assert isinstance(instance.interval_range, list) @@ -166,16 +168,16 @@ class TestCollectionBaseProperties: ) assert len(instance.interval_range) == 2 - def test_progress(self, mnt_filepath): + def test_progress(self, mnt_filepath: GRLDatasets): with pytest.raises(expected_exception=NotImplementedError) as cm: cm: ExceptionInfo instance = CollectionBase( start=AGO_15min, offset="3min", archive_path=mnt_filepath.data_src ) - x = instance.progress + _ = instance.progress assert "Must override" in str(cm.value) - def test_progress2(self, mnt_filepath): + def test_progress2(self, mnt_filepath: GRLDatasets): instance = CollectionBase( start=AGO_2HR, offset="15min", @@ -184,10 +186,10 @@ class TestCollectionBaseProperties: assert instance.df.empty with pytest.raises(expected_exception=NotImplementedError) as cm: - df = instance.progress + _ = instance.progress assert "Must override" in str(cm.value) - def test_items2(self, mnt_filepath): + def test_items2(self, mnt_filepath: GRLDatasets): """There can't be a test for this because the Items need a path whic isn't possible in the generic form """ @@ -197,7 +199,7 @@ class TestCollectionBaseProperties: with pytest.raises(expected_exception=NotImplementedError) as cm: cm: ExceptionInfo - items = instance.items + _ = instance.items assert "Must override" in str(cm.value) # item = items[-3] @@ -208,19 +210,19 @@ class TestCollectionBaseProperties: # assert str(df.product_id.dtype) == "object" # assert str(ddf.product_id.dtype) == "string" - def test_items3(self, mnt_filepath): + def test_items3(self, mnt_filepath: GRLDatasets): instance = CollectionBase( start=AGO_2HR, offset="15min", archive_path=mnt_filepath.data_src, ) with pytest.raises(expected_exception=NotImplementedError) as cm: - item = instance.items[0] + _ = instance.items[0] assert "Must override" in str(cm.value) class TestCollectionBaseMethodsCleanup: - def test_fetch_force_rr_latest(self, mnt_filepath): + def test_fetch_force_rr_latest(self, mnt_filepath: GRLDatasets): coll = CollectionBase(archive_path=mnt_filepath.data_src) with pytest.raises(expected_exception=Exception) as cm: @@ -228,7 +230,7 @@ class TestCollectionBaseMethodsCleanup: coll.fetch_force_rr_latest(sources=[]) assert "Must override" in str(cm.value) - def test_fetch_all_paths(self, mnt_filepath): + def test_fetch_all_paths(self, mnt_filepath: GRLDatasets): coll = CollectionBase(archive_path=mnt_filepath.data_src) with pytest.raises(expected_exception=NotImplementedError) as cm: @@ -242,16 +244,16 @@ class TestCollectionBaseMethodsCleanup: class TestCollectionBaseMethodsCleanup: @pytest.mark.skip - def test_cleanup_partials(self, mnt_filepath): + def test_cleanup_partials(self, mnt_filepath: GRLDatasets): instance = CollectionBase(archive_path=mnt_filepath.data_src) assert instance.cleanup_partials() is None # it doesn't return anything - def test_clear_tmp_archives(self, mnt_filepath): + def test_clear_tmp_archives(self, mnt_filepath: GRLDatasets): instance = CollectionBase(archive_path=mnt_filepath.data_src) assert instance.clear_tmp_archives() is None # it doesn't return anything @pytest.mark.skip - def test_clear_corrupt_archives(self, mnt_filepath): + def test_clear_corrupt_archives(self, mnt_filepath: GRLDatasets): """TODO: expand this so it actually has corrupt archives that we check to see if they're removed """ @@ -259,14 +261,14 @@ class TestCollectionBaseMethodsCleanup: assert instance.clear_corrupt_archives() is None # it doesn't return anything @pytest.mark.skip - def test_rebuild_symlinks(self, mnt_filepath): + def test_rebuild_symlinks(self, mnt_filepath: GRLDatasets): instance = CollectionBase(archive_path=mnt_filepath.data_src) assert instance.rebuild_symlinks() is None class TestCollectionBaseMethodsSourceTiming: - def test_get_item(self, mnt_filepath): + def test_get_item(self, mnt_filepath: GRLDatasets): instance = CollectionBase(archive_path=mnt_filepath.data_src) i = pd.Interval(left=1, right=2, closed="left") @@ -274,7 +276,7 @@ class TestCollectionBaseMethodsSourceTiming: instance.get_item(interval=i) assert "Must override" in str(cm.value) - def test_get_item_start(self, mnt_filepath): + def test_get_item_start(self, mnt_filepath: GRLDatasets): instance = CollectionBase(archive_path=mnt_filepath.data_src) dt = datetime.now(tz=UTC) @@ -284,7 +286,7 @@ class TestCollectionBaseMethodsSourceTiming: instance.get_item_start(start=start) assert "Must override" in str(cm.value) - def test_get_items(self, mnt_filepath): + def test_get_items(self, mnt_filepath: GRLDatasets): instance = CollectionBase(archive_path=mnt_filepath.data_src) dt = datetime.now(tz=UTC) @@ -293,21 +295,21 @@ class TestCollectionBaseMethodsSourceTiming: instance.get_items(since=dt) assert "Must override" in str(cm.value) - def test_get_items_from_year(self, mnt_filepath): + def test_get_items_from_year(self, mnt_filepath: GRLDatasets): instance = CollectionBase(archive_path=mnt_filepath.data_src) with pytest.raises(expected_exception=NotImplementedError) as cm: instance.get_items_from_year(year=2020) assert "Must override" in str(cm.value) - def test_get_items_last90(self, mnt_filepath): + def test_get_items_last90(self, mnt_filepath: GRLDatasets): instance = CollectionBase(archive_path=mnt_filepath.data_src) with pytest.raises(expected_exception=NotImplementedError) as cm: instance.get_items_last90() assert "Must override" in str(cm.value) - def test_get_items_last365(self, mnt_filepath): + def test_get_items_last365(self, mnt_filepath: GRLDatasets): instance = CollectionBase(archive_path=mnt_filepath.data_src) with pytest.raises(expected_exception=NotImplementedError) as cm: diff --git a/tests/incite/test_collection_base_item.py b/tests/incite/test_collection_base_item.py index 7a0a581..e09f54a 100644 --- a/tests/incite/test_collection_base_item.py +++ b/tests/incite/test_collection_base_item.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from datetime import UTC, datetime from os.path import join as pjoin from pathlib import Path @@ -8,7 +10,7 @@ import pandas as pd import pytest from pydantic import ValidationError -from generalresearch.incite.base import CollectionItemBase +from generalresearch.incite.base import CollectionItemBase, GRLDatasets class TestCollectionItemBase: @@ -40,20 +42,20 @@ class TestCollectionItemBaseProperties: def test_finish(self): instance = CollectionItemBase() - with pytest.raises(expected_exception=AttributeError) as cm: - res = instance.finish + with pytest.raises(expected_exception=AttributeError): + _ = instance.finish def test_interval(self): instance = CollectionItemBase() - with pytest.raises(expected_exception=AttributeError) as cm: - res = instance.interval + with pytest.raises(expected_exception=AttributeError): + _ = instance.interval def test_filename(self): instance = CollectionItemBase() with pytest.raises(expected_exception=NotImplementedError) as cm: - res = instance.filename + _ = instance.filename assert "Do not use CollectionItemBase directly" in str(cm.value) @@ -61,7 +63,7 @@ class TestCollectionItemBaseProperties: instance = CollectionItemBase() with pytest.raises(expected_exception=NotImplementedError) as cm: - res = instance.filename + _ = instance.filename assert "Do not use CollectionItemBase directly" in str(cm.value) @@ -69,27 +71,27 @@ class TestCollectionItemBaseProperties: instance = CollectionItemBase() with pytest.raises(expected_exception=NotImplementedError) as cm: - res = instance.filename + _ = instance.filename assert "Do not use CollectionItemBase directly" in str(cm.value) def test_path(self): instance = CollectionItemBase() - with pytest.raises(expected_exception=AttributeError) as cm: - res = instance.path + with pytest.raises(expected_exception=AttributeError): + _ = instance.path def test_partial_path(self): instance = CollectionItemBase() - with pytest.raises(expected_exception=AttributeError) as cm: - res = instance.partial_path + with pytest.raises(expected_exception=AttributeError): + _ = instance.partial_path def test_empty_path(self): instance = CollectionItemBase() - with pytest.raises(expected_exception=AttributeError) as cm: - res = instance.empty_path + with pytest.raises(expected_exception=AttributeError): + _ = instance.empty_path class TestCollectionItemBaseMethods: @@ -106,41 +108,41 @@ class TestCollectionItemBaseMethods: instance = CollectionItemBase() with pytest.raises(expected_exception=NotImplementedError) as cm: - res = instance.tmp_filename() + _ = instance.tmp_filename() assert "Do not use CollectionItemBase directly" in str(cm.value) def test_tmp_path(self): instance = CollectionItemBase() - with pytest.raises(expected_exception=AttributeError) as cm: - res = instance.tmp_path() + with pytest.raises(expected_exception=AttributeError): + instance.tmp_path() def test_is_empty(self): instance = CollectionItemBase() - with pytest.raises(expected_exception=AttributeError) as cm: - res = instance.is_empty() + with pytest.raises(expected_exception=AttributeError): + instance.is_empty() def test_has_empty(self): instance = CollectionItemBase() - with pytest.raises(expected_exception=AttributeError) as cm: - res = instance.has_empty() + with pytest.raises(expected_exception=AttributeError): + instance.has_empty() def test_has_partial_archive(self): instance = CollectionItemBase() - with pytest.raises(expected_exception=AttributeError) as cm: - res = instance.has_partial_archive() + with pytest.raises(expected_exception=AttributeError): + instance.has_partial_archive() @pytest.mark.parametrize("include_empty", [True, False]) - def test_has_archive(self, include_empty): + def test_has_archive(self, include_empty: bool): instance = CollectionItemBase() - with pytest.raises(expected_exception=AttributeError) as cm: - res = instance.has_archive(include_empty=include_empty) + with pytest.raises(expected_exception=AttributeError): + instance.has_archive(include_empty=include_empty) - def test_delete_archive_file(self, mnt_filepath): + def test_delete_archive_file(self, mnt_filepath: GRLDatasets): path1 = Path(pjoin(mnt_filepath.data_src, f"{uuid4().hex}.zip")) # Confirm it doesn't exist, and that delete_archive() doesn't throw @@ -155,7 +157,7 @@ class TestCollectionItemBaseMethods: CollectionItemBase.delete_archive(generic_path=path1) assert not path1.exists() - def test_delete_archive_dir(self, mnt_filepath): + def test_delete_archive_dir(self, mnt_filepath: GRLDatasets): path1 = Path(pjoin(mnt_filepath.data_src, f"{uuid4().hex}")) # Confirm it doesn't exist, and that delete_archive() doesn't throw @@ -174,20 +176,20 @@ class TestCollectionItemBaseMethods: def test_should_archive(self): instance = CollectionItemBase() - with pytest.raises(expected_exception=AttributeError) as cm: - res = instance.should_archive() + with pytest.raises(expected_exception=AttributeError): + _ = instance.should_archive() def test_set_empty(self): instance = CollectionItemBase() - with pytest.raises(expected_exception=AttributeError) as cm: - res = instance.set_empty() + with pytest.raises(expected_exception=AttributeError): + _ = instance.set_empty() def test_valid_archive(self): instance = CollectionItemBase() - with pytest.raises(expected_exception=AttributeError) as cm: - res = instance.valid_archive(generic_path=None, sample=None) + with pytest.raises(expected_exception=AttributeError): + _ = instance.valid_archive(generic_path=None, sample=None) class TestCollectionItemBaseMethodsORM: @@ -197,11 +199,11 @@ class TestCollectionItemBaseMethodsORM: pass @pytest.mark.parametrize("is_partial", [True, False]) - def test_to_archive(self, is_partial): + def test_to_archive(self, is_partial: bool): instance = CollectionItemBase() with pytest.raises(expected_exception=NotImplementedError) as cm: - res = instance.to_archive( + _ = instance.to_archive( ddf=dd.from_pandas(data=pd.DataFrame()), is_partial=is_partial ) assert "Must override" in str(cm.value) diff --git a/tests/incite/test_interval_idx.py b/tests/incite/test_interval_idx.py index 3034c21..03d29ea 100644 --- a/tests/incite/test_interval_idx.py +++ b/tests/incite/test_interval_idx.py @@ -1,4 +1,4 @@ -from datetime import datetime +from datetime import UTC, datetime import pandas as pd @@ -6,8 +6,8 @@ import pandas as pd class TestIntervalIndex: def test_init(self): - start = datetime(year=2000, month=1, day=1) - end = datetime(year=2000, month=1, day=10) + start = datetime(year=2000, month=1, day=1, tzinfo=UTC) + end = datetime(year=2000, month=1, day=10, tzinfo=UTC) iv_r: pd.IntervalIndex = pd.interval_range( start=start, end=end, freq="1d", closed="left" diff --git a/tests/managers/gr/test_business.py b/tests/managers/gr/test_business.py index 3490403..ed141b1 100644 --- a/tests/managers/gr/test_business.py +++ b/tests/managers/gr/test_business.py @@ -2,17 +2,36 @@ from uuid import uuid4 import pytest +from generalresearch.managers.gr.business import ( + BusinessAddressManager, + BusinessBankAccountManager, + BusinessManager, +) +from generalresearch.managers.gr.team import MembershipManager, TeamManager +from generalresearch.models.gr.authentication import GRUser +from generalresearch.models.gr.business import ( + Business, + BusinessAddress, + BusinessBankAccount, + TransferMethod, +) +from generalresearch.pg_helper import PostgresConfig + class TestBusinessBankAccountManager: - def test_init(self, business_bank_account_manager, gr_db): + def test_init( + self, + business_bank_account_manager: BusinessBankAccountManager, + gr_db: PostgresConfig, + ): assert business_bank_account_manager.pg_config == gr_db - def test_create(self, business: Business, business_bank_account_manager): - from generalresearch.models.gr.business import ( - BusinessBankAccount, - TransferMethod, - ) + def test_create( + self, + business: Business, + business_bank_account_manager: BusinessBankAccountManager, + ): instance = business_bank_account_manager.create( business_id=business.id, @@ -33,8 +52,9 @@ class TestBusinessBankAccountManager: class TestBusinessAddressManager: - def test_create(self, business: Business, business_address_manager): - from generalresearch.models.gr.business import BusinessAddress + def test_create( + self, business: Business, business_address_manager: BusinessAddressManager + ): res = business_address_manager.create(uuid=uuid4().hex, business_id=business.id) assert isinstance(res, BusinessAddress) @@ -43,14 +63,13 @@ class TestBusinessAddressManager: class TestBusinessManager: - def test_create(self, business_manager): - from generalresearch.models.gr.business import Business + def test_create(self, business_manager: BusinessManager): instance = business_manager.create_dummy() assert isinstance(instance, Business) assert isinstance(instance.id, int) - def test_get_or_create(self, business_manager): + def test_get_or_create(self, business_manager: BusinessManager): uuid_key = uuid4().hex assert business_manager.get_by_uuid(business_uuid=uuid_key) is None @@ -61,9 +80,10 @@ class TestBusinessManager: ) res = business_manager.get_by_uuid(business_uuid=uuid_key) + assert isinstance(res, Business) assert res.id == instance.id - def test_get_all(self, business_manager): + def test_get_all(self, business_manager: BusinessManager): res1 = business_manager.get_all() assert isinstance(res1, list) @@ -76,7 +96,11 @@ class TestBusinessManager: pass def test_get_by_user_id( - self, business_manager, gr_user, team_manager, membership_manager + self, + business_manager: BusinessManager, + gr_user: GRUser, + team_manager: TeamManager, + membership_manager: MembershipManager, ): res = business_manager.get_by_user_id(user_id=gr_user.id) assert len(res) == 0 @@ -93,7 +117,7 @@ class TestBusinessManager: # Create a Membership for the gr_user to the Team... but it doesn't # matter because the Team doesn't have any Business yet - m1 = membership_manager.create(team=t1, gr_user=gr_user) + _ = membership_manager.create(team=t1, gr_user=gr_user) res = business_manager.get_by_user_id(user_id=gr_user.id) assert len(res) == 0 @@ -113,15 +137,17 @@ class TestBusinessManager: def test_get_uuids_by_user_id(self): pass - def test_get_by_uuid(self, business: Business, business_manager): + def test_get_by_uuid(self, business: Business, business_manager: BusinessManager): instance = business_manager.get_by_uuid(business_uuid=business.uuid) + assert isinstance(instance, Business) assert business.id == instance.id - def test_get_by_id(self, business: Business, business_manager): + def test_get_by_id(self, business: Business, business_manager: BusinessManager): instance = business_manager.get_by_id(business_id=business.id) + assert isinstance(instance, Business) assert business.uuid == instance.uuid - def test_cache_key(self, business): + def test_cache_key(self, business: Business): assert "business:" in business.cache_key # def test_create_raise_on_duplicate(self): diff --git a/tests/managers/gr/test_team.py b/tests/managers/gr/test_team.py index 5e5c565..ae3e1bb 100644 --- a/tests/managers/gr/test_team.py +++ b/tests/managers/gr/test_team.py @@ -1,18 +1,29 @@ +from __future__ import annotations + +from collections.abc import Callable from uuid import uuid4 +from generalresearch.managers.gr.authentication import GRUserManager +from generalresearch.managers.gr.team import MembershipManager, TeamManager +from generalresearch.models.gr.authentication import GRUser +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 + class TestMembershipManager: - def test_init(self, membership_manager, gr_db): + def test_init(self, membership_manager: MembershipManager, gr_db: PostgresConfig): assert membership_manager.pg_config == gr_db class TestTeamManager: - def test_init(self, team_manager, gr_db): + def test_init(self, team_manager: TeamManager, gr_db: PostgresConfig): assert team_manager.pg_config == gr_db - def test_get_or_create(self, team_manager): + def test_get_or_create(self, team_manager: TeamManager): from generalresearch.models.gr.team import Team new_uuid = uuid4().hex @@ -24,7 +35,7 @@ class TestTeamManager: assert team.uuid == new_uuid assert team.name == "< Unknown >" - def test_get_all(self, team_manager): + def test_get_all(self, team_manager: TeamManager): res1 = team_manager.get_all() assert isinstance(res1, list) @@ -32,16 +43,20 @@ class TestTeamManager: res2 = team_manager.get_all() assert len(res1) == len(res2) - 1 - def test_create(self, team_manager): - from generalresearch.models.gr.team import Team + def test_create(self, team_manager: TeamManager): team: Team = team_manager.create_dummy() assert isinstance(team, Team) assert isinstance(team.id, int) - def test_add_user(self, team, team_manager, gr_um, gr_db, gr_redis_config): - from generalresearch.models.gr.authentication import GRUser - from generalresearch.models.gr.team import Membership + def test_add_user( + self, + team: Team, + team_manager: TeamManager, + gr_um: GRUserManager, + gr_db: PostgresConfig, + gr_redis_config: RedisConfig, + ): user: GRUser = gr_um.create_dummy() @@ -54,25 +69,23 @@ class TestTeamManager: assert len(team.gr_users) assert team.gr_users == [user] - def test_get_by_uuid(self, team_manager): - from generalresearch.models.gr.team import Team + def test_get_by_uuid(self, team_manager: TeamManager): team: Team = team_manager.create_dummy() instance = team_manager.get_by_uuid(team_uuid=team.uuid) assert team.id == instance.id - def test_get_by_id(self, team_manager): - from generalresearch.models.gr.team import Team + def test_get_by_id(self, team_manager: TeamManager): team: Team = team_manager.create_dummy() instance = team_manager.get_by_id(team_id=team.id) assert team.uuid == instance.uuid - def test_get_by_user(self, team, team_manager, gr_um): - from generalresearch.models.gr.authentication import GRUser - from generalresearch.models.gr.team import Team + def test_get_by_user( + self, team: Team, team_manager: TeamManager, gr_um: GRUserManager + ): user: GRUser = gr_um.create_dummy() team_manager.add_user(team=team, gr_user=user) @@ -86,15 +99,12 @@ class TestTeamManager: def test_get_by_user_duplicates( self, - gr_user_token, - gr_user, - membership, + gr_user: GRUser, product_factory: Callable[..., Product], - membership_factory, - team, - thl_web_rr: PostgresConfig, - gr_redis_config, - gr_db, + membership_factory: Callable[..., Membership], + team: Team, + gr_redis_config: RedisConfig, + gr_db: PostgresConfig, ): product_factory(team=team) membership_factory(team=team, gr_user=gr_user) diff --git a/tests/managers/network/test_label.py b/tests/managers/network/test_label.py index bfc7518..71efa95 100644 --- a/tests/managers/network/test_label.py +++ b/tests/managers/network/test_label.py @@ -27,7 +27,7 @@ def ip_label(utc_now) -> IPLabel: provider="GeoNodE", created_at=utc_now, ip=ip, - metadata=IPLabelMetadata(services=["RDP"]) + metadata=IPLabelMetadata(services=["RDP"]), ) @@ -181,7 +181,7 @@ def test_label_cidr_and_ipinfo( ip = fake.ipv6() ip_information_factory(ip=ip, geoname=ip_geoname) # We normalize for storage into ipinfo table - ip_norm, prefix = normalize_ip(ip) + ip_norm, _ = normalize_ip(ip) # Test with a larger network ip_48 = ipaddress.IPv6Network((ip, 48), strict=False) diff --git a/tests/managers/thl/test_contest/test_leaderboard.py b/tests/managers/thl/test_contest/test_leaderboard.py index 07d8d74..3a63075 100644 --- a/tests/managers/thl/test_contest/test_leaderboard.py +++ b/tests/managers/thl/test_contest/test_leaderboard.py @@ -5,7 +5,6 @@ from zoneinfo import ZoneInfo from generalresearch.currency import USDCent from generalresearch.managers.thl.contest_manager import ContestManager -from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager from generalresearch.managers.thl.user_manager.user_manager import UserManager from generalresearch.models.thl.contest.definitions import ( @@ -116,7 +115,9 @@ class TestLeaderboardContestCRUD: assert decision assert reason == ContestEndReason.ENDS_AT - contest_manager.end_contest_if_over(contest=contest, ledger_manager=thl_lm) + contest_manager.end_contest_if_over( + contest=contest, ledger_manager=thl_ledger_manager + ) c: LeaderboardContest = contest_manager.get(contest_uuid=contest.uuid) assert c.status == ContestStatus.COMPLETED diff --git a/tests/managers/thl/test_contest/test_milestone.py b/tests/managers/thl/test_contest/test_milestone.py index a2d575b..e3889bc 100644 --- a/tests/managers/thl/test_contest/test_milestone.py +++ b/tests/managers/thl/test_contest/test_milestone.py @@ -292,7 +292,10 @@ class TestMilestoneContestUserViews: assert len(cs) == 1 contest_manager.enter_milestone_contest( - contest_uuid=c.uuid, user=user, country_iso="us", ledger_manager=thl_lm + contest_uuid=c.uuid, + user=user, + country_iso="us", + ledger_manager=thl_ledger_manager, ) # User isn't eligible anymore diff --git a/tests/managers/thl/test_contest/test_raffle.py b/tests/managers/thl/test_contest/test_raffle.py index b435576..06d4676 100644 --- a/tests/managers/thl/test_contest/test_raffle.py +++ b/tests/managers/thl/test_contest/test_raffle.py @@ -14,6 +14,7 @@ from generalresearch.managers.thl.ledger_manager.exceptions import ( ) from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager from generalresearch.models.thl.contest import ( + Contest, ContestEndCondition, ContestEntryRule, ContestPrize, @@ -40,8 +41,6 @@ class TestRaffleContest: def test_should_end( self, contest: RaffleContest, - thl_ledger_manager: ThlLedgerManager, - contest_manager: ContestManager, ): # contest is active and has no entries should, msg = contest.should_end() @@ -67,7 +66,6 @@ class TestRaffleContestCRUD: self, contest_create: RaffleContestCreate, product_user_wallet_yes: Product, - thl_ledger_manager: ThlLedgerManager, contest_manager: ContestManager, ): c = contest_manager.create( @@ -329,7 +327,7 @@ class TestRaffleContestCRUD: contest_uuid=c.uuid, entry=entry, country_iso="us", - ledger_manager=thl_lm, + ledger_manager=thl_ledger_manager, ) assert "Entry would exceed max amount per user." in str(e.value) @@ -342,7 +340,7 @@ class TestRaffleContestCRUD: contest_uuid=c.uuid, entry=entry, country_iso="us", - ledger_manager=thl_lm, + ledger_manager=thl_ledger_manager, ) assert "Entry would exceed max amount per user per day." in str(e.value) @@ -354,7 +352,7 @@ class TestRaffleContestCRUD: contest_uuid=c.uuid, entry=entry, country_iso="us", - ledger_manager=thl_lm, + ledger_manager=thl_ledger_manager, ) # Then can't anymore @@ -366,7 +364,7 @@ class TestRaffleContestCRUD: contest_uuid=c.uuid, entry=entry, country_iso="us", - ledger_manager=thl_lm, + ledger_manager=thl_ledger_manager, ) assert "Entry would exceed max amount per user per day." in str(e.value) diff --git a/tests/managers/thl/test_ledger/test_lm_accounts.py b/tests/managers/thl/test_ledger/test_lm_accounts.py index 540bea8..7b65b2d 100644 --- a/tests/managers/thl/test_ledger/test_lm_accounts.py +++ b/tests/managers/thl/test_ledger/test_lm_accounts.py @@ -63,7 +63,7 @@ class TestLedgerAccountManagerNoResults: acct_id: UUIDStr, lm: LedgerManager, ): - qn = ":".join([currency, kind, acct_id]) + qn = f"{currency}:{kind}:{acct_id}" # (1) .get_many_ assert lm.get_account_many_(qualified_names=[qn], raise_on_error=False) == [] diff --git a/tests/managers/thl/test_ledger/test_lm_tx.py b/tests/managers/thl/test_ledger/test_lm_tx.py index 13495a7..ce609d6 100644 --- a/tests/managers/thl/test_ledger/test_lm_tx.py +++ b/tests/managers/thl/test_ledger/test_lm_tx.py @@ -24,8 +24,6 @@ class TestLedgerManagerCreateTx: """Confirm that the Permission values that are set on the Ledger Manger allow the Creation action to occur. """ - acct_uuid = uuid4().hex - # (1) With no Permissions defined test_lm = LedgerManager( pg_config=ledger_manager.pg_config, @@ -44,8 +42,6 @@ class TestLedgerManagerCreateTx: def test_create_assertions( self, - ledger_account_debit: LedgerAccount, - ledger_account_credit: LedgerAccount, ledger_manager: LedgerManager, ): with pytest.raises(expected_exception=ValueError) as excinfo: diff --git a/tests/managers/thl/test_ledger/test_thl_lm_accounts.py b/tests/managers/thl/test_ledger/test_thl_lm_accounts.py index dce9116..60eb71c 100644 --- a/tests/managers/thl/test_ledger/test_thl_lm_accounts.py +++ b/tests/managers/thl/test_ledger/test_thl_lm_accounts.py @@ -229,6 +229,7 @@ class TestThlLedgerManagerAccounts: # (1) known account and confirm it comes back res = ledger_manager.get_account(qualified_name=account1.qualified_name) + assert isinstance(res, LedgerAccount) assert account1.model_dump_json() == res.model_dump_json() # (2) known accounts and confirm they both come back @@ -291,6 +292,7 @@ class TestThlLedgerManagerAccounts: assert len(res) == 2 # Confirm an empty array comes back for all unknown qualified names + assert isinstance(ledger_manager.currency, LedgerCurrency) res = ledger_manager.get_accounts_if_exists( qualified_names=[ f"{ledger_manager.currency.value}:bp_wall:{uuid4().hex}" @@ -328,7 +330,7 @@ class TestThlLedgerManagerAccounts: product_uuids=product_uuids ) assert len(res) == len(product_uuids) - assert all([isinstance(i, LedgerAccount) for i in res]) + assert all(isinstance(i, LedgerAccount) for i in res) class TestLedgerAccountManager: @@ -351,10 +353,10 @@ class TestLedgerAccountManager: # First we want to validate that using the get_account method raises # an error for a random LedgerAccount which we know does not exist. with pytest.raises(LedgerAccountDoesntExistError): - lam.get_account(qualified_name=account.qualified_name) + ledger_account_manager.get_account(qualified_name=account.qualified_name) # Now that we know it doesn't exist, get_or_create for it - instance = lam.get_account_or_create(account=account) + instance = ledger_account_manager.get_account_or_create(account=account) # It should always return assert isinstance(instance, LedgerAccount) @@ -364,10 +366,11 @@ class TestLedgerAccountManager: self, user: User, thl_ledger_manager: ThlLedgerManager, - ledger_manager: LedgerManager, ledger_account_manager: LedgerAccountManager, ): + assert isinstance(user.product, Product) + with pytest.raises(LedgerAccountDoesntExistError): ledger_account_manager.get_account( qualified_name=f"test:bp_wallet:{user.product.id}" diff --git a/tests/managers/thl/test_ledger/test_thl_lm_bp_payout.py b/tests/managers/thl/test_ledger/test_thl_lm_bp_payout.py index e4a25a3..b518453 100644 --- a/tests/managers/thl/test_ledger/test_thl_lm_bp_payout.py +++ b/tests/managers/thl/test_ledger/test_thl_lm_bp_payout.py @@ -133,16 +133,17 @@ class TestThlLedgerManagerBPPayout: ) payoutevent_uuid = uuid4().hex - with caplog.at_level(logging.INFO): - with pytest.raises(LedgerTransactionConditionFailedError): - thl_ledger_manager.create_tx_bp_payout( - user.product, - amount=USDCent(10_000), - created=now + timedelta(minutes=2), - skip_one_per_day_check=True, - skip_wallet_balance_check=False, - payoutevent_uuid=payoutevent_uuid, - ) + with caplog.at_level(logging.INFO), pytest.raises( + LedgerTransactionConditionFailedError + ): + thl_ledger_manager.create_tx_bp_payout( + user.product, + amount=USDCent(10_000), + created=now + timedelta(minutes=2), + skip_one_per_day_check=True, + skip_wallet_balance_check=False, + payoutevent_uuid=payoutevent_uuid, + ) assert "failed condition check balance:" in caplog.text thl_ledger_manager.create_tx_bp_payout( @@ -197,17 +198,18 @@ class TestThlLedgerManagerBPPayout: assert balance == int(rand_amount) * -1 # Test some basic assertions - with caplog.at_level(logging.INFO): - with pytest.raises(expected_exception=Exception): - thl_ledger_manager.create_tx_bp_payout( - product=product, - amount=rand_amount, - payoutevent_uuid=uuid4().hex, - created=datetime.now(tz=UTC), - skip_wallet_balance_check=False, - skip_one_per_day_check=False, - skip_flag_check=False, - ) + with caplog.at_level(logging.INFO), pytest.raises( + expected_exception=ValueError + ): + thl_ledger_manager.create_tx_bp_payout( + product=product, + amount=rand_amount, + payoutevent_uuid=uuid4().hex, + created=datetime.now(tz=UTC), + skip_wallet_balance_check=False, + skip_one_per_day_check=False, + skip_flag_check=False, + ) assert "failed condition check >1 tx per day" in caplog.text def test_create_tx_redis_failure( @@ -291,7 +293,7 @@ class TestThlLedgerManagerBPPayout: # Will fail due to multiple per day payoutevent_uuid2 = uuid4().hex with pytest.raises(expected_exception=Exception) as e: - tx = thl_ledger_manager.create_tx_bp_payout( + thl_ledger_manager.create_tx_bp_payout( product=product, amount=rand_amount, payoutevent_uuid=payoutevent_uuid2, @@ -348,7 +350,7 @@ class TestThlLedgerManagerBPPayout: # Create TX will fail on lock exit, after the tx was created! with pytest.raises(expected_exception=Exception) as e: - tx = thl_ledger_manager.create_tx_bp_payout( + thl_ledger_manager.create_tx_bp_payout( product=product, amount=rand_amount, payoutevent_uuid=payoutevent_uuid, @@ -384,7 +386,9 @@ class TestPayoutEventManagerBPPayout: product, rand_amount, now, direction=Direction.CREDIT ) assert thl_ledger_manager.get_account_balance(bp_wallet_account) == rand_amount - brokerage_product_payout_event_manager.set_account_lookup_table(thl_lm=thl_lm) + brokerage_product_payout_event_manager.set_account_lookup_table( + thl_lm=thl_ledger_manager + ) pe = brokerage_product_payout_event_manager.create_bp_payout_event( thl_ledger_manager=thl_ledger_manager, @@ -557,7 +561,7 @@ class TestPayoutEventManagerBPPayout: # Will fail on lock exit, after the tx was created! # But it'll see that the tx was created and so everything will be fine Lock.release = broken_release - pe = brokerage_product_payout_event_manager.create_bp_payout_event( + brokerage_product_payout_event_manager.create_bp_payout_event( thl_ledger_manager=thl_ledger_manager, product=product, created=now, diff --git a/tests/managers/thl/test_ledger/test_thl_lm_tx.py b/tests/managers/thl/test_ledger/test_thl_lm_tx.py index 89adb0b..1860d6d 100644 --- a/tests/managers/thl/test_ledger/test_thl_lm_tx.py +++ b/tests/managers/thl/test_ledger/test_thl_lm_tx.py @@ -23,7 +23,6 @@ from generalresearch.models.thl.definitions import ( WALL_ALLOWED_STATUS_STATUS_CODE, ) from generalresearch.models.thl.ledger import ( - AccountType, Direction, LedgerAccount, TransactionType, @@ -287,17 +286,18 @@ class TestThlLedgerTxManager: assert balance == int(rand_amount) * -1 # Test some basic assertions - with caplog.at_level(logging.INFO): - with pytest.raises(expected_exception=Exception): - thl_ledger_manager.create_tx_bp_payout( - product=product, - amount=rand_amount, - payoutevent_uuid=uuid4().hex, - created=datetime.now(tz=UTC), - skip_wallet_balance_check=False, - skip_one_per_day_check=False, - skip_flag_check=False, - ) + with caplog.at_level(logging.INFO), pytest.raises( + expected_exception=ValueError + ): + thl_ledger_manager.create_tx_bp_payout( + product=product, + amount=rand_amount, + payoutevent_uuid=uuid4().hex, + created=datetime.now(tz=UTC), + skip_wallet_balance_check=False, + skip_one_per_day_check=False, + skip_flag_check=False, + ) assert "failed condition check >1 tx per day" in caplog.text def test_create_tx_bp_payout_( @@ -1794,7 +1794,7 @@ class TestThlLedgerManagerAdj: ) thl_ledger_manager.create_tx_bp_payment(session, created=wall1.started) - revenue = ththl_ledger_managerl_lm.get_account_task_complete_revenue() + revenue = thl_ledger_manager.get_account_task_complete_revenue() bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet( user.product ) diff --git a/tests/managers/thl/test_payout.py b/tests/managers/thl/test_payout.py index e6c597b..0f3f103 100644 --- a/tests/managers/thl/test_payout.py +++ b/tests/managers/thl/test_payout.py @@ -727,8 +727,6 @@ class TestBusinessPayoutEventManager: # {"uuid": bp_pe.uuid, "status": PayoutStatus.FAILED}, # ) - assert 1 == 0 - def test_ach_payment( self, mnt_filepath: GRLDatasets, diff --git a/tests/managers/thl/test_survey.py b/tests/managers/thl/test_survey.py index 2c2bf9d..c3ab162 100644 --- a/tests/managers/thl/test_survey.py +++ b/tests/managers/thl/test_survey.py @@ -11,9 +11,6 @@ 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 @@ -183,7 +180,8 @@ class TestSurvey: ] uqad = {} for uqa in uqas: - for k, _ in uqa.calc_answers.items(): + assert uqa.calc_answers + for k in uqa.calc_answers: if k in qualifying_questions: uqad[k] = uqa uqad[uqa.property_code] = uqa diff --git a/tests/managers/thl/test_user_manager/test_base.py b/tests/managers/thl/test_user_manager/test_base.py index 5822207..8cd83ad 100644 --- a/tests/managers/thl/test_user_manager/test_base.py +++ b/tests/managers/thl/test_user_manager/test_base.py @@ -21,7 +21,7 @@ from generalresearch.managers.thl.user_manager.user_manager import ( UserManager, ) from generalresearch.managers.thl.userhealth import AuditLogManager -from generalresearch.models.thl.product import Product, UserCreateConfig, product +from generalresearch.models.thl.product import Product, UserCreateConfig from generalresearch.models.thl.user import User from generalresearch.pg_helper import PostgresConfig diff --git a/tests/models/network/test_nmap.py b/tests/models/network/test_nmap.py index 5e9f4d0..db39997 100644 --- a/tests/models/network/test_nmap.py +++ b/tests/models/network/test_nmap.py @@ -8,7 +8,7 @@ from generalresearch.managers.network.tool_run import ToolRunManager from generalresearch.models.network.definitions import IPProtocol from generalresearch.models.network.nmap.execute import execute_nmap from generalresearch.models.network.nmap.result import NmapResult, PortState -from generalresearch.models.network.tool_run import NmapRun, Status, ToolClass, ToolName +from generalresearch.models.network.tool_run import NmapRun, ToolClass, ToolName fake = faker.Faker() diff --git a/tests/models/spectrum/test_survey.py b/tests/models/spectrum/test_survey.py index 7ddd407..f97860f 100644 --- a/tests/models/spectrum/test_survey.py +++ b/tests/models/spectrum/test_survey.py @@ -405,3 +405,47 @@ class TestSpectrumSurvey: assert (None, {"c", "d"}) == s.determine_eligibility_soft( {"a": True, "b": True, "c": None, "d": None} ) + + +def test_spectrum_something(spectrum_api_surveys_json: list[str]): + # make sure hashes for 111111 are in db + c1 = SpectrumCondition( + question_id="1001", + value_type=ConditionValueType.LIST, + values=["a", "b", "c"], + negate=False, + logical_operator=LogicalOperator.OR, + ) + c2 = SpectrumCondition( + question_id="1001", + value_type=ConditionValueType.LIST, + values=["a"], + negate=False, + logical_operator=LogicalOperator.OR, + ) + c3 = SpectrumCondition( + question_id="1002", + value_type=ConditionValueType.RANGE, + values=["18-24", "30-32"], + negate=False, + logical_operator=LogicalOperator.OR, + ) + c4 = SpectrumCondition( + question_id="212", + value_type=ConditionValueType.LIST, + values=["23", "24"], + negate=False, + logical_operator=LogicalOperator.OR, + ) + c5 = SpectrumCondition( + question_id="1031", + value_type=ConditionValueType.LIST, + values=["113", "114", "121"], + negate=False, + logical_operator=LogicalOperator.OR, + ) + _conditions = [c1, c2, c3, c4, c5] + + survey = SpectrumSurvey.model_validate_json(spectrum_api_surveys_json[0]) + assert c1.criterion_hash in survey.qualifications + assert c3.criterion_hash in survey.qualifications diff --git a/tests/models/thl/test_product.py b/tests/models/thl/test_product.py index adf276d..880799a 100644 --- a/tests/models/thl/test_product.py +++ b/tests/models/thl/test_product.py @@ -1013,7 +1013,7 @@ class TestProductCache: assert res is None with pytest.raises(expected_exception=AssertionError): product.set_cache( - thl_lm=thl_lm, + thl_lm=thl_ledger_manager, ds=mnt_filepath, client=client_no_amm, bp_pem=brokerage_product_payout_event_manager, @@ -1035,7 +1035,7 @@ class TestProductCache: # Now try again with everything in place product.set_cache( - thl_lm=thl_lm, + thl_lm=thl_ledger_manager, ds=mnt_filepath, client=client_no_amm, bp_pem=brokerage_product_payout_event_manager, -- cgit v1.2.3 From ba544e2ba31432aad4d2acaba3e1f90c27137ded Mon Sep 17 00:00:00 2001 From: Max Nanis Date: Fri, 28 Aug 2026 00:22:17 -0700 Subject: Ruff evening! --- generalresearch/managers/events.py | 40 +++---- generalresearch/managers/gr/authentication.py | 3 +- generalresearch/managers/innovate/survey.py | 2 +- generalresearch/managers/leaderboard/manager.py | 8 +- generalresearch/managers/morning/survey.py | 2 +- generalresearch/managers/network/label.py | 2 +- generalresearch/managers/network/mtr.py | 9 +- generalresearch/managers/network/nmap.py | 9 +- generalresearch/managers/network/rdns.py | 5 +- generalresearch/managers/sago/survey.py | 2 +- generalresearch/managers/spectrum/survey.py | 2 +- generalresearch/managers/thl/buyer.py | 4 +- generalresearch/managers/thl/cashout_method.py | 4 +- generalresearch/managers/thl/category.py | 7 +- generalresearch/managers/thl/contest_manager.py | 9 +- generalresearch/managers/thl/ipinfo.py | 8 +- .../managers/thl/ledger_manager/ledger.py | 13 ++- generalresearch/managers/thl/payout.py | 46 ++++---- generalresearch/managers/thl/product.py | 33 +++--- generalresearch/managers/thl/profiling/user_upk.py | 8 +- generalresearch/managers/thl/survey.py | 121 ++++++++++----------- generalresearch/managers/thl/survey_penalty.py | 1 - generalresearch/managers/thl/tango_api.py | 2 +- .../managers/thl/user_manager/__init__.py | 2 +- .../thl/user_manager/mysql_user_manager.py | 11 +- generalresearch/managers/thl/userhealth.py | 8 +- generalresearch/managers/thl/wall.py | 2 +- generalresearch/managers/thl/wallet/tango.py | 6 +- generalresearch/mariadb.py | 8 -- generalresearch/models/admin/__init__.py | 2 +- generalresearch/models/admin/request.py | 2 +- generalresearch/models/cint/question.py | 2 +- generalresearch/models/cint/survey.py | 15 +-- generalresearch/models/custom_types.py | 6 +- generalresearch/models/dynata/survey.py | 12 +- generalresearch/models/dynata/task_collection.py | 2 +- generalresearch/models/gr/authentication.py | 13 +-- generalresearch/models/gr/business.py | 64 +++++------ generalresearch/models/gr/team.py | 19 ++-- generalresearch/models/innovate/question.py | 10 +- generalresearch/models/innovate/survey.py | 39 ++++--- generalresearch/models/legacy/bucket.py | 30 +++-- generalresearch/models/legacy/questions.py | 32 ++---- generalresearch/models/morning/survey.py | 10 +- generalresearch/models/morning/task_collection.py | 2 +- generalresearch/models/network/label.py | 12 +- generalresearch/models/network/nmap/result.py | 2 +- generalresearch/models/network/rdns/command.py | 2 +- generalresearch/models/precision/question.py | 6 +- generalresearch/models/precision/survey.py | 2 +- generalresearch/models/prodege/question.py | 11 +- generalresearch/models/prodege/survey.py | 6 +- generalresearch/models/prodege/task_collection.py | 2 +- generalresearch/models/repdata/question.py | 4 +- generalresearch/models/repdata/survey.py | 7 +- generalresearch/models/repdata/task_collection.py | 2 +- generalresearch/models/sago/question.py | 5 +- generalresearch/models/sago/survey.py | 24 ++-- generalresearch/models/thl/contest/__init__.py | 2 +- generalresearch/models/thl/contest/contest.py | 2 +- .../models/thl/contest/contest_entry.py | 1 + generalresearch/models/thl/contest/raffle.py | 6 +- generalresearch/models/thl/demographics.py | 8 +- generalresearch/models/thl/finance.py | 2 +- generalresearch/models/thl/ledger.py | 4 +- generalresearch/models/thl/offerwall/__init__.py | 2 +- generalresearch/models/thl/offerwall/base.py | 6 +- generalresearch/models/thl/payout_format.py | 2 +- generalresearch/models/thl/product.py | 4 +- .../models/thl/profiling/marketplace.py | 7 +- generalresearch/models/thl/report_task.py | 2 +- generalresearch/models/thl/session.py | 11 +- generalresearch/models/thl/soft_pair.py | 2 +- generalresearch/models/thl/survey/penalty.py | 2 +- .../models/thl/survey/task_collection.py | 3 +- generalresearch/models/thl/task_status.py | 9 +- generalresearch/pg_helper.py | 6 +- generalresearch/sql_helper.py | 2 +- generalresearch/utils/grpc_logger.py | 6 +- generalresearch/wall_status_codes/fullcircle.py | 2 +- generalresearch/wall_status_codes/innovate.py | 2 +- generalresearch/wall_status_codes/lucid.py | 2 +- generalresearch/wall_status_codes/morning.py | 2 +- generalresearch/wall_status_codes/pollfish.py | 2 +- test_utils/managers/contest/conftest.py | 4 +- test_utils/models/contest/conftest.py | 37 ++++--- test_utils/spectrum/conftest.py | 77 ++++++++++--- .../incite/collections/test_df_collection_base.py | 17 ++- .../collections/test_df_collection_item_base.py | 19 +++- .../test_df_collection_thl_marketplaces.py | 14 ++- .../collections/test_df_collection_thl_web.py | 120 ++++++++++++++------ .../mergers/foundations/test_user_id_product.py | 21 +++- tests/incite/mergers/test_pop_ledger.py | 6 +- tests/incite/test_collection_base.py | 2 +- tests/managers/thl/test_contest/test_milestone.py | 11 +- .../test_ledger/test_thl_lm_tx__user_payouts.py | 22 ++-- tests/models/spectrum/test_survey.py | 44 +------- 97 files changed, 669 insertions(+), 546 deletions(-) (limited to 'tests/incite/collections') diff --git a/generalresearch/managers/events.py b/generalresearch/managers/events.py index f6c429e..efc8c0d 100644 --- a/generalresearch/managers/events.py +++ b/generalresearch/managers/events.py @@ -343,8 +343,8 @@ class TaskStatsManager(RedisManager): by_source=live_tasks_max_payout_by_source, ) - task_created_count_last_1h = dict() - task_created_count_last_24h = dict() + task_created_count_last_1h = {} + task_created_count_last_24h = {} for source in sources: task_created_count_last_1h[source] = pipe_res.pop(0) task_created_count_last_24h[source] = pipe_res.pop(0) @@ -381,25 +381,25 @@ class SessionStatsManager(RedisManager): older than 1 hr (in the 1 hr bucket) will expire. """ - # Must be ordered. Don't change this - global_keys = [ - "session_enters_last_1h", - "session_enters_last_24h", - "session_fails_last_1h", - "session_fails_last_24h", - "session_completes_last_1h", - "session_completes_last_24h", - "sum_payouts_last_1h", - "sum_payouts_last_24h", - "sum_user_payouts_last_1h", - "sum_user_payouts_last_24h", - # "session_fail_loi_sum_last_1h", - "session_fail_loi_sum_last_24h", - # "session_complete_loi_sum_last_1h", - "session_complete_loi_sum_last_24h", - ] - def __init__(self, *args, **kwargs): + # Must be ordered. Don't change this + self.global_keys = [ + "session_enters_last_1h", + "session_enters_last_24h", + "session_fails_last_1h", + "session_fails_last_24h", + "session_completes_last_1h", + "session_completes_last_24h", + "sum_payouts_last_1h", + "sum_payouts_last_24h", + "sum_user_payouts_last_1h", + "sum_user_payouts_last_24h", + # "session_fail_loi_sum_last_1h", + "session_fail_loi_sum_last_24h", + # "session_complete_loi_sum_last_1h", + "session_complete_loi_sum_last_24h", + ] + super().__init__(*args, **kwargs) self.SUM_HASH_LUA = self.redis_client.register_script(SUM_HASH_LUA_SCRIPT) diff --git a/generalresearch/managers/gr/authentication.py b/generalresearch/managers/gr/authentication.py index 80bee4b..721895e 100644 --- a/generalresearch/managers/gr/authentication.py +++ b/generalresearch/managers/gr/authentication.py @@ -270,7 +270,6 @@ class GRTokenManager(PostgresManager): ) conn.commit() - def get_by_user_id(self, user_id: PositiveInt) -> GRToken | None: # django authtoken_token table has (user_id) UNIQUE constraint # therefore, this will only return 0 or 1 GRTokens @@ -295,7 +294,7 @@ class GRTokenManager(PostgresManager): res = result[0] - for k, _ in res.items(): + for k in res: if isinstance(res[k], datetime): res[k] = res[k].replace(tzinfo=UTC) diff --git a/generalresearch/managers/innovate/survey.py b/generalresearch/managers/innovate/survey.py index cddfba2..f6d00a8 100644 --- a/generalresearch/managers/innovate/survey.py +++ b/generalresearch/managers/innovate/survey.py @@ -179,5 +179,5 @@ class InnovateSurveyManager(SurveyManager): if e.args[0] == 1062: existing_sns.add(sn) else: - raise e + raise self.update([surveys[sn] for sn in existing_sns]) diff --git a/generalresearch/managers/leaderboard/manager.py b/generalresearch/managers/leaderboard/manager.py index 07e3e2c..ed13cf2 100644 --- a/generalresearch/managers/leaderboard/manager.py +++ b/generalresearch/managers/leaderboard/manager.py @@ -45,9 +45,7 @@ class LeaderboardManager: self.country_iso = country_iso self.within_time_aware = None if within_time is None: - self.within_time_aware = datetime.now(tz=UTC).astimezone( - self.timezone - ) + self.within_time_aware = datetime.now(tz=UTC).astimezone(self.timezone) elif within_time.tzinfo is not None: self.within_time_aware = within_time.astimezone(self.timezone) else: @@ -123,7 +121,9 @@ class LeaderboardManager: user_idx = user_indices[0][0] user_row = user_indices[0][1] if user_row.rank == max([row.rank for row in rows]): - user_idx = [i for i, row in enumerate(rows) if row.rank == user_row.rank][0] + user_idx = next( + i for i, row in enumerate(rows) if row.rank == user_row.rank + ) start: int = max(user_idx - limit, 0) end: int = min(user_idx + limit + 1, len(rows)) diff --git a/generalresearch/managers/morning/survey.py b/generalresearch/managers/morning/survey.py index 0e29010..2488cd8 100644 --- a/generalresearch/managers/morning/survey.py +++ b/generalresearch/managers/morning/survey.py @@ -258,5 +258,5 @@ class MorningSurveyManager(SurveyManager): if e.args[0] == 1062: existing_sns.add(sn) else: - raise e + raise self.update([surveys[sn] for sn in existing_sns]) diff --git a/generalresearch/managers/network/label.py b/generalresearch/managers/network/label.py index 1f44862..f0ba9f7 100644 --- a/generalresearch/managers/network/label.py +++ b/generalresearch/managers/network/label.py @@ -30,7 +30,7 @@ class IPLabelManager(PostgresManager): params = ip_label.model_dump_postgres() with self.pg_config.make_connection() as conn, conn.cursor() as c: c.execute(query, params) - pk = c.fetchone()["id"] + _pk = c.fetchone()["id"] return ip_label def make_filter_str( diff --git a/generalresearch/managers/network/mtr.py b/generalresearch/managers/network/mtr.py index 19c5caf..54d74b7 100644 --- a/generalresearch/managers/network/mtr.py +++ b/generalresearch/managers/network/mtr.py @@ -42,8 +42,7 @@ class MTRRunManager(PostgresManager): if params_hops: c.executemany(query_hops, params_hops) else: - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute(query, params) - if params_hops: - c.executemany(query_hops, params_hops) + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute(query, params) + if params_hops: + c.executemany(query_hops, params_hops) diff --git a/generalresearch/managers/network/nmap.py b/generalresearch/managers/network/nmap.py index a8470c8..84d13ad 100644 --- a/generalresearch/managers/network/nmap.py +++ b/generalresearch/managers/network/nmap.py @@ -50,8 +50,7 @@ class NmapRunManager(PostgresManager): if nmap_run.ports: c.executemany(query_ports, params_ports) else: - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute(query, params) - if nmap_run.ports: - c.executemany(query_ports, params_ports) + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute(query, params) + if nmap_run.ports: + c.executemany(query_ports, params_ports) diff --git a/generalresearch/managers/network/rdns.py b/generalresearch/managers/network/rdns.py index 0b41a9a..c8ce913 100644 --- a/generalresearch/managers/network/rdns.py +++ b/generalresearch/managers/network/rdns.py @@ -28,6 +28,5 @@ class RDNSRunManager(PostgresManager): if c: c.execute(query, params) else: - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute(query, params) + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute(query, params) diff --git a/generalresearch/managers/sago/survey.py b/generalresearch/managers/sago/survey.py index 462d2ef..a13fbce 100644 --- a/generalresearch/managers/sago/survey.py +++ b/generalresearch/managers/sago/survey.py @@ -179,6 +179,6 @@ class SagoSurveyManager(SurveyManager): if e.args[0] == 1062: existing_sns.add(sn) else: - raise e + raise self.update([surveys[sn] for sn in existing_sns]) diff --git a/generalresearch/managers/spectrum/survey.py b/generalresearch/managers/spectrum/survey.py index 9b58d43..987ce7b 100644 --- a/generalresearch/managers/spectrum/survey.py +++ b/generalresearch/managers/spectrum/survey.py @@ -212,6 +212,6 @@ class SpectrumSurveyManager(SurveyManager): if e.args[0] == 1062: existing_sns.add(sn) else: - raise e + raise self.update([surveys[sn] for sn in existing_sns]) diff --git a/generalresearch/managers/thl/buyer.py b/generalresearch/managers/thl/buyer.py index 2cb582f..5aa2a01 100644 --- a/generalresearch/managers/thl/buyer.py +++ b/generalresearch/managers/thl/buyer.py @@ -18,8 +18,8 @@ class BuyerManager(PostgresManager): ): super().__init__(pg_config=pg_config, permissions=permissions) # self.buyer_pk: Dict[Buyer, int] = dict() - self.source_code_buyer: dict[str, Buyer] = dict() - self.source_code_pk: dict[str, int] = dict() + self.source_code_buyer: dict[str, Buyer] = {} + self.source_code_pk: dict[str, int] = {} self.populate_caches() def populate_caches(self): diff --git a/generalresearch/managers/thl/cashout_method.py b/generalresearch/managers/thl/cashout_method.py index 10282d4..f45e692 100644 --- a/generalresearch/managers/thl/cashout_method.py +++ b/generalresearch/managers/thl/cashout_method.py @@ -160,7 +160,7 @@ class CashoutMethodManager(PostgresManager): is_live: bool | None = True, ): filters = [] - params = dict() + params = {} if uuid is not None: params["uuid"] = uuid filters.append("id = %(uuid)s") @@ -292,7 +292,7 @@ class CashoutMethodManager(PostgresManager): x["type"] = PayoutType(x["provider"].upper()) if "data" not in x: - x["data"] = dict() + x["data"] = {} x["data"].update(x.pop("_data_")) x["data"]["type"] = x["type"] if user and x["type"] in {PayoutType.PAYPAL, PayoutType.CASH_IN_MAIL}: diff --git a/generalresearch/managers/thl/category.py b/generalresearch/managers/thl/category.py index 05ceb8f..e8a6aa6 100644 --- a/generalresearch/managers/thl/category.py +++ b/generalresearch/managers/thl/category.py @@ -9,8 +9,6 @@ from generalresearch.pg_helper import PostgresConfig class CategoryManager(PostgresManager): - categories = dict() - category_label_map = dict() def __init__( self, @@ -18,8 +16,9 @@ class CategoryManager(PostgresManager): permissions: Collection[Permission] | None = None, ): super().__init__(pg_config=pg_config, permissions=permissions) - self.categories: dict[UUIDStr, Category] = dict() - self.category_label_map: dict[str, Category] = dict() + self.categories: dict[UUIDStr, Category] = {} + self.category_label_map: dict[str, Category] = {} + self.populate_caches() def populate_caches(self): diff --git a/generalresearch/managers/thl/contest_manager.py b/generalresearch/managers/thl/contest_manager.py index 62146d7..64206e1 100644 --- a/generalresearch/managers/thl/contest_manager.py +++ b/generalresearch/managers/thl/contest_manager.py @@ -173,7 +173,7 @@ class ContestBaseManager(PostgresManager): except ValueError as e: if e.args[0] == "Contest not found": return None - raise e + raise @staticmethod def make_filter_str( @@ -187,7 +187,7 @@ class ContestBaseManager(PostgresManager): has_participants: bool | None = None, ) -> tuple[str, dict[str, Any]]: filters = [] - params = dict() + params = {} if product_id: params["product_id"] = product_id @@ -681,7 +681,7 @@ class RaffleContestManager(ContestBaseManager): raise ContestError(msg) if contest.entry_type == ContestEntryType.CASH: - tx = ledger_manager.create_tx_user_enter_contest( + ledger_manager.create_tx_user_enter_contest( contest_uuid=contest.uuid, contest_entry=entry ) @@ -827,7 +827,6 @@ class MilestoneContestManager(ContestBaseManager): ) self.end_milestone_contest(contest) - def enter_contest_db_work_milestone( self, contest: MilestoneUserView, user: User, incr: PositiveInt ) -> MilestoneEntry: @@ -1052,7 +1051,7 @@ class ContestManager( ) -> NonNegativeInt: contests_closed = 0 for contest in contests: - should_end, reason = contest.should_end() + should_end, _ = contest.should_end() if should_end: if hasattr(contest, "redis_client"): contest.redis_client = redis_client diff --git a/generalresearch/managers/thl/ipinfo.py b/generalresearch/managers/thl/ipinfo.py index e1143c2..9595cab 100644 --- a/generalresearch/managers/thl/ipinfo.py +++ b/generalresearch/managers/thl/ipinfo.py @@ -423,7 +423,9 @@ class IPInformationManager(PostgresManager): FROM thl_ipinformation WHERE updated >= NOW() - INTERVAL '12 hours' """ - denominator = list(pg_config.execute_sql_query(query=query))[0]["denominator"] + denominator = next(iter(pg_config.execute_sql_query(query=query)))[ + "denominator" + ] if denominator == 0: pass @@ -509,7 +511,7 @@ class GeoIpInfoManager(PostgresManagerWithRedis): res = [GeoIPInformation.model_validate_json(raw) for raw in res if raw] gs = {x.ip: x for x in res} - res2 = dict() + res2 = {} for ip, (normalized_ip, lookup_prefix) in ip_norm_lookup.items(): if normalized_ip not in gs: # try the non-normalized (remove me also 28 days from 2025-11-15) @@ -719,7 +721,7 @@ class GeoIpInfoManager(PostgresManagerWithRedis): gs = [GeoIPInformation.from_mysql(i) for i in res] gs = {g.ip: g for g in gs} - res2 = dict() + res2 = {} for ip, (normalized_ip, lookup_prefix) in ip_norm_lookup.items(): if normalized_ip not in gs: diff --git a/generalresearch/managers/thl/ledger_manager/ledger.py b/generalresearch/managers/thl/ledger_manager/ledger.py index 410f6ca..a5263f0 100644 --- a/generalresearch/managers/thl/ledger_manager/ledger.py +++ b/generalresearch/managers/thl/ledger_manager/ledger.py @@ -150,7 +150,7 @@ class LedgerTransactionManager(LedgerManagerBasePostgres): ), "LedgerTransactionManager has insufficient Permissions" if metadata is None: - metadata = dict() + metadata = {} if created is None: created = datetime.now(tz=UTC) @@ -429,7 +429,7 @@ class LedgerTransactionManager(LedgerManagerBasePostgres): ) } else: - metadata = dict() + metadata = {} entries = [ LedgerEntry( @@ -750,7 +750,7 @@ class LedgerMetadataManager(LedgerManagerBasePostgres): """ - tx_ids = set([tx.id for tx in transactions]) + tx_ids = {tx.id for tx in transactions} res = self.pg_config.execute_sql_query( query=""" SELECT @@ -782,7 +782,7 @@ class LedgerMetadataManager(LedgerManagerBasePostgres): from the database. """ - tx_ids = set([tx.id for tx in transactions]) + tx_ids = {tx.id for tx in transactions} res = self.pg_config.execute_sql_query( query=""" SELECT tx_meta.id @@ -792,7 +792,7 @@ class LedgerMetadataManager(LedgerManagerBasePostgres): params=[list(tx_ids)], ) - return set([i["id"] for i in res]) + return {i["id"] for i in res} class LedgerEntryManager(LedgerManagerBasePostgres): @@ -803,7 +803,7 @@ class LedgerEntryManager(LedgerManagerBasePostgres): def get_tx_entries_by_txs( self, transactions: list[LedgerTransaction] ) -> list[LedgerEntry]: - tx_ids = set([tx.id for tx in transactions]) + tx_ids = {tx.id for tx in transactions} tx_entries = self.pg_config.execute_sql_query( query=""" SELECT @@ -1141,4 +1141,5 @@ class LedgerManager( } for k, v in d.items(): v["total"] = (v["debit"] - v["credit"]) * k.normal_balance.value + return d diff --git a/generalresearch/managers/thl/payout.py b/generalresearch/managers/thl/payout.py index 8bc6843..f50e0d2 100644 --- a/generalresearch/managers/thl/payout.py +++ b/generalresearch/managers/thl/payout.py @@ -95,9 +95,9 @@ class PayoutEventManager(PostgresManagerWithRedis): with self.pg_config.make_connection() as conn: with conn.cursor() as c: c.execute(query=query, params=d) - assert c.rowcount == 1, ( - "Nothing was updated! Are you sure this payout_event exists?" - ) + assert ( + c.rowcount == 1 + ), "Nothing was updated! Are you sure this payout_event exists?" conn.commit() @@ -140,7 +140,7 @@ class UserPayoutEventManager(PayoutEventManager): # the purposes of returning to the user. pe = self.get_by_uuid(pe_uuid=pe_uuid) - transaction_info = dict() + transaction_info = {} order: dict[str, Any] = pe.order_data if pe.payout_type == PayoutType.TANGO and pe.status == PayoutStatus.COMPLETE: reward = order["reward"] @@ -411,7 +411,7 @@ class BrokerageProductPayoutEventManager(PayoutEventManager): *** IT IS ONLY FOR Brokerage Product PAYOUTS *** """ - params = dict() + params = {} filters = [] if ext_ref_id: # This is transaction id for tracking ACH/Wires with a banking institution @@ -630,9 +630,9 @@ class BusinessPayoutEventManager(PostgresManagerWithRedis): for bp_payout in d["bp_payouts"]: bp_payout["created"] = datetime.fromisoformat(bp_payout["created"]) bpe = BusinessPayoutEvent.model_validate(d) - assert bpe.bp_payouts is not None and len(bpe.bp_payouts) > 0, ( - "No BP payouts found for this Business Payout Event. This shouldn't happen!" - ) + assert ( + bpe.bp_payouts is not None and len(bpe.bp_payouts) > 0 + ), "No BP payouts found for this Business Payout Event. This shouldn't happen!" return bpe def filter_by( @@ -677,9 +677,9 @@ class BusinessPayoutEventManager(PostgresManagerWithRedis): for bp_payout in row["bp_payouts"]: bp_payout["created"] = datetime.fromisoformat(bp_payout["created"]) bpe = BusinessPayoutEvent.model_validate(row) - assert bpe.bp_payouts is not None and len(bpe.bp_payouts) > 0, ( - "No BP payouts found for this Business Payout Event. This shouldn't happen!" - ) + assert ( + bpe.bp_payouts is not None and len(bpe.bp_payouts) > 0 + ), "No BP payouts found for this Business Payout Event. This shouldn't happen!" bpes.append(bpe) return bpes @@ -696,9 +696,9 @@ class BusinessPayoutEventManager(PostgresManagerWithRedis): for bp_pe in bpe.bp_payouts ] txs = thl_lm.get_tx_ids_by_tags(tags=tags) - assert len(txs) == len(bpe.bp_payouts), ( - f"Expected {len(bpe.bp_payouts)} BP payouts but found {len(txs)}!" - ) + assert len(txs) == len( + bpe.bp_payouts + ), f"Expected {len(bpe.bp_payouts)} BP payouts but found {len(txs)}!" return True def resume_failed_business_payout( @@ -824,9 +824,9 @@ class BusinessPayoutEventManager(PostgresManagerWithRedis): shortfall: int = int(target_amount) - w_df["deduction"].sum() w_df["remaining_balance"] = w_df["available_balance"] - w_df["deduction"] - assert w_df[w_df["deduction"] > w_df["available_balance"]].empty, ( - "Trying to deduct more from an Product than what is available" - ) + assert w_df[ + w_df["deduction"] > w_df["available_balance"] + ].empty, "Trying to deduct more from an Product than what is available" return w_df @@ -898,9 +898,9 @@ class BusinessPayoutEventManager(PostgresManagerWithRedis): ) -> BusinessPayoutEvent: assert isinstance(bpe, BusinessPayoutEventCreate) assert bpe.bp_payouts, "Must provide at least one BP Payout" - assert {bp_pe.status for bp_pe in bpe.bp_payouts} == {PayoutStatus.PENDING}, ( - "All BP Payouts must be PENDING" - ) + assert {bp_pe.status for bp_pe in bpe.bp_payouts} == { + PayoutStatus.PENDING + }, "All BP Payouts must be PENDING" INSERT_SUPPLIER_PAYOUT = """ INSERT INTO supplier_payout ( business_id, created, amount, @@ -993,9 +993,9 @@ class BusinessPayoutEventManager(PostgresManagerWithRedis): # Can't pay any Products that don't have a remaining balance df = df[df["remaining_balance"] > 0].copy() - assert df.deduction.sum() == business.balance.recoup, ( - "recoup_proportional failure" - ) + assert ( + df.deduction.sum() == business.balance.recoup + ), "recoup_proportional failure" df["issue_amount"] = BusinessPayoutEventManager.distribute_amount( df=df, amount=amount diff --git a/generalresearch/managers/thl/product.py b/generalresearch/managers/thl/product.py index 54fa7c8..aac2979 100644 --- a/generalresearch/managers/thl/product.py +++ b/generalresearch/managers/thl/product.py @@ -33,6 +33,7 @@ if TYPE_CHECKING: ProfilingConfig, SessionConfig, SourcesConfig, + SupplyConfig, UserCreateConfig, UserHealthConfig, UserWalletConfig, @@ -167,15 +168,14 @@ class ProductManager(PostgresManager): if filter_uuids is None or len(filter_uuids) == 0: return [] - with self.pg_config.make_connection() as sql_connection: - with sql_connection.cursor() as c: - res = [] - for chunk in chunked(filter_uuids, 500): - res.extend( - self.fetch_uuids_( - c=c, filter_uuids=chunk, filter_column=filter_column - ) + with self.pg_config.make_connection() as sql_connection, sql_connection.cursor() as c: + res = [] + for chunk in chunked(filter_uuids, 500): + res.extend( + self.fetch_uuids_( + c=c, filter_uuids=chunk, filter_column=filter_column ) + ) return res def fetch_uuids_( @@ -258,9 +258,10 @@ class ProductManager(PostgresManager): for k, v in res1.items(): try: r.append(Product.model_validate(v)) - except ValidationError as e: + except ValidationError: logger.info(f"failed to parse product: {k}") - raise e + raise + return r def create( @@ -272,7 +273,7 @@ class ProductManager(PostgresManager): business_id: UUIDStr | None = None, harmonizer_domain: str | None = None, commission_pct: Decimal = Decimal("0.05"), - sources_config: SourcesConfig | SupplyConfigs | None = None, + sources_config: SourcesConfig | SupplyConfig | None = None, payout_config: PayoutConfig | None = None, session_config: SessionConfig | None = None, profiling_config: ProfilingConfig | None = None, @@ -360,10 +361,10 @@ class ProductManager(PostgresManager): insert_data["payments_enabled"] = instance.payments_enabled try: - insert_data["id_int"] = list(self.pg_config.execute_sql_query(query=""" + insert_data["id_int"] = next(iter(self.pg_config.execute_sql_query(query=""" SELECT COALESCE(MAX(id_int), 0) + 1 as id_int FROM userprofile_brokerageproduct - """))[0]["id_int"] + """)))["id_int"] instance.id_int = insert_data["id_int"] query = """ @@ -400,14 +401,14 @@ class ProductManager(PostgresManager): try: return self.get_by_uuid(product_uuid=instance.id) - except Exception: + except AssertionError: pass finally: self.cache_clear(instance.id) # If we couldn't find the Product, then go ahead and raise. capture_exception(e) - raise e + raise bpconfig = instance.model_dump( include={"sources_config", "user_wallet"}, mode="json" @@ -477,7 +478,7 @@ class ProductManager(PostgresManager): data["grs_domain"] = data.pop("harmonizer_domain") data = {k: v for k, v in data.items() if k in in_bp_keys} data["id"] = product_uuid - update_str = ", ".join(f"{k}=%({k})s" for k in data.keys()) + update_str = ", ".join(f"{k}=%({k})s" for k in data) self.pg_config.execute_write( f""" UPDATE userprofile_brokerageproduct diff --git a/generalresearch/managers/thl/profiling/user_upk.py b/generalresearch/managers/thl/profiling/user_upk.py index a2cddb3..6106037 100644 --- a/generalresearch/managers/thl/profiling/user_upk.py +++ b/generalresearch/managers/thl/profiling/user_upk.py @@ -158,7 +158,7 @@ class UserUpkManager(PostgresManagerWithRedis): country_isos = {x["country_iso"] for x in upk_ans_dict} assert len(country_isos) == 1 - country_iso = list(country_isos)[0] + country_iso = next(iter(country_isos)) for x in upk_ans_dict: x["pred"] = x["pred"].replace("gr:", "") x["obj"] = x["obj"].replace("gr:", "") @@ -304,15 +304,15 @@ class UserUpkManager(PostgresManagerWithRedis): def set_user_upk(self, upk_ans: list[UpkQuestionAnswer]) -> None: user_id = {x.user_id for x in upk_ans} assert len(user_id) == 1, "only run for 1 user at a time" - user_id = list(user_id)[0] + user_id = next(iter(user_id)) curr_upk = self.get_user_upk(user_id=user_id) curr_upk_simple = self.get_user_upk_simple(user_id=user_id) new_upk_simple = defaultdict(set) delete_items = set() - upk_multi = list() - delete_upk_multi = list() + upk_multi = [] + delete_upk_multi = [] for x in upk_ans: # For zero or more (multiple values) We want all values to equal these. # Might involve deleting values if they exist and are not in upk_ans diff --git a/generalresearch/managers/thl/survey.py b/generalresearch/managers/thl/survey.py index a9ec841..024ad38 100644 --- a/generalresearch/managers/thl/survey.py +++ b/generalresearch/managers/thl/survey.py @@ -134,7 +134,7 @@ class SurveyManager(PostgresManager): if len(survey_keys) == 0: return [] - params = dict() + params = {} survey_source_ids = defaultdict(set) for sk in survey_keys: @@ -354,59 +354,6 @@ class SurveyManager(PostgresManager): class SurveyStatManager(PostgresManager): - KEYS = [ - "survey_id", - "quota_id", - "country_iso", - "version", - "cpi", - "complete_too_fast_cutoff", - "prescreen_conv_alpha", - "prescreen_conv_beta", - "conv_alpha", - "conv_beta", - "dropoff_alpha", - "dropoff_beta", - "completion_time_mu", - "completion_time_sigma", - "mobile_eligible_alpha", - "mobile_eligible_beta", - "desktop_eligible_alpha", - "desktop_eligible_beta", - "tablet_eligible_alpha", - "tablet_eligible_beta", - "long_fail_rate", - "user_report_coeff", - "recon_likelihood", - "score_x0", - "score_x1", - "score", - "updated_at", - "survey_is_live", - "survey_survey_id", - "survey_source", - ] - - SURVEY_STATS_COL_MAP = { - "PRESCREEN_CONVERSION.alpha": "prescreen_conv_alpha", - "PRESCREEN_CONVERSION.beta": "prescreen_conv_beta", - "CONVERSION.alpha": "conv_alpha", - "CONVERSION.beta": "conv_beta", - "COMPLETION_TIME.mu": "completion_time_mu", - "COMPLETION_TIME.sigma": "completion_time_sigma", - "LONG_FAIL.value": "long_fail_rate", - "USER_REPORT_COEFF.value": "user_report_coeff", - "RECON_LIKELIHOOD.value": "recon_likelihood", - "DROPOFF_RATE.alpha": "dropoff_alpha", - "DROPOFF_RATE.beta": "dropoff_beta", - "IS_MOBILE_ELIGIBLE.alpha": "mobile_eligible_alpha", - "IS_MOBILE_ELIGIBLE.beta": "mobile_eligible_beta", - "IS_DESKTOP_ELIGIBLE.alpha": "desktop_eligible_alpha", - "IS_DESKTOP_ELIGIBLE.beta": "desktop_eligible_beta", - "IS_TABLET_ELIGIBLE.alpha": "tablet_eligible_alpha", - "IS_TABLET_ELIGIBLE.beta": "tablet_eligible_beta", - "cpi": "cpi", - } def __init__( self, @@ -419,6 +366,60 @@ class SurveyStatManager(PostgresManager): ) # self.ensure_surveystat_key_type() + self.KEYS = [ + "survey_id", + "quota_id", + "country_iso", + "version", + "cpi", + "complete_too_fast_cutoff", + "prescreen_conv_alpha", + "prescreen_conv_beta", + "conv_alpha", + "conv_beta", + "dropoff_alpha", + "dropoff_beta", + "completion_time_mu", + "completion_time_sigma", + "mobile_eligible_alpha", + "mobile_eligible_beta", + "desktop_eligible_alpha", + "desktop_eligible_beta", + "tablet_eligible_alpha", + "tablet_eligible_beta", + "long_fail_rate", + "user_report_coeff", + "recon_likelihood", + "score_x0", + "score_x1", + "score", + "updated_at", + "survey_is_live", + "survey_survey_id", + "survey_source", + ] + + self.SURVEY_STATS_COL_MAP = { + "PRESCREEN_CONVERSION.alpha": "prescreen_conv_alpha", + "PRESCREEN_CONVERSION.beta": "prescreen_conv_beta", + "CONVERSION.alpha": "conv_alpha", + "CONVERSION.beta": "conv_beta", + "COMPLETION_TIME.mu": "completion_time_mu", + "COMPLETION_TIME.sigma": "completion_time_sigma", + "LONG_FAIL.value": "long_fail_rate", + "USER_REPORT_COEFF.value": "user_report_coeff", + "RECON_LIKELIHOOD.value": "recon_likelihood", + "DROPOFF_RATE.alpha": "dropoff_alpha", + "DROPOFF_RATE.beta": "dropoff_beta", + "IS_MOBILE_ELIGIBLE.alpha": "mobile_eligible_alpha", + "IS_MOBILE_ELIGIBLE.beta": "mobile_eligible_beta", + "IS_DESKTOP_ELIGIBLE.alpha": "desktop_eligible_alpha", + "IS_DESKTOP_ELIGIBLE.beta": "desktop_eligible_beta", + "IS_TABLET_ELIGIBLE.alpha": "tablet_eligible_alpha", + "IS_TABLET_ELIGIBLE.beta": "tablet_eligible_beta", + "cpi": "cpi", + } + # # def ensure_surveystat_key_type(self): # SQL = """ @@ -570,12 +571,10 @@ class SurveyStatManager(PostgresManager): = (v.survey_id, v.quota_id, v.country_iso, v.version); """ params = [item for row in keys for item in row] - with self.pg_config.make_connection() as conn: - # self.register_surveystat_key(conn) - with conn.cursor() as c: - c.execute(query, params=params) - res = c.fetchall() - # print('\n'.join([x['QUERY PLAN'] for x in res])) + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute(query, params=params) + res = c.fetchall() + # print('\n'.join([x['QUERY PLAN'] for x in res])) return [SurveyStat.model_validate(x) for x in res] def update_surveystats_for_source( @@ -633,7 +632,7 @@ class SurveyStatManager(PostgresManager): country_iso: str | None = None, ) -> tuple[str, dict[str, Any]]: filters = [] - params = dict() + params = {} if updated_after is not None: params["updated_after"] = updated_after filters.append("ss.updated_at >= %(updated_after)s") diff --git a/generalresearch/managers/thl/survey_penalty.py b/generalresearch/managers/thl/survey_penalty.py index efaa930..bf914cb 100644 --- a/generalresearch/managers/thl/survey_penalty.py +++ b/generalresearch/managers/thl/survey_penalty.py @@ -61,7 +61,6 @@ class SurveyPenaltyManager(RedisManager): return f"{self.redis_prefix}:{uuid_id}" def set_penalties(self, penalties: list[Penalty]): - """ """ if len(penalties) > 1000: LOG.warning("SurveyPenaltyManager.set_penalties batch me!") assert len(penalties) < 10_000, "something is surely wrong" diff --git a/generalresearch/managers/thl/tango_api.py b/generalresearch/managers/thl/tango_api.py index 657224e..dab560e 100644 --- a/generalresearch/managers/thl/tango_api.py +++ b/generalresearch/managers/thl/tango_api.py @@ -122,7 +122,7 @@ class TangoClient: return self.get_order(reference_order_id) except TangoError as e: if "The order you requested cannot be found" not in e.args[0]: - raise e + raise return None def create_order(self, order: TangoOrderRequest) -> dict[str, Any]: diff --git a/generalresearch/managers/thl/user_manager/__init__.py b/generalresearch/managers/thl/user_manager/__init__.py index 0392edb..b3fa8f6 100644 --- a/generalresearch/managers/thl/user_manager/__init__.py +++ b/generalresearch/managers/thl/user_manager/__init__.py @@ -63,7 +63,7 @@ def parse_bp_trust_df(fp: str | Path) -> dict[str, Any]: "entrance_limit_value": convert_int, "median_daily_completes_7d": convert_int, } - bptrust = dict() + bptrust = {} with open(fp, newline="") as csvfile: reader = csv.reader(csvfile) diff --git a/generalresearch/managers/thl/user_manager/mysql_user_manager.py b/generalresearch/managers/thl/user_manager/mysql_user_manager.py index d2d0ffc..e0a7548 100644 --- a/generalresearch/managers/thl/user_manager/mysql_user_manager.py +++ b/generalresearch/managers/thl/user_manager/mysql_user_manager.py @@ -139,11 +139,10 @@ class MysqlUserManager: """) try: - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute(query=query, params=params) - user_id = c.fetchone()["id"] - except psycopg.IntegrityError as e: + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute(query=query, params=params) + user_id = c.fetchone()["id"] + except psycopg.IntegrityError: # Two machines/processes are trying to create this same (product_id, product_user_id) # at the same time. There's a unique index, so mysql will not let two be created. # The 2nd should get an IntegrityError, meaning this already exists, and we can just query it. @@ -160,7 +159,7 @@ class MysqlUserManager: else: # We specifically queried the NON read-replica, and we got an IntegrityError, so # something else must be wrong... - raise e + raise else: user = User( user_id=user_id, diff --git a/generalresearch/managers/thl/userhealth.py b/generalresearch/managers/thl/userhealth.py index babed04..26f08b4 100644 --- a/generalresearch/managers/thl/userhealth.py +++ b/generalresearch/managers/thl/userhealth.py @@ -221,7 +221,7 @@ class IPRecordManager(PostgresManagerWithRedis): "forwarded_ip5", "forwarded_ip6", ] - for col, ip in zip_longest( + for col, fwd_ip in zip_longest( fips_cols, [ forwarded_ip1, @@ -233,7 +233,7 @@ class IPRecordManager(PostgresManagerWithRedis): ], fillvalue=None, ): - data[col] = ipaddress.ip_address(ip).exploded if ip else ip + data[col] = ipaddress.ip_address(fwd_ip).exploded if fwd_ip else fwd_ip self.pg_config.execute_write( query=""" @@ -490,9 +490,7 @@ class AuditLogManager(PostgresManager): created_after: datetime | None = None, ) -> tuple[str, dict[str, Any]]: assert user_ids, "must pass at least 1 user_id" - assert all( - [isinstance(uid, int) for uid in user_ids] - ), "must pass user_id as int" + assert all(isinstance(uid, int) for uid in user_ids), "must pass user_id as int" if created_after is None: created_after = datetime.now(tz=UTC) - timedelta(days=7) diff --git a/generalresearch/managers/thl/wall.py b/generalresearch/managers/thl/wall.py index 03ca1c6..ac9fb62 100644 --- a/generalresearch/managers/thl/wall.py +++ b/generalresearch/managers/thl/wall.py @@ -484,7 +484,7 @@ class WallManager(PostgresManager): ORDER BY rs.source, rs.survey_id; """ - params = dict() + params = {} filters = [] # Instead of doing a big IN with a big set of tuples, since we know diff --git a/generalresearch/managers/thl/wallet/tango.py b/generalresearch/managers/thl/wallet/tango.py index 4abfc70..be8fd97 100644 --- a/generalresearch/managers/thl/wallet/tango.py +++ b/generalresearch/managers/thl/wallet/tango.py @@ -44,7 +44,7 @@ def complete_tango_order( tango_client=tango_client, ) - except Exception: + except AssertionError: # todo: its possible the order went through, but something else was wrong # we should try to retrieve the order by its ref_id and confirm it really # failed... @@ -70,8 +70,8 @@ def create_tango_order( """ Create a tango gift card order. Throws exception if anything is not right. - # https://integration-www.tangocard.com/raas_api_console/v2/ - # https://www.apimatic.io/apidocs/tangocard/v/2_3_4#/python + - https://integration-www.tangocard.com/raas_api_console/v2/ + - https://www.apimatic.io/apidocs/tangocard/v/2_3_4#/python :param utid: Card identifier :param amount: requested card value in USD diff --git a/generalresearch/mariadb.py b/generalresearch/mariadb.py index 8bcd8ee..5d43f97 100644 --- a/generalresearch/mariadb.py +++ b/generalresearch/mariadb.py @@ -32,11 +32,3 @@ def example(): for m in zip(c.metadata["field"], c.metadata["ext_type_or_format"]): # here we can just check if the field's ext_field_flag == 'UUID' (2) print(m[0], ext_field_flags_rev[m[1]]) - - -def get_column_types(): - # How does django do this? - res = """ - SELECT column_name, data_type - FROM information_schema.columns - WHERE table_name = 'morning_userpid' AND table_schema = DATABASE()""" diff --git a/generalresearch/models/admin/__init__.py b/generalresearch/models/admin/__init__.py index ad6302b..344c34a 100644 --- a/generalresearch/models/admin/__init__.py +++ b/generalresearch/models/admin/__init__.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime import pandas as pd from dateutil import relativedelta diff --git a/generalresearch/models/admin/request.py b/generalresearch/models/admin/request.py index 5fdc784..6112786 100644 --- a/generalresearch/models/admin/request.py +++ b/generalresearch/models/admin/request.py @@ -119,7 +119,7 @@ class ReportRequest(BaseModel): @property def end_naive(self) -> datetime: - return datetime.now(tz=None) + return datetime.now(tz=None) # noqa @property def ts_start(self) -> pd.Timestamp: diff --git a/generalresearch/models/cint/question.py b/generalresearch/models/cint/question.py index f8a287a..4c0e52f 100644 --- a/generalresearch/models/cint/question.py +++ b/generalresearch/models/cint/question.py @@ -44,7 +44,7 @@ class CintQuestionType(StrEnum): # This seems to be invalid as there are no options??? "Grid": None, } - return API_TYPE_MAP[a] if a in API_TYPE_MAP else None + return API_TYPE_MAP.get(a) class CintUserQuestionAnswer(MarketplaceUserQuestionAnswer): diff --git a/generalresearch/models/cint/survey.py b/generalresearch/models/cint/survey.py index cfc91ef..fde4559 100644 --- a/generalresearch/models/cint/survey.py +++ b/generalresearch/models/cint/survey.py @@ -12,6 +12,7 @@ from pydantic import ( ConfigDict, Field, NonNegativeInt, + ValidationError, computed_field, model_validator, ) @@ -67,7 +68,7 @@ class CintQuota(BaseModel): condition_hashes: list[str] | None = Field(min_length=1, default=None) def __hash__(self): - return hash(tuple((tuple(self.condition_hashes), self.quota_id))) + return hash((tuple(self.condition_hashes), self.quota_id)) @model_validator(mode="after") def validate_condition_len(self) -> Self: @@ -317,7 +318,7 @@ class CintSurvey(MarketplaceTask): def from_api(cls, d: dict[str, Any]) -> Self | None: try: return cls._from_api(d) - except Exception as e: + except ValidationError as e: logger.warning(f"Unable to parse survey: {d}. {e}") return None @@ -370,8 +371,8 @@ class CintSurvey(MarketplaceTask): d["mobile_conversion"] = None d["revenue_per_click"] = None - d["conditions"] = dict() - d.setdefault("survey_qualifications", list()) + d["conditions"] = {} + d.setdefault("survey_qualifications", []) qualifications = [CintCondition.from_api(q) for q in d["survey_qualifications"]] for q in qualifications: d["conditions"][q.criterion_hash] = q @@ -416,7 +417,7 @@ class CintSurvey(MarketplaceTask): return d @classmethod - def from_mysql(cls, d: Dict[str, Any]) -> Self: + def from_mysql(cls, d: dict[str, Any]) -> Self: d["created_at"] = d["created_at"].replace(tzinfo=UTC) d["last_updated"] = d["last_updated"].replace(tzinfo=UTC) d["qualifications"] = json.loads(d["qualifications"]) @@ -465,7 +466,7 @@ class CintSurvey(MarketplaceTask): ) -> tuple[bool | None, set[str]]: # Many surveys have 0 quotas. Quotas are exclusionary. # They can NOT match a quota where currently_open=0 - total_quota = [q for q in self.quotas if q.quota_type == "total"][0] + total_quota = next(q for q in self.quotas if q.quota_type == "total") if not total_quota.is_open: return False, set() quotas = [q for q in self.quotas if q.quota_type != "total"] @@ -474,7 +475,7 @@ class CintSurvey(MarketplaceTask): quota_eval = { quota: quota.matches_soft(criteria_evaluation) for quota in quotas } - evals = set(g[0] for g in quota_eval.values()) + evals = {g[0] for g in quota_eval.values()} if any(m[0] is True and not q.is_open for q, m in quota_eval.items()): # matched a full quota return False, set() diff --git a/generalresearch/models/custom_types.py b/generalresearch/models/custom_types.py index c200b34..5e4db3e 100644 --- a/generalresearch/models/custom_types.py +++ b/generalresearch/models/custom_types.py @@ -98,7 +98,7 @@ LanguageISOLike = Annotated[ def check_valid_uuid(v: str) -> str: try: assert UUID(v).hex == v - except Exception: + except (ValueError, AssertionError): raise ValueError("Invalid UUID") return v @@ -106,7 +106,7 @@ def check_valid_uuid(v: str) -> str: def is_valid_uuid(v: str) -> bool: try: assert UUID(v).hex == v - except Exception: + except (ValueError, AssertionError): return False return True @@ -165,7 +165,7 @@ CoercedStr = Annotated[str, BeforeValidator(coerce_int_to_str)] # Serializers that can transform a collection of str into a comma separated # str bidirectionally -to_comma_sep_str = PlainSerializer(lambda x: ",".join(sorted(list(x))), return_type=str) +to_comma_sep_str = PlainSerializer(lambda x: ",".join(sorted(x)), return_type=str) enum_to_comma_sep_str = PlainSerializer( lambda x: ",".join(sorted([str(y.value) for y in x])), return_type=str ) diff --git a/generalresearch/models/dynata/survey.py b/generalresearch/models/dynata/survey.py index 5a9f763..942ab4f 100644 --- a/generalresearch/models/dynata/survey.py +++ b/generalresearch/models/dynata/survey.py @@ -168,7 +168,7 @@ class DynataQuota(BaseModel): status: DynataStatus = Field() def __hash__(self): - return hash(tuple((tuple(self.condition_hashes), self.count, self.status))) + return hash((tuple(self.condition_hashes), self.count, self.status)) @property def is_open(self) -> bool: @@ -244,7 +244,7 @@ class DynataQuotaGroup(RootModel): ) -> tuple[bool | None, set[str]]: # Qualify for ANY quota object within a quota group obj_evals = {obj: obj.passes_soft(criteria_evaluation) for obj in self.root} - evals = set(v[0] for v in obj_evals.values()) + evals = {v[0] for v in obj_evals.values()} # If we match 1 obj, then the others don't matter if any(evals): return True, set() @@ -319,7 +319,7 @@ class DynataFilterGroup(RootModel): ) -> tuple[bool | None, set[str]]: # Passes back "passes" (T/F/none) and a list of unknown criterion hashes obj_evals = {obj: obj.passes_soft(criteria_evaluation) for obj in self.root} - evals = set(v[0] for v in obj_evals.values()) + evals = {v[0] for v in obj_evals.values()} # If we match 1 obj, then the others don't matter if any(evals): return True, set() @@ -548,7 +548,7 @@ class DynataSurvey(MarketplaceTask): return d @classmethod - def from_db(cls, d: Dict[str, Any]) -> Self: + def from_db(cls, d: dict[str, Any]) -> Self: d["created"] = d["created"].replace(tzinfo=UTC) d["last_updated"] = d["last_updated"].replace(tzinfo=UTC) d["filters"] = json.loads(d["filters"]) @@ -578,7 +578,7 @@ class DynataSurvey(MarketplaceTask): group_eval = { group: group.passes_soft(criteria_evaluation) for group in self.filters } - evals = set(g[0] for g in group_eval.values()) + evals = {g[0] for g in group_eval.values()} if False in evals: return False, set() elif None in evals: @@ -614,7 +614,7 @@ class DynataSurvey(MarketplaceTask): group_eval = { quota: quota.passes_soft(criteria_evaluation) for quota in self.quotas } - evals = set(g[0] for g in group_eval.values()) + evals = {g[0] for g in group_eval.values()} if False in evals: return False, set() elif None in evals: diff --git a/generalresearch/models/dynata/task_collection.py b/generalresearch/models/dynata/task_collection.py index 71cf3db..2b82bfd 100644 --- a/generalresearch/models/dynata/task_collection.py +++ b/generalresearch/models/dynata/task_collection.py @@ -54,7 +54,7 @@ DynataTaskCollectionSchema = DataFrameSchema( class DynataTaskCollection(TaskCollection): - items: List[DynataSurvey] + items: list[DynataSurvey] _schema = DynataTaskCollectionSchema def to_row(self, s: DynataSurvey) -> dict[str, Any]: diff --git a/generalresearch/models/gr/authentication.py b/generalresearch/models/gr/authentication.py index f9644fe..25f65fa 100644 --- a/generalresearch/models/gr/authentication.py +++ b/generalresearch/models/gr/authentication.py @@ -4,7 +4,7 @@ import binascii import json import os from datetime import UTC, datetime -from typing import TYPE_CHECKING, Any, Self +from typing import TYPE_CHECKING, Any from pydantic import ( AnyHttpUrl, @@ -283,16 +283,15 @@ class GRUser(BaseModel): ex=ex_secs, ) - # --- ORM --- @classmethod - def from_postgresql(cls, d: dict) -> Self: + def from_postgresql(cls, d: dict[str, Any]) -> GRUser: d["date_joined"] = d["date_joined"].replace(tzinfo=UTC) return GRUser.model_validate(d) @classmethod - def from_redis(cls, d: str | dict[str, Any]) -> Self: + def from_redis(cls, d: str | dict[str, Any]) -> GRUser: if isinstance(d, str): d = json.loads(d) assert isinstance(d, dict) @@ -357,13 +356,13 @@ class GRToken(BaseModel): # --- Properties --- @property - def auth_header(self, key_name: str = "Authorization") -> dict[str, str]: - return {key_name: self.key} + def auth_header(self) -> dict[str, str]: + return {"Authorization": self.key} # --- ORM --- @classmethod - def from_redis(cls, d: str | dict[str, Any]) -> Self: + def from_redis(cls, d: str | dict[str, Any]) -> GRToken: if isinstance(d, str): d = json.loads(d) assert isinstance(d, dict) diff --git a/generalresearch/models/gr/business.py b/generalresearch/models/gr/business.py index 74b5c29..064c200 100644 --- a/generalresearch/models/gr/business.py +++ b/generalresearch/models/gr/business.py @@ -6,14 +6,15 @@ import os from datetime import UTC, datetime from enum import Enum, StrEnum from pathlib import Path -from typing import TYPE_CHECKING, Self +from typing import TYPE_CHECKING from uuid import uuid4 import pandas as pd +import pyarrow as pa from dask.distributed import Client from psycopg.cursor import Cursor from psycopg.rows import dict_row -from pydantic import BaseModel, ConfigDict, Field, PositiveInt +from pydantic import BaseModel, ConfigDict, Field, PositiveInt, ValidationError from pydantic.json_schema import SkipJsonSchema from pydantic_extra_types.phone_numbers import PhoneNumber @@ -210,9 +211,11 @@ class Business(BaseModel): payouts: list[BusinessPayoutEvent] | None = Field( default=None, name="Business Payouts", - description="These are the ACH or Wire payments that were sent to the" - "Business as a single amount, summed for all the Business" - "child Products", + description=( + "These are the ACH or Wire payments that were sent to the" + "Business as a single amount, summed for all the Business" + "child Products" + ), ) pop_financial: list[POPFinancial] | None = Field(default=None) @@ -237,18 +240,19 @@ class Business(BaseModel): # --- Prefetch --- def prefetch_addresses(self, pg_config: PostgresConfig) -> None: - with pg_config.make_connection() as conn: - with conn.cursor(row_factory=dict_row) as c: - c.execute( - query=""" + with pg_config.make_connection() as conn, conn.cursor( + row_factory=dict_row + ) as c: + c.execute( + query=""" SELECT * FROM common_businessaddress AS ba WHERE ba.business_id = %s LIMIT 1 """, - params=[self.id], - ) - res = c.fetchall() + params=[self.id], + ) + res = c.fetchall() if len(res) == 0: self.addresses = [] @@ -258,22 +262,23 @@ class Business(BaseModel): def prefetch_teams(self, pg_config: PostgresConfig) -> None: from generalresearch.models.gr.team import Team - with pg_config.make_connection() as conn: - with conn.cursor(row_factory=dict_row) as c: - c: Cursor + with pg_config.make_connection() as conn, conn.cursor( + row_factory=dict_row + ) as c: + c: Cursor - c.execute( - query=""" + c.execute( + query=""" SELECT t.* FROM common_team AS t INNER JOIN common_team_businesses AS tb ON tb.team_id = t.id WHERE tb.business_id = %s """, - params=(self.id,), - ) + params=(self.id,), + ) - res = c.fetchall() + res = c.fetchall() if len(res) == 0: self.teams = [] @@ -542,11 +547,10 @@ class Business(BaseModel): ) try: - test = pd.read_parquet(path, engine="pyarrow") - except Exception as e: + _ = pd.read_parquet(path, engine="pyarrow") + except (pa.ArrowException, OSError, ValueError) as e: raise OSError(f"Parquet verification failed: {e}") - def prebuild_enriched_wall_parquet( self, thl_pg_config: PostgresConfig, @@ -586,11 +590,10 @@ class Business(BaseModel): ) try: - test = pd.read_parquet(path, engine="pyarrow") - except Exception as e: + _ = pd.read_parquet(path, engine="pyarrow") + except (pa.ArrowException, OSError, ValueError) as e: raise OSError(f"Parquet verification failed: {e}") - @classmethod def required_fields(cls) -> list[str]: return [ @@ -651,7 +654,7 @@ class Business(BaseModel): client=client, pop_ledger=pop_ledger, ) - self.prebuild_payouts(thl_pg_config=thl_web_rr, thl_lm=thl_lm, bpem=bpem) + self.prebuild_payouts(bpem=bpem) self.prebuild_pop_financial( thl_pg_config=thl_web_rr, thl_lm=thl_lm, @@ -713,7 +716,7 @@ class Business(BaseModel): uuid: UUIDStr, fields: list[str], gr_redis_config: RedisConfig, - ) -> Self | None: + ) -> Business | None: keys: list[str] = Business.required_fields() + fields if "pop_financial" in keys: @@ -724,7 +727,7 @@ class Business(BaseModel): rc = gr_redis_config.create_redis_client() try: - res: list = rc.hmget(name=f"business:{uuid}", keys=keys) + res: list[str | bytes | None] = rc.hmget(name=f"business:{uuid}", keys=keys) d = { val: json.loads(res[idx]) if res[idx] is not None else None for idx, val in enumerate(keys) @@ -742,6 +745,5 @@ class Business(BaseModel): result["pop_financial"] = pop_financial return Business.model_validate(result) - except Exception as e: - logging.exception(e) + except ValidationError: return None diff --git a/generalresearch/models/gr/team.py b/generalresearch/models/gr/team.py index 78a9ba9..4752bea 100644 --- a/generalresearch/models/gr/team.py +++ b/generalresearch/models/gr/team.py @@ -5,16 +5,18 @@ import os from datetime import UTC, datetime from enum import Enum from pathlib import Path -from typing import TYPE_CHECKING, Self +from typing import TYPE_CHECKING from uuid import uuid4 import pandas as pd +import pyarrow as pa from dask.distributed import Client from pydantic import ( BaseModel, ConfigDict, Field, PositiveInt, + ValidationError, field_validator, ) from pydantic.json_schema import SkipJsonSchema @@ -191,10 +193,9 @@ class Team(BaseModel): try: _ = pd.read_parquet(path, engine="pyarrow") - except Exception as e: + except (pa.ArrowException, OSError, ValueError) as e: raise OSError(f"Parquet verification failed: {e}") - def prebuild_enriched_wall_parquet( self, thl_pg_config: PostgresConfig, @@ -235,10 +236,9 @@ class Team(BaseModel): try: _ = pd.read_parquet(path, engine="pyarrow") - except Exception as e: + except (pa.ArrowException, OSError, ValueError) as e: raise OSError(f"Parquet verification failed: {e}") - @classmethod def required_fields(cls) -> list[str]: return [ @@ -281,8 +281,6 @@ class Team(BaseModel): enriched_session: EnrichedSessionMerge | None = None, enriched_wall: EnrichedWallMerge | None = None, ) -> None: - ex_secs = 60 * 60 * 24 * 3 # 3 days - self.prefetch_products(thl_pg_config=thl_web_rr) self.prefetch_gr_users(pg_config=pg_config, redis_config=redis_config) self.prefetch_businesses(pg_config=pg_config, redis_config=redis_config) @@ -323,7 +321,6 @@ class Team(BaseModel): enriched_wall=enriched_wall, ) - # --- ORM --- @classmethod @@ -332,14 +329,14 @@ class Team(BaseModel): uuid: UUIDStr, fields: list[str], gr_redis_config: RedisConfig, - ) -> Self | None: + ) -> Team | None: keys: list = Team.required_fields() + fields rc = gr_redis_config.create_redis_client() try: - res: list = rc.hmget(name=f"team:{uuid}", keys=keys) + res: list[str | bytes | None] = rc.hmget(name=f"team:{uuid}", keys=keys) d = {val: json.loads(res[idx]) for idx, val in enumerate(keys)} return Team.model_validate(d) - except Exception: + except ValidationError: return None diff --git a/generalresearch/models/innovate/question.py b/generalresearch/models/innovate/question.py index 89310a2..f5a4846 100644 --- a/generalresearch/models/innovate/question.py +++ b/generalresearch/models/innovate/question.py @@ -6,7 +6,7 @@ import logging from enum import StrEnum from typing import TYPE_CHECKING, Any, Literal -from pydantic import BaseModel, Field, field_validator, model_validator +from pydantic import BaseModel, Field, ValidationError, field_validator, model_validator from generalresearch.models import Source from generalresearch.models.innovate import InnovateQuestionID @@ -71,7 +71,7 @@ class InnovateQuestionType(StrEnum): @classmethod def from_api(cls, a: int): API_TYPE_MAP = cls.get_api_map() - return API_TYPE_MAP[a] if a in API_TYPE_MAP else None + return API_TYPE_MAP.get(a) class InnovateQuestion(MarketplaceQuestion): @@ -141,7 +141,7 @@ class InnovateQuestion(MarketplaceQuestion): @classmethod def from_api( - cls, d: dict, country_iso: str, language_iso: str + cls, d: dict[str, Any], country_iso: str, language_iso: str ) -> InnovateQuestion | None: """ :param d: Raw response from API @@ -151,13 +151,13 @@ class InnovateQuestion(MarketplaceQuestion): """ try: return cls._from_api(d, country_iso, language_iso) - except Exception as e: + except ValidationError as e: logger.warning(f"Unable to parse question: {d}. {e}") return None @classmethod def _from_api( - cls, d: dict, country_iso: str, language_iso: str + cls, d: dict[str, Any], country_iso: str, language_iso: str ) -> InnovateQuestion: # Question AGE returns options even though its marked as a text entry (but only in some locales) d["QuestionKey"] = d["QuestionKey"].lower() diff --git a/generalresearch/models/innovate/survey.py b/generalresearch/models/innovate/survey.py index 3c37fe3..d07f960 100644 --- a/generalresearch/models/innovate/survey.py +++ b/generalresearch/models/innovate/survey.py @@ -17,6 +17,7 @@ from pydantic import ( BaseModel, ConfigDict, Field, + ValidationError, computed_field, model_validator, ) @@ -69,7 +70,7 @@ class InnovateCondition(MarketplaceCondition): d["logical_operator"] = LogicalOperator.OR d["value_type"] = ConditionValueType.LIST d["negate"] = False - d["values"] = list(set(x.strip().lower() for x in d["values"])) + d["values"] = list({x.strip().lower() for x in d["values"]}) return cls.model_validate(d) @@ -88,7 +89,7 @@ class InnovateQuota(BaseModel): condition_hashes: list[str] = Field(min_length=0, default_factory=list) def __hash__(self): - return hash(tuple((tuple(self.condition_hashes), self.remaining_count))) + return hash((tuple(self.condition_hashes), self.remaining_count)) @property def is_open(self) -> bool: @@ -99,7 +100,7 @@ class InnovateQuota(BaseModel): ) @classmethod - def from_api(cls, d: dict): + def from_api(cls, d: dict[str, Any]): return cls.model_validate(d) def passes(self, criteria_evaluation: dict[str, bool | None]) -> bool: @@ -263,13 +264,13 @@ class InnovateSurvey(MarketplaceTask): def from_api(cls, d: dict[str, Any]) -> InnovateSurvey | None: try: return cls._from_api(d) - except Exception as e: + except ValidationError as e: logger.warning(f"Unable to parse survey: {d}. {e}") return None @classmethod def _from_api(cls, d: dict[str, Any]) -> InnovateSurvey: - d["conditions"] = dict() + d["conditions"] = {} # If we haven't hit the "detail" endpoint, we won't get this d.setdefault("qualifications", []) @@ -317,11 +318,14 @@ class InnovateSurvey(MarketplaceTask): # Fancy repr that abbreviates exclude_pids and excluded_surveys repr_args = list(self.__repr_args__()) for n, (k, v) in enumerate(repr_args): - if k in {"exclude_pids", "include_pids", "excluded_surveys"}: - if v and len(v) > 6: - v = sorted(v) - v = v[:3] + ["…"] + v[-3:] - repr_args[n] = (k, v) + if ( + k in {"exclude_pids", "include_pids", "excluded_surveys"} + and v + and len(v) > 6 + ): + v = sorted(v) + v = v[:3] + ["…"] + v[-3:] + repr_args[n] = (k, v) join_str = ", " repr_str = join_str.join( repr(v) if a is None else f"{a}={v!r}" for a, v in repr_args @@ -380,14 +384,21 @@ class InnovateSurvey(MarketplaceTask): """ assert isinstance(att_survey_ids, set), "must pass a set" assert isinstance(att_job_ids, set), "must pass a set" + if self.survey_id in att_survey_ids: return False - if self.duplicate_check_level == InnovateDuplicateCheckLevel.JOB: - if self.job_id in att_job_ids: - return False + + if ( + self.duplicate_check_level == InnovateDuplicateCheckLevel.JOB + and self.job_id in att_job_ids + ): + return False + if self.duplicate_check_level == InnovateDuplicateCheckLevel.EXCLUDED_SURVEYS: + assert self.excluded_surveys is not None if self.excluded_surveys.intersection(att_survey_ids): return False + return True def passes_qualifications( @@ -431,7 +442,7 @@ class InnovateSurvey(MarketplaceTask): quota_eval = { quota: quota.matches_soft(criteria_evaluation) for quota in self.quotas } - evals = set(g[0] for g in quota_eval.values()) + evals = {g[0] for g in quota_eval.values()} if any(m[0] is True and not q.is_open for q, m in quota_eval.items()): # matched a full quota return False, set() diff --git a/generalresearch/models/legacy/bucket.py b/generalresearch/models/legacy/bucket.py index 812241d..3705f4f 100644 --- a/generalresearch/models/legacy/bucket.py +++ b/generalresearch/models/legacy/bucket.py @@ -120,8 +120,10 @@ class BucketBase(BaseModel): ) uri: HttpsUrl = Field( examples=[ - "https://task.generalresearch.com/api/v1/52d3f63b2709/797df4136c604a6c8599818296aae6d1/?i" - "=5ba2fe5010cc4d078fc3cc0b0cc264c3&b=test&66482fb=e7baf5e" + ( + "https://task.generalresearch.com/api/v1/52d3f63b2709/797df4136c604a6c8599818296aae6d1/?i" + "=5ba2fe5010cc4d078fc3cc0b0cc264c3&b=test&66482fb=e7baf5e" + ) ], description="The URL to send a respondent into. Must not edit this URL in any way", ) @@ -465,12 +467,12 @@ class PayoutSummaryDecimal(StatisticalSummary): class PayoutSummary(StatisticalSummary): """Payouts are in Integer USD Cents""" - min: int = Field(gt=0, le=10000) - max: int = Field(gt=0, le=10000) - q1: int = Field(gt=0, le=10000) - q2: int = Field(gt=0, le=10000) - q3: int = Field(gt=0, le=10000) - mean: int | None = Field(gt=0, le=10000, default=None) + min: int = Field(gt=0, le=10_000) + max: int = Field(gt=0, le=10_000) + q1: int = Field(gt=0, le=10_000) + q2: int = Field(gt=0, le=10_000) + q3: int = Field(gt=0, le=10_000) + mean: int | None = Field(gt=0, le=10_000, default=None) model_config = { "json_schema_extra": { @@ -724,8 +726,10 @@ class OneShotOfferwallBucket(BaseModel): ) uri: HttpsUrl = Field( examples=[ - "https://task.generalresearch.com/api/v1/52d3f63b2709/797df4136c604a6c8599818296aae6d1/?i" - "=5ba2fe5010cc4d078fc3cc0b0cc264c3&b=test&66482fb=e7baf5e" + ( + "https://task.generalresearch.com/api/v1/52d3f63b2709/797df4136c604a6c8599818296aae6d1/?i" + "=5ba2fe5010cc4d078fc3cc0b0cc264c3&b=test&66482fb=e7baf5e" + ) ], description="The URL to send a respondent into. Must not edit this URL in any way", ) @@ -759,8 +763,10 @@ class WXETOfferwallBucket(BaseModel): ) uri: HttpsUrl = Field( examples=[ - "https://task.generalresearch.com/api/v1/52d3f63b2709/797df4136c604a6c8599818296aae6d1/?i" - "=5ba2fe5010cc4d078fc3cc0b0cc264c3&b=test&66482fb=e7baf5e" + ( + "https://task.generalresearch.com/api/v1/52d3f63b2709/797df4136c604a6c8599818296aae6d1/?i" + "=5ba2fe5010cc4d078fc3cc0b0cc264c3&b=test&66482fb=e7baf5e" + ) ], description="The URL to send a respondent into. Must not edit this URL in any way", ) diff --git a/generalresearch/models/legacy/questions.py b/generalresearch/models/legacy/questions.py index 4651ab0..9f37837 100644 --- a/generalresearch/models/legacy/questions.py +++ b/generalresearch/models/legacy/questions.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Annotated, Any, Self +from typing import TYPE_CHECKING, Annotated, Any from pydantic import ( BaseModel, @@ -87,25 +87,17 @@ class UserQuestionAnswerIn(BaseModel): fingerprint_tz = "a91cb1dea814480dba12d9b7b48696dd" fingerprint_fingerprint = "1d1e2e8380ac474b87fb4e4c569b48df" - if self.question_id in { - user_agent_qid, - fingerprint_langs, - fingerprint_tz, - fingerprint_fingerprint, - }: - if len(self.answer) != 1: - raise ValueError("Too many answer values provided") - - return self - - @model_validator(mode="after") - def user_agent_check(self) -> Self: - # TODO: where / how do I want to pass in this Werz user_agent stuff? - user_agent_qid = "2fbedb2b9f7647b09ff5e52fa119cc5e" - - if self.question_id == user_agent_qid: - val = self.answer[0] - # assert val == request.user_agent.to_header(): + if ( + self.question_id + in { + user_agent_qid, + fingerprint_langs, + fingerprint_tz, + fingerprint_fingerprint, + } + and len(self.answer) != 1 + ): + raise ValueError("Too many answer values provided") return self diff --git a/generalresearch/models/morning/survey.py b/generalresearch/models/morning/survey.py index 5011698..91d1bce 100644 --- a/generalresearch/models/morning/survey.py +++ b/generalresearch/models/morning/survey.py @@ -200,7 +200,7 @@ class MorningQuota(MorningStatistics, MarketplaceTask): data["country_isos"] = [data["country_iso"]] if isinstance(data["language_isos"], str): data["language_isos"] = set(data["language_isos"].split(",")) - data["language_iso"] = sorted(data["language_isos"])[0] + data["language_iso"] = min(data["language_isos"]) return data @property @@ -276,11 +276,11 @@ class MorningQuota(MorningStatistics, MarketplaceTask): self, criteria_evaluation: dict[str, bool | None] ) -> tuple[bool | None, list[str]]: # Passes back "matches" (T/F/none) and a list of unknown criterion hashes - unknowns = list() + unknowns = [] for c in self.condition_hashes: eval_value = criteria_evaluation.get(c) if eval_value is False: - return False, list() + return False, [] if eval_value is None: unknowns.append(c) if unknowns: @@ -359,7 +359,7 @@ class MorningBid(MorningTaskStatistics): @property def language_iso_any(self): - return sorted(self.language_isos)[0] + return min(self.language_isos) @property def locale(self): @@ -417,7 +417,7 @@ class MorningBid(MorningTaskStatistics): if "conditions" in data: return data - data["conditions"] = dict() + data["conditions"] = {} for quota in data["quotas"]: if "qualifications" in quota: quota_conditions = [ diff --git a/generalresearch/models/morning/task_collection.py b/generalresearch/models/morning/task_collection.py index eb4cbd1..9303a2f 100644 --- a/generalresearch/models/morning/task_collection.py +++ b/generalresearch/models/morning/task_collection.py @@ -108,7 +108,7 @@ class MorningTaskCollection(TaskCollection): ] quota_fields = list(quota_columns.keys()) rows = [] - bid_dict = dict() + bid_dict = {} for k in bid_fields: bid_dict[k] = getattr(bid, k) bid_dict["bid.id"] = bid.id diff --git a/generalresearch/models/network/label.py b/generalresearch/models/network/label.py index e4ddd18..60a6e58 100644 --- a/generalresearch/models/network/label.py +++ b/generalresearch/models/network/label.py @@ -2,6 +2,7 @@ from __future__ import annotations import ipaddress from enum import StrEnum +from ipaddress import IPv4Network, IPv6Network from pydantic import ( BaseModel, @@ -84,12 +85,13 @@ class IPLabel(BaseModel): @field_validator("ip", mode="before") @classmethod - def normalize_and_validate_network(cls, v): - net = ipaddress.ip_network(v, strict=False) + def normalize_and_validate_network( + cls, v: IPvAnyNetwork + ) -> IPv4Network | IPv6Network | None: + net = ipaddress.ip_network(address=v, strict=False) - if isinstance(net, ipaddress.IPv6Network): - if net.prefixlen > 64: - raise ValueError("IPv6 network must be /64 or larger") + if isinstance(net, ipaddress.IPv6Network) and net.prefixlen > 64: + raise ValueError("IPv6 network must be /64 or larger") return net diff --git a/generalresearch/models/network/nmap/result.py b/generalresearch/models/network/nmap/result.py index 55c2109..4552e15 100644 --- a/generalresearch/models/network/nmap/result.py +++ b/generalresearch/models/network/nmap/result.py @@ -411,7 +411,7 @@ class NmapResult(BaseModel): def model_dump_postgres(self): # Writes for the network_portscan table - d = dict() + d = {} data = self.model_dump( mode="json", include={ diff --git a/generalresearch/models/network/rdns/command.py b/generalresearch/models/network/rdns/command.py index e88a84d..bccead0 100644 --- a/generalresearch/models/network/rdns/command.py +++ b/generalresearch/models/network/rdns/command.py @@ -20,7 +20,7 @@ def run_rdns(config: RDNSRunCommand) -> RDNSResult: def build_rdns_command(ip: str) -> str: # e.g. dig +noall +answer -x 1.2.3.4 - return " ".join(["dig", "+noall", "+answer", "-x", ip]) + return f"dig +noall +answer -x {ip}" def get_dig_version() -> str: diff --git a/generalresearch/models/precision/question.py b/generalresearch/models/precision/question.py index cc90aa9..f532998 100644 --- a/generalresearch/models/precision/question.py +++ b/generalresearch/models/precision/question.py @@ -54,15 +54,15 @@ class PrecisionQuestionType(StrEnum): TEXT_ENTRY = "t" @classmethod - def from_api(cls, a: int): - API_TYPE_MAP = { + def from_api(cls, a: int) -> PrecisionQuestionType | None: + api_type_map: dict[str, PrecisionQuestionType] = { "Drop Down": PrecisionQuestionType.SINGLE_SELECT, "Multi Select": PrecisionQuestionType.MULTI_SELECT, "Single Select": PrecisionQuestionType.SINGLE_SELECT, "Single Select Matrix": PrecisionQuestionType.SINGLE_SELECT, "Vertical Question": PrecisionQuestionType.SINGLE_SELECT, } - return API_TYPE_MAP[a] if a in API_TYPE_MAP else None + return api_type_map.get(a, None) class PrecisionUserQuestionAnswer(MarketplaceUserQuestionAnswer): diff --git a/generalresearch/models/precision/survey.py b/generalresearch/models/precision/survey.py index b27b8c4..bf9e83e 100644 --- a/generalresearch/models/precision/survey.py +++ b/generalresearch/models/precision/survey.py @@ -312,7 +312,7 @@ class PrecisionSurvey(MarketplaceTask): quota_eval = { quota: quota.matches_soft(criteria_evaluation) for quota in self.quotas } - evals = set(g[0] for g in quota_eval.values()) + evals = {g[0] for g in quota_eval.values()} if any(m[0] is True and not q.is_open for q, m in quota_eval.items()): # matched a full quota return False, set() diff --git a/generalresearch/models/prodege/question.py b/generalresearch/models/prodege/question.py index 574c4fd..58aed67 100644 --- a/generalresearch/models/prodege/question.py +++ b/generalresearch/models/prodege/question.py @@ -8,7 +8,14 @@ from enum import StrEnum from functools import cached_property from typing import TYPE_CHECKING, Any, Literal -from pydantic import BaseModel, ConfigDict, Field, PositiveInt, model_validator +from pydantic import ( + BaseModel, + ConfigDict, + Field, + PositiveInt, + ValidationError, + model_validator, +) from generalresearch.locales import Localelator from generalresearch.models import MAX_INT32, Source @@ -143,7 +150,7 @@ class ProdegeQuestion(MarketplaceQuestion): """ try: return cls._from_api(d, country_iso) - except Exception as e: + except ValidationError as e: logger.warning(f"Unable to parse question: {d}. {e}") return None diff --git a/generalresearch/models/prodege/survey.py b/generalresearch/models/prodege/survey.py index 7ab6df6..5d0369a 100644 --- a/generalresearch/models/prodege/survey.py +++ b/generalresearch/models/prodege/survey.py @@ -539,7 +539,7 @@ class ProdegeSurvey(MarketplaceTask): d["country_isos"] = [ locale_helper.get_country_iso(d.pop("country_code").lower()) ] - d["country_iso"] = sorted(d["country_isos"])[0] + d["country_iso"] = min(d["country_isos"]) # No languages are returned anywhere for anything d["language_isos"] = [ locale_helper.get_default_lang_from_country(d["country_isos"][0]) @@ -552,7 +552,7 @@ class ProdegeSurvey(MarketplaceTask): d["past_participation"] = ProdegePastParticipation.from_api( d["past_participation"] ) - d["conditions"] = dict() + d["conditions"] = {} for quota in d["quotas"]: quota["condition_hashes"] = [] for c in quota["targeting_criteria"]: @@ -563,7 +563,7 @@ class ProdegeSurvey(MarketplaceTask): d["quotas"] = [ProdegeQuota.from_api(q) for q in d["quotas"]] countries = {q.country_iso for q in d["quotas"] if q.country_iso} if countries: - d["country_iso"] = sorted(countries)[0] + d["country_iso"] = min(countries) d["country_isos"] = countries d["language_iso"] = locale_helper.get_default_lang_from_country( d["country_iso"] diff --git a/generalresearch/models/prodege/task_collection.py b/generalresearch/models/prodege/task_collection.py index 4544050..9f6a81b 100644 --- a/generalresearch/models/prodege/task_collection.py +++ b/generalresearch/models/prodege/task_collection.py @@ -76,7 +76,7 @@ class ProdegeTaskCollection(TaskCollection): "used_question_ids", "all_hashes", ] - d = dict() + d = {} for k in fields: d[k] = getattr(s, k) d["cpi"] = float(d["cpi"]) diff --git a/generalresearch/models/repdata/question.py b/generalresearch/models/repdata/question.py index 9dda97f..4fa2d22 100644 --- a/generalresearch/models/repdata/question.py +++ b/generalresearch/models/repdata/question.py @@ -12,6 +12,7 @@ from pydantic import ( ConfigDict, Field, PositiveInt, + ValidationError, field_validator, model_validator, ) @@ -142,6 +143,7 @@ class RepDataQuestion(MarketplaceQuestion): @property def internal_id(self) -> str: + assert self.lucid_id return self.lucid_id @field_validator("question_id", mode="before") @@ -167,7 +169,7 @@ class RepDataQuestion(MarketplaceQuestion): """ try: return cls._from_api(d, country_iso, language_iso) - except Exception as e: + except ValidationError as e: logger.warning(f"Unable to parse question: {d}. {e}") return None diff --git a/generalresearch/models/repdata/survey.py b/generalresearch/models/repdata/survey.py index 43a592c..c5b0730 100644 --- a/generalresearch/models/repdata/survey.py +++ b/generalresearch/models/repdata/survey.py @@ -13,6 +13,7 @@ from pydantic import ( BaseModel, ConfigDict, Field, + ValidationError, computed_field, field_validator, model_validator, @@ -459,7 +460,7 @@ class RepDataSurvey(BaseModel): @property def all_conditions(self) -> list[RepDataCondition]: - cs = list() + cs = [] for stream in self.streams: cs.extend(stream.all_conditions) # dedupe by criterion_hash @@ -477,7 +478,7 @@ class RepDataSurvey(BaseModel): """ try: return cls._from_api(survey_response) - except Exception as e: + except ValidationError as e: survey_id = survey_response.get("survey_id") or survey_response.get( "SurveyNumber" ) @@ -485,7 +486,7 @@ class RepDataSurvey(BaseModel): return None @classmethod - def _from_api(cls, survey_response) -> RepDataSurvey: + def _from_api(cls, survey_response: dict[str, Any]) -> RepDataSurvey: d = survey_response.copy() d["country_iso"] = locale_helper.get_country_iso(d["SurveyCountry"].lower()) d["language_iso"] = locale_helper.get_language_iso(d["SurveyLanguage"].lower()) diff --git a/generalresearch/models/repdata/task_collection.py b/generalresearch/models/repdata/task_collection.py index d625349..5b9a4ba 100644 --- a/generalresearch/models/repdata/task_collection.py +++ b/generalresearch/models/repdata/task_collection.py @@ -110,7 +110,7 @@ class RepDataTaskCollection(TaskCollection): "remaining_count", ] rows = [] - d = dict() + d = {} for k in survey_fields: d[k] = getattr(s, k) d["allowed_devices"] = s.allowed_devices_str diff --git a/generalresearch/models/sago/question.py b/generalresearch/models/sago/question.py index 474543d..291214f 100644 --- a/generalresearch/models/sago/question.py +++ b/generalresearch/models/sago/question.py @@ -13,6 +13,7 @@ from pydantic import ( ConfigDict, Field, PositiveInt, + ValidationError, field_validator, model_validator, ) @@ -86,7 +87,7 @@ class SagoQuestionType(StrEnum): 6: SagoQuestionType.TEXT_ENTRY, 7: SagoQuestionType.TEXT_ENTRY, } - return API_TYPE_MAP[a] if a in API_TYPE_MAP else None + return API_TYPE_MAP.get(a, None) class SagoUserQuestionAnswer(BaseModel): @@ -182,7 +183,7 @@ class SagoQuestion(MarketplaceQuestion): """ try: return cls._from_api(d, country_iso, language_iso) - except Exception as e: + except ValidationError as e: logger.warning(f"Unable to parse question: {d}. {e}") return None diff --git a/generalresearch/models/sago/survey.py b/generalresearch/models/sago/survey.py index 83aad8c..8550cd3 100644 --- a/generalresearch/models/sago/survey.py +++ b/generalresearch/models/sago/survey.py @@ -8,7 +8,14 @@ from functools import cached_property from typing import Annotated, Any, Literal, Self from more_itertools import flatten -from pydantic import BaseModel, ConfigDict, Field, computed_field, model_validator +from pydantic import ( + BaseModel, + ConfigDict, + Field, + ValidationError, + computed_field, + model_validator, +) from generalresearch.locales import Localelator from generalresearch.models import LogicalOperator, Source @@ -71,7 +78,7 @@ class SagoQuota(BaseModel): # There is no explicit status. The quota is closed if the count is 0 def __hash__(self) -> int: - return hash(tuple((tuple(self.condition_hashes), self.remaining_count))) + return hash((tuple(self.condition_hashes), self.remaining_count)) @property def is_open(self) -> bool: @@ -261,7 +268,7 @@ class SagoSurvey(MarketplaceTask): def from_api(cls, d: dict[str, Any]) -> SagoSurvey | None: try: return cls._from_api(d) - except Exception as e: + except ValidationError as e: logger.warning(f"Unable to parse survey: {d}. {e}") return None @@ -273,11 +280,10 @@ class SagoSurvey(MarketplaceTask): # Fancy repr that abbreviates ip_exclusions and survey_exclusions repr_args = list(self.__repr_args__()) for n, (k, v) in enumerate(repr_args): - if k in {"ip_exclusions", "survey_exclusions"}: - if v and len(v) > 6: - v = sorted(v) - v = v[:3] + ["…"] + v[-3:] - repr_args[n] = (k, v) + if k in {"ip_exclusions", "survey_exclusions"} and v and len(v) > 6: + v = sorted(v) + v = v[:3] + ["…"] + v[-3:] + repr_args[n] = (k, v) join_str = ", " repr_str = join_str.join( repr(v) if a is None else f"{a}={v!r}" for a, v in repr_args @@ -362,7 +368,7 @@ class SagoSurvey(MarketplaceTask): quota_eval = { quota: quota.matches_soft(criteria_evaluation) for quota in self.quotas } - evals = set(g[0] for g in quota_eval.values()) + evals = {g[0] for g in quota_eval.values()} if any(m[0] is True and not q.is_open for q, m in quota_eval.items()): # matched a full quota return False, set() diff --git a/generalresearch/models/thl/contest/__init__.py b/generalresearch/models/thl/contest/__init__.py index 842d85e..65b28f4 100644 --- a/generalresearch/models/thl/contest/__init__.py +++ b/generalresearch/models/thl/contest/__init__.py @@ -1,7 +1,7 @@ from __future__ import annotations from datetime import UTC, datetime -from typing import Self +from typing import Any, Self from uuid import uuid4 from pydantic import ( diff --git a/generalresearch/models/thl/contest/contest.py b/generalresearch/models/thl/contest/contest.py index bd0fc04..5814bef 100644 --- a/generalresearch/models/thl/contest/contest.py +++ b/generalresearch/models/thl/contest/contest.py @@ -186,7 +186,7 @@ class Contest(ContestBase): @classmethod def model_validate_mysql(cls, data: dict[str, Any]) -> Self: - data = {k: v for k, v in data.items() if k in cls.model_fields.keys()} + data = {k: v for k, v in data.items() if k in cls.model_fields} if isinstance(data["end_condition"], dict): data["end_condition"] = ContestEndCondition.model_validate( data["end_condition"] diff --git a/generalresearch/models/thl/contest/contest_entry.py b/generalresearch/models/thl/contest/contest_entry.py index bb3aef4..b5f0ac3 100644 --- a/generalresearch/models/thl/contest/contest_entry.py +++ b/generalresearch/models/thl/contest/contest_entry.py @@ -1,6 +1,7 @@ from __future__ import annotations from datetime import UTC, datetime +from typing import Any from uuid import uuid4 from pydantic import ( diff --git a/generalresearch/models/thl/contest/raffle.py b/generalresearch/models/thl/contest/raffle.py index 16a0a47..072f011 100644 --- a/generalresearch/models/thl/contest/raffle.py +++ b/generalresearch/models/thl/contest/raffle.py @@ -203,9 +203,7 @@ class RaffleContest(RaffleContestCreate, Contest): c = self.end_condition if c.target_entry_amount and self.current_amount >= c.target_entry_amount: return True - if c.ends_at and datetime.now(tz=UTC) >= c.ends_at: - return True - return False + return bool(c.ends_at and datetime.now(tz=UTC) >= c.ends_at) def model_dump_mysql(self) -> dict[str, Any]: d = super().model_dump_mysql() @@ -213,7 +211,7 @@ class RaffleContest(RaffleContestCreate, Contest): return d @classmethod - def model_validate_mysql(cls, data: dict) -> Self: + def model_validate_mysql(cls, data: dict[str, Any]) -> Self: data["entry_rule"] = ContestEntryRule.model_validate(data["entry_rule"]) return super().model_validate_mysql(data) diff --git a/generalresearch/models/thl/demographics.py b/generalresearch/models/thl/demographics.py index b6a8be1..c11f8b2 100644 --- a/generalresearch/models/thl/demographics.py +++ b/generalresearch/models/thl/demographics.py @@ -76,7 +76,7 @@ class AgeGroup(Enum): return self.label -def calculate_demographic_metrics(opps: list[MarketplaceTask]) -> list: +def calculate_demographic_metrics(opps: list[MarketplaceTask]) -> list[dict[str, Any]]: """ Measurement: marketplace_survey_demographics tags: source (marketplace) @@ -86,7 +86,7 @@ def calculate_demographic_metrics(opps: list[MarketplaceTask]) -> list: """ source = {opp.source for opp in opps} assert len(source) == 1 - source = list(source)[0] + source = next(iter(source)) survey_cpi = defaultdict(list) target_open = defaultdict(int) for opp in opps: @@ -100,7 +100,7 @@ def calculate_demographic_metrics(opps: list[MarketplaceTask]) -> list: survey_counter = {k: len(v) for k, v in survey_cpi.items()} survey_counter = {k: {"count": v} for k, v in survey_counter.items() if v} - grp_stats = dict() + grp_stats = {} for grp, costs in survey_cpi.items(): stats = { "cost_min": np.min(costs), @@ -155,7 +155,7 @@ def calculate_used_question_metrics( """ source = {opp.source for opp in opps} assert len(source) == 1 - source = list(source)[0] + source = next(iter(source)) country_q_counter = defaultdict(Counter) for opp in opps: for q in opp.used_question_ids: diff --git a/generalresearch/models/thl/finance.py b/generalresearch/models/thl/finance.py index 8c94390..79a74a7 100644 --- a/generalresearch/models/thl/finance.py +++ b/generalresearch/models/thl/finance.py @@ -557,7 +557,7 @@ class BusinessBalances(BaseModel): they all explicitly are set """ - if any([pb.product_id is None for pb in v]): + if any(pb.product_id is None for pb in v): raise ValueError("'product_id' must be set for BusinessBalance children.") return v diff --git a/generalresearch/models/thl/ledger.py b/generalresearch/models/thl/ledger.py index 19dde20..a9fbbb1 100644 --- a/generalresearch/models/thl/ledger.py +++ b/generalresearch/models/thl/ledger.py @@ -1,7 +1,7 @@ from __future__ import annotations from datetime import UTC, datetime -from enum import StrEnum +from enum import IntEnum, StrEnum from typing import Annotated, Any, Literal, Self from uuid import uuid4 @@ -36,7 +36,7 @@ from generalresearch.models.thl.payout_format import ( from generalresearch.utils.enum import ReprEnumMeta -class Direction(int, Enum, metaclass=ReprEnumMeta): +class Direction(IntEnum, metaclass=ReprEnumMeta): """Entries on the debit side will increase debit normal accounts, while entries on the credit side will decrease them. Conversely, entries on the credit side will increase credit normal accounts, while entries on diff --git a/generalresearch/models/thl/offerwall/__init__.py b/generalresearch/models/thl/offerwall/__init__.py index e7c8e03..0c3d51d 100644 --- a/generalresearch/models/thl/offerwall/__init__.py +++ b/generalresearch/models/thl/offerwall/__init__.py @@ -267,7 +267,7 @@ class OfferWallRequest(BaseModel): # We need this so thl-core can refresh an offerwall in order to continue # a session d = self.model_dump(mode="json") - kwargs = dict() + kwargs = {} keys = [ "n_bins", "min_bin_size", diff --git a/generalresearch/models/thl/offerwall/base.py b/generalresearch/models/thl/offerwall/base.py index 33489df..33b9847 100644 --- a/generalresearch/models/thl/offerwall/base.py +++ b/generalresearch/models/thl/offerwall/base.py @@ -398,8 +398,10 @@ class OfferwallBucket(BaseModel): ) uri: HttpsUrl | None = Field( examples=[ - "https://task.generalresearch.com/api/v1/52d3f63b2709/797df4136c604a6c8599818296aae6d1/?i" - "=5ba2fe5010cc4d078fc3cc0b0cc264c3&b=test&66482fb=e7baf5e" + ( + "https://task.generalresearch.com/api/v1/52d3f63b2709/797df4136c604a6c8599818296aae6d1/?i" + "=5ba2fe5010cc4d078fc3cc0b0cc264c3&b=test&66482fb=e7baf5e" + ) ], description="The URL to send a respondent into. Must not edit this URL in any way", default=None, diff --git a/generalresearch/models/thl/payout_format.py b/generalresearch/models/thl/payout_format.py index 4d616b6..d29c9de 100644 --- a/generalresearch/models/thl/payout_format.py +++ b/generalresearch/models/thl/payout_format.py @@ -70,7 +70,7 @@ def format_payout_format(payout_format: str, payout_int: int) -> str: except TypeError: # "{payout()*1:}" - TypeError: 'int' object is not callable raise ValueError("Invalid type reference.") - except Exception: + except Exception: # noqa raise ValueError("Invalid payout transformation") formatstr = f"{{:{formatstr}}}" diff --git a/generalresearch/models/thl/product.py b/generalresearch/models/thl/product.py index a7ecd55..65ed177 100644 --- a/generalresearch/models/thl/product.py +++ b/generalresearch/models/thl/product.py @@ -957,7 +957,7 @@ class Product(BaseModel, validate_assignment=True): @field_validator("harmonizer_domain", mode="before") def harmonizer_domain_https(cls, s: str | None): # in the db, this has no scheme. accept both with a default of https:// - if s is not None and not (s.startswith("https://") or s.startswith("http://")): + if s is not None and not (s.startswith(("https://", "http://"))): s = f"https://{s}" return s @@ -1371,7 +1371,7 @@ class Product(BaseModel, validate_assignment=True): if self.payout_config.payout_transformation is None: return None payout_xform_func = self.get_payout_transformation_func() - kwargs = dict() + kwargs = {} if "user_wallet_balance" in inspect.signature(payout_xform_func).parameters: kwargs["user_wallet_balance"] = user_wallet_balance user_payout: Decimal = payout_xform_func(bp_payout, **kwargs) diff --git a/generalresearch/models/thl/profiling/marketplace.py b/generalresearch/models/thl/profiling/marketplace.py index ad4ce80..0129e38 100644 --- a/generalresearch/models/thl/profiling/marketplace.py +++ b/generalresearch/models/thl/profiling/marketplace.py @@ -82,10 +82,9 @@ class MarketplaceQuestion(BaseModel, ABC): # question has more than 6. repr_args = list(self.__repr_args__()) for n, (k, v) in enumerate(repr_args): - if k == "options": - if v and len(v) > 6: - v = v[:3] + ["..."] + v[-3:] - repr_args[n] = ("options", v) + if k == "options" and v and len(v) > 6: + v = v[:3] + ["..."] + v[-3:] + repr_args[n] = ("options", v) join_str = ", " repr_str = join_str.join( repr(v) if a is None else f"{a}={v!r}" for a, v in repr_args diff --git a/generalresearch/models/thl/report_task.py b/generalresearch/models/thl/report_task.py index d29599d..299ba90 100644 --- a/generalresearch/models/thl/report_task.py +++ b/generalresearch/models/thl/report_task.py @@ -28,7 +28,7 @@ def prioritize_report_values( return None report_values = list(set(report_values)) random.shuffle(report_values) - return sorted(report_values, key=lambda x: REPORT_PRIORITY[x])[-1] + return max(report_values, key=lambda x: REPORT_PRIORITY[x]) class ReportTask(BaseModel): diff --git a/generalresearch/models/thl/session.py b/generalresearch/models/thl/session.py index fe7194a..e4e264f 100644 --- a/generalresearch/models/thl/session.py +++ b/generalresearch/models/thl/session.py @@ -234,10 +234,7 @@ class WallBase(BaseModel): return self.is_visible() and self.status == Status.COMPLETE def allow_session(self) -> bool: - if self.status == Status.COMPLETE: - return False - - return True + return self.status != Status.COMPLETE def update(self, **kwargs) -> None: """ @@ -969,10 +966,7 @@ class Session(BaseModel): return True # Hard limit of 40 wall events per session - if len(self.wall_events) >= 40: - return True - - return False + return len(self.wall_events) >= 40 def determine_payments( self, @@ -985,6 +979,7 @@ class Session(BaseModel): ) product = self.user.product + assert product # Handle brokerage product payouts bp_pay: Decimal = product.determine_bp_payment(thl_net) commission_amount: Decimal = thl_net - bp_pay diff --git a/generalresearch/models/thl/soft_pair.py b/generalresearch/models/thl/soft_pair.py index f3b2b6f..7c2f36e 100644 --- a/generalresearch/models/thl/soft_pair.py +++ b/generalresearch/models/thl/soft_pair.py @@ -50,7 +50,7 @@ class SoftPairResult: return ( self.survey_id + ":" - + ";".join(sorted(set([c.question_id for c in self.conditions]))) + + ";".join(sorted({c.question_id for c in self.conditions})) ) else: return None diff --git a/generalresearch/models/thl/survey/penalty.py b/generalresearch/models/thl/survey/penalty.py index 04f8e20..755d25c 100644 --- a/generalresearch/models/thl/survey/penalty.py +++ b/generalresearch/models/thl/survey/penalty.py @@ -56,7 +56,7 @@ class TeamSurveyPenalty(SurveyPenalty): Penalty = Annotated[ - Union[BPSurveyPenalty, TeamSurveyPenalty], + BPSurveyPenalty | TeamSurveyPenalty, Field(discriminator="kind"), ] PenaltyListAdapter = TypeAdapter(list[Penalty]) diff --git a/generalresearch/models/thl/survey/task_collection.py b/generalresearch/models/thl/survey/task_collection.py index b80166f..d8db0d1 100644 --- a/generalresearch/models/thl/survey/task_collection.py +++ b/generalresearch/models/thl/survey/task_collection.py @@ -38,7 +38,8 @@ class TaskCollection(BaseModel): except pa.errors.SchemaErrors as exc: idx = exc.failure_cases["index"] if len(idx) >= len(df) * 0.10: - raise exc + raise + logger.info(f"{self.__repr_name__()}:handle_df:{json.dumps(exc.message)}") df.drop(index=list(idx), inplace=True) # we need to redo the validation after removing failing rows! diff --git a/generalresearch/models/thl/task_status.py b/generalresearch/models/thl/task_status.py index de767d6..7719b18 100644 --- a/generalresearch/models/thl/task_status.py +++ b/generalresearch/models/thl/task_status.py @@ -224,11 +224,12 @@ class TaskStatusResponse(BaseModel): return v or 0 @field_validator("kwargs", mode="after") - def sanitize_kwargs(cls, v: dict | None) -> dict | None: + def sanitize_kwargs(cls, v: dict[str, Any] | None) -> dict[str, Any] | None: if v and "clicked_timestamp" in v: try: - clicked_timestamp = datetime.strptime( - v["clicked_timestamp"], "%Y-%m-%d %H:%M:%S.%f" + clicked_timestamp = datetime.strptime( # noqa + date_string=v["clicked_timestamp"], + format="%Y-%m-%d %H:%M:%S.%f", ) v["clicked_timestamp"] = ( clicked_timestamp.isoformat(timespec="microseconds") + "Z" @@ -238,7 +239,7 @@ class TaskStatusResponse(BaseModel): return v @model_validator(mode="before") - def transform_user_payout(cls, d): + def transform_user_payout(cls, d: dict[str, Any]): # If the user_payout is None and there is a payout_format, make the user_payout 0 if d.get("user_payout") is None and d.get("payout_format"): d["user_payout"] = 0 diff --git a/generalresearch/pg_helper.py b/generalresearch/pg_helper.py index 1d5d30b..a397247 100644 --- a/generalresearch/pg_helper.py +++ b/generalresearch/pg_helper.py @@ -108,10 +108,8 @@ class PostgresConfig: def execute_write(self, query, params=None) -> int: cmd = query.lstrip().upper() - assert ( - cmd.startswith("INSERT") - or cmd.startswith("UPDATE") - or cmd.startswith("DELETE") + assert cmd.startswith( + ("INSERT", "UPDATE", "DELETE") ), "Supports INSERT/UPDATE only" with self.make_connection() as conn: diff --git a/generalresearch/sql_helper.py b/generalresearch/sql_helper.py index 08b660d..ae2b8d8 100644 --- a/generalresearch/sql_helper.py +++ b/generalresearch/sql_helper.py @@ -315,7 +315,7 @@ class SqlHelper(SqlConnector): field_names = ["`" + x + "`" for x in field_names] field_name_str = ",".join(field_names) if filter_d: - lookup_vals = " AND ".join([f"`{fn}`=%({fn})s" for fn in filter_d.keys()]) + lookup_vals = " AND ".join([f"`{fn}`=%({fn})s" for fn in filter_d]) lookup_str = f" WHERE {lookup_vals}" else: lookup_str = "" diff --git a/generalresearch/utils/grpc_logger.py b/generalresearch/utils/grpc_logger.py index 59f7471..8f2f454 100644 --- a/generalresearch/utils/grpc_logger.py +++ b/generalresearch/utils/grpc_logger.py @@ -33,9 +33,11 @@ try: response = handler_func(request, context) code = context.code() or grpc.StatusCode.OK return response - except Exception as e: + + except Exception: code = context.code() or grpc.StatusCode.INTERNAL - raise e + raise + finally: duration_ms = int((time.time() - start_time) * 1000) peer = context.peer() or "unknown" diff --git a/generalresearch/wall_status_codes/fullcircle.py b/generalresearch/wall_status_codes/fullcircle.py index aeaa4c7..cd9fdff 100644 --- a/generalresearch/wall_status_codes/fullcircle.py +++ b/generalresearch/wall_status_codes/fullcircle.py @@ -29,7 +29,7 @@ status_codes_ext_map: dict[StatusCode1, list[str]] = { StatusCode1.PS_FAIL: [], StatusCode1.PS_OVERQUOTA: [], } -ext_status_code_map: dict[str, StatusCode1] = dict() +ext_status_code_map: dict[str, StatusCode1] = {} for k, v in status_codes_ext_map.items(): k: StatusCode1 v: list[str] diff --git a/generalresearch/wall_status_codes/innovate.py b/generalresearch/wall_status_codes/innovate.py index 936ee6c..e3d2468 100644 --- a/generalresearch/wall_status_codes/innovate.py +++ b/generalresearch/wall_status_codes/innovate.py @@ -38,7 +38,7 @@ status_codes_ext_map: dict[StatusCode1, list[str]] = { StatusCode1.PS_FAIL: ["5"], StatusCode1.PS_OVERQUOTA: ["7"], } -ext_status_code_map = dict() +ext_status_code_map = {} for k, v in status_codes_ext_map.items(): for vv in v: ext_status_code_map[status_codes_ext_map.get(vv, vv)] = k diff --git a/generalresearch/wall_status_codes/lucid.py b/generalresearch/wall_status_codes/lucid.py index 3cc1b5e..c4c5e90 100644 --- a/generalresearch/wall_status_codes/lucid.py +++ b/generalresearch/wall_status_codes/lucid.py @@ -102,7 +102,7 @@ status_codes_ext_map: dict[StatusCode1, list[str]] = { StatusCode1.PS_OVERQUOTA: ["40", "41", "42"], } -ext_status_code_map: dict[str, StatusCode1] = dict() +ext_status_code_map: dict[str, StatusCode1] = {} for k, v in status_codes_ext_map.items(): k: StatusCode1 v: list[str] diff --git a/generalresearch/wall_status_codes/morning.py b/generalresearch/wall_status_codes/morning.py index ffa6be2..6f63b82 100644 --- a/generalresearch/wall_status_codes/morning.py +++ b/generalresearch/wall_status_codes/morning.py @@ -97,7 +97,7 @@ status_codes_ext_map: dict[StatusCode1, list[str]] = { "quota_invalid_for_bid", ], } -ext_status_code_map: dict[str, StatusCode1] = dict() +ext_status_code_map: dict[str, StatusCode1] = {} for k, v in status_codes_ext_map.items(): k: StatusCode1 v: list[str] diff --git a/generalresearch/wall_status_codes/pollfish.py b/generalresearch/wall_status_codes/pollfish.py index a5c6e25..e1ad12a 100644 --- a/generalresearch/wall_status_codes/pollfish.py +++ b/generalresearch/wall_status_codes/pollfish.py @@ -58,7 +58,7 @@ status_codes_ext_map: dict[StatusCode1, list[str]] = { ], StatusCode1.PS_OVERQUOTA: ["quota_full", "survey_closed", "survey_expired"], } -ext_status_code_map = dict() +ext_status_code_map = {} for k, v in status_codes_ext_map.items(): k: StatusCode1 v: list[str] diff --git a/test_utils/managers/contest/conftest.py b/test_utils/managers/contest/conftest.py index 67935e7..a9375f6 100644 --- a/test_utils/managers/contest/conftest.py +++ b/test_utils/managers/contest/conftest.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import pytest from generalresearch.managers.base import Permission @@ -11,8 +13,6 @@ def contest_manager(thl_web_rw: PostgresConfig) -> ContestManager: assert thl_web_rw.dsn.path assert "/unittest-" in thl_web_rw.dsn.path - from generalresearch.managers.thl.contest_manager import ContestManager - return ContestManager( pg_config=thl_web_rw, permissions=[ diff --git a/test_utils/models/contest/conftest.py b/test_utils/models/contest/conftest.py index d9c8a6b..84930b8 100644 --- a/test_utils/models/contest/conftest.py +++ b/test_utils/models/contest/conftest.py @@ -38,6 +38,26 @@ from generalresearch.models.thl.user import User # === Managers === +# --- Factories --- + + +@pytest.fixture(scope="function") +def raffle_contest_factory( + product_user_wallet_yes: Product, + raffle_contest_create: RaffleContestCreate, + contest_manager: ContestManager, +) -> Callable[..., RaffleContest]: + + def _inner(**kwargs): + raffle_contest_create.update(**kwargs) + return contest_manager.create( + product_id=product_user_wallet_yes.uuid, + contest_create=raffle_contest_create, + ) + + return _inner + + # === Models === @@ -82,23 +102,6 @@ def raffle_contest( ) -@pytest.fixture(scope="function") -def raffle_contest_factory( - product_user_wallet_yes: Product, - raffle_contest_create: RaffleContestCreate, - contest_manager: ContestManager, -) -> Callable[..., RaffleContest]: - - def _inner(**kwargs): - raffle_contest_create.update(**kwargs) - return contest_manager.create( - product_id=product_user_wallet_yes.uuid, - contest_create=raffle_contest_create, - ) - - return _inner - - @pytest.fixture def milestone_contest_create() -> MilestoneContestCreate: from generalresearch.models.thl.contest import ( diff --git a/test_utils/spectrum/conftest.py b/test_utils/spectrum/conftest.py index 7cd9321..a8ce9d9 100644 --- a/test_utils/spectrum/conftest.py +++ b/test_utils/spectrum/conftest.py @@ -1,32 +1,32 @@ from __future__ import annotations -import logging import time from datetime import UTC, datetime from decimal import Decimal -from typing import TYPE_CHECKING, Any +from typing import Any import pytest +from generalresearch.config import GRLBaseSettings from generalresearch.managers.spectrum.survey import ( SpectrumCriteriaManager, SpectrumSurveyManager, ) -from generalresearch.models.spectrum.survey import SpectrumSurvey +from generalresearch.models import ( + LogicalOperator, +) +from generalresearch.models.spectrum.survey import ( + SpectrumCondition, + SpectrumSurvey, +) +from generalresearch.models.thl.survey.condition import ConditionValueType from generalresearch.sql_helper import SqlHelper -from .surveys_json import CONDITIONS, SURVEYS_JSON - -if TYPE_CHECKING: - from generalresearch.config import GRLBaseSettings - @pytest.fixture(scope="session") def spectrum_rw(settings: GRLBaseSettings) -> SqlHelper: - logging.info(f"{settings.spectrum_rw_db=}") - assert settings.spectrum_rw_db is not None - assert "/unittest-" in settings.spectrum_rw_db.path + assert "/unittest-" in str(settings.spectrum_rw_db.path) return SqlHelper( dsn=settings.spectrum_rw_db, @@ -38,27 +38,36 @@ def spectrum_rw(settings: GRLBaseSettings) -> SqlHelper: @pytest.fixture(scope="session") def spectrum_criteria_manager(spectrum_rw: SqlHelper) -> SpectrumCriteriaManager: + assert spectrum_rw.dsn + assert spectrum_rw.dsn.path assert "/unittest-" in spectrum_rw.dsn.path return SpectrumCriteriaManager(spectrum_rw) @pytest.fixture(scope="session") def spectrum_survey_manager(spectrum_rw: SqlHelper) -> SpectrumSurveyManager: + assert spectrum_rw.dsn + assert spectrum_rw.dsn.path assert "/unittest-" in spectrum_rw.dsn.path return SpectrumSurveyManager(spectrum_rw) @pytest.fixture(scope="session") def setup_spectrum_surveys( - spectrum_rw: SqlHelper, spectrum_survey_manager, spectrum_criteria_manager + spectrum_rw: SqlHelper, + spectrum_survey_manager: SpectrumSurveyManager, + spectrum_criteria_manager: SpectrumCriteriaManager, + spectrum_conditions: list[SpectrumCondition], + spectrum_api_surveys_json: list[str], ) -> None: now = datetime.now(UTC) # make sure these example surveys exist in db - surveys = [SpectrumSurvey.model_validate_json(x) for x in SURVEYS_JSON] + surveys = [SpectrumSurvey.model_validate_json(x) for x in spectrum_api_surveys_json] for s in surveys: s.modified_api = datetime.now(tz=UTC) + spectrum_survey_manager.create_or_update(surveys) - spectrum_criteria_manager.update(CONDITIONS) + spectrum_criteria_manager.update(spectrum_conditions) # and make sure they have allocation for 687 spectrum_rw.execute_sql_query( @@ -198,6 +207,46 @@ def spectrum_api_surveys_json() -> list[str]: ] +def spectrum_conditions() -> list[SpectrumCondition]: + # make sure hashes for 111111 are in db + c1 = SpectrumCondition( + question_id="1001", + value_type=ConditionValueType.LIST, + values=["a", "b", "c"], + negate=False, + logical_operator=LogicalOperator.OR, + ) + c2 = SpectrumCondition( + question_id="1001", + value_type=ConditionValueType.LIST, + values=["a"], + negate=False, + logical_operator=LogicalOperator.OR, + ) + c3 = SpectrumCondition( + question_id="1002", + value_type=ConditionValueType.RANGE, + values=["18-24", "30-32"], + negate=False, + logical_operator=LogicalOperator.OR, + ) + c4 = SpectrumCondition( + question_id="212", + value_type=ConditionValueType.LIST, + values=["23", "24"], + negate=False, + logical_operator=LogicalOperator.OR, + ) + c5 = SpectrumCondition( + question_id="1031", + value_type=ConditionValueType.LIST, + values=["113", "114", "121"], + negate=False, + logical_operator=LogicalOperator.OR, + ) + return [c1, c2, c3, c4, c5] + + @pytest.fixture(scope="session") def spectrum_api_survey_json() -> dict[str, Any]: return { diff --git a/tests/incite/collections/test_df_collection_base.py b/tests/incite/collections/test_df_collection_base.py index b9f0181..c236700 100644 --- a/tests/incite/collections/test_df_collection_base.py +++ b/tests/incite/collections/test_df_collection_base.py @@ -9,6 +9,7 @@ from generalresearch.incite.collections import ( DFCollection, DFCollectionType, ) +from generalresearch.pg_helper import PostgresConfig if TYPE_CHECKING: from generalresearch.incite.base import GRLDatasets @@ -45,7 +46,9 @@ class TestDFCollectionBase: class TestDFCollectionBaseProperties: @pytest.mark.skip - def test_df_collection_items(self, mnt_filepath: GRLDatasets, df_coll_type): + def test_df_collection_items( + self, mnt_filepath: GRLDatasets, df_coll_type: DFCollectionType + ): instance = DFCollection( data_type=df_coll_type, start=datetime(year=1800, month=1, day=1, tzinfo=UTC), @@ -57,7 +60,9 @@ class TestDFCollectionBaseProperties: assert len(instance.interval_range) == len(instance.items) assert len(instance.items) == 366 - def test_df_collection_progress(self, mnt_filepath: GRLDatasets, df_coll_type): + def test_df_collection_progress( + self, mnt_filepath: GRLDatasets, df_coll_type: DFCollectionType + ): instance = DFCollection( data_type=df_coll_type, start=datetime(year=1800, month=1, day=1, tzinfo=UTC), @@ -70,7 +75,9 @@ class TestDFCollectionBaseProperties: assert isinstance(instance.progress, pd.DataFrame) assert instance.progress.shape == (366, 6) - def test_df_collection_schema(self, mnt_filepath: GRLDatasets, df_coll_type): + def test_df_collection_schema( + self, mnt_filepath: GRLDatasets, df_coll_type: DFCollectionType + ): instance1 = DFCollection( data_type=DFCollectionType.WALL, archive_path=mnt_filepath.data_src ) @@ -87,9 +94,9 @@ class TestDFCollectionBaseProperties: class TestDFCollectionBaseMethods: @pytest.mark.skip - def test_initial_load(self, mnt_filepath: GRLDatasets, thl_web_rr): + def test_initial_load(self, mnt_filepath: GRLDatasets, thl_web_rr: PostgresConfig): instance = DFCollection( - pg_config=thl_web_rr: PostgresConfig, + pg_config=thl_web_rr, data_type=DFCollectionType.USER, start=datetime(year=2022, month=1, day=1, minute=0, tzinfo=UTC), finished=datetime(year=2022, month=1, day=1, minute=5, tzinfo=UTC), diff --git a/tests/incite/collections/test_df_collection_item_base.py b/tests/incite/collections/test_df_collection_item_base.py index 9a2ecf3..e0171c2 100644 --- a/tests/incite/collections/test_df_collection_item_base.py +++ b/tests/incite/collections/test_df_collection_item_base.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from datetime import UTC, datetime from typing import TYPE_CHECKING @@ -19,7 +21,7 @@ df_collection_types = [e for e in DFCollectionType if e is not DFCollectionType. @pytest.mark.parametrize("df_coll_type", df_collection_types) class TestDFCollectionItemBase: - def test_init(self, mnt_filepath: GRLDatasets, df_coll_type): + def test_init(self, mnt_filepath: GRLDatasets, df_coll_type: DFCollectionType): collection = DFCollection( data_type=df_coll_type, offset="100d", @@ -38,14 +40,16 @@ class TestDFCollectionItemBase: class TestDFCollectionItemProperties: @pytest.mark.skip - def test_filename(self, df_coll_type): + def test_filename(self, df_coll_type: DFCollectionType): pass @pytest.mark.parametrize("df_coll_type", df_collection_types) class TestDFCollectionItemMethods: - def test_has_mysql_false(self, mnt_filepath: GRLDatasets, df_coll_type): + def test_has_mysql_false( + self, mnt_filepath: GRLDatasets, df_coll_type: DFCollectionType + ): collection = DFCollection( data_type=df_coll_type, offset="100d", @@ -58,7 +62,10 @@ class TestDFCollectionItemMethods: assert not instance1.has_mysql() def test_has_mysql_true( - self, thl_web_rr: PostgresConfig, mnt_filepath: GRLDatasets, df_coll_type + self, + thl_web_rr: PostgresConfig, + mnt_filepath: GRLDatasets, + df_coll_type: DFCollectionType, ): collection = DFCollection( data_type=df_coll_type, @@ -66,7 +73,7 @@ class TestDFCollectionItemMethods: start=datetime(year=1800, month=1, day=1, tzinfo=UTC), finished=datetime(year=1900, month=1, day=1, tzinfo=UTC), archive_path=mnt_filepath.archive_path(enum_type=df_coll_type), - pg_config=thl_web_rr: PostgresConfig, + pg_config=thl_web_rr, ) # Has RR, assume unittest server is online @@ -74,5 +81,5 @@ class TestDFCollectionItemMethods: assert instance2.has_mysql() @pytest.mark.skip - def test_update_partial_archive(self, df_coll_type): + def test_update_partial_archive(self, df_coll_type: DFCollectionType): pass diff --git a/tests/incite/collections/test_df_collection_thl_marketplaces.py b/tests/incite/collections/test_df_collection_thl_marketplaces.py index d2d3ce4..b4b5b00 100644 --- a/tests/incite/collections/test_df_collection_thl_marketplaces.py +++ b/tests/incite/collections/test_df_collection_thl_marketplaces.py @@ -4,6 +4,7 @@ from itertools import product import pytest from pandera.pandas import Column, DataFrameSchema, Index +from generalresearch.incite.base import GRLDatasets from generalresearch.incite.collections import DFCollection, DFCollectionType from generalresearch.incite.collections.thl_marketplaces import ( InnovateSurveyHistoryCollection, @@ -11,6 +12,7 @@ from generalresearch.incite.collections.thl_marketplaces import ( SagoSurveyHistoryCollection, SpectrumSurveyTimeseriesCollection, ) +from generalresearch.pg_helper import PostgresConfig def combo_object(): @@ -29,7 +31,13 @@ def combo_object(): @pytest.mark.parametrize("df_coll, offset", combo_object()) class TestDFCollection_thl_marketplaces: - def test_init(self, mnt_filepath, df_coll, offset, spectrum_rw): + def test_init( + self, + mnt_filepath: GRLDatasets, + df_coll: DFCollection, + offset: str, + spectrum_rw: PostgresConfig, + ): assert issubclass(df_coll, DFCollection) # This is stupid, but we need to pull the default from the @@ -38,7 +46,7 @@ class TestDFCollection_thl_marketplaces: assert isinstance(data_type, DFCollectionType) # (1) Can't be totally empty, needs a path... - with pytest.raises(expected_exception=Exception) as cm: + with pytest.raises(expected_exception=Exception): instance = df_coll() # (2) Confirm it only needs the archive_path @@ -61,7 +69,7 @@ class TestDFCollection_thl_marketplaces: assert isinstance(instance._schema, DataFrameSchema) assert isinstance(instance._schema.index, Index) - for c in instance._schema.columns.keys(): + for c in instance._schema.columns: assert isinstance(c, str) col = instance._schema.columns[c] assert isinstance(col, Column) diff --git a/tests/incite/collections/test_df_collection_thl_web.py b/tests/incite/collections/test_df_collection_thl_web.py index bcdeb83..6d509bc 100644 --- a/tests/incite/collections/test_df_collection_thl_web.py +++ b/tests/incite/collections/test_df_collection_thl_web.py @@ -3,19 +3,16 @@ from __future__ import annotations from collections.abc import Generator from datetime import datetime from itertools import product -from typing import TYPE_CHECKING import dask.dataframe as dd import pandas as pd import pytest from pandera.pandas import DataFrameSchema -from generalresearch.incite.collections import DFCollection, DFCollectionType - -if TYPE_CHECKING: - from generalresearch.incite.collections import ( - DFCollectionType, - ) +from generalresearch.incite.collections import ( + DFCollection, + DFCollectionType, +) def combo_object() -> Generator[tuple]: @@ -39,7 +36,10 @@ def combo_object() -> Generator[tuple]: class TestDFCollection_thl_web: def test_init( - self, df_collection_data_type: DFCollectionType, offset: str, df_collection + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, ): assert isinstance(df_collection_data_type, DFCollectionType) assert isinstance(df_collection, DFCollection) @@ -50,12 +50,12 @@ class TestDFCollection_thl_web: ) class TestDFCollection_thl_web_Properties: - def test_items(self, df_collection): + def test_items(self, df_collection: DFCollection): assert isinstance(df_collection.items, list) for i in df_collection.items: assert i._collection == df_collection - def test__schema(self, df_collection): + def test__schema(self, df_collection: DFCollection): assert isinstance(df_collection._schema, DataFrameSchema) @@ -65,16 +65,16 @@ class TestDFCollection_thl_web_Properties: class TestDFCollection_thl_web_BaseProperties: @pytest.mark.skip - def test__interval_range(self, df_collection): + def test__interval_range(self, df_collection: DFCollection): pass - def test_interval_start(self, df_collection): + def test_interval_start(self, df_collection: DFCollection): assert isinstance(df_collection.interval_start, datetime) - def test_interval_range(self, df_collection): + def test_interval_range(self, df_collection: DFCollection): assert isinstance(df_collection.interval_range, list) - def test_progress(self, df_collection): + def test_progress(self, df_collection: DFCollection): assert isinstance(df_collection.progress, pd.DataFrame) @@ -84,17 +84,21 @@ class TestDFCollection_thl_web_BaseProperties: class TestDFCollection_thl_web_Methods: @pytest.mark.skip - def test_initial_loads(self, df_collection_data_type, df_collection, offset): + def test_initial_loads( + self, df_collection_data_type, df_collection: DFCollection, offset: str + ): pass @pytest.mark.skip def test_fetch_force_rr_latest( - self, df_collection_data_type, df_collection, offset: str + self, df_collection_data_type, df_collection: DFCollection, offset: str ): pass @pytest.mark.skip - def test_force_rr_latest(self, df_collection_data_type, df_collection, offset): + def test_force_rr_latest( + self, df_collection_data_type, df_collection: DFCollection, offset: str + ): pass @@ -103,63 +107,108 @@ class TestDFCollection_thl_web_Methods: ) class TestDFCollection_thl_web_BaseMethods: - def test_fetch_all_paths(self, df_collection_data_type, offset: str, df_collection): + def test_fetch_all_paths( + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, + ): res = df_collection.fetch_all_paths( items=None, force_rr_latest=False, include_partial=False ) assert isinstance(res, list) @pytest.mark.skip - def test_ddf(self, df_collection_data_type, offset: str, df_collection): + def test_ddf( + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, + ): res = df_collection.ddf() assert isinstance(res, dd.DataFrame) # -- cleanup -- @pytest.mark.skip def test_schedule_cleanup( - self, df_collection_data_type, offset: str, df_collection + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, ): pass @pytest.mark.skip - def test_cleanup(self, df_collection_data_type, offset: str, df_collection): + def test_cleanup( + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, + ): pass @pytest.mark.skip def test_cleanup_partials( - self, df_collection_data_type, offset: str, df_collection + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, ): pass @pytest.mark.skip def test_clear_tmp_archives( - self, df_collection_data_type, offset: str, df_collection + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, ): pass @pytest.mark.skip def test_clear_corrupt_archives( - self, df_collection_data_type, offset: str, df_collection + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, ): pass @pytest.mark.skip def test_rebuild_symlinks( - self, df_collection_data_type, offset: str, df_collection + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, ): pass # -- Source timing -- @pytest.mark.skip - def test_get_item(self, df_collection_data_type, offset: str, df_collection): + def test_get_item( + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, + ): pass @pytest.mark.skip - def test_get_item_start(self, df_collection_data_type, offset: str, df_collection): + def test_get_item_start( + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, + ): pass @pytest.mark.skip - def test_get_items(self, df_collection_data_type, offset: str, df_collection): + def test_get_items( + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, + ): # If we get all the items from the start of the collection, it # should include all the items! res1 = df_collection.items @@ -168,18 +217,27 @@ class TestDFCollection_thl_web_BaseMethods: @pytest.mark.skip def test_get_items_from_year( - self, df_collection_data_type, offset: str, df_collection + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, ): pass @pytest.mark.skip def test_get_items_last90( - self, df_collection_data_type, offset: str, df_collection + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, ): pass @pytest.mark.skip def test_get_items_last365( - self, df_collection_data_type, offset: str, df_collection + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, ): pass diff --git a/tests/incite/mergers/foundations/test_user_id_product.py b/tests/incite/mergers/foundations/test_user_id_product.py index 10802e5..7367056 100644 --- a/tests/incite/mergers/foundations/test_user_id_product.py +++ b/tests/incite/mergers/foundations/test_user_id_product.py @@ -1,11 +1,15 @@ +from __future__ import annotations + from datetime import UTC, datetime, timedelta from itertools import product import pandas as pd import pytest +from dask.distributed import Client as DaskClient # noinspection PyUnresolvedReferences from generalresearch.incite.mergers.foundations.user_id_product import ( + UserIdProductMerge, UserIdProductMergeItem, ) @@ -23,14 +27,21 @@ from generalresearch.incite.mergers.foundations.user_id_product import ( class TestUserIDProduct: @pytest.mark.skip - def test_base(self, client_no_amm, user_id_product_merge): + def test_base( + self, client_no_amm: DaskClient, user_id_product_merge: UserIdProductMerge + ): ddf = user_id_product_merge.ddf() df = client_no_amm.compute(collections=ddf, sync=True) assert isinstance(df, pd.DataFrame) assert not df.empty @pytest.mark.skip - def test_base_item(self, client_no_amm, user_id_product_merge, user_collection): + def test_base_item( + self, + client_no_amm: DaskClient, + user_id_product_merge: UserIdProductMerge, + user_collection, + ): assert len(user_id_product_merge.items) == 1 for item in user_id_product_merge.items: @@ -40,7 +51,7 @@ class TestUserIDProduct: try: modified_time1 = path.stat().st_mtime - except Exception: + except OSError: modified_time1 = 0 user_id_product_merge.build(client=client_no_amm, user_coll=user_collection) @@ -49,7 +60,9 @@ class TestUserIDProduct: assert modified_time2 > modified_time1 @pytest.mark.skip - def test_read(self, client_no_amm, user_id_product_merge): + def test_read( + self, client_no_amm: DaskClient, user_id_product_merge: UserIdProductMerge + ): users_ddf = user_id_product_merge.ddf() df = client_no_amm.compute(collections=users_ddf, sync=True) diff --git a/tests/incite/mergers/test_pop_ledger.py b/tests/incite/mergers/test_pop_ledger.py index 529a641..2146344 100644 --- a/tests/incite/mergers/test_pop_ledger.py +++ b/tests/incite/mergers/test_pop_ledger.py @@ -86,9 +86,7 @@ class TestMergePOPLedger: # -- - user_wallet_account: LedgerAccount = ( - thl_ledger_manager.get_account_or_create_user_wallet(user=u) - ) + thl_ledger_manager.get_account_or_create_user_wallet(user=u) cash_account: LedgerAccount = thl_ledger_manager.get_account_cash() rev_account: LedgerAccount = ( thl_ledger_manager.get_account_task_complete_revenue() @@ -295,7 +293,7 @@ class TestMergePOPLedger: assert isinstance(df.index, pd.Index) assert isinstance(df.index, pd.DatetimeIndex) - bp_account_balance = thl_ledger_manager.get_account_balance(account=bp_account) + thl_ledger_manager.get_account_balance(account=bp_account) # Initial sum initial_sum = df.sum().sum() diff --git a/tests/incite/test_collection_base.py b/tests/incite/test_collection_base.py index d6ce2b1..577eda9 100644 --- a/tests/incite/test_collection_base.py +++ b/tests/incite/test_collection_base.py @@ -241,7 +241,7 @@ class TestCollectionBaseMethodsCleanup: assert "Must override" in str(cm.value) -class TestCollectionBaseMethodsCleanup: +class TestCollectionBaseMethodsCleanup2: @pytest.mark.skip def test_cleanup_partials(self, mnt_filepath: GRLDatasets): diff --git a/tests/managers/thl/test_contest/test_milestone.py b/tests/managers/thl/test_contest/test_milestone.py index e3889bc..aa738e1 100644 --- a/tests/managers/thl/test_contest/test_milestone.py +++ b/tests/managers/thl/test_contest/test_milestone.py @@ -15,6 +15,7 @@ from generalresearch.models.thl.contest.milestone import ( MilestoneContestCreate, MilestoneUserView, ) +from generalresearch.models.thl.contest.raffle import RaffleContest from generalresearch.models.thl.product import Product from generalresearch.models.thl.user import User @@ -241,7 +242,7 @@ class TestMilestoneContestUserViews: def test_list_user_eligible_country( self, user_with_wallet: User, - contest_factory: Callable[..., Contest], + raffle_contest_factory: Callable[..., Contest], thl_ledger_manager: ThlLedgerManager, contest_manager: ContestManager, ): @@ -252,7 +253,7 @@ class TestMilestoneContestUserViews: assert len(cs) == 0 # Create a contest. It'll be in the US/CA - contest_factory(country_isos={"us", "ca"}) + raffle_contest_factory(country_isos={"us", "ca"}) # Not eligible in mexico cs = contest_manager.get_many_by_user_eligible( @@ -265,7 +266,7 @@ class TestMilestoneContestUserViews: assert len(cs) == 1 # Create another, any country - contest_factory(country_isos=None) + raffle_contest_factory(country_isos=None) cs = contest_manager.get_many_by_user_eligible( user=user_with_wallet, country_iso="mx" ) @@ -278,12 +279,12 @@ class TestMilestoneContestUserViews: def test_list_user_eligible( self, user_with_money: User, - contest_factory: Callable[..., Contest], + raffle_contest_factory: Callable[..., RaffleContest], thl_ledger_manager: ThlLedgerManager, contest_manager: ContestManager, ): # User reaches milestone after 1 complete - c = contest_factory(target_amount=1) + c = raffle_contest_factory(target_amount=1) user = user_with_money cs = contest_manager.get_many_by_user_eligible( diff --git a/tests/managers/thl/test_ledger/test_thl_lm_tx__user_payouts.py b/tests/managers/thl/test_ledger/test_thl_lm_tx__user_payouts.py index 5fb6935..82dc143 100644 --- a/tests/managers/thl/test_ledger/test_thl_lm_tx__user_payouts.py +++ b/tests/managers/thl/test_ledger/test_thl_lm_tx__user_payouts.py @@ -277,11 +277,12 @@ class TestLedgerManagerAMT: thl_ledger_manager.create_tx_user_payout_cancelled( user, payout_event=pe, skip_flag_check=True ) - with pytest.raises(expected_exception=LedgerTransactionConditionFailedError): - with caplog.at_level(logging.WARNING): - thl_ledger_manager.create_tx_user_payout_complete( - user, payout_event=pe, skip_flag_check=True - ) + with pytest.raises( + expected_exception=LedgerTransactionConditionFailedError + ), caplog.at_level(logging.WARNING): + thl_ledger_manager.create_tx_user_payout_complete( + user, payout_event=pe, skip_flag_check=True + ) assert "trying to complete payout that was already cancelled" in caplog.text cash = thl_ledger_manager.get_account_cash() @@ -319,11 +320,12 @@ class TestLedgerManagerAMT: thl_ledger_manager.create_tx_user_payout_complete( user, payout_event=pe2, skip_flag_check=True ) - with pytest.raises(expected_exception=LedgerTransactionConditionFailedError): - with caplog.at_level(logging.WARNING): - thl_ledger_manager.create_tx_user_payout_cancelled( - user, payout_event=pe2, skip_flag_check=True - ) + with pytest.raises( + expected_exception=LedgerTransactionConditionFailedError + ), caplog.at_level(logging.WARNING): + thl_ledger_manager.create_tx_user_payout_cancelled( + user, payout_event=pe2, skip_flag_check=True + ) assert "trying to cancel payout that was already completed" in caplog.text diff --git a/tests/models/spectrum/test_survey.py b/tests/models/spectrum/test_survey.py index f97860f..bad6857 100644 --- a/tests/models/spectrum/test_survey.py +++ b/tests/models/spectrum/test_survey.py @@ -407,44 +407,12 @@ class TestSpectrumSurvey: ) -def test_spectrum_something(spectrum_api_surveys_json: list[str]): - # make sure hashes for 111111 are in db - c1 = SpectrumCondition( - question_id="1001", - value_type=ConditionValueType.LIST, - values=["a", "b", "c"], - negate=False, - logical_operator=LogicalOperator.OR, - ) - c2 = SpectrumCondition( - question_id="1001", - value_type=ConditionValueType.LIST, - values=["a"], - negate=False, - logical_operator=LogicalOperator.OR, - ) - c3 = SpectrumCondition( - question_id="1002", - value_type=ConditionValueType.RANGE, - values=["18-24", "30-32"], - negate=False, - logical_operator=LogicalOperator.OR, - ) - c4 = SpectrumCondition( - question_id="212", - value_type=ConditionValueType.LIST, - values=["23", "24"], - negate=False, - logical_operator=LogicalOperator.OR, - ) - c5 = SpectrumCondition( - question_id="1031", - value_type=ConditionValueType.LIST, - values=["113", "114", "121"], - negate=False, - logical_operator=LogicalOperator.OR, - ) - _conditions = [c1, c2, c3, c4, c5] +def test_spectrum_something( + spectrum_conditions: list[SpectrumCondition], spectrum_api_surveys_json: list[str] +): + + c1 = spectrum_conditions[0] + c3 = spectrum_conditions[2] survey = SpectrumSurvey.model_validate_json(spectrum_api_surveys_json[0]) assert c1.criterion_hash in survey.qualifications -- cgit v1.2.3 From 8219bf4814a7aa374a500516f2254f39085f0357 Mon Sep 17 00:00:00 2001 From: Max Nanis Date: Fri, 28 Aug 2026 17:28:10 -0700 Subject: Ruff afternoon --- generalresearch/__init__.py | 19 +++- generalresearch/config.py | 6 +- generalresearch/currency.py | 26 ++--- generalresearch/grliq/managers/forensic_data.py | 84 +++++++------- generalresearch/grliq/managers/forensic_summary.py | 8 +- generalresearch/grliq/models/events.py | 7 +- generalresearch/grliq/models/forensic_data.py | 19 ++-- generalresearch/grliq/models/forensic_summary.py | 7 +- generalresearch/grliq/models/useragents.py | 4 +- generalresearch/healing_ppe.py | 2 +- generalresearch/incite/base.py | 28 ++--- generalresearch/incite/collections/__init__.py | 123 +++++--------------- generalresearch/incite/exceptions.py | 19 ++++ generalresearch/incite/mergers/__init__.py | 26 ++--- .../incite/mergers/foundations/__init__.py | 18 +-- .../incite/mergers/foundations/enriched_session.py | 26 ++--- .../mergers/foundations/enriched_task_adjust.py | 9 +- .../incite/mergers/foundations/enriched_wall.py | 2 +- generalresearch/incite/mergers/ym_survey_wall.py | 4 +- generalresearch/incite/mergers/ym_wall_summary.py | 13 ++- generalresearch/incite/schemas/admin_responses.py | 73 +++++++----- generalresearch/locales/__init__.py | 4 - generalresearch/locales/setup_json.py | 124 ++++++++++----------- generalresearch/logging.py | 5 +- generalresearch/managers/cint/survey.py | 2 +- generalresearch/managers/criteria.py | 20 ++-- generalresearch/managers/dynata/survey.py | 64 +++++------ generalresearch/managers/events.py | 57 +++++----- generalresearch/managers/gr/authentication.py | 6 +- generalresearch/managers/innovate/survey.py | 83 +++++++------- generalresearch/managers/morning/survey.py | 87 ++++++++------- generalresearch/managers/network/mtr.py | 7 +- generalresearch/managers/network/nmap.py | 2 +- generalresearch/managers/network/rdns.py | 5 +- generalresearch/managers/precision/survey.py | 51 ++++----- generalresearch/managers/prodege/survey.py | 51 ++++----- generalresearch/managers/repdata/survey.py | 81 +++++++------- generalresearch/managers/sago/survey.py | 65 +++++------ generalresearch/managers/spectrum/survey.py | 63 +++++------ generalresearch/managers/thl/task_adjustment.py | 8 +- .../thl/user_manager/mysql_user_manager.py | 2 +- generalresearch/managers/thl/userhealth.py | 4 +- generalresearch/managers/thl/wall.py | 2 +- generalresearch/models/gr/business.py | 5 +- generalresearch/models/legacy/bucket.py | 6 + generalresearch/pg_helper.py | 11 +- generalresearch/sql_helper.py | 5 +- .../collections/test_df_collection_item_thl_web.py | 12 +- .../test_df_collection_thl_marketplaces.py | 2 +- tests/managers/thl/test_contest/test_milestone.py | 2 +- 50 files changed, 677 insertions(+), 682 deletions(-) create mode 100644 generalresearch/incite/exceptions.py (limited to 'tests/incite/collections') diff --git a/generalresearch/__init__.py b/generalresearch/__init__.py index 3b2ec3d..f27385c 100644 --- a/generalresearch/__init__.py +++ b/generalresearch/__init__.py @@ -1,20 +1,27 @@ +import logging import threading import time from collections.abc import Callable from functools import wraps -from typing import Any, Optional +from typing import ParamSpec, TypeVar from decorator import decorator from wrapt import FunctionWrapper, ObjectProxy +P = ParamSpec("P") +R = TypeVar("R") + +ExceptionType = type[BaseException] +ExceptionsArg = ExceptionType | tuple[ExceptionType, ...] + def retry( - exceptions, + exceptions: ExceptionsArg, tries: int = 4, delay: float = 0.5, backoff: int = 2, - logger: Any | None = None, -) -> Callable: + logger: logging.Logger | None = None, +) -> Callable[[Callable[P, R]], Callable[P, R]]: """ https://www.calazan.com/retry-decorator-for-python-3/ Retry calling the decorated function using an exponential backoff. @@ -29,10 +36,10 @@ def retry( logger: Logger to use. If None, print. """ - def deco_retry(f): + def deco_retry(f: Callable[P, R]) -> Callable[P, R]: @wraps(f) - def f_retry(*args, **kwargs): + def f_retry(*args: P.args, **kwargs: P.kwargs) -> R: mtries, mdelay = tries, delay while mtries > 1: try: diff --git a/generalresearch/config.py b/generalresearch/config.py index c6f41e8..c414069 100644 --- a/generalresearch/config.py +++ b/generalresearch/config.py @@ -16,9 +16,9 @@ def is_debug() -> bool: import os is_developer: bool = os.getenv("USER") in {"nanis", "gstupp"} - is_pytest1: bool = bool(os.getenv("PYTEST_TEST", False)) - is_pytest2: bool = bool(os.getenv("PYTEST_CURRENT_TEST", False)) - is_pytest3: bool = bool(os.getenv("PYTEST_VERSION", False)) + is_pytest1: bool = bool(os.getenv("PYTEST_TEST")) + is_pytest2: bool = bool(os.getenv("PYTEST_CURRENT_TEST")) + is_pytest3: bool = bool(os.getenv("PYTEST_VERSION")) is_debugging1: bool = os.getenv("DEBUG", "").lower() in ("1", "true", "yes") is_debugging2: bool = os.getenv("PYTHON_DEBUG", "").lower() in ("1", "true", "yes") is_jenkins: bool = bool(os.getenv("JENKINS_HOME")) or bool(os.getenv("JENKINS_URL")) diff --git a/generalresearch/currency.py b/generalresearch/currency.py index 9402e04..716cb0f 100644 --- a/generalresearch/currency.py +++ b/generalresearch/currency.py @@ -25,7 +25,7 @@ def format_usd_cent(usd_cent: int) -> str: class USDCent(int): - def __new__(cls, value, *args, **kwargs): + def __new__(cls, value: int, *args, **kwargs): if isinstance(value, float): warnings.warn( @@ -42,17 +42,17 @@ class USDCent(int): return super(cls, cls).__new__(cls, value) - def __add__(self, other): + def __add__(self, other: Any): assert isinstance(other, USDCent) res = super().__add__(other) return self.__class__(res) - def __sub__(self, other): + def __sub__(self, other: Any): assert isinstance(other, USDCent) res = super().__sub__(other) return self.__class__(res) - def __mul__(self, other): + def __mul__(self, other: Any): assert isinstance(other, USDCent) res = super().__mul__(other) return self.__class__(res) @@ -61,14 +61,14 @@ class USDCent(int): res = super().__abs__() return self.__class__(res) - def __truediv__(self, other): + def __truediv__(self): raise ValueError("Division not allowed for USDCent") def __str__(self): - return "%d" % int(self) + return f"{int(self):d}" def __repr__(self): - return "USDCent(%d)" % int(self) + return f"USDCent({int(self)})" @classmethod def __get_pydantic_core_schema__( @@ -110,17 +110,17 @@ class USDMill(int): return super(cls, cls).__new__(cls, value) - def __add__(self, other): + def __add__(self, other: Any): assert isinstance(other, USDMill) res = super().__add__(other) return self.__class__(res) - def __sub__(self, other): + def __sub__(self, other: Any): assert isinstance(other, USDMill) res = super().__sub__(other) return self.__class__(res) - def __mul__(self, other): + def __mul__(self, other: Any): assert isinstance(other, USDMill) res = super().__mul__(other) return self.__class__(res) @@ -129,14 +129,14 @@ class USDMill(int): res = super().__abs__() return self.__class__(res) - def __truediv__(self, other): + def __truediv__(self): raise ValueError("Division not allowed for USDMill") def __str__(self): - return "%d" % int(self) + return f"{int(self):d}" def __repr__(self): - return "USDMill(%d)" % int(self) + return f"USDMill({int(self)})" @classmethod def __get_pydantic_core_schema__( diff --git a/generalresearch/grliq/managers/forensic_data.py b/generalresearch/grliq/managers/forensic_data.py index 093f7ae..c1eac37 100644 --- a/generalresearch/grliq/managers/forensic_data.py +++ b/generalresearch/grliq/managers/forensic_data.py @@ -103,14 +103,13 @@ class GrlIqDataManager: is_attempt_allowed = %(is_attempt_allowed)s WHERE uuid = %(uuid)s """) - with self.postgres_config.make_connection() as conn: - with conn.cursor() as c: - c.execute(query, data) - if c.rowcount != 1: - raise ValueError( - f"Expected 1 row to be updated, but {c.rowcount} rows were affected." - ) - conn.commit() + with self.postgres_config.make_connection() as conn, conn.cursor() as c: + c.execute(query, data) + if c.rowcount != 1: + raise ValueError( + f"Expected 1 row to be updated, but {c.rowcount} rows were affected." + ) + conn.commit() def update_fingerprint(self, iq_data: GrlIqData) -> None: # We should only run this if we modified the fingerprint algorithm @@ -123,14 +122,13 @@ class GrlIqDataManager: SET fingerprint = %(fingerprint)s WHERE uuid = %(uuid)s """) - with self.postgres_config.make_connection() as conn: - with conn.cursor() as c: - c.execute(query, data) - if c.rowcount != 1: - raise ValueError( - f"Expected 1 row to be updated, but {c.rowcount} rows were affected." - ) - conn.commit() + with self.postgres_config.make_connection() as conn, conn.cursor() as c: + c.execute(query, data) + if c.rowcount != 1: + raise ValueError( + f"Expected 1 row to be updated, but {c.rowcount} rows were affected." + ) + conn.commit() def update_data(self, iq_data: GrlIqData) -> None: # We should only run this if we structured new fields and want to @@ -141,14 +139,13 @@ class GrlIqDataManager: SET data = %(data)s WHERE id = %(id)s """) - with self.postgres_config.make_connection() as conn: - with conn.cursor() as c: - c.execute(query, data) - if c.rowcount != 1: - raise ValueError( - f"Expected 1 row to be updated, but {c.rowcount} rows were affected." - ) - conn.commit() + with self.postgres_config.make_connection() as conn, conn.cursor() as c: + c.execute(query, data) + if c.rowcount != 1: + raise ValueError( + f"Expected 1 row to be updated, but {c.rowcount} rows were affected." + ) + conn.commit() def get_data_if_exists( self, forensic_uuid: UUIDStr, load_events: bool = False @@ -472,7 +469,7 @@ class GrlIqDataManager: filters.append("d.fingerprint = ANY(%(fingerprints)s)") if product_ids and len(product_ids) == 1: - product_id = list(product_ids)[0] + product_id = next(iter(product_ids)) product_ids = None if product_ids: @@ -536,9 +533,9 @@ class GrlIqDataManager: [f"(%(bp_{i})s, %(bpuid_{i})s)" for i in range(len(users))] ) filters.append(f"(d.product_id, d.product_user_id) IN ({user_args})") - for i, user in enumerate(users): - params[f"bp_{i}"] = user.product_id - params[f"bpuid_{i}"] = user.product_user_id + for i, _user in enumerate(users): + params[f"bp_{i}"] = _user.product_id + params[f"bpuid_{i}"] = _user.product_user_id if phase: params["phase"] = phase.value @@ -595,24 +592,19 @@ class GrlIqDataManager: ) if only_product_id: - try: - with self.postgres_config.make_connection() as conn: - with conn.cursor() as c: - c.execute( - query=""" - SELECT count AS c - FROM grliq_forensicdata_product_counts - WHERE product_id = %s - LIMIT 1 - """, - params=(product_id,), - ) - res = c.fetchone() - if res and res["c"] >= 0: - return int(res["c"]) - - except Exception: - pass + with self.postgres_config.make_connection() as conn, conn.cursor() as c: + c.execute( + query=""" + SELECT count AS c + FROM grliq_forensicdata_product_counts + WHERE product_id = %s + LIMIT 1 + """, + params=(product_id,), + ) + res = c.fetchone() + if res and res["c"] >= 0: + return int(res["c"]) query = f""" SELECT COUNT(1) AS c diff --git a/generalresearch/grliq/managers/forensic_summary.py b/generalresearch/grliq/managers/forensic_summary.py index 98e00d8..b86e1f5 100644 --- a/generalresearch/grliq/managers/forensic_summary.py +++ b/generalresearch/grliq/managers/forensic_summary.py @@ -109,7 +109,7 @@ def calculate_timing_summary( for k, v in country_distributions.items() } - out = dict() + out = {} for country_iso, median_rtts in country_median_rtts.items(): country_stats = country_distributions[country_iso] z_scores = [ @@ -158,12 +158,14 @@ def run_user_forensic_summary( ) session_uuids = {x["session_uuid"] for x in res} - timing_res: list[dict] = iq_em.filter_distinct_timing(session_uuids=session_uuids) + timing_res: list[dict[str, Any]] = iq_em.filter_distinct_timing( + session_uuids=session_uuids + ) country_timing_data_summary = ( calculate_timing_summary(redis_config=redis_config, timing_res=timing_res) if timing_res - else dict() + else {} ) s = UserForensicSummary( diff --git a/generalresearch/grliq/models/events.py b/generalresearch/grliq/models/events.py index 69b67e5..995a6fa 100644 --- a/generalresearch/grliq/models/events.py +++ b/generalresearch/grliq/models/events.py @@ -177,12 +177,7 @@ class TimingData(BaseModel): @property def server_location(self) -> str: - # TODO: when we have more locations ... - return ( - "fremont_ca" - if self.server_hostname in {"grliq-web-0", "grliq-web-1"} - else "fremont_ca" - ) + return "fremont_ca" @property def has_data(self): diff --git a/generalresearch/grliq/models/forensic_data.py b/generalresearch/grliq/models/forensic_data.py index 666cb81..69b1760 100644 --- a/generalresearch/grliq/models/forensic_data.py +++ b/generalresearch/grliq/models/forensic_data.py @@ -474,9 +474,11 @@ class GrlIqData(BaseModel): description="Bit-packed string for font support. Each element is 32 bits, with each bit representing T/F for " "font support.", examples=[ - "72|768|262144|1073741824|0|0|540672|73728|7340032|1342177280|117446656|256|16|0|543|4290797636" - "|1677723648|4168998400|0|1048576|262144|268500994|1342177280|262144|125829376|37888000|0|435363842|0" - "|2147483648|109543424|1880099872|268435471" + ( + "72|768|262144|1073741824|0|0|540672|73728|7340032|1342177280|117446656|256|16|0|543|4290797636" + "|1677723648|4168998400|0|1048576|262144|268500994|1342177280|262144|125829376|37888000|0|435363842|0" + "|2147483648|109543424|1880099872|268435471" + ) ], ) @@ -558,19 +560,21 @@ class GrlIqData(BaseModel): @cached_property def audio_codecs_named(self) -> dict[str, bool]: + assert self.audio_codecs return dict( zip( AUDIO_CODEC_NAMES, - [True if x == "3" else False for x in self.audio_codecs.split(",")], + [x == "3" for x in self.audio_codecs.split(",")], ) ) @cached_property def video_codecs_named(self) -> dict[str, bool]: + assert self.video_codecs return dict( zip( VIDEO_CODEC_NAMES, - [True if x == "3" else False for x in self.video_codecs.split(",")], + [x == "3" for x in self.video_codecs.split(",")], ) ) @@ -779,9 +783,8 @@ class GrlIqData(BaseModel): minutes=90 ), "expired session" - def model_dump_sql(self, **kwargs) -> dict[str, Any]: - d = dict() + d = {} d["uuid"] = self.uuid d["session_uuid"] = self.mid d["created_at"] = self.created_at @@ -801,7 +804,7 @@ class GrlIqData(BaseModel): return d @classmethod - def from_db(cls, d: dict[str, Any]) -> Self: + def from_db(cls, d: dict[str, Any]) -> GrlIqData: res = GrlIqData.model_validate(d["data"]) if d.get("category_result"): diff --git a/generalresearch/grliq/models/forensic_summary.py b/generalresearch/grliq/models/forensic_summary.py index f5b0f25..6b0e065 100644 --- a/generalresearch/grliq/models/forensic_summary.py +++ b/generalresearch/grliq/models/forensic_summary.py @@ -251,10 +251,9 @@ class CountryRTTDistribution(BaseModel): Render a boxplot from the RTT percentiles. """ try: - # annoying pycharm error import matplotlib.pyplot as plt - except ImportError as e: - raise e + except ImportError: + return p = self.rtt_percentiles data = { @@ -266,7 +265,7 @@ class CountryRTTDistribution(BaseModel): "fliers": [p[0]] + ([p[100]] if p[100] > p[95] else []), } - fig, ax = plt.subplots(figsize=(4, 1.5)) + _, ax = plt.subplots(figsize=(4, 1.5)) ax.bxp([data], showfliers=True, vert=False) ax.set_title(f"RTT Boxplot for {self.country_iso}") ax.set_xlabel("RTT (ms)") diff --git a/generalresearch/grliq/models/useragents.py b/generalresearch/grliq/models/useragents.py index de63a67..3d5e5ce 100644 --- a/generalresearch/grliq/models/useragents.py +++ b/generalresearch/grliq/models/useragents.py @@ -132,7 +132,7 @@ class BrowserInfo(BaseModel): class DeviceInfo(BaseModel): family: DeviceModelFamily = Field() - brand: DeviceBrand = Field() + brand: DeviceBrand | None = Field(default=None) model: DeviceModelFamily = Field() @field_validator("family", "model", mode="before") @@ -186,7 +186,7 @@ class GrlUserAgent(BaseModel): def ua_string_values(self) -> dict[str, str]: # Returns the raw parsed string values for each of these. To be used # for db filtering, identifying trends, etc. - d = dict() + d = {} d["ua_browser_family"] = self.ua_parsed.browser.family d["ua_browser_version"] = self.ua_parsed.browser.version_string d["ua_os_family"] = self.ua_parsed.os.family diff --git a/generalresearch/healing_ppe.py b/generalresearch/healing_ppe.py index 254a893..dfee689 100644 --- a/generalresearch/healing_ppe.py +++ b/generalresearch/healing_ppe.py @@ -73,7 +73,7 @@ def test(): time.sleep(0.5) # Kill a process in the pool - pid = list(pool._processes.keys())[0] + pid = next(iter(pool._processes.keys())) os.kill(pid, signal.SIGKILL) time.sleep(0.5) diff --git a/generalresearch/incite/base.py b/generalresearch/incite/base.py index 9d504bd..060df4e 100644 --- a/generalresearch/incite/base.py +++ b/generalresearch/incite/base.py @@ -25,6 +25,7 @@ from uuid import uuid4 import dask import dask.dataframe as dd import pandas as pd +import pandera as pa import pyarrow.parquet as pq from distributed import Client as DaskClient from pandera.pandas import DataFrameSchema @@ -45,6 +46,7 @@ from pydantic.json_schema import SkipJsonSchema from sentry_sdk import capture_exception from generalresearch.config import is_debug +from generalresearch.incite.collections import DFCollectionItem from generalresearch.incite.schemas import ( ARCHIVE_AFTER, empty_dataframe_from_schema, @@ -61,7 +63,7 @@ if TYPE_CHECKING: Collection = DFCollection | MergeCollection logging.basicConfig() -LOG = logging.getLogger() +LOG = logging.getLogger(f"{__name__}.incite") # Item = Union["DFCollectionItem", "MergeCollectionItem"] Item = Any @@ -133,6 +135,7 @@ class GRLDatasets(BaseModel): from generalresearch.incite.mergers import MergeType folder = "mergers" if isinstance(enum_type, MergeType) else "raw/df-collections" + assert self.incite is not None return Path( pjoin(self.data_src, self.incite.point, folder, str(enum_type.value)) ) @@ -203,9 +206,6 @@ class CollectionBase(BaseModel): @model_validator(mode="after") def check_model_after(self) -> Self: - if self.offset is None or self.start is None: - return self - offset_total_sec = pd.Timedelta(self.offset).total_seconds() start_total_sec = (datetime.now(tz=UTC) - self.start).total_seconds() @@ -230,7 +230,7 @@ class CollectionBase(BaseModel): return v try: pd.Timedelta(v) - except Exception as e: + except (ValueError, TypeError) as e: capture_exception(error=e) raise ValueError( "Invalid offset alias provided. Please review: " @@ -554,7 +554,7 @@ class CollectionBase(BaseModel): try: pq.ParquetDataset(highest_version).read().to_pandas() - except Exception: + except (pa.ArrowInvalid, pa.ArrowIOError, FileNotFoundError): # If the most recent version isn't valid, we don't want to # create a symlink to it. # TODO: We could try to be smart and iterate down the most recent @@ -764,11 +764,12 @@ class CollectionItemBase(BaseModel): # regex = re.compile(r'\.parquet\.[0-9a-f]{32}', re.I) builds = [] for fn in os.listdir(coll.archive_path): - if fn.startswith(self.filename): - - # Don't include the "broken link" or mmfsymlink text file - if fn != self.filename and fn != self.partial_filename: - builds.append(fn) + if ( + fn.startswith(self.filename) + and fn != self.filename + and fn != self.partial_filename + ): + builds.append(fn) if len(builds) == 0: return None @@ -872,7 +873,7 @@ class CollectionItemBase(BaseModel): raise ValueError("Unknown path type.") df = parquet.read().to_pandas() - except Exception: + except (pa.ArrowInvalid, pa.ArrowIOError, OSError): LOG.warning(f"Invalid archive {path=}") df = None @@ -891,8 +892,7 @@ class CollectionItemBase(BaseModel): try: schema: DataFrameSchema = self._collection._schema return schema.validate(check_obj=df, lazy=True, sample=sample) - except Exception as e: - LOG.exception(e) + except pa.errors.SchemaErrors as e: capture_exception(error=e) return None diff --git a/generalresearch/incite/collections/__init__.py b/generalresearch/incite/collections/__init__.py index 42c3d31..f1e1b77 100644 --- a/generalresearch/incite/collections/__init__.py +++ b/generalresearch/incite/collections/__init__.py @@ -1,6 +1,5 @@ from __future__ import annotations -import logging import os import subprocess import time @@ -12,16 +11,18 @@ from typing import Any import dask import dask.dataframe as dd import pandas as pd +import pyarrow as pa import pyarrow.parquet as pq +from dask.distributed import Client as DaskClient from dask.distributed import Future -from distributed import Client, as_completed +from distributed import as_completed from more_itertools import chunked from pandera.pandas import DataFrameSchema from psycopg import Cursor from pydantic import Field, FilePath, ValidationInfo, field_validator from sentry_sdk import capture_exception -from generalresearch.incite.base import CollectionBase, CollectionItemBase +from generalresearch.incite.base import LOG, CollectionBase, CollectionItemBase from generalresearch.incite.schemas import ( ARCHIVE_AFTER, ORDER_KEY, @@ -49,9 +50,6 @@ from generalresearch.incite.schemas.thl_web import ( UserHealthIPHistoryWSSchema, ) from generalresearch.pg_helper import PostgresConfig -from generalresearch.sql_helper import SqlHelper - -LOG = logging.getLogger("incite") DT_STR = "%Y-%m-%d %H:%M:%S" @@ -106,18 +104,6 @@ class DFCollectionItem(CollectionItemBase): # --- Methods --- - def has_mysql(self) -> bool: - if self._collection.sql_helper is None: - return False - - connected = True - try: - self._collection.sql_helper.execute_sql_query("""SELECT 1;""") - except: - connected = False - - return connected - def has_postgres(self) -> bool: if self._collection.pg_config is None: return False @@ -125,7 +111,7 @@ class DFCollectionItem(CollectionItemBase): connected = True try: self._collection.pg_config.execute_sql_query("""SELECT 1;""") - except: + except AssertionError: connected = False return connected @@ -148,7 +134,7 @@ class DFCollectionItem(CollectionItemBase): since = max([since, self.start]) # don't allow to query before the item's start df = df[df[order_key] < since].copy() - _df = self.from_mysql(since=since) + _df = self.from_db(since=since) if _df is not None: df = pd.concat([df, _df]) @@ -161,7 +147,7 @@ class DFCollectionItem(CollectionItemBase): return True def create_partial_archive(self) -> bool: - _df = self.from_mysql() + _df = self.from_db() if _df is None: # Returned no rows, but the period is not closed, so we # don't want to mark as empty. Do nothing. @@ -172,71 +158,15 @@ class DFCollectionItem(CollectionItemBase): def to_dict(self) -> dict[str, Any]: return self._to_dict() - def from_mysql(self, since: datetime | None = None) -> pd.DataFrame | None: + def from_db(self, since: datetime | None = None) -> pd.DataFrame | None: if self._collection.data_type == DFCollectionType.LEDGER: assert since is None, "Shouldn't pass since for Ledger item" assert self._collection.pg_config is not None return self.from_postgres_ledger() else: - if self._collection.sql_helper: - return self.from_mysql_standard(since=since) - else: - return self.from_postgres_standard(since=since) - - def from_mysql_standard(self, since: datetime | None = None) -> pd.DataFrame | None: - - assert ( - self._collection.data_type != DFCollectionType.LEDGER - ), "Can't call from_mysql_standard for Ledger DFCollectionItem" - - start, finish = self.start, self.finish - LOG.debug( - f"{self._collection.data_type.value}.from_mysql(" - f"start={start.strftime(DT_STR)}, " - f"finish={finish.strftime(DT_STR)})" - ) - coll = self._collection - schema = coll._schema - sql_helper = coll.sql_helper - - start = since or start - order_key = schema.metadata[ORDER_KEY] - cols = list(schema.columns.keys()) + [schema.index.name] - cols_str = ",".join(map(sql_helper._quote, cols)) - db_name = sql_helper.db - - try: - res = sql_helper.execute_sql_query( - query=f""" - SELECT {cols_str} - FROM `{db_name}`.`{coll.data_type.value}` - WHERE `{order_key}` >= %s AND `{order_key}` < %s; - """, - params=[start, finish], - ) - except Exception as e: - capture_exception(error=e) - LOG.error(f"_from_mysql Exception: {e}") - return None - - if not res: - LOG.warning("_from_mysql query returned nothing") - # Return an empty df.DataFrame with the correct columns - return empty_dataframe_from_schema(coll._schema) - - df = pd.DataFrame.from_records(res).set_index(coll._schema.index.name) - df = self.validate_df(df=df) - - if df is None: - LOG.warning("_from_mysql query results failed validation") - # Schema validation can fail... - return None - - return df + return self.from_db_standard(since=since) - def from_postgres_standard( - self, since: datetime | None = None - ) -> pd.DataFrame | None: + def from_db_standard(self, since: datetime | None = None) -> pd.DataFrame | None: assert ( self._collection.data_type != DFCollectionType.LEDGER ), "Can't call from_postgres_standard for Ledger DFCollectionItem" @@ -250,6 +180,7 @@ class DFCollectionItem(CollectionItemBase): coll = self._collection schema = coll._schema pg_config = coll.pg_config + assert pg_config, "Must provide PostgresConfig" start = since or start order_key = schema.metadata[ORDER_KEY] @@ -265,7 +196,7 @@ class DFCollectionItem(CollectionItemBase): """, params=[start, finish], ) - except Exception as e: + except AssertionError as e: capture_exception(error=e) LOG.error(f"_from_postgres Exception: {e}") return None @@ -298,13 +229,14 @@ class DFCollectionItem(CollectionItemBase): ) coll = self._collection + assert coll.pg_config, "Must provide PostgresConfig" pg_config: PostgresConfig = coll.pg_config limit = 20000 offset = 0 res = [] while True: - logging.info( + LOG.info( f"{self._collection.data_type.value}.from_postgres_ledger({limit=}, {offset=})" ) chunk = pg_config.execute_sql_query( @@ -399,7 +331,7 @@ class DFCollectionItem(CollectionItemBase): """ assert isinstance(ddf, dd.DataFrame), "must pass dask df" - client: Client | None = self._collection._client + client: DaskClient | None = self._collection._client # client = None if client: @@ -455,7 +387,9 @@ class DFCollectionItem(CollectionItemBase): tmp_path = self.tmp_path() try: schema = self._collection._schema - partition = schema.metadata.get(PARTITION_ON, None) + assert schema + assert schema.metadata + partition = schema.metadata.get(PARTITION_ON) ddf.to_parquet( path=tmp_path, @@ -466,7 +400,7 @@ class DFCollectionItem(CollectionItemBase): compression="brotli", ) - except Exception as e: + except (pa.ArrowInvalid, pa.ArrowIOError, OSError) as e: LOG.exception(e) self.delete_archive(tmp_path) return False @@ -516,7 +450,8 @@ class DFCollectionItem(CollectionItemBase): collection = self._collection schema = collection._schema - client: Client | None = collection._client + assert schema + client: DaskClient | None = collection._client next_numbered_path = self.next_numbered_path(self.partial_path) partial_path = self.partial_path @@ -544,7 +479,8 @@ class DFCollectionItem(CollectionItemBase): return False try: - partition = schema.metadata.get(PARTITION_ON, None) + assert schema.metadata + partition = schema.metadata.get(PARTITION_ON) ddf.to_parquet( path=next_numbered_path, partition_on=partition, @@ -553,7 +489,7 @@ class DFCollectionItem(CollectionItemBase): write_metadata_file=True, compression="brotli", ) - except Exception as e: + except (pa.ArrowInvalid, pa.ArrowIOError, OSError) as e: LOG.exception(e) self.delete_archive(next_numbered_path) return False @@ -572,7 +508,7 @@ class DFCollectionItem(CollectionItemBase): assert self.should_archive(), "not ready to archive!" - df: pd.DataFrame | None = self.from_mysql() + df: pd.DataFrame | None = self.from_db() if df is None: self.set_empty() @@ -592,7 +528,6 @@ class DFCollection(CollectionBase): # --- Private --- pg_config: PostgresConfig | None = Field(default=None) - sql_helper: SqlHelper | None = Field(default=None) def __repr__(self): res = self.signature() + "\n" @@ -620,7 +555,7 @@ class DFCollection(CollectionBase): return res @field_validator("data_type") - def check_data_type(cls, data_type, info: ValidationInfo): + def check_data_type(cls, data_type: DFCollectionType | None, info: ValidationInfo): if data_type is None: raise ValueError("Must explicitly provide a data_type") @@ -647,7 +582,7 @@ class DFCollection(CollectionBase): def initial_load( self, - client: Client | None = None, + client: DaskClient | None = None, sync: bool = True, since: datetime | None = None, client_resources: dict[str, Any] | None = None, @@ -685,7 +620,7 @@ class DFCollection(CollectionBase): if sync: fs = client.compute(fs, sync=False, priority=2, resources=client_resources) - ac = as_completed(fs, timeout=timeout) + _ = as_completed(fs, timeout=timeout) return fs else: @@ -734,7 +669,7 @@ class DFCollection(CollectionBase): def force_rr_latest( self, - client: Client, + client: DaskClient, client_resources: dict[str, Any] | None = None, sync: bool = True, ) -> list[Future]: diff --git a/generalresearch/incite/exceptions.py b/generalresearch/incite/exceptions.py new file mode 100644 index 0000000..112978f --- /dev/null +++ b/generalresearch/incite/exceptions.py @@ -0,0 +1,19 @@ +class BuildItemsError(Exception): + + def __init__(self, message: str): + self.message = message + super().__init__(self.message) + + +class BuildError(Exception): + + def __init__(self, message: str): + self.message = message + super().__init__(self.message) + + +class FetchError(Exception): + + def __init__(self, message: str): + self.message = message + super().__init__(self.message) diff --git a/generalresearch/incite/mergers/__init__.py b/generalresearch/incite/mergers/__init__.py index 45d3e2d..810bedc 100644 --- a/generalresearch/incite/mergers/__init__.py +++ b/generalresearch/incite/mergers/__init__.py @@ -108,13 +108,12 @@ class MergeCollectionItem(CollectionItemBase): client: Client, ddf: dd.DataFrame, is_partial: bool = False, - client_resources=None, ) -> bool: assert is_partial is False, "use to_archive_symlink" return self._to_archive(client=client, ddf=ddf, client_resources=None) def _to_archive( - self, client: Client, ddf: dd.DataFrame, client_resources=None + self, client: Client, ddf: dd.DataFrame | None, client_resources=None ) -> bool: """ For archiving an item. Will write an empty file if ddf is empty. @@ -125,12 +124,15 @@ class MergeCollectionItem(CollectionItemBase): if ddf is None: return False - row_len = client.compute(collections=ddf.shape[0], sync=True) + row_len: int = client.compute(collections=ddf.shape[0], sync=True) + assert row_len assert row_len > 0, "empty ddf" tmp_path = self.tmp_path() schema = self._collection._schema - partition = schema.metadata.get(PARTITION_ON, None) + assert schema.metadata + + partition = schema.metadata.get(PARTITION_ON) f = ddf.to_parquet( compute=False, path=tmp_path, @@ -178,8 +180,7 @@ class MergeCollectionItem(CollectionItemBase): collection = self._collection LOG.warning(f"{collection.merge_type.value}.to_archive_symlink()") - if not isinstance(ddf, dd.DataFrame): - raise ValueError("must pass a dask df") + assert isinstance(ddf, dd.DataFrame), "must pass a dask df" # We should validate before or after!!! # _validate_df(self.compute(ddf), coll._schema) @@ -212,13 +213,10 @@ class MergeCollectionItem(CollectionItemBase): else: subprocess.call(["ln", "-sfnT", target, path.as_posix()]) - if validate_after: - if not self.valid_archive(self.path): - LOG.error( - f"{collection.merge_type.value} failed validation: {self.path}" - ) - self.delete_archive(self.path) - return False + if validate_after and not self.valid_archive(self.path): + LOG.error(f"{collection.merge_type.value} failed validation: {self.path}") + self.delete_archive(self.path) + return False return True # todo: unclear what the common interface should be here ... ? @@ -256,7 +254,7 @@ class MergeCollection(CollectionBase): return self @field_validator("merge_type") - def check_merge_type(cls, merge_type, info: ValidationInfo): + def check_merge_type(cls, merge_type: MergeType | None, info: ValidationInfo): if merge_type is None: raise ValueError("Must explicitly provide a merge_type") diff --git a/generalresearch/incite/mergers/foundations/__init__.py b/generalresearch/incite/mergers/foundations/__init__.py index d3100fa..561e759 100644 --- a/generalresearch/incite/mergers/foundations/__init__.py +++ b/generalresearch/incite/mergers/foundations/__init__.py @@ -68,11 +68,9 @@ def lookup_product_and_team_id( assert len(user_ids) <= 1000, "you should chunk this bro" res: list[dict[str, Any]] = [] - with pg_config.make_connection() as conn: - try: - with conn.cursor() as c: - c.execute( - query=""" + with pg_config.make_connection() as conn, conn.cursor() as c: + c.execute( + query=""" SELECT u.id AS user_id, u.product_id, bp.team_id @@ -81,13 +79,9 @@ def lookup_product_and_team_id( ON bp.id = u.product_id WHERE u.id = ANY(%s); """, - params=[list(user_ids)], - ) - res.extend(c.fetchall()) - - except Exception as e: - LOG.exception(f"lookup_product_and_team_id: {e}") - raise + params=[list(user_ids)], + ) + res.extend(c.fetchall()) return res diff --git a/generalresearch/incite/mergers/foundations/enriched_session.py b/generalresearch/incite/mergers/foundations/enriched_session.py index 8a707ca..049b1bc 100644 --- a/generalresearch/incite/mergers/foundations/enriched_session.py +++ b/generalresearch/incite/mergers/foundations/enriched_session.py @@ -6,8 +6,8 @@ from typing import TYPE_CHECKING, Any, Literal import dask.dataframe as dd import pandas as pd +from dask.distributed import Client as DaskClient from dask.distributed import as_completed -from distributed import Client from more_itertools import chunked, flatten from generalresearch.incite.collections.thl_web import ( @@ -46,7 +46,7 @@ class EnrichedSessionMergeItem(MergeCollectionItem): session_coll: SessionDFCollection, wall_coll: WallDFCollection, pg_config: PostgresConfig, - client: Client | None = None, + client: DaskClient | None = None, client_resources: dict[str, Any] | None = None, ) -> None: @@ -140,9 +140,9 @@ class EnrichedSessionMergeItem(MergeCollectionItem): try: results = client.gather(list(futures)) - except Exception as e: + except Exception: client.cancel(futures, asynchronous=False, force=True) - raise e + raise dfp = pd.DataFrame( list(flatten(results)), columns=["user_id", "product_id", "team_id"] @@ -154,18 +154,13 @@ class EnrichedSessionMergeItem(MergeCollectionItem): df = df[df["started"].between(start, end)] is_missing = df[["product_id"]].isna().sum().sum() > 0 - session_is_partial = any([w.should_archive() is False for w in session_items]) + session_is_partial = any(w.should_archive() is False for w in session_items) session_is_missing = any( - [ - w.should_archive() is True and w.has_archive() is False - for w in session_items - ] + w.should_archive() is True and w.has_archive() is False + for w in session_items ) wall_is_missing = any( - [ - w.should_archive() is True and w.has_archive() is False - for w in wall_items - ] + w.should_archive() is True and w.has_archive() is False for w in wall_items ) is_partial = ( is_missing or session_is_partial or session_is_missing or wall_is_missing @@ -203,7 +198,7 @@ class EnrichedSessionMerge(MergeCollection): def build( self, - client: Client, + client: DaskClient, session_coll: SessionDFCollection, wall_coll: WallDFCollection, pg_config: PostgresConfig, @@ -232,7 +227,7 @@ class EnrichedSessionMerge(MergeCollection): def to_admin_response( self, rr: ReportRequest, - client: Client, + client: DaskClient, product_ids: list[UUIDStr] | None = None, user: User | None = None, ) -> pd.DataFrame: @@ -243,6 +238,7 @@ class EnrichedSessionMerge(MergeCollection): filters = [] if user: + assert product_ids assert ( len(product_ids) <= 1 ), "Can't search more than 1 Product ID for a specific User" diff --git a/generalresearch/incite/mergers/foundations/enriched_task_adjust.py b/generalresearch/incite/mergers/foundations/enriched_task_adjust.py index 234cd5b..e8a3654 100644 --- a/generalresearch/incite/mergers/foundations/enriched_task_adjust.py +++ b/generalresearch/incite/mergers/foundations/enriched_task_adjust.py @@ -11,6 +11,7 @@ from sentry_sdk import capture_exception from generalresearch.incite.collections.thl_web import ( TaskAdjustmentDFCollection, ) +from generalresearch.incite.exceptions import BuildError, BuildItemsError from generalresearch.incite.mergers import ( MergeCollection, MergeCollectionItem, @@ -61,7 +62,7 @@ class EnrichedTaskAdjustMergeItem(MergeCollectionItem): ] if len(task_adj_coll_items) == 0: - raise Exception("TaskAdjColl item collection failed") + raise BuildItemsError("TaskAdjColl item collection failed") ddf: dd.DataFrame | None = task_adj_coll.ddf( items=task_adj_coll_items, @@ -83,6 +84,8 @@ class EnrichedTaskAdjustMergeItem(MergeCollectionItem): ("started", "<", end), ], ) + + assert isinstance(ddf, pd.DataFrame) # Naked compute... don't log # LOG.info(f"TaskAdjustmentDetailMergeCollectionItem.rows: {len(ddf.index)}") @@ -91,7 +94,7 @@ class EnrichedTaskAdjustMergeItem(MergeCollectionItem): ew_items = [ew for ew in enriched_wall.items if ew.interval.overlaps(ir)] if len(ew_items) == 0: - raise Exception( + raise BuildItemsError( "EnrichedWall item collection failed for EnrichedTaskAdjColl" ) @@ -209,5 +212,5 @@ class EnrichedTaskAdjustMerge(MergeCollection): enriched_wall=enriched_wall, pg_config=pg_config, ) - except Exception as e: + except BuildError as e: capture_exception(error=e) diff --git a/generalresearch/incite/mergers/foundations/enriched_wall.py b/generalresearch/incite/mergers/foundations/enriched_wall.py index b2ac7bb..a74a556 100644 --- a/generalresearch/incite/mergers/foundations/enriched_wall.py +++ b/generalresearch/incite/mergers/foundations/enriched_wall.py @@ -148,7 +148,7 @@ class EnrichedWallMergeItem(MergeCollectionItem): is_missing = False df = df.dropna(subset=["product_id", "session_id"], how="any") - wall_is_partial = any([w.should_archive() is False for w in wall_items]) + wall_is_partial = any(w.should_archive() is False for w in wall_items) is_partial = is_missing or wall_is_partial # Lots of downstream issues with this... diff --git a/generalresearch/incite/mergers/ym_survey_wall.py b/generalresearch/incite/mergers/ym_survey_wall.py index 9750b57..a99e8ec 100644 --- a/generalresearch/incite/mergers/ym_survey_wall.py +++ b/generalresearch/incite/mergers/ym_survey_wall.py @@ -10,6 +10,7 @@ from distributed import Client from sentry_sdk import capture_exception from generalresearch.incite.collections.thl_web import WallDFCollection +from generalresearch.incite.exceptions import BuildError from generalresearch.incite.mergers import ( MergeCollection, MergeCollectionItem, @@ -109,7 +110,6 @@ class YMSurveyWallMergeCollectionItem(MergeCollectionItem): LOG.warning("YMSurveyWallMerge failed validation") - class YMSurveyWallMerge(MergeCollection): merge_type: Literal[MergeType.YM_SURVEY_WALL] = MergeType.YM_SURVEY_WALL collection_item_class: Literal[YMSurveyWallMergeCollectionItem] = ( @@ -143,7 +143,7 @@ class YMSurveyWallMerge(MergeCollection): wall_coll=wall_coll, enriched_session=enriched_session, ) - except Exception as e: + except BuildError as e: capture_exception(error=e) item.delete_dangling_partials(keep_latest=2, target_path=item.path) diff --git a/generalresearch/incite/mergers/ym_wall_summary.py b/generalresearch/incite/mergers/ym_wall_summary.py index 37fc3b9..69ef5c5 100644 --- a/generalresearch/incite/mergers/ym_wall_summary.py +++ b/generalresearch/incite/mergers/ym_wall_summary.py @@ -12,6 +12,7 @@ from generalresearch.incite.collections.thl_web import ( SessionDFCollection, WallDFCollection, ) +from generalresearch.incite.exceptions import FetchError from generalresearch.incite.mergers import ( MergeCollection, MergeCollectionItem, @@ -41,6 +42,8 @@ class YMWallSummaryMergeItem(MergeCollectionItem): ddf = wall_collection.ddf( items=wall_items, force_rr_latest=False, include_partial=True ) + assert isinstance(ddf, pd.DataFrame) + ddf = ddf[ddf["started"].between(start, end)] # Then we need the sessions for these wall events. They'll have started @@ -88,12 +91,14 @@ class YMWallSummaryMerge(MergeCollection): @field_validator("offset") def check_offset_ym_wall_summary(cls, v: str | None): # the offset MUST be on a whole day, no hourly + assert v assert v.endswith("D"), "offset must be in days" return v @field_validator("start") def check_start_ym_wall_summary(cls, v: datetime | None): # the start MUST be start on midnight exactly + assert v assert v.time() == time(0, 0, 0, 0), "start must no have a time component" return v @@ -117,7 +122,7 @@ class YMWallSummaryMerge(MergeCollection): # item every time build is run even if it isn't closed # if item.should_archive(): item.fetch(wall_collection, session_collection, user_id_product) - except Exception as e: + except FetchError as e: capture_exception(e) @staticmethod @@ -176,10 +181,10 @@ class YMWallSummaryMerge(MergeCollection): # df.to_parquet(str(self.archive_path) + ".all.parquet") pass - def get_counts(self, product_id): + def get_counts(self, product_id: str): # examples... product_id = "" - df = dd.read_parquet( + _ = dd.read_parquet( str(self.archive_path) + ".all.parquet", filters=[ ("product_id", "=", product_id), @@ -187,7 +192,7 @@ class YMWallSummaryMerge(MergeCollection): ], ).compute() country_iso = "de" - df = dd.read_parquet( + _ = dd.read_parquet( str(self.archive_path) + ".all.parquet", filters=[ ("product_id", "=", product_id), diff --git a/generalresearch/incite/schemas/admin_responses.py b/generalresearch/incite/schemas/admin_responses.py index bb2852f..e65c6e2 100644 --- a/generalresearch/incite/schemas/admin_responses.py +++ b/generalresearch/incite/schemas/admin_responses.py @@ -1,5 +1,9 @@ -from datetime import datetime +from __future__ import annotations +from collections.abc import Callable +from datetime import UTC, datetime + +import pandas as pd from pandera.pandas import ( Check, Column, @@ -14,6 +18,14 @@ BIG_INT32 = 2_147_483_647 SIX_HOUR_SECONDS = 6 * 60 * 6 ROUNDING = 2 +_fillna: Callable[[pd.Series], pd.Series] = lambda s: s.fillna(value=0.00) +_clip: Callable[[pd.Series], pd.Series] = lambda s: s.clip( + lower=0, upper=SIX_HOUR_SECONDS +) +_round: Callable[[pd.Series], pd.Series] = lambda s: s.round(decimals=ROUNDING) +_tz_localize_none: Callable[[pd.Series], pd.Series] = lambda i: i.dt.tz_localize(None) + + AdminPOPSchema = DataFrameSchema( # Generic: used for Session or Wall index=MultiIndex( @@ -25,10 +37,15 @@ AdminPOPSchema = DataFrameSchema( Index( name="index0", dtype=Timestamp, - parsers=[Parser(lambda i: i.dt.tz_localize(None))], + parsers=[Parser(_tz_localize_none)], checks=[ Check.less_than( - max_value=datetime(year=datetime.now().year + 1, month=1, day=1) + max_value=datetime( + year=datetime.now(tz=UTC).year + 1, + month=1, + day=1, + tzinfo=UTC, + ) ) ], ), @@ -44,85 +61,85 @@ AdminPOPSchema = DataFrameSchema( "elapsed_avg": Column( dtype=float, parsers=[ - Parser(lambda s: s.fillna(value=0.00)), - Parser(lambda s: s.clip(lower=0, upper=SIX_HOUR_SECONDS)), - Parser(lambda s: s.round(decimals=ROUNDING)), + Parser(_fillna), + Parser(_clip), + Parser(_round), ], checks=Check.between(min_value=0, max_value=SIX_HOUR_SECONDS), ), "elapsed_total": Column( dtype=int, parsers=[ - Parser(lambda s: s.fillna(value=0)), + Parser(_fillna), ], checks=Check.between(min_value=0, max_value=BIG_INT32), ), "payout_avg": Column( dtype=float, parsers=[ - Parser(lambda s: s.fillna(value=0.00)), - Parser(lambda s: s.round(decimals=ROUNDING)), + Parser(_fillna), + Parser(_round), ], checks=Check.between(min_value=0, max_value=100), ), "payout_total": Column( dtype=float, parsers=[ - Parser(lambda s: s.fillna(value=0.00)), - Parser(lambda s: s.round(decimals=ROUNDING)), + Parser(_fillna), + Parser(_round), ], checks=Check.between(min_value=0, max_value=BIG_INT32), ), "entrances": Column( dtype=int, parsers=[ - Parser(lambda s: s.fillna(value=0)), + Parser(_fillna), ], checks=Check.between(min_value=0, max_value=BIG_INT32), ), "completes": Column( dtype=int, parsers=[ - Parser(lambda s: s.fillna(value=0)), + Parser(_fillna), ], checks=Check.between(min_value=0, max_value=BIG_INT32), ), "users": Column( dtype=int, parsers=[ - Parser(lambda s: s.fillna(value=0)), + Parser(_fillna), ], checks=Check.between(min_value=0, max_value=BIG_INT32), ), "conversion": Column( dtype=float, parsers=[ - Parser(lambda s: s.fillna(value=0.00)), - Parser(lambda s: s.round(decimals=ROUNDING)), + Parser(_fillna), + Parser(_round), ], checks=Check.between(min_value=0.00, max_value=1.00), ), "epc": Column( dtype=float, parsers=[ - Parser(lambda s: s.fillna(value=0.00)), - Parser(lambda s: s.round(decimals=ROUNDING)), + Parser(_fillna), + Parser(_round), ], checks=Check.between(min_value=0, max_value=100), ), "eph": Column( dtype=float, parsers=[ - Parser(lambda s: s.fillna(value=0.00)), - Parser(lambda s: s.round(decimals=ROUNDING)), + Parser(_fillna), + Parser(_round), ], checks=Check.between(min_value=0, max_value=BIG_INT32), ), "cpc": Column( dtype=float, parsers=[ - Parser(lambda s: s.fillna(value=0.00)), - Parser(lambda s: s.round(decimals=ROUNDING)), + Parser(_fillna), + Parser(_round), ], checks=Check.between(min_value=0, max_value=250), ), @@ -140,21 +157,21 @@ AdminPOPWallSchema = DataFrameSchema( "buyers": Column( dtype=int, parsers=[ - Parser(lambda s: s.fillna(value=0)), + Parser(_fillna), ], checks=Check.between(min_value=0, max_value=BIG_INT32), ), "surveys": Column( dtype=int, parsers=[ - Parser(lambda s: s.fillna(value=0)), + Parser(_fillna), ], checks=Check.between(min_value=0, max_value=BIG_INT32), ), "sessions": Column( dtype=int, parsers=[ - Parser(lambda s: s.fillna(value=0)), + Parser(_fillna), ], checks=Check.between(min_value=0, max_value=BIG_INT32), ), @@ -169,15 +186,15 @@ AdminPOPSessionSchema = DataFrameSchema( "attempts_avg": Column( dtype=float, parsers=[ - Parser(lambda s: s.fillna(value=0.00)), - Parser(lambda s: s.round(decimals=ROUNDING)), + Parser(_fillna), + Parser(_round), ], checks=Check.between(min_value=0, max_value=25), ), "attempts_total": Column( dtype=int, parsers=[ - Parser(lambda s: s.fillna(value=0)), + Parser(_fillna), ], checks=Check.between(min_value=0, max_value=BIG_INT32), ), diff --git a/generalresearch/locales/__init__.py b/generalresearch/locales/__init__.py index 38c1832..4bb10d0 100644 --- a/generalresearch/locales/__init__.py +++ b/generalresearch/locales/__init__.py @@ -19,10 +19,6 @@ class Localelator: EVERYTHING IS LOWERCASE!!! (except this comment) """ - lang_alpha2_to_alpha3b = dict() - lang_alpha3_to_alpha3b = dict() - languages = set() - def __init__(self): d = json.loads(pkgutil.get_data(__name__, "iso639-3.json")) self.lang_alpha2_to_alpha3b = {x["alpha_2"]: x["alpha_3b"] for x in d} diff --git a/generalresearch/locales/setup_json.py b/generalresearch/locales/setup_json.py index 57caa00..71beb09 100644 --- a/generalresearch/locales/setup_json.py +++ b/generalresearch/locales/setup_json.py @@ -1,62 +1,62 @@ -import json - - -def country_default_lang(): - """ - Some marketplaces have no language specified. Surveys are in the "default - language for that country", whatever that means. This helper is meant to - provide a reasonable guess as to what language it is. - - Derived from: http://download.geonames.org/export/dump/countryInfo.txt - """ - raise ValueError("no need to run this, I already ran it.") - import pandas as pd - - from generalresearch.locales import Localelator - - l = Localelator() - - df = pd.read_csv( - "http://download.geonames.org/export/dump/countryInfo.txt", - sep="\t", - skiprows=49, - ) - df["default_lang"] = df.Languages.str.split(",").str[0].str.split("-").str[0] - df.default_lang = df.default_lang.fillna("en") - df.default_lang = df.default_lang.map( - lambda x: l.get_language_iso(x) if x in l.languages else "eng" - ) - df["#ISO"] = df["#ISO"].str.lower() - df["country_iso"] = df["#ISO"].map( - lambda x: l.get_country_iso(x) if x in l.countries else None - ) - df = df[df.country_iso.notnull()] - d = df.set_index("country_iso").default_lang.to_dict() - with open("country_default_lang.json", "w") as f: - json.dump(d, f, indent=2) - return d - - -def setup_json(): - # pycountry is 30mb, which makes using this package on AWS lambda problematic. - # These JSONs are stolen from pycountry and adapted. - - raise ValueError("no need to run this, I already ran it.") - - # languages - d = json.load(open("iso639-3.json")) - d["639-3"] = [x for x in d["639-3"] if "alpha_2" in x] - for x in d["639-3"]: - x["alpha_3b"] = x.pop("bibliographic", None) or x["alpha_3"] - del x["scope"] - del x["type"] - with open("iso639-3.json", "w") as f: - json.dump(d["639-3"], f, indent=2) - - # countries - d = json.load(open("iso3166-1.json"))["3166-1"] - for x in d: - x["alpha_2"] = x["alpha_2"].lower() - x["alpha_3"] = x["alpha_3"].lower() - with open("iso3166-1.json", "w") as f: - json.dump(d, f, indent=2) +# import json + + +# def country_default_lang(): +# """ +# Some marketplaces have no language specified. Surveys are in the "default +# language for that country", whatever that means. This helper is meant to +# provide a reasonable guess as to what language it is. + +# Derived from: http://download.geonames.org/export/dump/countryInfo.txt +# """ +# raise ValueError("no need to run this, I already ran it.") +# import pandas as pd + +# from generalresearch.locales import Localelator + +# l = Localelator() + +# df = pd.read_csv( +# "http://download.geonames.org/export/dump/countryInfo.txt", +# sep="\t", +# skiprows=49, +# ) +# df["default_lang"] = df.Languages.str.split(",").str[0].str.split("-").str[0] +# df.default_lang = df.default_lang.fillna("en") +# df.default_lang = df.default_lang.map( +# lambda x: l.get_language_iso(x) if x in l.languages else "eng" +# ) +# df["#ISO"] = df["#ISO"].str.lower() +# df["country_iso"] = df["#ISO"].map( +# lambda x: l.get_country_iso(x) if x in l.countries else None +# ) +# df = df[df.country_iso.notnull()] +# d = df.set_index("country_iso").default_lang.to_dict() +# with open("country_default_lang.json", "w") as f: +# json.dump(d, f, indent=2) +# return d + + +# def setup_json(): +# # pycountry is 30mb, which makes using this package on AWS lambda problematic. +# # These JSONs are stolen from pycountry and adapted. + +# raise ValueError("no need to run this, I already ran it.") + +# # languages +# d = json.load(open("iso639-3.json")) +# d["639-3"] = [x for x in d["639-3"] if "alpha_2" in x] +# for x in d["639-3"]: +# x["alpha_3b"] = x.pop("bibliographic", None) or x["alpha_3"] +# del x["scope"] +# del x["type"] +# with open("iso639-3.json", "w") as f: +# json.dump(d["639-3"], f, indent=2) + +# # countries +# d = json.load(open("iso3166-1.json"))["3166-1"] +# for x in d: +# x["alpha_2"] = x["alpha_2"].lower() +# x["alpha_3"] = x["alpha_3"].lower() +# with open("iso3166-1.json", "w") as f: +# json.dump(d, f, indent=2) diff --git a/generalresearch/logging.py b/generalresearch/logging.py index 9b72e0b..40f2174 100644 --- a/generalresearch/logging.py +++ b/generalresearch/logging.py @@ -1,6 +1,7 @@ import decimal import json from datetime import date +from typing import Any class ThlJsonEncoder(json.JSONEncoder): @@ -11,11 +12,11 @@ class ThlJsonEncoder(json.JSONEncoder): datetime/date to isoformat """ - def default(self, o): + def default(self, o: Any) -> Any: if isinstance(o, decimal.Decimal): return str(o) if isinstance(o, set): - return sorted(list(o)) + return sorted(o) if isinstance(o, date): return o.isoformat() return super().default(o) diff --git a/generalresearch/managers/cint/survey.py b/generalresearch/managers/cint/survey.py index 819ae3d..686a964 100644 --- a/generalresearch/managers/cint/survey.py +++ b/generalresearch/managers/cint/survey.py @@ -141,5 +141,5 @@ class CintSurveyManager(SurveyManager): if e.args[0] == 1062: existing_sns.add(sn) else: - raise e + raise self.update([surveys[sn] for sn in existing_sns]) diff --git a/generalresearch/managers/criteria.py b/generalresearch/managers/criteria.py index c13b8ac..b8ae6ac 100644 --- a/generalresearch/managers/criteria.py +++ b/generalresearch/managers/criteria.py @@ -9,20 +9,21 @@ from more_itertools import chunked from generalresearch.managers.base import SqlManager from generalresearch.models.thl.survey import MarketplaceCondition +DB_FIELDS = [ + "hash", + "question_id", + "logical_operator", + "values", + "value_type", + "negate", +] + class CriteriaManager(SqlManager, ABC): """ Using the terms "criteria" & "condition" interchangeably! """ - DB_FIELDS = [ - "hash", - "question_id", - "logical_operator", - "values", - "value_type", - "negate", - ] CONDITION_MODEL = None TABLE_NAME = "" @@ -59,7 +60,7 @@ class CriteriaManager(SqlManager, ABC): def update(self, conditions: Collection[MarketplaceCondition]) -> None: # Add any new hashes into the DB - this_hashes = set([condition.criterion_hash for condition in conditions]) + this_hashes = {condition.criterion_hash for condition in conditions} known_hashes = self.filter_exists(this_hashes) new_hashes = this_hashes - known_hashes @@ -95,7 +96,6 @@ class CriteriaManager(SqlManager, ABC): ) conn.commit() - @property def mysql_fields(self) -> str: return ", ".join([f"`{k}`" for k in self.DB_FIELDS]) diff --git a/generalresearch/managers/dynata/survey.py b/generalresearch/managers/dynata/survey.py index 7643dc4..21ec42e 100644 --- a/generalresearch/managers/dynata/survey.py +++ b/generalresearch/managers/dynata/survey.py @@ -13,6 +13,34 @@ from generalresearch.models.dynata.survey import DynataCondition, DynataSurvey logger = logging.getLogger() +SURVEY_FIELDS = [ + "survey_id", + "status", + "is_live", + "client_id", + "bid_loi", + "bid_ir", + "country_iso", + "language_iso", + "cpi", + "expected_count", + "project_id", + "group_id", + "calculation_type", + "days_in_field", + "order_number", + "requirements", + "allowed_devices", + "category_exclusions", + "project_exclusions", + "live_link", + "category_ids", + "filters", + "quotas", + "used_question_ids", + "created", +] + class DynataCriteriaManager(CriteriaManager): CONDITION_MODEL = DynataCondition @@ -20,33 +48,6 @@ class DynataCriteriaManager(CriteriaManager): class DynataSurveyManager(SurveyManager): - SURVEY_FIELDS = [ - "survey_id", - "status", - "is_live", - "client_id", - "bid_loi", - "bid_ir", - "country_iso", - "language_iso", - "cpi", - "expected_count", - "project_id", - "group_id", - "calculation_type", - "days_in_field", - "order_number", - "requirements", - "allowed_devices", - "category_exclusions", - "project_exclusions", - "live_link", - "category_ids", - "filters", - "quotas", - "used_question_ids", - "created", - ] def get_survey_library( self, @@ -107,7 +108,7 @@ class DynataSurveyManager(SurveyManager): conn: pymysql.Connection = self.sql_helper.make_connection() conn.autocommit(True) c = conn.cursor() - create_fields = ["id"] + self.SURVEY_FIELDS + ["last_updated"] + create_fields = ["id"] + SURVEY_FIELDS + ["last_updated"] fields_str = ", ".join([f"`{x}`" for x in create_fields]) values_str = ", ".join([f"%({x})s" for x in create_fields]) @@ -124,10 +125,10 @@ class DynataSurveyManager(SurveyManager): def update(self, surveys: list[DynataSurvey]) -> bool: now = datetime.now(tz=UTC) - update_fields = self.SURVEY_FIELDS + ["last_updated"] + update_fields = SURVEY_FIELDS + ["last_updated"] data = [survey.to_mysql() for survey in surveys] - survey_data = [[d[k] for k in self.SURVEY_FIELDS] + [now] for d in data] + survey_data = [[d[k] for k in SURVEY_FIELDS] + [now] for d in data] self.sql_helper.bulk_update("dynata_survey", update_fields, survey_data) return True @@ -154,5 +155,6 @@ class DynataSurveyManager(SurveyManager): if e.args[0] == 1062: existing_sns.add(sn) else: - raise e + raise + self.update([surveys[sn] for sn in existing_sns]) diff --git a/generalresearch/managers/events.py b/generalresearch/managers/events.py index efc8c0d..6779104 100644 --- a/generalresearch/managers/events.py +++ b/generalresearch/managers/events.py @@ -1,16 +1,16 @@ from __future__ import annotations -import logging import math import socket import threading import time from datetime import UTC, datetime, timedelta from decimal import Decimal -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Any from redis.client import PubSub, Redis +from generalresearch.incite.base import LOG from generalresearch.managers.base import RedisManager from generalresearch.models import Source from generalresearch.models.custom_types import UUIDStr @@ -216,16 +216,18 @@ class UserStatsManager(RedisManager): class TaskStatsManager(RedisManager): - task_stats = [ - "task_created_count_last_1h", - "task_created_count_last_24h", - "live_task_count", - "live_tasks_max_payout", - "TaskStatsManager:latest", - ] def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) + + self.task_stats = [ + "task_created_count_last_1h", + "task_created_count_last_24h", + "live_task_count", + "live_tasks_max_payout", + "TaskStatsManager:latest", + ] + self.SUM_HASH_LUA = self.redis_client.register_script(SUM_HASH_LUA_SCRIPT) self.MAX_HASH_LUA = self.redis_client.register_script(MAX_HASH_LUA_SCRIPT) @@ -435,8 +437,8 @@ class SessionStatsManager(RedisManager): pipe.hincrby(name, key, 1) pipe.hexpire(name, ttl, key, nx=True) # BP-specific tracker - pipe.hincrby(name + ":" + user.product_id, key, 1) - pipe.hexpire(name + ":" + user.product_id, ttl, key, nx=True) + pipe.hincrby(f"{name}:{user.product_id}", key, 1) + pipe.hexpire(f"{name}:{user.product_id}", ttl, key, nx=True) # We're not returning this, but keep the sums, so we can # calculate the avg @@ -444,8 +446,8 @@ class SessionStatsManager(RedisManager): value = round(session.elapsed.total_seconds()) pipe.hincrby(name, key, value) pipe.hexpire(name, ttl, key, nx=True) - pipe.hincrby(name + ":" + user.product_id, key, value) - pipe.hexpire(name + ":" + user.product_id, ttl, key, nx=True) + pipe.hincrby(f"{name}:{user.product_id}", key, value) + pipe.hexpire(f"{name}:{user.product_id}", ttl, key, nx=True) pipe.execute() @@ -476,31 +478,31 @@ class SessionStatsManager(RedisManager): pipe.hincrby(name, key, 1) pipe.hexpire(name, ttl, key, nx=True) # BP-specific tracker - pipe.hincrby(name + ":" + user.product_id, key, 1) - pipe.hexpire(name + ":" + user.product_id, ttl, key, nx=True) + pipe.hincrby(f"{name}:{user.product_id}", key, 1) + pipe.hexpire(f"{name}:{user.product_id}", ttl, key, nx=True) name = "sum_payouts_" + name_postfix amount = round(session.payout * 100) pipe.hincrby(name, key, amount) pipe.hexpire(name, ttl, key, nx=True) - pipe.hincrby(name + ":" + user.product_id, key, amount) - pipe.hexpire(name + ":" + user.product_id, ttl, key, nx=True) + pipe.hincrby(f"{name}:{user.product_id}", key, amount) + pipe.hexpire(f"{name}:{user.product_id}", ttl, key, nx=True) if session.user_payout: name = "sum_user_payouts_" + name_postfix amount = round(session.user_payout * 100) pipe.hincrby(name, key, amount) pipe.hexpire(name, ttl, key, nx=True) - pipe.hincrby(name + ":" + user.product_id, key, amount) - pipe.hexpire(name + ":" + user.product_id, ttl, key, nx=True) + pipe.hincrby(f"{name}:{user.product_id}", key, amount) + pipe.hexpire(f"{name}:{user.product_id}", ttl, key, nx=True) # We're not returning this, but keep the sums, so we can calculate the avg name = "session_complete_loi_sum_" + name_postfix value = round(session.elapsed.total_seconds()) pipe.hincrby(name, key, value) pipe.hexpire(name, ttl, key, nx=True) - pipe.hincrby(name + ":" + user.product_id, key, value) - pipe.hexpire(name + ":" + user.product_id, ttl, key, nx=True) + pipe.hincrby(f"{name}:{user.product_id}", key, value) + pipe.hexpire(f"{name}:{user.product_id}", ttl, key, nx=True) pipe.execute() @@ -563,6 +565,7 @@ class SessionStatsManager(RedisManager): res["session_avg_user_payout_last_24h"] = None res["session_complete_avg_loi_last_24h"] = None res["session_fail_avg_loi_last_24h"] = None + if res["session_completes_last_24h"]: res["session_avg_payout_last_24h"] = math.ceil( res["sum_payouts_last_24h"] / res["session_completes_last_24h"] @@ -630,7 +633,7 @@ class EventManager(StatsManager): def get_active_subscribers(self) -> set[UUIDStr]: res = self.redis_client.pubsub_channels(f"{self.cache_prefix}:event-channel:*") - product_ids = {x.rsplit(":", 1)[-1] for x in res} + product_ids = {str(x.rsplit(":", 1)[-1]) for x in res} return product_ids def stats_worker(self): @@ -638,7 +641,7 @@ class EventManager(StatsManager): try: self.stats_worker_task() except Exception as e: - logging.exception(e) + LOG.exception(e) finally: time.sleep(60) @@ -654,14 +657,14 @@ class EventManager(StatsManager): lock_key = f"{self.cache_prefix}:event-channel-lock" res = self.redis_client.set(lock_key, 1, ex=120, nx=True) if not res: - logging.debug("failed to acquire stats_worker_task lock") + LOG.debug("failed to acquire stats_worker_task lock") return - logging.info("Acquired stats_worker_task lock") + LOG.info("Acquired stats_worker_task lock") for product_id in self.get_active_subscribers(): if time.monotonic() - now > 120: - logging.exception("stats_worker_task is taking too long") + LOG.exception("stats_worker_task is taking too long") break channel = self.get_channel_name(product_id) msg = self.get_stats_message(product_id=product_id) @@ -680,7 +683,7 @@ class EventManager(StatsManager): return - def make_influx_point(self, channel: str, numsub: int): + def make_influx_point(self, channel: str, numsub: int) -> dict[str, Any]: return { "measurement": "redis_pubsub_subscribers", "tags": {"hostname": socket.gethostname(), "channel": channel}, diff --git a/generalresearch/managers/gr/authentication.py b/generalresearch/managers/gr/authentication.py index 721895e..851b88a 100644 --- a/generalresearch/managers/gr/authentication.py +++ b/generalresearch/managers/gr/authentication.py @@ -160,7 +160,7 @@ class GRUserManager(PostgresManagerWithRedis): res = thl_pg_config.execute_sql_query( query=""" - SELECT bp.id + SELECT bp.id::uuid as uuid FROM userprofile_brokerageproduct AS bp WHERE bp.business_id = ANY(%s) """, @@ -234,10 +234,10 @@ class GRTokenManager(PostgresManager): res = c.fetchall() if len(res) == 0: - raise Exception(f"No GRUser with token of '{api_key}'") + raise ValueError(f"No GRUser with token of '{api_key}'") if len(res) > 1: - raise Exception(f"Too many GRUsers found with token of '{api_key}'") + raise ValueError(f"Too many GRUsers found with token of '{api_key}'") item = res[0] diff --git a/generalresearch/managers/innovate/survey.py b/generalresearch/managers/innovate/survey.py index f6d00a8..a4e36c1 100644 --- a/generalresearch/managers/innovate/survey.py +++ b/generalresearch/managers/innovate/survey.py @@ -16,6 +16,44 @@ from generalresearch.models.innovate.survey import ( logger = logging.getLogger() +SURVEY_FIELDS = [ + "survey_id", + "status", + "country_iso", + "language_iso", + "cpi", + "buyer_id", + "job_id", + "survey_name", + "desired_count", + "remaining_count", + "supplier_completes_achieved", + "global_completes", + "global_starts", + "global_median_loi", + "global_conversion", + "bid_loi", + "bid_ir", + "allowed_devices", + "entry_link", + "category", + "requires_pii", + "excluded_surveys", + "duplicate_check_level", + "exclude_pids", + "include_pids", + "is_revenue_sharing", + "group_type", + "off_hour_traffic", + "qualifications", + "quotas", + "used_question_ids", + "is_live", + "modified_api", + "created_api", + "expected_end_date", +] + class InnovateCriteriaManager(CriteriaManager): CONDITION_MODEL = InnovateCondition @@ -23,43 +61,6 @@ class InnovateCriteriaManager(CriteriaManager): class InnovateSurveyManager(SurveyManager): - SURVEY_FIELDS = [ - "survey_id", - "status", - "country_iso", - "language_iso", - "cpi", - "buyer_id", - "job_id", - "survey_name", - "desired_count", - "remaining_count", - "supplier_completes_achieved", - "global_completes", - "global_starts", - "global_median_loi", - "global_conversion", - "bid_loi", - "bid_ir", - "allowed_devices", - "entry_link", - "category", - "requires_pii", - "excluded_surveys", - "duplicate_check_level", - "exclude_pids", - "include_pids", - "is_revenue_sharing", - "group_type", - "off_hour_traffic", - "qualifications", - "quotas", - "used_question_ids", - "is_live", - "modified_api", - "created_api", - "expected_end_date", - ] def get_survey_library( self, @@ -104,7 +105,7 @@ class InnovateSurveyManager(SurveyManager): assert filters, "Must set at least 1 filter" filter_str = " AND ".join(filters) filter_str = "WHERE " + filter_str if filter_str else "" - fields = set(self.SURVEY_FIELDS) | {"created", "updated"} + fields = set(SURVEY_FIELDS) | {"created", "updated"} if exclude_fields: fields -= exclude_fields @@ -126,7 +127,7 @@ class InnovateSurveyManager(SurveyManager): conn: pymysql.Connection = self.sql_helper.make_connection() conn.autocommit(True) c = conn.cursor() - create_fields = self.SURVEY_FIELDS + ["created", "updated"] + create_fields = SURVEY_FIELDS + ["created", "updated"] fields_str = ", ".join([f"`{x}`" for x in create_fields]) values_str = ", ".join([f"%({x})s" for x in create_fields]) @@ -143,10 +144,10 @@ class InnovateSurveyManager(SurveyManager): def update(self, surveys: list[InnovateSurvey]) -> bool: now = datetime.now(tz=UTC) - update_fields = self.SURVEY_FIELDS + ["updated"] + update_fields = SURVEY_FIELDS + ["updated"] data = [survey.to_mysql() for survey in surveys] - survey_data = [[d[k] for k in self.SURVEY_FIELDS] + [now] for d in data] + survey_data = [[d[k] for k in SURVEY_FIELDS] + [now] for d in data] self.sql_helper.bulk_update( table_name="innovate_survey", field_names=update_fields, diff --git a/generalresearch/managers/morning/survey.py b/generalresearch/managers/morning/survey.py index 2488cd8..b43cd71 100644 --- a/generalresearch/managers/morning/survey.py +++ b/generalresearch/managers/morning/survey.py @@ -14,6 +14,50 @@ from generalresearch.models.morning.survey import MorningBid, MorningCondition logger = logging.getLogger() +STAT_FIELDS = [ + "obs_median_loi", + "qualified_conversion", + "num_available", + "num_completes", + "num_failures", + "num_in_progress", + "num_over_quotas", + "num_qualified", + "num_quality_terminations", + "num_timeouts", +] +STAT_EXTENDED_FIELDS = ["system_conversion", "num_entrants", "num_screenouts"] +BID_FIELDS = ( + [ + "id", + "status", + "country_iso", + "language_isos", + "buyer_account_id", + "buyer_id", + "name", + "supplier_exclusive", + "survey_type", + "timeout", + "topic_id", + "bid_loi", + "exclusions", + "used_question_ids", + "expected_end", + "created_api", + "is_live", + ] + + STAT_FIELDS + + STAT_EXTENDED_FIELDS +) +QUOTA_FIELDS = [ + "id", + "cpi", + "condition_hashes", +] + STAT_FIELDS +BID_DB_SOURCE = "`thl-morning`.`morning_surveybid`" +QUOTA_DB_SOURCE = "`thl-morning`.`morning_surveyquota`" + class MorningCriteriaManager(CriteriaManager): CONDITION_MODEL = MorningCondition @@ -21,49 +65,6 @@ class MorningCriteriaManager(CriteriaManager): class MorningSurveyManager(SurveyManager): - STAT_FIELDS = [ - "obs_median_loi", - "qualified_conversion", - "num_available", - "num_completes", - "num_failures", - "num_in_progress", - "num_over_quotas", - "num_qualified", - "num_quality_terminations", - "num_timeouts", - ] - STAT_EXTENDED_FIELDS = ["system_conversion", "num_entrants", "num_screenouts"] - BID_FIELDS = ( - [ - "id", - "status", - "country_iso", - "language_isos", - "buyer_account_id", - "buyer_id", - "name", - "supplier_exclusive", - "survey_type", - "timeout", - "topic_id", - "bid_loi", - "exclusions", - "used_question_ids", - "expected_end", - "created_api", - "is_live", - ] - + STAT_FIELDS - + STAT_EXTENDED_FIELDS - ) - QUOTA_FIELDS = [ - "id", - "cpi", - "condition_hashes", - ] + STAT_FIELDS - BID_DB_SOURCE = "`thl-morning`.`morning_surveybid`" - QUOTA_DB_SOURCE = "`thl-morning`.`morning_surveyquota`" def get_survey_library( self, diff --git a/generalresearch/managers/network/mtr.py b/generalresearch/managers/network/mtr.py index 54d74b7..179b8a9 100644 --- a/generalresearch/managers/network/mtr.py +++ b/generalresearch/managers/network/mtr.py @@ -41,8 +41,9 @@ class MTRRunManager(PostgresManager): c.execute(query, params) if params_hops: c.executemany(query_hops, params_hops) + else: - with self.pg_config.make_connection() as conn, conn.cursor() as c: - c.execute(query, params) + with self.pg_config.make_connection() as conn, conn.cursor() as _c: + _c.execute(query, params) if params_hops: - c.executemany(query_hops, params_hops) + _c.executemany(query_hops, params_hops) diff --git a/generalresearch/managers/network/nmap.py b/generalresearch/managers/network/nmap.py index 84d13ad..574bce1 100644 --- a/generalresearch/managers/network/nmap.py +++ b/generalresearch/managers/network/nmap.py @@ -50,7 +50,7 @@ class NmapRunManager(PostgresManager): if nmap_run.ports: c.executemany(query_ports, params_ports) else: - with self.pg_config.make_connection() as conn, conn.cursor() as c: + with self.pg_config.make_connection() as conn, conn.cursor(): c.execute(query, params) if nmap_run.ports: c.executemany(query_ports, params_ports) diff --git a/generalresearch/managers/network/rdns.py b/generalresearch/managers/network/rdns.py index c8ce913..95a1381 100644 --- a/generalresearch/managers/network/rdns.py +++ b/generalresearch/managers/network/rdns.py @@ -27,6 +27,7 @@ class RDNSRunManager(PostgresManager): params = run.model_dump_postgres() if c: c.execute(query, params) + else: - with self.pg_config.make_connection() as conn, conn.cursor() as c: - c.execute(query, params) + with self.pg_config.make_connection() as conn, conn.cursor() as _c: + _c.execute(query, params) diff --git a/generalresearch/managers/precision/survey.py b/generalresearch/managers/precision/survey.py index c13dca8..cc28287 100644 --- a/generalresearch/managers/precision/survey.py +++ b/generalresearch/managers/precision/survey.py @@ -16,6 +16,30 @@ from generalresearch.models.precision.survey import ( logger = logging.getLogger() +SURVEY_FIELDS = [ + # 'country_iso', 'language_iso', # these come from join table + "survey_id", + "is_live", + "status", + "cpi", + "group_id", + "name", + "survey_guid", + "buyer_id", + "category_id", + "bid_loi", + "bid_ir", + "global_conversion", + "desired_count", + "achieved_count", + "allowed_devices", + "entry_link", + "excluded_surveys", + "quotas", + "used_question_ids", + "expected_end_date", +] + class PrecisionCriteriaManager(CriteriaManager): CONDITION_MODEL = PrecisionCondition @@ -23,29 +47,6 @@ class PrecisionCriteriaManager(CriteriaManager): class PrecisionSurveyManager(SurveyManager): - SURVEY_FIELDS = [ - # 'country_iso', 'language_iso', # these come from join table - "survey_id", - "is_live", - "status", - "cpi", - "group_id", - "name", - "survey_guid", - "buyer_id", - "category_id", - "bid_loi", - "bid_ir", - "global_conversion", - "desired_count", - "achieved_count", - "allowed_devices", - "entry_link", - "excluded_surveys", - "quotas", - "used_question_ids", - "expected_end_date", - ] def get_survey_library( self, @@ -109,7 +110,7 @@ class PrecisionSurveyManager(SurveyManager): conn: pymysql.Connection = self.sql_helper.make_connection() conn.autocommit(False) c = conn.cursor() - create_fields = self.SURVEY_FIELDS + ["created", "updated"] + create_fields = SURVEY_FIELDS + ["created", "updated"] fields_str = ", ".join([f"`{x}`" for x in create_fields]) values_str = ", ".join([f"%({x})s" for x in create_fields]) @@ -241,5 +242,5 @@ class PrecisionSurveyManager(SurveyManager): if e.args[0] == 1062: existing_sns.add(sn) else: - raise e + raise self.update([surveys[sn] for sn in existing_sns]) diff --git a/generalresearch/managers/prodege/survey.py b/generalresearch/managers/prodege/survey.py index 983572c..ef98d7b 100644 --- a/generalresearch/managers/prodege/survey.py +++ b/generalresearch/managers/prodege/survey.py @@ -9,6 +9,31 @@ from generalresearch.managers.criteria import CriteriaManager from generalresearch.managers.survey import SurveyManager from generalresearch.models.prodege.survey import ProdegeCondition, ProdegeSurvey +SURVEY_FIELDS = [ + "survey_id", + "survey_name", + "status", + "country_iso", + "language_iso", + "cpi", + "desired_count", + "remaining_count", + "achieved_completes", + "bid_loi", + "bid_ir", + "actual_loi", + "actual_ir", + "conversion_rate", + "entrance_url", + "max_clicks_settings", + "past_participation", + "include_psids", + "exclude_psids", + "quotas", + "used_question_ids", + "is_live", +] + class ProdegeCriteriaManager(CriteriaManager): CONDITION_MODEL = ProdegeCondition @@ -16,30 +41,6 @@ class ProdegeCriteriaManager(CriteriaManager): class ProdegeSurveyManager(SurveyManager): - SURVEY_FIELDS = [ - "survey_id", - "survey_name", - "status", - "country_iso", - "language_iso", - "cpi", - "desired_count", - "remaining_count", - "achieved_completes", - "bid_loi", - "bid_ir", - "actual_loi", - "actual_ir", - "conversion_rate", - "entrance_url", - "max_clicks_settings", - "past_participation", - "include_psids", - "exclude_psids", - "quotas", - "used_question_ids", - "is_live", - ] def get_survey_library( self, @@ -98,7 +99,7 @@ class ProdegeSurveyManager(SurveyManager): conn: pymysql.Connection = self.sql_helper.make_connection() conn.autocommit(True) c = conn.cursor() - create_fields = self.SURVEY_FIELDS + ["created", "updated"] + create_fields = SURVEY_FIELDS + ["created", "updated"] fields_str = ", ".join([f"`{x}`" for x in create_fields]) values_str = ", ".join([f"%({x})s" for x in create_fields]) diff --git a/generalresearch/managers/repdata/survey.py b/generalresearch/managers/repdata/survey.py index fe6f621..ab66374 100644 --- a/generalresearch/managers/repdata/survey.py +++ b/generalresearch/managers/repdata/survey.py @@ -15,6 +15,40 @@ from generalresearch.models.repdata.survey import ( RepDataSurveyHashed, ) +SURVEY_FIELDS = [ + "survey_id", + "survey_uuid", + "survey_name", + "project_uuid", + "survey_status", + "country_iso", + "language_iso", + "estimated_loi", + "estimated_ir", + "collects_pii", + "allowed_devices", +] +STREAM_FIELDS = [ + "stream_id", + "stream_uuid", + "stream_name", + "stream_status", + "calculation_type", + "qualification_hashes", + "hashed_quotas", + "expected_count", + "cpi", + "days_in_field", + "actual_ir", + "actual_loi", + "actual_conversion", + "actual_complete_count", + "actual_count", + "used_question_ids", + "survey_id", + "remaining_count", +] + class RepDataCriteriaManager(CriteriaManager): CONDITION_MODEL = RepDataCondition @@ -22,39 +56,6 @@ class RepDataCriteriaManager(CriteriaManager): class RepDataSurveyManager(SurveyManager): - SURVEY_FIELDS = [ - "survey_id", - "survey_uuid", - "survey_name", - "project_uuid", - "survey_status", - "country_iso", - "language_iso", - "estimated_loi", - "estimated_ir", - "collects_pii", - "allowed_devices", - ] - STREAM_FIELDS = [ - "stream_id", - "stream_uuid", - "stream_name", - "stream_status", - "calculation_type", - "qualification_hashes", - "hashed_quotas", - "expected_count", - "cpi", - "days_in_field", - "actual_ir", - "actual_loi", - "actual_conversion", - "actual_complete_count", - "actual_count", - "used_question_ids", - "survey_id", - "remaining_count", - ] def get_survey_library( self, @@ -127,7 +128,7 @@ class RepDataSurveyManager(SurveyManager): conn: pymysql.Connection = self.sql_helper.make_connection() conn.autocommit(True) c = conn.cursor() - create_fields = self.SURVEY_FIELDS + ["created", "last_updated"] + create_fields = SURVEY_FIELDS + ["created", "last_updated"] fields_str = ", ".join([f"`{x}`" for x in create_fields]) values_str = ", ".join([f"%({x})s" for x in create_fields]) @@ -141,10 +142,10 @@ class RepDataSurveyManager(SurveyManager): args=survey_data, ) - fields_str = ", ".join([f"`{x}`" for x in self.STREAM_FIELDS]) - values_str = ", ".join([f"%({x})s" for x in self.STREAM_FIELDS]) + fields_str = ", ".join([f"`{x}`" for x in STREAM_FIELDS]) + values_str = ", ".join([f"%({x})s" for x in STREAM_FIELDS]) stream_data = [ - {k: v for k, v in stream.items() if k in self.STREAM_FIELDS} + {k: v for k, v in stream.items() if k in STREAM_FIELDS} for stream in d["streams"] ] for sd in stream_data: @@ -161,10 +162,10 @@ class RepDataSurveyManager(SurveyManager): def update(self, surveys: list[RepDataSurveyHashed]) -> bool: now = datetime.now(tz=UTC) - update_fields = self.SURVEY_FIELDS + ["last_updated"] + update_fields = SURVEY_FIELDS + ["last_updated"] data = [survey.to_mysql() for survey in surveys] - survey_data = [[d[k] for k in self.SURVEY_FIELDS] + [now] for d in data] + survey_data = [[d[k] for k in SURVEY_FIELDS] + [now] for d in data] self.sql_helper.bulk_update( table_name="repdata_survey", field_names=update_fields, @@ -175,7 +176,7 @@ class RepDataSurveyManager(SurveyManager): for d in data: for stream in d["streams"]: stream["survey_id"] = d["survey_id"] - stream_data.append([stream[k] for k in self.STREAM_FIELDS]) + stream_data.append([stream[k] for k in STREAM_FIELDS]) self.sql_helper.bulk_update( table_name="repdata_surveystream", diff --git a/generalresearch/managers/sago/survey.py b/generalresearch/managers/sago/survey.py index a13fbce..9228528 100644 --- a/generalresearch/managers/sago/survey.py +++ b/generalresearch/managers/sago/survey.py @@ -13,6 +13,31 @@ from generalresearch.models.sago.survey import SagoCondition, SagoSurvey logger = logging.getLogger() +SURVEY_FIELDS = [ + "survey_id", + "is_live", + "status", + "country_iso", + "language_iso", + "cpi", + "buyer_id", + "account_id", + "study_type_id", + "industry_id", + "allowed_devices", + "collects_pii", + "bid_loi", + "bid_ir", + "live_link", + "survey_exclusions", + "ip_exclusions", + "remaining_count", + "qualifications", + "quotas", + "used_question_ids", + "modified_api", +] + class SagoCriteriaManager(CriteriaManager): CONDITION_MODEL = SagoCondition @@ -20,30 +45,6 @@ class SagoCriteriaManager(CriteriaManager): class SagoSurveyManager(SurveyManager): - SURVEY_FIELDS = [ - "survey_id", - "is_live", - "status", - "country_iso", - "language_iso", - "cpi", - "buyer_id", - "account_id", - "study_type_id", - "industry_id", - "allowed_devices", - "collects_pii", - "bid_loi", - "bid_ir", - "live_link", - "survey_exclusions", - "ip_exclusions", - "remaining_count", - "qualifications", - "quotas", - "used_question_ids", - "modified_api", - ] def get_survey_library( self, @@ -85,7 +86,7 @@ class SagoSurveyManager(SurveyManager): assert filters, "Must set at least 1 filter" filter_str = " AND ".join(filters) filter_str = "WHERE " + filter_str if filter_str else "" - fields = set(self.SURVEY_FIELDS) | {"created", "updated"} + fields = set(SURVEY_FIELDS) | {"created", "updated"} if exclude_fields: fields -= exclude_fields fields_str = ", ".join([f"`{v}`" for v in fields]) @@ -106,7 +107,7 @@ class SagoSurveyManager(SurveyManager): conn: pymysql.Connection = self.sql_helper.make_connection() conn.autocommit(True) c = conn.cursor() - create_fields = self.SURVEY_FIELDS + ["created", "updated"] + create_fields = SURVEY_FIELDS + ["created", "updated"] fields_str = ", ".join([f"`{x}`" for x in create_fields]) values_str = ", ".join([f"%({x})s" for x in create_fields]) @@ -123,10 +124,10 @@ class SagoSurveyManager(SurveyManager): def update(self, surveys: list[SagoSurvey]) -> bool: now = datetime.now(tz=UTC) - update_fields = self.SURVEY_FIELDS + ["updated"] + update_fields = SURVEY_FIELDS + ["updated"] data = [survey.to_mysql() for survey in surveys] - survey_data = [[d[k] for k in self.SURVEY_FIELDS] + [now] for d in data] + survey_data = [[d[k] for k in SURVEY_FIELDS] + [now] for d in data] self.sql_helper.bulk_update("sago_survey", update_fields, survey_data) return True @@ -156,8 +157,8 @@ class SagoSurveyManager(SurveyManager): return True def create_or_update(self, surveys: list[SagoSurvey]) -> None: - surveys = {s.survey_id: s for s in surveys} - sns = set(surveys.keys()) + _surveys = {s.survey_id: s for s in surveys} + sns = set(_surveys.keys()) existing_sns = { x["survey_id"] for x in self.sql_helper.execute_sql_query( @@ -171,7 +172,7 @@ class SagoSurveyManager(SurveyManager): } create_sns = sns - existing_sns for sn in create_sns: - survey = surveys[sn] + survey = _surveys[sn] try: self.create(survey) except IntegrityError as e: @@ -181,4 +182,4 @@ class SagoSurveyManager(SurveyManager): else: raise - self.update([surveys[sn] for sn in existing_sns]) + self.update([_surveys[sn] for sn in existing_sns]) diff --git a/generalresearch/managers/spectrum/survey.py b/generalresearch/managers/spectrum/survey.py index 987ce7b..5059716 100644 --- a/generalresearch/managers/spectrum/survey.py +++ b/generalresearch/managers/spectrum/survey.py @@ -16,6 +16,37 @@ from generalresearch.models.spectrum.survey import ( logger = logging.getLogger() +SURVEY_FIELDS = [ + "survey_id", + "survey_name", + "status", + "country_iso", + "language_iso", + "cpi", + "field_end_date", + "category_code", + "calculation_type", + "requires_pii", + "buyer_id", + "survey_exclusions", + "exclusion_period", + "bid_loi", + "bid_ir", + "last_block_loi", + "last_block_ir", + "overall_ir", + "overall_loi", + "project_last_complete_date", + "include_psids", + "exclude_psids", + "qualifications", + "quotas", + "used_question_ids", + "is_live", + "modified_api", + "created_api", +] + class SpectrumCriteriaManager(CriteriaManager): CONDITION_MODEL = SpectrumCondition @@ -23,36 +54,6 @@ class SpectrumCriteriaManager(CriteriaManager): class SpectrumSurveyManager(SurveyManager): - SURVEY_FIELDS = [ - "survey_id", - "survey_name", - "status", - "country_iso", - "language_iso", - "cpi", - "field_end_date", - "category_code", - "calculation_type", - "requires_pii", - "buyer_id", - "survey_exclusions", - "exclusion_period", - "bid_loi", - "bid_ir", - "last_block_loi", - "last_block_ir", - "overall_ir", - "overall_loi", - "project_last_complete_date", - "include_psids", - "exclude_psids", - "qualifications", - "quotas", - "used_question_ids", - "is_live", - "modified_api", - "created_api", - ] def get_survey_library( self, @@ -115,7 +116,7 @@ class SpectrumSurveyManager(SurveyManager): conn: pymysql.Connection = self.sql_helper.make_connection() conn.autocommit(True) c = conn.cursor() - create_fields = self.SURVEY_FIELDS + ["updated"] + create_fields = SURVEY_FIELDS + ["updated"] fields_str = ", ".join([f"`{x}`" for x in create_fields]) values_str = ", ".join([f"%({x})s" for x in create_fields]) diff --git a/generalresearch/managers/thl/task_adjustment.py b/generalresearch/managers/thl/task_adjustment.py index 60bade7..2d89334 100644 --- a/generalresearch/managers/thl/task_adjustment.py +++ b/generalresearch/managers/thl/task_adjustment.py @@ -24,6 +24,9 @@ from generalresearch.models.thl.session import ( ) from generalresearch.models.thl.task_adjustment import TaskAdjustmentEvent +logging.basicConfig() +logger = logging.getLogger(__name__) + class TaskAdjustmentManager(PostgresManager): @@ -129,8 +132,9 @@ class TaskAdjustmentManager(PostgresManager): user.prefetch_product(self.pg_config) if adjusted_status == WallAdjustedStatus.ADJUSTED_TO_FAIL: + assert wall.cpi amount_usd = wall.cpi * -1 - adjusted_cpi = 0 + adjusted_cpi = Decimal(0) elif adjusted_status == WallAdjustedStatus.ADJUSTED_TO_COMPLETE: amount_usd = wall.cpi adjusted_cpi = wall.cpi @@ -169,7 +173,7 @@ class TaskAdjustmentManager(PostgresManager): new_adjusted_cpi=new_adjusted_cpi, ) except AssertionError as e: - logging.warning(e) + logger.warning(e) return event = TaskAdjustmentEvent( diff --git a/generalresearch/managers/thl/user_manager/mysql_user_manager.py b/generalresearch/managers/thl/user_manager/mysql_user_manager.py index e0a7548..af65d65 100644 --- a/generalresearch/managers/thl/user_manager/mysql_user_manager.py +++ b/generalresearch/managers/thl/user_manager/mysql_user_manager.py @@ -172,7 +172,7 @@ class MysqlUserManager: return user - @lru_cache(maxsize=5000) + @lru_cache(maxsize=5_000) def product_id_exists(self, product_id: str): # 'id' is the primary key, there can only be 0 or 1 query = """ diff --git a/generalresearch/managers/thl/userhealth.py b/generalresearch/managers/thl/userhealth.py index 26f08b4..f28fe0c 100644 --- a/generalresearch/managers/thl/userhealth.py +++ b/generalresearch/managers/thl/userhealth.py @@ -374,10 +374,10 @@ class AuditLogManager(PostgresManager): ) if len(res) == 0: - raise Exception(f"No AuditLog with id of '{auditlog_id}'") + raise ValueError(f"No AuditLog with id of '{auditlog_id}'") if len(res) > 1: - raise Exception(f"Too many AuditLog found with id of '{auditlog_id}'") + raise ValueError(f"Too many AuditLog found with id of '{auditlog_id}'") return AuditLog.from_mysql(res[0]) diff --git a/generalresearch/managers/thl/wall.py b/generalresearch/managers/thl/wall.py index ac9fb62..774db9e 100644 --- a/generalresearch/managers/thl/wall.py +++ b/generalresearch/managers/thl/wall.py @@ -579,7 +579,7 @@ class WallCacheManager(PostgresManagerWithRedis): # b as second element and a as third element" attempts = sorted(attempts, key=lambda x: x.started) json_res = [attempt.model_dump_json() for attempt in attempts] - res = self.redis_client.lpush(redis_key, *json_res) + _ = self.redis_client.lpush(redis_key, *json_res) self.redis_client.expire(redis_key, time=60 * 60 * 24) # So this doesn't grow forever, keep only the most recent 5k diff --git a/generalresearch/models/gr/business.py b/generalresearch/models/gr/business.py index 064c200..11a5770 100644 --- a/generalresearch/models/gr/business.py +++ b/generalresearch/models/gr/business.py @@ -38,6 +38,9 @@ from generalresearch.redis_helper import RedisConfig from generalresearch.utils.aggregation import group_by_year from generalresearch.utils.enum import ReprEnumMeta +logging.basicConfig() +logger = logging.getLogger(__name__) + if TYPE_CHECKING: from generalresearch.incite.base import GRLDatasets from generalresearch.incite.mergers.foundations.enriched_session import ( @@ -324,7 +327,7 @@ class Business(BaseModel): for product_uuid in self.product_uuids: if product_uuid not in bp_account: refresh = True - logging.exception( + logger.exception( f"Business {self.uuid} does not have a BP Wallet Account for Product {product_uuid}. Creating..." ) product = product_lookup[product_uuid] diff --git a/generalresearch/models/legacy/bucket.py b/generalresearch/models/legacy/bucket.py index 3705f4f..2650b90 100644 --- a/generalresearch/models/legacy/bucket.py +++ b/generalresearch/models/legacy/bucket.py @@ -441,6 +441,12 @@ class DurationSummary(StatisticalSummary): @classmethod def from_bucket(cls, bucket: Bucket) -> DurationSummary: + assert bucket.loi_min + assert bucket.loi_max + assert bucket.loi_q1 + assert bucket.loi_q2 + assert bucket.loi_q3 + return cls( min=bucket.loi_min.total_seconds(), max=bucket.loi_max.total_seconds(), diff --git a/generalresearch/pg_helper.py b/generalresearch/pg_helper.py index a397247..b1ac556 100644 --- a/generalresearch/pg_helper.py +++ b/generalresearch/pg_helper.py @@ -3,6 +3,7 @@ from __future__ import annotations from datetime import UTC import psycopg +from psycopg.abc import Query from psycopg.rows import RowFactory, dict_row from psycopg.types.datetime import TimestampLoader from psycopg.types.net import InetLoader @@ -74,7 +75,9 @@ class PostgresConfig: self.row_factory = row_factory @property - def db(self): + def db(self) -> str: + assert self.dsn + assert self.dsn.path return self.dsn.path[1:] def make_connection(self) -> psycopg.Connection: @@ -98,15 +101,15 @@ class PostgresConfig: conn.adapters.register_loader("inet", InetHostLoader) return conn - def execute_sql_query(self, query, params=None): + def execute_sql_query(self, query: Query, params=None): # This is only intended for SELECT queries - assert "SELECT" in query.upper(), "Supports SELECTs only" + assert "SELECT" in str(query).upper(), "Supports SELECTs only" with self.make_connection() as conn, conn.cursor() as c: c.execute(query=query, params=params) return c.fetchall() - def execute_write(self, query, params=None) -> int: + def execute_write(self, query: Query, params=None) -> int: cmd = query.lstrip().upper() assert cmd.startswith( ("INSERT", "UPDATE", "DELETE") diff --git a/generalresearch/sql_helper.py b/generalresearch/sql_helper.py index ae2b8d8..ef53c3b 100644 --- a/generalresearch/sql_helper.py +++ b/generalresearch/sql_helper.py @@ -14,6 +14,9 @@ ListOrTupleOfListOrTuple = ( DataBaseDsn = MySQLDsn | MariaDBDsn | PostgresDsn | None +logging.basicConfig() +logger = logging.getLogger(__name__) + class MultipleObjectsReturned(Exception): pass @@ -131,7 +134,7 @@ class SqlHelper(SqlConnector): ) -> list[dict[str, Any]]: for param in params if params else []: if isinstance(param, (tuple, list, set)) and len(param) == 0: - logging.warning("param is empty. not executing query") + logger.warning("param is empty. not executing query") return [] connection = self.make_connection() c = connection.cursor() diff --git a/tests/incite/collections/test_df_collection_item_thl_web.py b/tests/incite/collections/test_df_collection_item_thl_web.py index edf90f7..5f9a3f6 100644 --- a/tests/incite/collections/test_df_collection_item_thl_web.py +++ b/tests/incite/collections/test_df_collection_item_thl_web.py @@ -250,7 +250,7 @@ class TestDFCollectionItemMethod: # Unlike .from_mysql_ledger(), .from_mysql_standard() will return # back and empty df with the correct columns in place delete_df_collection(coll=df_collection) - df = item.from_mysql() + df = item.from_db() if df_collection.data_type == DFCollectionType.LEDGER: assert df is None else: @@ -260,7 +260,7 @@ class TestDFCollectionItemMethod: incite_item_factory(user=u1, item=item) - df = item.from_mysql() + df = item.from_db() assert isinstance(df, pd.DataFrame) assert not df.empty assert set(df.columns) == set(df_collection._schema.columns.keys()) @@ -401,7 +401,7 @@ class TestDFCollectionItemMethod: # Load up the data that we'll be using for various to_archive # methods. - df = item.from_mysql() + df = item.from_db() ddf = dd.from_pandas(df, npartitions=1) # (1) Write the basic archive, the issue is that because it's @@ -444,7 +444,7 @@ class TestDFCollectionItemMethod: # Load up the data that we'll be using for various to_archive # methods. Will always be empty pd.DataFrames for now... - df = item.from_mysql() + df = item.from_db() ddf = dd.from_pandas(df, npartitions=1) # (1) Confirm a missing ddf (shouldn't bc of type hint) should @@ -876,7 +876,7 @@ class TestDFCollectionItemFunctionalTest: # Load up the data that we'll be using for various to_archive # methods. Will always be empty pd.DataFrames for now... - df = item.from_mysql() + df = item.from_db() ddf = dd.from_pandas(df, npartitions=1) assert isinstance(ddf, dd.DataFrame) @@ -946,7 +946,7 @@ class TestDFCollectionItemFunctionalTest: for item in df_collection.items: assert not item.has_empty() - df: pd.DataFrame = item.from_mysql() + df: pd.DataFrame = item.from_db() # We do this check b/c the Ledger returns back None and # I don't want it to fail when we go to make a ddf diff --git a/tests/incite/collections/test_df_collection_thl_marketplaces.py b/tests/incite/collections/test_df_collection_thl_marketplaces.py index b4b5b00..6a7e5c9 100644 --- a/tests/incite/collections/test_df_collection_thl_marketplaces.py +++ b/tests/incite/collections/test_df_collection_thl_marketplaces.py @@ -46,7 +46,7 @@ class TestDFCollection_thl_marketplaces: assert isinstance(data_type, DFCollectionType) # (1) Can't be totally empty, needs a path... - with pytest.raises(expected_exception=Exception): + with pytest.raises(expected_exception=ValueError): instance = df_coll() # (2) Confirm it only needs the archive_path diff --git a/tests/managers/thl/test_contest/test_milestone.py b/tests/managers/thl/test_contest/test_milestone.py index aa738e1..e29ba4c 100644 --- a/tests/managers/thl/test_contest/test_milestone.py +++ b/tests/managers/thl/test_contest/test_milestone.py @@ -242,7 +242,7 @@ class TestMilestoneContestUserViews: def test_list_user_eligible_country( self, user_with_wallet: User, - raffle_contest_factory: Callable[..., Contest], + raffle_contest_factory: Callable[..., RaffleContest], thl_ledger_manager: ThlLedgerManager, contest_manager: ContestManager, ): -- cgit v1.2.3 From 89ed44f466dc9a93d6f85931fb6eea0e9cbd27f6 Mon Sep 17 00:00:00 2001 From: Max Nanis Date: Tue, 1 Sep 2026 10:43:29 -0700 Subject: working through circular imports, removing __init__ loaded Types / Definitions --- generalresearch/grliq/managers/event_plotter.py | 5 +- generalresearch/grliq/managers/forensic_events.py | 8 +- generalresearch/grliq/managers/forensic_results.py | 8 +- generalresearch/grliq/managers/forensic_summary.py | 16 +-- .../mergers/foundations/enriched_task_adjust.py | 2 +- .../schemas/mergers/foundations/enriched_wall.py | 2 +- .../incite/schemas/mergers/ym_wall_summary.py | 2 +- generalresearch/incite/schemas/thl_web.py | 2 +- generalresearch/managers/cint/user_pid.py | 2 +- generalresearch/managers/dynata/user_pid.py | 2 +- generalresearch/managers/events.py | 2 +- generalresearch/managers/innovate/user_pid.py | 2 +- generalresearch/managers/marketplace/user_pid.py | 2 +- generalresearch/managers/morning/user_pid.py | 2 +- generalresearch/managers/network/label.py | 6 +- generalresearch/managers/precision/user_pid.py | 2 +- generalresearch/managers/prodege/user_pid.py | 2 +- generalresearch/managers/repdata/user_pid.py | 2 +- generalresearch/managers/sago/user_pid.py | 2 +- generalresearch/managers/spectrum/user_pid.py | 2 +- generalresearch/managers/thl/buyer.py | 2 +- generalresearch/managers/thl/cashout_method.py | 2 +- generalresearch/managers/thl/contest_manager.py | 4 +- .../managers/thl/ledger_manager/ledger.py | 2 +- .../managers/thl/ledger_manager/thl_ledger.py | 2 +- generalresearch/managers/thl/payout.py | 2 +- generalresearch/managers/thl/session.py | 4 +- generalresearch/managers/thl/survey.py | 2 +- generalresearch/managers/thl/task_adjustment.py | 2 +- generalresearch/managers/thl/wall.py | 4 +- generalresearch/managers/thl/wallet/__init__.py | 2 +- generalresearch/managers/utils.py | 16 +++ generalresearch/models/cint/question.py | 2 +- generalresearch/models/cint/survey.py | 2 +- generalresearch/models/custom_types.py | 2 +- generalresearch/models/definitions.py | 114 +++++++++++++++++++++ generalresearch/models/device.py | 2 +- generalresearch/models/dynata/question.py | 2 +- generalresearch/models/dynata/survey.py | 4 +- generalresearch/models/dynata/task_collection.py | 2 +- generalresearch/models/events.py | 2 +- generalresearch/models/gr/business.py | 22 ++-- generalresearch/models/innovate/question.py | 2 +- generalresearch/models/innovate/survey.py | 8 +- generalresearch/models/legacy/bucket.py | 2 +- generalresearch/models/legacy/questions.py | 2 +- generalresearch/models/lucid/question.py | 2 +- generalresearch/models/lucid/survey.py | 2 +- generalresearch/models/morning/question.py | 2 +- generalresearch/models/morning/survey.py | 2 +- generalresearch/models/pollfish/question.py | 2 +- generalresearch/models/precision/question.py | 2 +- generalresearch/models/precision/survey.py | 2 +- generalresearch/models/prodege/question.py | 2 +- generalresearch/models/prodege/survey.py | 6 +- generalresearch/models/repdata/question.py | 2 +- generalresearch/models/repdata/survey.py | 2 +- generalresearch/models/repdata/task_collection.py | 2 +- generalresearch/models/sago/question.py | 2 +- generalresearch/models/sago/survey.py | 2 +- generalresearch/models/spectrum/question.py | 2 +- generalresearch/models/spectrum/survey.py | 2 +- generalresearch/models/spectrum/task_collection.py | 2 +- generalresearch/models/thl/__init__.py | 4 +- generalresearch/models/thl/category.py | 5 +- .../models/thl/contest/contest_entry.py | 13 +-- generalresearch/models/thl/contest/raffle.py | 4 +- generalresearch/models/thl/finance.py | 15 +-- generalresearch/models/thl/ledger.py | 2 +- generalresearch/models/thl/offerwall/__init__.py | 2 +- generalresearch/models/thl/offerwall/base.py | 2 +- generalresearch/models/thl/offerwall/cache.py | 2 +- generalresearch/models/thl/payout.py | 2 +- generalresearch/models/thl/product.py | 4 +- .../models/thl/profiling/marketplace.py | 4 +- .../models/thl/profiling/upk_question.py | 8 +- .../models/thl/profiling/upk_question_answer.py | 2 +- generalresearch/models/thl/profiling/user_info.py | 2 +- .../models/thl/profiling/user_question_answer.py | 8 +- generalresearch/models/thl/session.py | 4 +- generalresearch/models/thl/soft_pair.py | 2 +- generalresearch/models/thl/survey/__init__.py | 2 +- generalresearch/models/thl/survey/buyer.py | 2 +- generalresearch/models/thl/survey/condition.py | 2 +- generalresearch/models/thl/survey/model.py | 2 +- generalresearch/models/thl/survey/penalty.py | 2 +- generalresearch/models/thl/task_adjustment.py | 4 +- generalresearch/models/thl/user.py | 4 +- generalresearch/models/thl/user_profile.py | 2 +- generalresearch/models/thl/user_quality_event.py | 2 +- generalresearch/models/thl/user_streak.py | 2 +- .../models/thl/wallet/cashout_method.py | 4 +- generalresearch/models/thl/wallet/definitions.py | 87 ++++++++++++++++ generalresearch/models/thl/wallet/payout.py | 2 +- generalresearch/schemas/survey_stats.py | 2 +- generalresearch/wall_status_codes/__init__.py | 2 +- test_utils/conftest.py | 17 +-- test_utils/grliq/conftest.py | 12 ++- test_utils/incite/collections/conftest.py | 4 +- test_utils/incite/mergers/conftest.py | 55 +++++----- test_utils/managers/cashout_methods.py | 2 +- test_utils/managers/conftest.py | 55 +++++----- test_utils/managers/contest/conftest.py | 6 +- test_utils/managers/gr/conftest.py | 7 +- test_utils/managers/ledger/conftest.py | 14 ++- test_utils/managers/thl/conftest.py | 51 ++++----- test_utils/managers/upk/conftest.py | 9 +- test_utils/models/conftest.py | 4 +- test_utils/models/contest/conftest.py | 29 +++--- test_utils/models/gr/conftest.py | 38 +++---- test_utils/models/ledger/conftest.py | 4 +- test_utils/models/network/conftest.py | 5 +- test_utils/models/thl/conftest.py | 80 ++++++++------- test_utils/models/upk/conftest.py | 3 +- test_utils/spectrum/conftest.py | 8 +- .../incite/collections/test_df_collection_base.py | 2 +- .../collections/test_df_collection_item_base.py | 5 +- .../collections/test_df_collection_item_thl_web.py | 20 ++-- .../test_df_collection_thl_marketplaces.py | 9 +- .../collections/test_df_collection_thl_web.py | 2 +- .../mergers/foundations/test_enriched_session.py | 31 +++--- .../foundations/test_enriched_task_adjust.py | 30 +++--- .../mergers/foundations/test_enriched_wall.py | 29 +++--- .../mergers/foundations/test_user_id_product.py | 9 +- tests/incite/mergers/test_merge_collection.py | 7 +- tests/incite/mergers/test_merge_collection_item.py | 13 ++- tests/incite/mergers/test_pop_ledger.py | 21 ++-- tests/incite/mergers/test_ym_survey_merge.py | 24 +++-- tests/incite/test_collection_base.py | 6 +- tests/incite/test_collection_base_item.py | 6 +- tests/managers/gr/test_business.py | 19 ++-- tests/managers/gr/test_team.py | 15 +-- tests/managers/leaderboard.py | 5 +- tests/managers/network/test_label.py | 17 ++- tests/managers/test_events.py | 12 ++- tests/managers/test_lucid.py | 6 +- tests/managers/thl/test_buyer.py | 7 +- tests/managers/thl/test_cashout_method.py | 19 ++-- tests/managers/thl/test_category.py | 7 +- .../managers/thl/test_contest/test_leaderboard.py | 19 ++-- tests/managers/thl/test_contest/test_milestone.py | 20 ++-- tests/managers/thl/test_contest/test_raffle.py | 21 ++-- tests/managers/thl/test_harmonized_uqa.py | 7 +- tests/managers/thl/test_ipinfo.py | 7 +- tests/managers/thl/test_ledger/test_lm_accounts.py | 10 +- tests/managers/thl/test_ledger/test_lm_tx.py | 7 +- .../managers/thl/test_ledger/test_lm_tx_entries.py | 10 +- tests/managers/thl/test_ledger/test_lm_tx_locks.py | 15 +-- .../thl/test_ledger/test_lm_tx_metadata.py | 11 +- .../thl/test_ledger/test_thl_lm_accounts.py | 15 +-- .../thl/test_ledger/test_thl_lm_bp_payout.py | 22 ++-- tests/managers/thl/test_ledger/test_thl_lm_tx.py | 28 +++-- .../test_ledger/test_thl_lm_tx__user_payouts.py | 13 ++- tests/managers/thl/test_ledger/test_thl_pem.py | 21 ++-- tests/managers/thl/test_ledger/test_user_txs.py | 19 ++-- tests/managers/thl/test_ledger/test_wallet.py | 9 +- tests/managers/thl/test_payout.py | 53 +++++----- tests/managers/thl/test_product.py | 9 +- tests/managers/thl/test_product_prod.py | 5 +- tests/managers/thl/test_profiling/test_question.py | 7 +- tests/managers/thl/test_profiling/test_schema.py | 9 +- tests/managers/thl/test_profiling/test_user_upk.py | 6 +- tests/managers/thl/test_session_manager.py | 15 +-- tests/managers/thl/test_survey.py | 18 ++-- tests/managers/thl/test_survey_penalty.py | 7 +- tests/managers/thl/test_task_adjustment.py | 21 ++-- tests/managers/thl/test_task_status.py | 18 ++-- tests/managers/thl/test_user_manager/test_base.py | 18 ++-- tests/managers/thl/test_user_manager/test_mysql.py | 11 +- tests/managers/thl/test_user_manager/test_redis.py | 10 +- .../thl/test_user_manager/test_user_fetch.py | 8 +- .../thl/test_user_manager/test_user_metadata.py | 13 ++- tests/managers/thl/test_user_streak.py | 11 +- tests/managers/thl/test_userhealth.py | 19 +++- tests/managers/thl/test_wall_manager.py | 13 ++- tests/models/custom_types/test_aware_datetime.py | 4 +- tests/models/custom_types/test_dsn.py | 4 +- tests/models/custom_types/test_uuid_str.py | 4 +- tests/models/dynata/test_survey.py | 3 + tests/models/gr/test_authentication.py | 14 ++- tests/models/gr/test_base.py | 4 +- tests/models/gr/test_business.py | 40 ++++---- tests/models/gr/test_team.py | 38 ++++--- tests/models/innovate/test_question.py | 2 +- .../models/legacy/test_offerwall_parse_response.py | 2 +- .../models/legacy/test_user_question_answer_in.py | 9 +- tests/models/network/test_mtr.py | 6 +- tests/models/network/test_nmap.py | 8 +- tests/models/network/test_nmap_parser.py | 9 +- tests/models/network/test_rdns.py | 6 +- tests/models/spectrum/test_question.py | 2 +- tests/models/spectrum/test_survey.py | 2 +- tests/models/spectrum/test_survey_manager.py | 12 ++- tests/models/test_device.py | 2 +- tests/models/test_finance.py | 19 ++-- tests/models/thl/test_adjustments.py | 17 +-- tests/models/thl/test_buyer.py | 2 +- tests/models/thl/test_contest/test_contest.py | 6 +- .../thl/test_contest/test_leaderboard_contest.py | 7 +- .../models/thl/test_contest/test_raffle_contest.py | 7 +- tests/models/thl/test_marketplace_condition.py | 6 +- tests/models/thl/test_payout.py | 8 +- tests/models/thl/test_payout_format.py | 8 +- tests/models/thl/test_product.py | 37 ++++--- tests/models/thl/test_product_userwalletconfig.py | 2 +- tests/models/thl/test_soft_pair.py | 2 +- tests/models/thl/test_user.py | 9 +- tests/models/thl/test_user_metadata.py | 2 +- tests/models/thl/test_wall.py | 2 +- tests/models/thl/test_wall_session.py | 2 +- tests/test_postgres.py | 5 +- 211 files changed, 1318 insertions(+), 772 deletions(-) create mode 100644 generalresearch/managers/utils.py create mode 100644 generalresearch/models/definitions.py create mode 100644 generalresearch/models/thl/wallet/definitions.py (limited to 'tests/incite/collections') diff --git a/generalresearch/grliq/managers/event_plotter.py b/generalresearch/grliq/managers/event_plotter.py index 94b70ef..61cc52c 100644 --- a/generalresearch/grliq/managers/event_plotter.py +++ b/generalresearch/grliq/managers/event_plotter.py @@ -1,12 +1,15 @@ import html import webbrowser +from typing import TYPE_CHECKING import numpy as np from more_itertools import windowed from scipy.spatial.distance import euclidean from generalresearch.grliq.managers.colormap import turbo_colormap_data -from generalresearch.grliq.models.events import KeyboardEvent, MouseEvent + +if TYPE_CHECKING: + from generalresearch.grliq.models.events import KeyboardEvent, MouseEvent def make_events_svg( diff --git a/generalresearch/grliq/managers/forensic_events.py b/generalresearch/grliq/managers/forensic_events.py index 93da481..a97a9c2 100644 --- a/generalresearch/grliq/managers/forensic_events.py +++ b/generalresearch/grliq/managers/forensic_events.py @@ -1,7 +1,7 @@ import json from collections.abc import Collection from datetime import datetime -from typing import Any +from typing import TYPE_CHECKING, Any from uuid import uuid4 from psycopg import sql @@ -14,8 +14,10 @@ from generalresearch.grliq.models.events import ( PointerMove, TimingData, ) -from generalresearch.models.custom_types import UUIDStr -from generalresearch.pg_helper import PostgresConfig + +if TYPE_CHECKING: + from generalresearch.models.custom_types import UUIDStr + from generalresearch.pg_helper import PostgresConfig class GrlIqEventManager: diff --git a/generalresearch/grliq/managers/forensic_results.py b/generalresearch/grliq/managers/forensic_results.py index 158e582..93b1cdc 100644 --- a/generalresearch/grliq/managers/forensic_results.py +++ b/generalresearch/grliq/managers/forensic_results.py @@ -1,14 +1,16 @@ from collections.abc import Collection from datetime import datetime -from typing import Any +from typing import TYPE_CHECKING, Any from generalresearch.grliq.models.forensic_result import ( GrlIqForensicCategoryResult, Phase, ) from generalresearch.grliq.models.useragents import GrlUserAgent -from generalresearch.models.thl.user import User -from generalresearch.pg_helper import PostgresConfig + +if TYPE_CHECKING: + from generalresearch.models.thl.user import User + from generalresearch.pg_helper import PostgresConfig class GrlIqCategoryResultsReader: diff --git a/generalresearch/grliq/managers/forensic_summary.py b/generalresearch/grliq/managers/forensic_summary.py index b86e1f5..c222075 100644 --- a/generalresearch/grliq/managers/forensic_summary.py +++ b/generalresearch/grliq/managers/forensic_summary.py @@ -3,14 +3,10 @@ from __future__ import annotations import statistics from collections import defaultdict from datetime import UTC, datetime, timedelta -from typing import Any +from typing import TYPE_CHECKING, Any import numpy as np -from generalresearch.grliq.managers.forensic_data import GrlIqDataManager -from generalresearch.grliq.managers.forensic_events import ( - GrlIqEventManager, -) from generalresearch.grliq.models.forensic_result import ( GrlIqCheckerResults, GrlIqForensicCategoryResult, @@ -22,8 +18,14 @@ from generalresearch.grliq.models.forensic_summary import ( TimingDataCountrySummary, UserForensicSummary, ) -from generalresearch.models.thl.user import User -from generalresearch.redis_helper import RedisConfig + +if TYPE_CHECKING: + from generalresearch.grliq.managers.forensic_data import GrlIqDataManager + from generalresearch.grliq.managers.forensic_events import ( + GrlIqEventManager, + ) + from generalresearch.models.thl.user import User + from generalresearch.redis_helper import RedisConfig def calculate_category_summary( diff --git a/generalresearch/incite/schemas/mergers/foundations/enriched_task_adjust.py b/generalresearch/incite/schemas/mergers/foundations/enriched_task_adjust.py index ead42d9..ac9a35a 100644 --- a/generalresearch/incite/schemas/mergers/foundations/enriched_task_adjust.py +++ b/generalresearch/incite/schemas/mergers/foundations/enriched_task_adjust.py @@ -4,7 +4,7 @@ from pandera.pandas import Check, Column, DataFrameSchema, Index from generalresearch.incite.schemas import ARCHIVE_AFTER, ORDER_KEY from generalresearch.incite.schemas.thl_web import THLTaskAdjustmentSchema from generalresearch.locales import Localelator -from generalresearch.models import DeviceType, Source +from generalresearch.models.definitions import DeviceType, Source from generalresearch.models.thl.definitions import ( WallAdjustedStatus, ) diff --git a/generalresearch/incite/schemas/mergers/foundations/enriched_wall.py b/generalresearch/incite/schemas/mergers/foundations/enriched_wall.py index 1443f28..71d0eab 100644 --- a/generalresearch/incite/schemas/mergers/foundations/enriched_wall.py +++ b/generalresearch/incite/schemas/mergers/foundations/enriched_wall.py @@ -5,7 +5,7 @@ from pandera.pandas import Check, Column, DataFrameSchema, Index from generalresearch.incite.schemas import ARCHIVE_AFTER, PARTITION_ON from generalresearch.locales import Localelator -from generalresearch.models import DeviceType, Source +from generalresearch.models.definitions import DeviceType, Source from generalresearch.models.thl.definitions import ( ReportValue, Status, diff --git a/generalresearch/incite/schemas/mergers/ym_wall_summary.py b/generalresearch/incite/schemas/mergers/ym_wall_summary.py index 16cfc2f..737b925 100644 --- a/generalresearch/incite/schemas/mergers/ym_wall_summary.py +++ b/generalresearch/incite/schemas/mergers/ym_wall_summary.py @@ -6,7 +6,7 @@ from pandera.pandas import Check, Column, DataFrameSchema, Index from generalresearch.incite.schemas import ARCHIVE_AFTER from generalresearch.locales import Localelator -from generalresearch.models import Source +from generalresearch.models.definitions import Source COUNTRY_ISOS: set[str] = Localelator().get_all_countries() kosovo = "xk" diff --git a/generalresearch/incite/schemas/thl_web.py b/generalresearch/incite/schemas/thl_web.py index 30c7076..36ee8e9 100644 --- a/generalresearch/incite/schemas/thl_web.py +++ b/generalresearch/incite/schemas/thl_web.py @@ -6,7 +6,7 @@ from pandera.pandas import Check, Column, DataFrameSchema, Index, MultiIndex from generalresearch.incite.schemas import ARCHIVE_AFTER, ORDER_KEY from generalresearch.locales import Localelator -from generalresearch.models import DeviceType, Source +from generalresearch.models.definitions import DeviceType, Source from generalresearch.models.thl.definitions import ( ReportValue, SessionAdjustedStatus, diff --git a/generalresearch/managers/cint/user_pid.py b/generalresearch/managers/cint/user_pid.py index 4f749a0..0265823 100644 --- a/generalresearch/managers/cint/user_pid.py +++ b/generalresearch/managers/cint/user_pid.py @@ -1,5 +1,5 @@ from generalresearch.managers.marketplace.user_pid import UserPidManager -from generalresearch.models import Source +from generalresearch.models.definitions import Source class CintUserPidManager(UserPidManager): diff --git a/generalresearch/managers/dynata/user_pid.py b/generalresearch/managers/dynata/user_pid.py index aefed34..67ff968 100644 --- a/generalresearch/managers/dynata/user_pid.py +++ b/generalresearch/managers/dynata/user_pid.py @@ -1,5 +1,5 @@ from generalresearch.managers.marketplace.user_pid import UserPidManager -from generalresearch.models import Source +from generalresearch.models.definitions import Source class DynataUserPidManager(UserPidManager): diff --git a/generalresearch/managers/events.py b/generalresearch/managers/events.py index c43a020..30cec0c 100644 --- a/generalresearch/managers/events.py +++ b/generalresearch/managers/events.py @@ -12,7 +12,7 @@ from redis.client import PubSub, Redis from generalresearch.incite.base import LOG from generalresearch.managers.base import RedisManager -from generalresearch.models import Source +from generalresearch.models.definitions import Source from generalresearch.models.events import ( AggregateBySource, EventEnvelope, diff --git a/generalresearch/managers/innovate/user_pid.py b/generalresearch/managers/innovate/user_pid.py index 100b0ca..7544c89 100644 --- a/generalresearch/managers/innovate/user_pid.py +++ b/generalresearch/managers/innovate/user_pid.py @@ -1,5 +1,5 @@ from generalresearch.managers.marketplace.user_pid import UserPidManager -from generalresearch.models import Source +from generalresearch.models.definitions import Source class InnovateUserPidManager(UserPidManager): diff --git a/generalresearch/managers/marketplace/user_pid.py b/generalresearch/managers/marketplace/user_pid.py index fe24d38..00dae8a 100644 --- a/generalresearch/managers/marketplace/user_pid.py +++ b/generalresearch/managers/marketplace/user_pid.py @@ -6,7 +6,7 @@ from typing import TYPE_CHECKING from uuid import UUID from generalresearch.managers.base import SqlManager -from generalresearch.models import Source +from generalresearch.models.definitions import Source if TYPE_CHECKING: from generalresearch.sql_helper import SqlHelper diff --git a/generalresearch/managers/morning/user_pid.py b/generalresearch/managers/morning/user_pid.py index 78de3bd..5896734 100644 --- a/generalresearch/managers/morning/user_pid.py +++ b/generalresearch/managers/morning/user_pid.py @@ -1,5 +1,5 @@ from generalresearch.managers.marketplace.user_pid import UserPidManager -from generalresearch.models import Source +from generalresearch.models.definitions import Source class MorningUserPidManager(UserPidManager): diff --git a/generalresearch/managers/network/label.py b/generalresearch/managers/network/label.py index cec59ad..aed5ff6 100644 --- a/generalresearch/managers/network/label.py +++ b/generalresearch/managers/network/label.py @@ -5,11 +5,12 @@ from datetime import UTC, datetime, timedelta from typing import TYPE_CHECKING from psycopg import sql -from pydantic import IPvAnyNetwork, TypeAdapter +from pydantic import TypeAdapter from generalresearch.managers.base import PostgresManager from generalresearch.models.custom_types import ( - AwareDatetimeISO, + IPvAnyAddressStr, + IPvAnyNetwork, IPvAnyNetworkStr, ) from generalresearch.models.network.label import IPLabel @@ -17,7 +18,6 @@ from generalresearch.models.network.label import IPLabel if TYPE_CHECKING: from generalresearch.models.custom_types import ( AwareDatetimeISO, - IPvAnyNetworkStr, ) from generalresearch.models.network.label import IPLabelKind, IPLabelSource diff --git a/generalresearch/managers/precision/user_pid.py b/generalresearch/managers/precision/user_pid.py index 50e97e6..ed2d58d 100644 --- a/generalresearch/managers/precision/user_pid.py +++ b/generalresearch/managers/precision/user_pid.py @@ -1,5 +1,5 @@ from generalresearch.managers.marketplace.user_pid import UserPidManager -from generalresearch.models import Source +from generalresearch.models.definitions import Source class PrecisionUserPidManager(UserPidManager): diff --git a/generalresearch/managers/prodege/user_pid.py b/generalresearch/managers/prodege/user_pid.py index 7c92e28..c18c109 100644 --- a/generalresearch/managers/prodege/user_pid.py +++ b/generalresearch/managers/prodege/user_pid.py @@ -1,5 +1,5 @@ from generalresearch.managers.marketplace.user_pid import UserPidManager -from generalresearch.models import Source +from generalresearch.models.definitions import Source class ProdegeUserPidManager(UserPidManager): diff --git a/generalresearch/managers/repdata/user_pid.py b/generalresearch/managers/repdata/user_pid.py index 9d53897..5fdeccf 100644 --- a/generalresearch/managers/repdata/user_pid.py +++ b/generalresearch/managers/repdata/user_pid.py @@ -1,5 +1,5 @@ from generalresearch.managers.marketplace.user_pid import UserPidManager -from generalresearch.models import Source +from generalresearch.models.definitions import Source class RepdataUserPidManager(UserPidManager): diff --git a/generalresearch/managers/sago/user_pid.py b/generalresearch/managers/sago/user_pid.py index 311abb7..b7ce771 100644 --- a/generalresearch/managers/sago/user_pid.py +++ b/generalresearch/managers/sago/user_pid.py @@ -1,5 +1,5 @@ from generalresearch.managers.marketplace.user_pid import UserPidManager -from generalresearch.models import Source +from generalresearch.models.definitions import Source class SagoUserPidManager(UserPidManager): diff --git a/generalresearch/managers/spectrum/user_pid.py b/generalresearch/managers/spectrum/user_pid.py index 495e73c..980c28d 100644 --- a/generalresearch/managers/spectrum/user_pid.py +++ b/generalresearch/managers/spectrum/user_pid.py @@ -1,5 +1,5 @@ from generalresearch.managers.marketplace.user_pid import UserPidManager -from generalresearch.models import Source +from generalresearch.models.definitions import Source class SpectrumUserPidManager(UserPidManager): diff --git a/generalresearch/managers/thl/buyer.py b/generalresearch/managers/thl/buyer.py index 1e20e2f..38214c6 100644 --- a/generalresearch/managers/thl/buyer.py +++ b/generalresearch/managers/thl/buyer.py @@ -8,7 +8,7 @@ from generalresearch.managers.base import Permission, PostgresManager from generalresearch.models.thl.survey.buyer import Buyer if TYPE_CHECKING: - from generalresearch.models import Source + from generalresearch.models.definitions import Source from generalresearch.pg_helper import PostgresConfig diff --git a/generalresearch/managers/thl/cashout_method.py b/generalresearch/managers/thl/cashout_method.py index e701da3..c12c920 100644 --- a/generalresearch/managers/thl/cashout_method.py +++ b/generalresearch/managers/thl/cashout_method.py @@ -9,10 +9,10 @@ from uuid import UUID, uuid4 from pydantic import NonNegativeInt from generalresearch.managers.base import PostgresManager -from generalresearch.models.thl.wallet import PayoutType from generalresearch.models.thl.wallet.cashout_method import ( CashoutMethod, ) +from generalresearch.models.thl.wallet.definitions import PayoutType if TYPE_CHECKING: from generalresearch.models.thl.user import User diff --git a/generalresearch/managers/thl/contest_manager.py b/generalresearch/managers/thl/contest_manager.py index 3f85d31..68b2cf0 100644 --- a/generalresearch/managers/thl/contest_manager.py +++ b/generalresearch/managers/thl/contest_manager.py @@ -14,7 +14,9 @@ from generalresearch.models.thl.contest import ( ContestPrize, ContestWinner, ) +from generalresearch.models.thl.contest.contest_entry import ContestEntry from generalresearch.models.thl.contest.definitions import ( + ContestEntryType, ContestStatus, ContestType, ) @@ -35,8 +37,6 @@ from generalresearch.models.thl.contest.milestone import ( MilestoneUserView, ) from generalresearch.models.thl.contest.raffle import ( - ContestEntry, - ContestEntryType, RaffleContest, RaffleUserView, ) diff --git a/generalresearch/managers/thl/ledger_manager/ledger.py b/generalresearch/managers/thl/ledger_manager/ledger.py index 6cb4b28..f2455d4 100644 --- a/generalresearch/managers/thl/ledger_manager/ledger.py +++ b/generalresearch/managers/thl/ledger_manager/ledger.py @@ -13,7 +13,6 @@ from pydantic import AwareDatetime, NonNegativeInt, PositiveInt from redis.exceptions import LockError, LockNotOwnedError from generalresearch.currency import LedgerCurrency -from generalresearch.managers import parse_order_by from generalresearch.managers.base import ( Permission, PostgresManager, @@ -28,6 +27,7 @@ from generalresearch.managers.thl.ledger_manager.exceptions import ( LedgerTransactionFlagAlreadyExistsError, LedgerTransactionReleaseLockError, ) +from generalresearch.managers.utils import parse_order_by from generalresearch.models.custom_types import check_valid_uuid from generalresearch.models.thl.ledger import ( LedgerAccount, diff --git a/generalresearch/managers/thl/ledger_manager/thl_ledger.py b/generalresearch/managers/thl/ledger_manager/thl_ledger.py index 7aed619..bd27acf 100644 --- a/generalresearch/managers/thl/ledger_manager/thl_ledger.py +++ b/generalresearch/managers/thl/ledger_manager/thl_ledger.py @@ -52,7 +52,7 @@ from generalresearch.models.thl.ledger import ( ) from generalresearch.models.thl.product import Product from generalresearch.models.thl.session import Status -from generalresearch.models.thl.wallet import PayoutType +from generalresearch.models.thl.wallet.definitions import PayoutType if TYPE_CHECKING: from generalresearch.models.custom_types import UUIDStr diff --git a/generalresearch/managers/thl/payout.py b/generalresearch/managers/thl/payout.py index 2914ba4..1749783 100644 --- a/generalresearch/managers/thl/payout.py +++ b/generalresearch/managers/thl/payout.py @@ -31,11 +31,11 @@ from generalresearch.models.thl.payout import ( PayoutEvent, UserPayoutEvent, ) -from generalresearch.models.thl.wallet import PayoutType from generalresearch.models.thl.wallet.cashout_method import ( CashMailOrderData, CashoutRequestInfo, ) +from generalresearch.models.thl.wallet.definitions import PayoutType if TYPE_CHECKING: from generalresearch.managers.thl.ledger_manager.thl_ledger import ( diff --git a/generalresearch/managers/thl/session.py b/generalresearch/managers/thl/session.py index 7f17252..41d3893 100644 --- a/generalresearch/managers/thl/session.py +++ b/generalresearch/managers/thl/session.py @@ -10,12 +10,12 @@ from faker import Faker from psycopg import sql from pydantic import NonNegativeInt, PositiveInt -from generalresearch.managers import parse_order_by from generalresearch.managers.base import ( Permission, PostgresManager, ) from generalresearch.managers.thl.product import ProductManager +from generalresearch.managers.utils import parse_order_by from generalresearch.models.legacy.bucket import Bucket from generalresearch.models.thl.session import ( Session, @@ -28,8 +28,8 @@ from generalresearch.models.thl.task_status import ( from generalresearch.models.thl.user import User if TYPE_CHECKING: - from generalresearch.models import DeviceType from generalresearch.models.custom_types import UUIDStr + from generalresearch.models.definitions import DeviceType from generalresearch.models.thl.definitions import ( SessionStatusCode2, Status, diff --git a/generalresearch/managers/thl/survey.py b/generalresearch/managers/thl/survey.py index eacb345..92777e5 100644 --- a/generalresearch/managers/thl/survey.py +++ b/generalresearch/managers/thl/survey.py @@ -13,7 +13,7 @@ from pydantic import NonNegativeInt from generalresearch.managers.base import Permission, PostgresManager from generalresearch.managers.thl.buyer import BuyerManager from generalresearch.managers.thl.category import CategoryManager -from generalresearch.models import Source +from generalresearch.models.definitions import Source from generalresearch.models.thl.survey.model import ( Survey, SurveyStat, diff --git a/generalresearch/managers/thl/task_adjustment.py b/generalresearch/managers/thl/task_adjustment.py index e3f382d..d0d83cb 100644 --- a/generalresearch/managers/thl/task_adjustment.py +++ b/generalresearch/managers/thl/task_adjustment.py @@ -6,12 +6,12 @@ from decimal import Decimal from functools import cached_property from typing import TYPE_CHECKING -from generalresearch.managers import parse_order_by from generalresearch.managers.base import ( PostgresManager, ) from generalresearch.managers.thl.session import SessionManager from generalresearch.managers.thl.wall import WallManager +from generalresearch.managers.utils import parse_order_by from generalresearch.models.thl.definitions import ( Status, WallAdjustedStatus, diff --git a/generalresearch/managers/thl/wall.py b/generalresearch/managers/thl/wall.py index b9dc94d..83697f5 100644 --- a/generalresearch/managers/thl/wall.py +++ b/generalresearch/managers/thl/wall.py @@ -14,12 +14,12 @@ from psycopg import sql from psycopg.rows import dict_row from pydantic import AwareDatetime, PositiveInt -from generalresearch.managers import parse_order_by from generalresearch.managers.base import ( PostgresManager, PostgresManagerWithRedis, ) -from generalresearch.models import Source +from generalresearch.managers.utils import parse_order_by +from generalresearch.models.definitions import Source from generalresearch.models.thl.definitions import ( WallAdjustedStatus, ) diff --git a/generalresearch/managers/thl/wallet/__init__.py b/generalresearch/managers/thl/wallet/__init__.py index 457483f..f805872 100644 --- a/generalresearch/managers/thl/wallet/__init__.py +++ b/generalresearch/managers/thl/wallet/__init__.py @@ -6,7 +6,7 @@ from generalresearch.managers.thl.wallet.approve import ( approve_paypal_order, ) from generalresearch.models.thl.definitions import PayoutStatus -from generalresearch.models.thl.wallet import PayoutType +from generalresearch.models.thl.wallet.definitions import PayoutType if TYPE_CHECKING: from generalresearch.managers.thl.ledger_manager.thl_ledger import ( diff --git a/generalresearch/managers/utils.py b/generalresearch/managers/utils.py new file mode 100644 index 0000000..bc745fd --- /dev/null +++ b/generalresearch/managers/utils.py @@ -0,0 +1,16 @@ +def parse_order_by(order_by_str: str) -> str: + """ + Converts django-rest-framework ordering str to mysql clause + :param order_by_str: e.g. 'created,-name' + :return: mysql clause e.g. ORDER BY created ASC, name DESC + """ + fields = order_by_str.split(",") + + order_clause = [] + for field in fields: + if field.startswith("-"): + order_clause.append(f"{field[1:]} DESC") + else: + order_clause.append(f"{field} ASC") + + return "ORDER BY " + ", ".join(order_clause) diff --git a/generalresearch/models/cint/question.py b/generalresearch/models/cint/question.py index 44efd13..ab46653 100644 --- a/generalresearch/models/cint/question.py +++ b/generalresearch/models/cint/question.py @@ -8,7 +8,7 @@ from uuid import UUID from pydantic import BaseModel, Field, field_validator, model_validator -from generalresearch.models import Source, string_utils +from generalresearch.models.definitions import Source, string_utils from generalresearch.models.thl.profiling.marketplace import ( MarketplaceQuestion, MarketplaceUserQuestionAnswer, diff --git a/generalresearch/models/cint/survey.py b/generalresearch/models/cint/survey.py index 8c8f882..ebba09e 100644 --- a/generalresearch/models/cint/survey.py +++ b/generalresearch/models/cint/survey.py @@ -18,7 +18,7 @@ from pydantic import ( ) from generalresearch.locales import Localelator -from generalresearch.models import Source, TaskCalculationType +from generalresearch.models.definitions import Source, TaskCalculationType from generalresearch.models.thl.demographics import Gender from generalresearch.models.thl.survey import MarketplaceTask from generalresearch.models.thl.survey.condition import ( diff --git a/generalresearch/models/custom_types.py b/generalresearch/models/custom_types.py index 5e4db3e..680a99c 100644 --- a/generalresearch/models/custom_types.py +++ b/generalresearch/models/custom_types.py @@ -20,7 +20,7 @@ from pydantic.functional_validators import AfterValidator, BeforeValidator from pydantic.networks import IPvAnyNetwork, UrlConstraints from pydantic_core import MultiHostHost, Url -from generalresearch.models import DeviceType, Source +from generalresearch.models.definitions import DeviceType, Source HOSTNAME_REGEX = re.compile( r"^[a-zA-Z0-9]([a-zA-Z0-9-]{0,61}[a-zA-Z0-9])?(\.[a-zA-Z0-9]([a-zA-Z0-9-]{0,61}[a-zA-Z0-9])?)*$" diff --git a/generalresearch/models/definitions.py b/generalresearch/models/definitions.py new file mode 100644 index 0000000..c0348d7 --- /dev/null +++ b/generalresearch/models/definitions.py @@ -0,0 +1,114 @@ +from __future__ import annotations + +from enum import IntEnum, StrEnum + +from generalresearch.utils.enum import ReprEnumMeta + + +class Source(StrEnum, metaclass=ReprEnumMeta): + # The external marketplace, or the source of the survey / work. + # Max length of the value is 2. + GRS = "g" + CINT = "c" + DALIA = "a" # deprecated + DYNATA = "d" + ETX = "et" + FULL_CIRCLE = "f" + INNOVATE = "i" + LUCID = "l" + MORNING_CONSULT = "m" + OPEN_LABS = "n" + POLLFISH = "o" + PRECISION = "e" + PRODEGE_USER = "r" # deprecated + PRODEGE = "pr" # using 'r' for vendor_wall + PULLEY = "p" # deprecated + REPDATA = "rd" # using 'q' for vendor_wall + SAGO = "h" + SPECTRUM = "s" + TESTING = "t" # Used internally for testing + TESTING2 = "u" # Used internally for testing + WXET = "w" + + +class DebitKey(IntEnum, metaclass=ReprEnumMeta): + # The debit key for marketplaces + CINT = 8 + DALIA = 9 + DYNATA = 6 + # ETX = None + FULL_CIRCLE = 15 + INNOVATE = 7 + LUCID = 0 + MORNING_CONSULT = 12 + # OPEN_LABS = None + POLLFISH = 13 + PRECISION = 14 + PRODEGE = 11 + SAGO = 10 + SPECTRUM = 5 + # WXET = None + + +class DeviceType(IntEnum, metaclass=ReprEnumMeta): + UNKNOWN = 0 + MOBILE = 1 + DESKTOP = 2 + TABLET = 3 + + +class LogicalOperator(StrEnum, metaclass=ReprEnumMeta): + OR = "OR" + AND = "AND" + # There is currently no use case for NOT. See MarketplaceCondition.explain_not + NOT = "NOT" + + +class TaskStatus(StrEnum, metaclass=ReprEnumMeta): + # A survey is live if it is open and, given all conditions are met, is + # possible to send in traffic. All other statuses are just variants of + # NOT Live (not accepting traffic) + LIVE = "LIVE" + + # This is a generic NOT Live status. A marketplace may use other more + # specific statuses but in practice they don't matter because all we care + # about is if the task is LIVE. + NOT_LIVE = "NOT_LIVE" + + # We need a status to mark if a survey we thought was live does not come + # back from the API, we'll mark it as NOT_FOUND. + NOT_FOUND = "NOT_FOUND" + + +class TaskCalculationType(StrEnum): + COMPLETES = "COMPLETES" + STARTS = "STARTS" + + @classmethod + def from_api(cls, v: str) -> TaskCalculationType: + return { + "complete": cls.COMPLETES, + "completes": cls.COMPLETES, + "survey start": cls.STARTS, + "survey starts": cls.STARTS, + "start": cls.STARTS, + "prescreens": cls.STARTS, + "prescreen": cls.STARTS, + }[v.lower()] + + @classmethod + def prodege_from_api(cls, v: int) -> TaskCalculationType: + return {1: cls.COMPLETES, 2: cls.STARTS}[v] + + @classmethod + def innovate_from_api(cls, v: int) -> TaskCalculationType: + return {0: cls.COMPLETES, 1: cls.STARTS}[v] + + +class URLQueryKey(StrEnum, metaclass=ReprEnumMeta): + PRODUCT_ID = "39057c8b" + PRODUCT_USER_ID = "c184efc0" + SESSION_ID = "0bb50182" + + +MAX_INT32 = 2**31 diff --git a/generalresearch/models/device.py b/generalresearch/models/device.py index cc15eee..432c897 100644 --- a/generalresearch/models/device.py +++ b/generalresearch/models/device.py @@ -1,6 +1,6 @@ from user_agents import parse as parse_ua -from generalresearch.models import DeviceType +from generalresearch.models.definitions import DeviceType def parse_device_from_useragent(user_agent: str) -> DeviceType: diff --git a/generalresearch/models/dynata/question.py b/generalresearch/models/dynata/question.py index 1ed560a..60c7366 100644 --- a/generalresearch/models/dynata/question.py +++ b/generalresearch/models/dynata/question.py @@ -11,7 +11,7 @@ from typing import TYPE_CHECKING, Any, Literal from pydantic import BaseModel, Field, PositiveInt, field_validator, model_validator -from generalresearch.models import MAX_INT32, Source +from generalresearch.models.definitions import MAX_INT32, Source from generalresearch.models.thl.profiling.marketplace import MarketplaceQuestion if TYPE_CHECKING: diff --git a/generalresearch/models/dynata/survey.py b/generalresearch/models/dynata/survey.py index 70e3659..4174d31 100644 --- a/generalresearch/models/dynata/survey.py +++ b/generalresearch/models/dynata/survey.py @@ -19,7 +19,7 @@ from pydantic import ( ) from generalresearch.locales import Localelator -from generalresearch.models import Source +from generalresearch.models.definitions import Source from generalresearch.models.dynata import DynataStatus from generalresearch.models.thl.demographics import ( Gender, @@ -31,7 +31,6 @@ from generalresearch.models.thl.survey.condition import ( ) if TYPE_CHECKING: - from generalresearch.models import TaskCalculationType from generalresearch.models.custom_types import ( AlphaNumStr, AlphaNumStrSet, @@ -39,6 +38,7 @@ if TYPE_CHECKING: CoercedStr, DeviceTypes, ) + from generalresearch.models.definitions import TaskCalculationType logging.basicConfig() logger = logging.getLogger() diff --git a/generalresearch/models/dynata/task_collection.py b/generalresearch/models/dynata/task_collection.py index 94868bb..c6cdc19 100644 --- a/generalresearch/models/dynata/task_collection.py +++ b/generalresearch/models/dynata/task_collection.py @@ -6,7 +6,7 @@ import pandas as pd from pandera.pandas import Check, Column, DataFrameSchema, Index from generalresearch.locales import Localelator -from generalresearch.models import TaskCalculationType +from generalresearch.models.definitions import TaskCalculationType from generalresearch.models.dynata import DynataStatus from generalresearch.models.thl.survey.task_collection import ( TaskCollection, diff --git a/generalresearch/models/events.py b/generalresearch/models/events.py index 8d059f9..34f6be8 100644 --- a/generalresearch/models/events.py +++ b/generalresearch/models/events.py @@ -14,12 +14,12 @@ from pydantic import ( ) if TYPE_CHECKING: - from generalresearch.models import Source from generalresearch.models.custom_types import ( AwareDatetimeISO, CountryISOLike, UUIDStr, ) + from generalresearch.models.definitions import Source from generalresearch.models.thl.definitions import ( SessionStatusCode2, Status, diff --git a/generalresearch/models/gr/business.py b/generalresearch/models/gr/business.py index 146d690..e11c54d 100644 --- a/generalresearch/models/gr/business.py +++ b/generalresearch/models/gr/business.py @@ -11,7 +11,7 @@ from uuid import uuid4 import pandas as pd import pyarrow as pa -from dask.distributed import Client +from dask.distributed import Client as DaskClient from psycopg.cursor import Cursor from psycopg.rows import dict_row from pydantic import BaseModel, ConfigDict, Field, PositiveInt, ValidationError @@ -24,6 +24,11 @@ from generalresearch.incite.schemas.mergers.pop_ledger import ( numerical_col_names, ) from generalresearch.models.admin.request import ReportRequest, ReportType +from generalresearch.models.custom_types import ( + AwareDatetime, + UUIDStr, + UUIDStrCoerce, +) from generalresearch.models.gr.team import Team from generalresearch.models.thl.finance import BusinessBalances, POPFinancial from generalresearch.models.thl.ledger import OrderBy @@ -32,11 +37,6 @@ from generalresearch.utils.enum import ReprEnumMeta if TYPE_CHECKING: from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge - from generalresearch.models.custom_types import ( - AwareDatetime, - UUIDStr, - UUIDStrCoerce, - ) from generalresearch.models.thl.ledger import LedgerAccount from generalresearch.models.thl.payout import BusinessPayoutEvent from generalresearch.pg_helper import PostgresConfig @@ -354,7 +354,7 @@ class Business(BaseModel): thl_pg_config: PostgresConfig, lm: LedgerManager, ds: GRLDatasets, - client: Client, + client: DaskClient, pop_ledger: PopLedgerMerge | None = None, at_timestamp: AwareDatetime | None = None, ) -> None: @@ -464,7 +464,7 @@ class Business(BaseModel): thl_pg_config: PostgresConfig, thl_lm: ThlLedgerManager, ds: GRLDatasets, - client: Client, + client: DaskClient, pop_ledger: PopLedgerMerge | None = None, ) -> None: """This is very similar to the Product POP Financial endpoint; however, @@ -518,7 +518,7 @@ class Business(BaseModel): self, thl_pg_config: PostgresConfig, ds: GRLDatasets, - client: Client, + client: DaskClient, mnt_gr_api: Path, enriched_session: EnrichedSessionMerge | None = None, ) -> None: @@ -561,7 +561,7 @@ class Business(BaseModel): self, thl_pg_config: PostgresConfig, ds: GRLDatasets, - client: Client, + client: DaskClient, mnt_gr_api: Path, enriched_wall: EnrichedWallMerge | None = None, ) -> None: @@ -633,7 +633,7 @@ class Business(BaseModel): pg_config: PostgresConfig, thl_web_rr: PostgresConfig, redis_config: RedisConfig, - client: Client, + client: DaskClient, ds: GRLDatasets, lm: LedgerManager, thl_lm: ThlLedgerManager, diff --git a/generalresearch/models/innovate/question.py b/generalresearch/models/innovate/question.py index fc89524..6423399 100644 --- a/generalresearch/models/innovate/question.py +++ b/generalresearch/models/innovate/question.py @@ -8,7 +8,7 @@ from typing import TYPE_CHECKING, Any, Literal from pydantic import BaseModel, Field, ValidationError, field_validator, model_validator -from generalresearch.models import Source +from generalresearch.models.definitions import Source from generalresearch.models.thl.profiling.marketplace import ( MarketplaceQuestion, MarketplaceUserQuestionAnswer, diff --git a/generalresearch/models/innovate/survey.py b/generalresearch/models/innovate/survey.py index 60921df..e718dda 100644 --- a/generalresearch/models/innovate/survey.py +++ b/generalresearch/models/innovate/survey.py @@ -24,7 +24,7 @@ from pydantic import ( ) from generalresearch.locales import Localelator -from generalresearch.models import ( +from generalresearch.models.definitions import ( LogicalOperator, Source, ) @@ -41,15 +41,15 @@ from generalresearch.models.thl.survey.condition import ( ) if TYPE_CHECKING: - from generalresearch.models import ( - TaskCalculationType, - ) from generalresearch.models.custom_types import ( AlphaNumStrSet, AwareDatetimeISO, CoercedStr, DeviceTypes, ) + from generalresearch.models.definitions import ( + TaskCalculationType, + ) from generalresearch.models.innovate.question import InnovateQuestionID logging.basicConfig() diff --git a/generalresearch/models/legacy/bucket.py b/generalresearch/models/legacy/bucket.py index f20a769..5f53b89 100644 --- a/generalresearch/models/legacy/bucket.py +++ b/generalresearch/models/legacy/bucket.py @@ -15,7 +15,7 @@ from pydantic import ( model_validator, ) -from generalresearch.models import Source +from generalresearch.models.definitions import Source from generalresearch.models.thl.stats import StatisticalSummary if TYPE_CHECKING: diff --git a/generalresearch/models/legacy/questions.py b/generalresearch/models/legacy/questions.py index bebd28f..c333804 100644 --- a/generalresearch/models/legacy/questions.py +++ b/generalresearch/models/legacy/questions.py @@ -219,7 +219,7 @@ class UserQuestionAnswers(BaseModel): self.user = res def prefetch_wall(self, wm: WallManager) -> None: - from generalresearch.models import Source + from generalresearch.models.definitions import Source res: Wall | None = wm.get_from_uuid_if_exists(wall_uuid=self.session_id) diff --git a/generalresearch/models/lucid/question.py b/generalresearch/models/lucid/question.py index 98f535b..c1b9e52 100644 --- a/generalresearch/models/lucid/question.py +++ b/generalresearch/models/lucid/question.py @@ -6,7 +6,7 @@ from typing import TYPE_CHECKING, Any, Literal, Self from pydantic import BaseModel, Field, field_validator, model_validator -from generalresearch.models import Source +from generalresearch.models.definitions import Source from generalresearch.models.thl.profiling.marketplace import ( MarketplaceQuestion, ) diff --git a/generalresearch/models/lucid/survey.py b/generalresearch/models/lucid/survey.py index 0f03e31..a04e529 100644 --- a/generalresearch/models/lucid/survey.py +++ b/generalresearch/models/lucid/survey.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING, Any, Self from pydantic import BaseModel, ConfigDict, Field, NonNegativeInt -from generalresearch.models import Source +from generalresearch.models.definitions import Source from generalresearch.models.thl.survey.condition import ( ConditionValueType, MarketplaceCondition, diff --git a/generalresearch/models/morning/question.py b/generalresearch/models/morning/question.py index 748fcc6..909992f 100644 --- a/generalresearch/models/morning/question.py +++ b/generalresearch/models/morning/question.py @@ -6,7 +6,7 @@ from uuid import UUID from pydantic import BaseModel, Field, field_validator, model_validator from generalresearch.locales import Localelator -from generalresearch.models import Source +from generalresearch.models.definitions import Source from generalresearch.models.thl.profiling.marketplace import ( MarketplaceQuestion, MarketplaceUserQuestionAnswer, diff --git a/generalresearch/models/morning/survey.py b/generalresearch/models/morning/survey.py index 1e217f6..25accb6 100644 --- a/generalresearch/models/morning/survey.py +++ b/generalresearch/models/morning/survey.py @@ -25,7 +25,7 @@ from pydantic import ( ) from generalresearch.locales import Localelator -from generalresearch.models import Source +from generalresearch.models.definitions import Source from generalresearch.models.morning import MorningStatus from generalresearch.models.thl.demographics import Gender from generalresearch.models.thl.survey import MarketplaceTask diff --git a/generalresearch/models/pollfish/question.py b/generalresearch/models/pollfish/question.py index 3b658fd..f0c733c 100644 --- a/generalresearch/models/pollfish/question.py +++ b/generalresearch/models/pollfish/question.py @@ -8,7 +8,7 @@ from typing import TYPE_CHECKING, Any, Literal, Self from pydantic import BaseModel, Field, model_validator -from generalresearch.models import Source +from generalresearch.models.definitions import Source from generalresearch.models.thl.profiling.marketplace import MarketplaceQuestion if TYPE_CHECKING: diff --git a/generalresearch/models/precision/question.py b/generalresearch/models/precision/question.py index ba17361..6ed6bbd 100644 --- a/generalresearch/models/precision/question.py +++ b/generalresearch/models/precision/question.py @@ -8,7 +8,7 @@ from typing import TYPE_CHECKING, Any, Literal from pydantic import BaseModel, Field, ValidationError, field_validator, model_validator -from generalresearch.models import Source, string_utils +from generalresearch.models.definitions import Source, string_utils from generalresearch.models.thl.profiling.marketplace import ( MarketplaceQuestion, MarketplaceUserQuestionAnswer, diff --git a/generalresearch/models/precision/survey.py b/generalresearch/models/precision/survey.py index fa30882..a9e34e6 100644 --- a/generalresearch/models/precision/survey.py +++ b/generalresearch/models/precision/survey.py @@ -15,7 +15,7 @@ from pydantic import ( model_validator, ) -from generalresearch.models import Source +from generalresearch.models.definitions import Source from generalresearch.models.precision import PrecisionStatus from generalresearch.models.thl.demographics import Gender from generalresearch.models.thl.survey import MarketplaceTask diff --git a/generalresearch/models/prodege/question.py b/generalresearch/models/prodege/question.py index c43b51a..b963785 100644 --- a/generalresearch/models/prodege/question.py +++ b/generalresearch/models/prodege/question.py @@ -18,7 +18,7 @@ from pydantic import ( ) from generalresearch.locales import Localelator -from generalresearch.models import MAX_INT32, Source +from generalresearch.models.definitions import MAX_INT32, Source from generalresearch.models.thl.profiling.marketplace import MarketplaceQuestion if TYPE_CHECKING: diff --git a/generalresearch/models/prodege/survey.py b/generalresearch/models/prodege/survey.py index 7e56a9c..e3c765e 100644 --- a/generalresearch/models/prodege/survey.py +++ b/generalresearch/models/prodege/survey.py @@ -20,7 +20,11 @@ from pydantic import ( ) from generalresearch.locales import Localelator -from generalresearch.models import LogicalOperator, Source, TaskCalculationType +from generalresearch.models.definitions import ( + LogicalOperator, + Source, + TaskCalculationType, +) from generalresearch.models.prodege import ( ProdegePastParticipationType, ProdegeStatus, diff --git a/generalresearch/models/repdata/question.py b/generalresearch/models/repdata/question.py index 8cb1fa7..a578741 100644 --- a/generalresearch/models/repdata/question.py +++ b/generalresearch/models/repdata/question.py @@ -17,7 +17,7 @@ from pydantic import ( model_validator, ) -from generalresearch.models import MAX_INT32, Source +from generalresearch.models.definitions import MAX_INT32, Source from generalresearch.models.thl.profiling.marketplace import MarketplaceQuestion if TYPE_CHECKING: diff --git a/generalresearch/models/repdata/survey.py b/generalresearch/models/repdata/survey.py index cea61ed..fc1b649 100644 --- a/generalresearch/models/repdata/survey.py +++ b/generalresearch/models/repdata/survey.py @@ -21,7 +21,7 @@ from pydantic import ( from generalresearch.grpc import timestamp_from_datetime from generalresearch.locales import Localelator -from generalresearch.models import ( +from generalresearch.models.definitions import ( DeviceType, LogicalOperator, Source, diff --git a/generalresearch/models/repdata/task_collection.py b/generalresearch/models/repdata/task_collection.py index 04d79bd..f2cb63b 100644 --- a/generalresearch/models/repdata/task_collection.py +++ b/generalresearch/models/repdata/task_collection.py @@ -6,7 +6,7 @@ import pandas as pd from pandera.pandas import Check, Column, DataFrameSchema, Index from generalresearch.locales import Localelator -from generalresearch.models import TaskCalculationType +from generalresearch.models.definitions import TaskCalculationType from generalresearch.models.repdata import RepDataStatus from generalresearch.models.thl.survey.task_collection import ( TaskCollection, diff --git a/generalresearch/models/sago/question.py b/generalresearch/models/sago/question.py index cf9ea19..bb51d31 100644 --- a/generalresearch/models/sago/question.py +++ b/generalresearch/models/sago/question.py @@ -18,7 +18,7 @@ from pydantic import ( model_validator, ) -from generalresearch.models import MAX_INT32, Source, string_utils +from generalresearch.models.definitions import MAX_INT32, Source, string_utils from generalresearch.models.thl.profiling.marketplace import MarketplaceQuestion if TYPE_CHECKING: diff --git a/generalresearch/models/sago/survey.py b/generalresearch/models/sago/survey.py index c2f886a..c9bf431 100644 --- a/generalresearch/models/sago/survey.py +++ b/generalresearch/models/sago/survey.py @@ -18,7 +18,7 @@ from pydantic import ( ) from generalresearch.locales import Localelator -from generalresearch.models import LogicalOperator, Source +from generalresearch.models.definitions import LogicalOperator, Source from generalresearch.models.sago import SagoStatus from generalresearch.models.thl.demographics import Gender from generalresearch.models.thl.survey import MarketplaceTask diff --git a/generalresearch/models/spectrum/question.py b/generalresearch/models/spectrum/question.py index 89fbeb3..9c9bfa0 100644 --- a/generalresearch/models/spectrum/question.py +++ b/generalresearch/models/spectrum/question.py @@ -18,7 +18,7 @@ from pydantic import ( model_validator, ) -from generalresearch.models import MAX_INT32, Source, string_utils +from generalresearch.models.definitions import MAX_INT32, Source, string_utils from generalresearch.models.thl.profiling.marketplace import ( MarketplaceQuestion, ) diff --git a/generalresearch/models/spectrum/survey.py b/generalresearch/models/spectrum/survey.py index 4daa00b..a02c510 100644 --- a/generalresearch/models/spectrum/survey.py +++ b/generalresearch/models/spectrum/survey.py @@ -10,7 +10,7 @@ from more_itertools import flatten from pydantic import BaseModel, ConfigDict, Field, computed_field, model_validator from generalresearch.locales import Localelator -from generalresearch.models import Source, TaskCalculationType +from generalresearch.models.definitions import Source, TaskCalculationType from generalresearch.models.spectrum import SpectrumStatus from generalresearch.models.thl.demographics import Gender from generalresearch.models.thl.survey import MarketplaceTask diff --git a/generalresearch/models/spectrum/task_collection.py b/generalresearch/models/spectrum/task_collection.py index d909292..8e49434 100644 --- a/generalresearch/models/spectrum/task_collection.py +++ b/generalresearch/models/spectrum/task_collection.py @@ -6,7 +6,7 @@ import pandas as pd from pandera.pandas import Check, Column, DataFrameSchema, Index from generalresearch.locales import Localelator -from generalresearch.models import TaskCalculationType +from generalresearch.models.definitions import TaskCalculationType from generalresearch.models.spectrum import SpectrumStatus from generalresearch.models.thl.survey.task_collection import ( TaskCollection, diff --git a/generalresearch/models/thl/__init__.py b/generalresearch/models/thl/__init__.py index cb04b29..7f2b8a9 100644 --- a/generalresearch/models/thl/__init__.py +++ b/generalresearch/models/thl/__init__.py @@ -8,7 +8,7 @@ from decimal import Decimal # BrokerageProductPayoutEvent, # PayoutEvent, # ) -# from generalresearch.models.thl.product import Product +from generalresearch.models.thl.product import Product # _ = ( # Product, @@ -18,7 +18,7 @@ from decimal import Decimal # POPFinancial, # ) -# Product.model_rebuild() +Product.model_rebuild() # PayoutEvent.model_rebuild() # BrokerageProductPayoutEvent.model_rebuild() diff --git a/generalresearch/models/thl/category.py b/generalresearch/models/thl/category.py index 32841a5..ebfc840 100644 --- a/generalresearch/models/thl/category.py +++ b/generalresearch/models/thl/category.py @@ -1,12 +1,11 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Any, Self +from typing import Any, Self from uuid import uuid4 from pydantic import BaseModel, Field, PositiveInt, model_validator -if TYPE_CHECKING: - from generalresearch.models.custom_types import UUIDStr +from generalresearch.models.custom_types import UUIDStr class Category(BaseModel, frozen=True): diff --git a/generalresearch/models/thl/contest/contest_entry.py b/generalresearch/models/thl/contest/contest_entry.py index a57b2df..17b288b 100644 --- a/generalresearch/models/thl/contest/contest_entry.py +++ b/generalresearch/models/thl/contest/contest_entry.py @@ -12,10 +12,12 @@ from pydantic import ( ) from generalresearch.currency import USDCent +from generalresearch.models.thl.contest.definitions import ( + ContestEntryType, +) if TYPE_CHECKING: from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr - from generalresearch.models.thl.contest.definitions import ContestEntryType from generalresearch.models.thl.user import User @@ -59,10 +61,7 @@ class ContestEntry(BaseModel): @model_validator(mode="before") @classmethod - def validate_amount_type(cls, data: dict) -> dict: - from generalresearch.models.thl.contest.definitions import ( - ContestEntryType, - ) + def validate_amount_type(cls, data: dict[str, Any]) -> dict[str, Any]: amount = data.get("amount") entry_type = data.get("entry_type") @@ -71,6 +70,7 @@ class ContestEntry(BaseModel): assert isinstance(amount, int) and not isinstance( amount, USDCent ), "amount must be int in ContestEntryType.COUNT" + elif entry_type == ContestEntryType.CASH: # This may be coming from the DB, in which case it is an int. data["amount"] = USDCent(data["amount"]) @@ -79,9 +79,6 @@ class ContestEntry(BaseModel): @computed_field() def amount_str(self) -> str: - from generalresearch.models.thl.contest.definitions import ( - ContestEntryType, - ) if self.entry_type == ContestEntryType.COUNT: return str(self.amount) diff --git a/generalresearch/models/thl/contest/raffle.py b/generalresearch/models/thl/contest/raffle.py index b944740..21bc481 100644 --- a/generalresearch/models/thl/contest/raffle.py +++ b/generalresearch/models/thl/contest/raffle.py @@ -26,11 +26,9 @@ from generalresearch.models.thl.contest.contest import ( ContestBase, ContestUserView, ) -from generalresearch.models.thl.contest.contest_entry import ( - ContestEntryType, -) from generalresearch.models.thl.contest.definitions import ( ContestEndReason, + ContestEntryType, ContestPrizeKind, ContestStatus, ContestType, diff --git a/generalresearch/models/thl/finance.py b/generalresearch/models/thl/finance.py index 4b750da..9e7d2c3 100644 --- a/generalresearch/models/thl/finance.py +++ b/generalresearch/models/thl/finance.py @@ -27,8 +27,9 @@ adjustment_example = random.randint(-1_000, 50 * 100) if TYPE_CHECKING: from generalresearch.currency import USDCent + from generalresearch.managers.thl.product import ProductManager from generalresearch.models.thl.ledger import LedgerAccount - from generalresearch.pg_helper import PostgresConfig + from generalresearch.models.thl.product import Product class AdjustmentType(BaseModel): @@ -516,7 +517,7 @@ class ProductBalances(BaseModel): if isinstance(input_data, pd.Series): return ProductBalances.model_validate(input_data.to_dict()) - elif isinstance(input_data, pd.DataFrame): + else: assert isinstance(input_data.index, pd.DatetimeIndex), "Invalid input data" # The pop merge is grouped by 1min intervals. Therefore, if we take @@ -529,9 +530,6 @@ class ProductBalances(BaseModel): pb.last_event = pq_last_event_close.to_pydatetime() return pb - else: - raise NotImplementedError("Can't handle this input") - def __str__(self) -> str: return ( f"Product: {self.product_id or '—'}\n" @@ -834,19 +832,17 @@ class BusinessBalances(BaseModel): def from_pandas( input_data: pd.DataFrame, accounts: list[LedgerAccount], - thl_pg_config: PostgresConfig, + product_manager: ProductManager, ) -> BusinessBalances: LOG.debug(f"BusinessBalances.from_pandas(input_data={input_data.shape})") from generalresearch.incite.schemas.mergers.pop_ledger import ( numerical_col_names, ) - from generalresearch.managers.thl.product import ProductManager from generalresearch.models.thl.ledger import ( AccountType, Direction, ) - from generalresearch.models.thl.product import Product # Validate the input accounts assert len(accounts) > 0, "Must provide accounts" @@ -872,8 +868,7 @@ class BusinessBalances(BaseModel): # Sort the ProductBalances so that they're always in a consistent # sorted order. - pm = ProductManager(pg_config=thl_pg_config) - products: list[Product] = pm.get_by_uuids( + products: list[Product] = product_manager.get_by_uuids( product_uuids=[pb.product_id for pb in product_balances] ) sorted_products_uuids = [ diff --git a/generalresearch/models/thl/ledger.py b/generalresearch/models/thl/ledger.py index c38e83b..fbfb6bb 100644 --- a/generalresearch/models/thl/ledger.py +++ b/generalresearch/models/thl/ledger.py @@ -354,7 +354,7 @@ class LedgerTransaction(BaseModel): def to_user_tx( self, user_account: LedgerAccount, product_id: str, payout_format: str ): - from generalresearch.models.thl.wallet import PayoutType + from generalresearch.models.thl.wallet.definitions import PayoutType d = self.model_dump(include={"created"}) d["tx_type"] = self.metadata.get("tx_type") diff --git a/generalresearch/models/thl/offerwall/__init__.py b/generalresearch/models/thl/offerwall/__init__.py index 0c3d51d..599cc1d 100644 --- a/generalresearch/models/thl/offerwall/__init__.py +++ b/generalresearch/models/thl/offerwall/__init__.py @@ -14,8 +14,8 @@ from pydantic import ( model_validator, ) -from generalresearch.models import Source from generalresearch.models.custom_types import IPvAnyAddressStr +from generalresearch.models.definitions import Source from generalresearch.models.thl.locales import ( CountryISO, LanguageISO, diff --git a/generalresearch/models/thl/offerwall/base.py b/generalresearch/models/thl/offerwall/base.py index 1d41ef2..fb0bc77 100644 --- a/generalresearch/models/thl/offerwall/base.py +++ b/generalresearch/models/thl/offerwall/base.py @@ -19,7 +19,7 @@ from pydantic import ( model_validator, ) -from generalresearch.models import Source +from generalresearch.models.definitions import Source from generalresearch.models.legacy.bucket import ( Bucket as LegacyBucket, ) diff --git a/generalresearch/models/thl/offerwall/cache.py b/generalresearch/models/thl/offerwall/cache.py index aa18014..2a733c9 100644 --- a/generalresearch/models/thl/offerwall/cache.py +++ b/generalresearch/models/thl/offerwall/cache.py @@ -6,8 +6,8 @@ from typing import TYPE_CHECKING, Any from pydantic import BaseModel, Field if TYPE_CHECKING: - from generalresearch.models import Source from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr + from generalresearch.models.definitions import Source from generalresearch.models.thl.offerwall import OfferWallRequest from generalresearch.models.thl.offerwall.base import ( OfferwallBase, diff --git a/generalresearch/models/thl/payout.py b/generalresearch/models/thl/payout.py index 9902af3..128723b 100644 --- a/generalresearch/models/thl/payout.py +++ b/generalresearch/models/thl/payout.py @@ -18,7 +18,7 @@ from pydantic.json_schema import SkipJsonSchema from generalresearch.currency import USDCent from generalresearch.models.thl.definitions import PayoutStatus -from generalresearch.models.thl.wallet import PayoutType +from generalresearch.models.thl.wallet.definitions import PayoutType if TYPE_CHECKING: from generalresearch.models.custom_types import ( diff --git a/generalresearch/models/thl/product.py b/generalresearch/models/thl/product.py index 988b72d..346a98b 100644 --- a/generalresearch/models/thl/product.py +++ b/generalresearch/models/thl/product.py @@ -38,13 +38,13 @@ from pydantic.json_schema import SkipJsonSchema from generalresearch.currency import USDCent from generalresearch.decorators import LOG -from generalresearch.models import Source from generalresearch.models.custom_types import ( AwareDatetimeISO, CountryISOLike, HttpsUrlStr, UUIDStr, ) +from generalresearch.models.definitions import Source from generalresearch.models.thl.finance import ( POPFinancial, ProductBalances, @@ -63,7 +63,7 @@ from generalresearch.models.thl.payout_format import ( examples as payout_format_examples, ) from generalresearch.models.thl.supplier_tag import SupplierTag -from generalresearch.models.thl.wallet import PayoutType +from generalresearch.models.thl.wallet.definitions import PayoutType from generalresearch.models.utils import decimal_to_usd_cents from generalresearch.redis_helper import RedisConfig diff --git a/generalresearch/models/thl/profiling/marketplace.py b/generalresearch/models/thl/profiling/marketplace.py index 0c1e39b..23501e3 100644 --- a/generalresearch/models/thl/profiling/marketplace.py +++ b/generalresearch/models/thl/profiling/marketplace.py @@ -7,16 +7,16 @@ from typing import TYPE_CHECKING, Any from pydantic import BaseModel, ConfigDict, Field, PositiveInt, computed_field -from generalresearch.models import MAX_INT32 +from generalresearch.models.definitions import MAX_INT32 if TYPE_CHECKING: - from generalresearch.models import Source from generalresearch.models.custom_types import ( AwareDatetimeISO, CountryISOLike, LanguageISOLike, UUIDStr, ) + from generalresearch.models.definitions import Source from generalresearch.models.thl.locales import CountryISO, LanguageISO diff --git a/generalresearch/models/thl/profiling/upk_question.py b/generalresearch/models/thl/profiling/upk_question.py index a73683c..9c7383a 100644 --- a/generalresearch/models/thl/profiling/upk_question.py +++ b/generalresearch/models/thl/profiling/upk_question.py @@ -5,7 +5,7 @@ import json import re from enum import StrEnum from functools import cached_property -from typing import TYPE_CHECKING, Annotated, Any, Literal +from typing import Annotated, Any, Literal from pydantic import ( BaseModel, @@ -17,12 +17,10 @@ from pydantic import ( model_validator, ) -from generalresearch.models import Source +from generalresearch.models.custom_types import UUIDStr +from generalresearch.models.definitions import Source from generalresearch.models.thl.category import Category -if TYPE_CHECKING: - from generalresearch.models.custom_types import UUIDStr - class UPKImportance(BaseModel): task_count: int | None = Field( diff --git a/generalresearch/models/thl/profiling/upk_question_answer.py b/generalresearch/models/thl/profiling/upk_question_answer.py index 41895b1..4d07970 100644 --- a/generalresearch/models/thl/profiling/upk_question_answer.py +++ b/generalresearch/models/thl/profiling/upk_question_answer.py @@ -13,7 +13,7 @@ from pydantic import ( model_validator, ) -from generalresearch.models import MAX_INT32 +from generalresearch.models.definitions import MAX_INT32 from generalresearch.models.thl.profiling.upk_property import ( Cardinality, PropertyType, diff --git a/generalresearch/models/thl/profiling/user_info.py b/generalresearch/models/thl/profiling/user_info.py index c82e2d2..40b4b17 100644 --- a/generalresearch/models/thl/profiling/user_info.py +++ b/generalresearch/models/thl/profiling/user_info.py @@ -6,8 +6,8 @@ from pydantic import BaseModel, ConfigDict, Field from pydantic.json_schema import SkipJsonSchema if TYPE_CHECKING: - from generalresearch.models import Source from generalresearch.models.custom_types import AwareDatetimeISO + from generalresearch.models.definitions import Source from generalresearch.models.thl.profiling.user_question_answer import ( MarketplaceResearchProfileQuestion, ) diff --git a/generalresearch/models/thl/profiling/user_question_answer.py b/generalresearch/models/thl/profiling/user_question_answer.py index b1868b3..a7c2194 100644 --- a/generalresearch/models/thl/profiling/user_question_answer.py +++ b/generalresearch/models/thl/profiling/user_question_answer.py @@ -14,12 +14,12 @@ from pydantic import ( model_validator, ) -from generalresearch.models import MAX_INT32 +from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr +from generalresearch.models.definitions import MAX_INT32 +from generalresearch.models.thl.locales import CountryISO, LanguageISO if TYPE_CHECKING: - from generalresearch.models import Source - from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr - from generalresearch.models.thl.locales import CountryISO, LanguageISO + from generalresearch.models.definitions import Source from generalresearch.models.thl.profiling.upk_question import UpkQuestion diff --git a/generalresearch/models/thl/session.py b/generalresearch/models/thl/session.py index 871e5c4..404cff7 100644 --- a/generalresearch/models/thl/session.py +++ b/generalresearch/models/thl/session.py @@ -18,7 +18,7 @@ from pydantic import ( model_validator, ) -from generalresearch.models import Source +from generalresearch.models.definitions import Source from generalresearch.models.thl import ( decimal_to_int_cents, int_cents_to_decimal, @@ -37,13 +37,13 @@ if TYPE_CHECKING: from generalresearch.managers.thl.ledger_manager.thl_ledger import ( ThlLedgerManager, ) - from generalresearch.models import DeviceType from generalresearch.models.custom_types import ( AwareDatetimeISO, EnumNameSerializer, IPvAnyAddressStr, UUIDStr, ) + from generalresearch.models.definitions import DeviceType from generalresearch.models.legacy.bucket import Bucket from generalresearch.models.thl.definitions import ( ReportValue, diff --git a/generalresearch/models/thl/soft_pair.py b/generalresearch/models/thl/soft_pair.py index 313d374..c0bf2dd 100644 --- a/generalresearch/models/thl/soft_pair.py +++ b/generalresearch/models/thl/soft_pair.py @@ -5,7 +5,7 @@ from enum import Enum from typing import TYPE_CHECKING if TYPE_CHECKING: - from generalresearch.models import Source + from generalresearch.models.definitions import Source from generalresearch.models.dynata.survey import DynataCondition from generalresearch.models.thl.survey.condition import ( MarketplaceCondition, diff --git a/generalresearch/models/thl/survey/__init__.py b/generalresearch/models/thl/survey/__init__.py index d0f2b33..76f819e 100644 --- a/generalresearch/models/thl/survey/__init__.py +++ b/generalresearch/models/thl/survey/__init__.py @@ -18,7 +18,7 @@ from generalresearch.models.thl.survey.condition import ( ) if TYPE_CHECKING: - from generalresearch.models import Source + from generalresearch.models.definitions import Source from generalresearch.models.thl.locales import ( CountryISO, CountryISOs, diff --git a/generalresearch/models/thl/survey/buyer.py b/generalresearch/models/thl/survey/buyer.py index 26846d3..ef309d1 100644 --- a/generalresearch/models/thl/survey/buyer.py +++ b/generalresearch/models/thl/survey/buyer.py @@ -16,7 +16,7 @@ from pydantic import ( ) from scipy.stats import beta as beta_dist -from generalresearch.models import Source +from generalresearch.models.definitions import Source if TYPE_CHECKING: from generalresearch.models.custom_types import ( diff --git a/generalresearch/models/thl/survey/condition.py b/generalresearch/models/thl/survey/condition.py index 514ee64..90cf27b 100644 --- a/generalresearch/models/thl/survey/condition.py +++ b/generalresearch/models/thl/survey/condition.py @@ -17,7 +17,7 @@ from pydantic import ( model_validator, ) -from generalresearch.models import LogicalOperator +from generalresearch.models.definitions import LogicalOperator MarketplaceConditionHash = Annotated[ str, StringConstraints(min_length=7, max_length=7, pattern=r"^[a-f0-9]+$") diff --git a/generalresearch/models/thl/survey/model.py b/generalresearch/models/thl/survey/model.py index 9fa3d8e..8986e4d 100644 --- a/generalresearch/models/thl/survey/model.py +++ b/generalresearch/models/thl/survey/model.py @@ -21,7 +21,6 @@ from generalresearch.models.thl.definitions import StatusCode1 from generalresearch.models.thl.pagination import Page if TYPE_CHECKING: - from generalresearch.models import Source from generalresearch.models.custom_types import ( AwareDatetimeISO, CountryISOLike, @@ -29,6 +28,7 @@ if TYPE_CHECKING: PropertyCode, SurveyKey, ) + from generalresearch.models.definitions import Source from generalresearch.models.thl.category import Category from generalresearch.models.thl.definitions import Status diff --git a/generalresearch/models/thl/survey/penalty.py b/generalresearch/models/thl/survey/penalty.py index 54edb94..25e07cf 100644 --- a/generalresearch/models/thl/survey/penalty.py +++ b/generalresearch/models/thl/survey/penalty.py @@ -7,11 +7,11 @@ from typing import TYPE_CHECKING, Annotated, Literal from pydantic import BaseModel, ConfigDict, Field, TypeAdapter if TYPE_CHECKING: - from generalresearch.models import Source from generalresearch.models.custom_types import ( AwareDatetimeISO, UUIDStr, ) + from generalresearch.models.definitions import Source class SurveyPenalty(BaseModel, abc.ABC): diff --git a/generalresearch/models/thl/task_adjustment.py b/generalresearch/models/thl/task_adjustment.py index fa5592e..fee2007 100644 --- a/generalresearch/models/thl/task_adjustment.py +++ b/generalresearch/models/thl/task_adjustment.py @@ -7,14 +7,14 @@ from uuid import uuid4 from pydantic import BaseModel, ConfigDict, Field, PositiveInt, model_validator -from generalresearch.models import MAX_INT32 +from generalresearch.models.definitions import MAX_INT32 from generalresearch.models.thl.definitions import ( WallAdjustedStatus, ) if TYPE_CHECKING: - from generalresearch.models import Source from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr + from generalresearch.models.definitions import Source class TaskAdjustmentEvent(BaseModel): diff --git a/generalresearch/models/thl/user.py b/generalresearch/models/thl/user.py index 302aa72..1f88dc6 100644 --- a/generalresearch/models/thl/user.py +++ b/generalresearch/models/thl/user.py @@ -20,7 +20,7 @@ from pydantic import ( ) from sentry_sdk import set_tag, set_user -from generalresearch.models import MAX_INT32 +from generalresearch.models.definitions import MAX_INT32 if TYPE_CHECKING: from generalresearch.managers.thl.ledger_manager.thl_ledger import ( @@ -253,7 +253,7 @@ class User(BaseModel): # # Delete from db.thl-marketplaces # We need DELETE credentials for all these... - # from generalresearch.models import Source + # from generalresearch.models.definitions import Source # mp_db_table = { # Source.SPECTRUM: "`thl-spectrum`.`spectrum_marketresearchprofilequestion`", # Source.INNOVATE: "`thl-innovate`.`innovate_marketresearchprofilequestion`", diff --git a/generalresearch/models/thl/user_profile.py b/generalresearch/models/thl/user_profile.py index 2dc19b7..c47c6f2 100644 --- a/generalresearch/models/thl/user_profile.py +++ b/generalresearch/models/thl/user_profile.py @@ -13,7 +13,7 @@ from pydantic import ( ) from pydantic.json_schema import SkipJsonSchema -from generalresearch.models import MAX_INT32, Source +from generalresearch.models.definitions import MAX_INT32, Source if TYPE_CHECKING: from generalresearch.models.custom_types import UUIDStr diff --git a/generalresearch/models/thl/user_quality_event.py b/generalresearch/models/thl/user_quality_event.py index 5438740..8c2e25f 100644 --- a/generalresearch/models/thl/user_quality_event.py +++ b/generalresearch/models/thl/user_quality_event.py @@ -7,7 +7,7 @@ from typing import TYPE_CHECKING, Literal from pydantic import BaseModel, Field, PositiveInt -from generalresearch.models import MAX_INT32, Source +from generalresearch.models.definitions import MAX_INT32, Source from generalresearch.utils.enum import ReprEnumMeta if TYPE_CHECKING: diff --git a/generalresearch/models/thl/user_streak.py b/generalresearch/models/thl/user_streak.py index 6cd853a..4c09d13 100644 --- a/generalresearch/models/thl/user_streak.py +++ b/generalresearch/models/thl/user_streak.py @@ -19,7 +19,7 @@ from pydantic import ( from pydantic.json_schema import SkipJsonSchema from generalresearch.managers.leaderboard import country_timezone -from generalresearch.models import MAX_INT32 +from generalresearch.models.definitions import MAX_INT32 if TYPE_CHECKING: from generalresearch.models.thl.locales import CountryISO diff --git a/generalresearch/models/thl/wallet/cashout_method.py b/generalresearch/models/thl/wallet/cashout_method.py index 1db85e8..9383c36 100644 --- a/generalresearch/models/thl/wallet/cashout_method.py +++ b/generalresearch/models/thl/wallet/cashout_method.py @@ -19,7 +19,7 @@ from pydantic import ( from generalresearch.models.legacy.api_status import StatusResponse from generalresearch.models.thl.definitions import PayoutStatus -from generalresearch.models.thl.wallet import PayoutType +from generalresearch.models.thl.wallet.definitions import PayoutType from generalresearch.utils.enum import ReprEnumMeta if TYPE_CHECKING: @@ -31,7 +31,7 @@ if TYPE_CHECKING: ) from generalresearch.models.thl.locales import CountryISO from generalresearch.models.thl.user import BPUIDStr, User - from generalresearch.models.thl.wallet import Currency + from generalresearch.models.thl.wallet.definitions import Currency logger = logging.getLogger() diff --git a/generalresearch/models/thl/wallet/definitions.py b/generalresearch/models/thl/wallet/definitions.py new file mode 100644 index 0000000..2d1eb8d --- /dev/null +++ b/generalresearch/models/thl/wallet/definitions.py @@ -0,0 +1,87 @@ +from enum import StrEnum + +from generalresearch.utils.enum import ReprEnumMeta + + +class PayoutType(StrEnum, metaclass=ReprEnumMeta): + """ + The method in which the requested payout is delivered. + """ + + # The max size of the db field that holds this value is 14, so please + # don't add new values longer than that! + + # User is paid out to their personal PayPal email address + PAYPAL = "PAYPAL" + # User is paid out via a Tango Gift Card + TANGO = "TANGO" + # DWOLLA + DWOLLA = "DWOLLA" + # A payment is made to a bank account using ACH + ACH = "ACH" + # A payment is made to a bank account using ACH + WIRE = "WIRE" + # A payment is made in cash and mailed to the user. + CASH_IN_MAIL = "CASH_IN_MAIL" + # A payment is made as a prize with some monetary value + PRIZE = "PRIZE" + + # This is used to designate either AMT_BONUS or AMT_HIT + AMT = "AMT" + # Amazon Mechanical Turk as a Bonus + AMT_BONUS = "AMT_BONUS" + # Amazon Mechanical Turk for a HIT + AMT_HIT = "AMT_ASSIGNMENT" + AMT_ASSIGNMENT = "AMT_ASSIGNMENT" + + +class Currency(StrEnum): + # United States Dollar + USD = "USD" + # Canadian Dollar + CAD = "CAD" + # British Pound Sterling + GBP = "GBP" + # Euro + EUR = "EUR" + # Indian Rupee + INR = "INR" + # Australian Dollar + AUD = "AUD" + # Polish Zloty + PLN = "PLN" + # Swedish Krona + SEK = "SEK" + # Singapore Dollar + SGD = "SGD" + # Mexican Peso + MXN = "MXN" + + +CURRENCY_FORMATTER = { + "USD": lambda x: f"${x / 100:,.2f}", + "CAD": lambda x: f"${x / 100:,.2f} CAD", + "GBP": lambda x: f"{x / 100:,.2f} £", + "EUR": lambda x: f"€{x / 100:,.2f}", + "INR": lambda x: f"₹{x / 100:,.2f}", + "AUD": lambda x: f"${x / 100:,.2f} AUD", + "PLN": lambda x: f"{x / 100:,.2f} zł", + "SEK": lambda x: f"{x / 100:,.2f} kr", + "SGD": lambda x: f"${x / 100:,.2f} SGD", + "MXN": lambda x: f"${x / 100:,.2f} MXN", +} + +# The max value user can redeem in one go in foreign currencies. should be < $250 +# in order to avoid exchange rate issues +CURRENCY_MAX_VALUE = { + "USD": 250, + "CAD": 200, + "GBP": 100, + "EUR": 100, + "INR": 10000, + "AUD": 200, + "PLN": 500, + "SEK": 1000, + "SGD": 200, + "MXN": 4000, +} diff --git a/generalresearch/models/thl/wallet/payout.py b/generalresearch/models/thl/wallet/payout.py index 7301b31..79c50e1 100644 --- a/generalresearch/models/thl/wallet/payout.py +++ b/generalresearch/models/thl/wallet/payout.py @@ -16,7 +16,7 @@ from pydantic import ( from generalresearch.currency import USDCent from generalresearch.models.thl.definitions import PayoutStatus -from generalresearch.models.thl.wallet import PayoutType +from generalresearch.models.thl.wallet.definitions import PayoutType if TYPE_CHECKING: from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr diff --git a/generalresearch/schemas/survey_stats.py b/generalresearch/schemas/survey_stats.py index b3acf34..dd592d4 100644 --- a/generalresearch/schemas/survey_stats.py +++ b/generalresearch/schemas/survey_stats.py @@ -2,7 +2,7 @@ import pandas as pd from pandera.pandas import Check, Column, DataFrameSchema, Index from generalresearch.locales import Localelator -from generalresearch.models import Source +from generalresearch.models.definitions import Source COUNTRY_ISOS = Localelator().get_all_countries() kosovo = "xk" diff --git a/generalresearch/wall_status_codes/__init__.py b/generalresearch/wall_status_codes/__init__.py index 37f3960..cca1a19 100644 --- a/generalresearch/wall_status_codes/__init__.py +++ b/generalresearch/wall_status_codes/__init__.py @@ -1,6 +1,6 @@ from typing import TYPE_CHECKING -from generalresearch.models import Source +from generalresearch.models.definitions import Source from generalresearch.models.thl.definitions import Status, StatusCode1 from generalresearch.wall_status_codes import ( cint, diff --git a/test_utils/conftest.py b/test_utils/conftest.py index f55fe11..397d98f 100644 --- a/test_utils/conftest.py +++ b/test_utils/conftest.py @@ -9,6 +9,7 @@ from collections.abc import Callable, Generator from datetime import UTC, datetime, timedelta from os.path import join as pjoin from pathlib import Path +from typing import TYPE_CHECKING from uuid import uuid4 import pytest @@ -17,23 +18,13 @@ from dotenv import load_dotenv from pydantic import MariaDBDsn, PostgresDsn, TypeAdapter from pytest import TempPathFactory -from generalresearch.config import GRLBaseSettings from generalresearch.currency import USDCent from generalresearch.models.custom_types import InternalHostname, PostgresDict -from generalresearch.pg_helper import PostgresConfig from generalresearch.sql_helper import SqlHelper -# -- redis notes from jenkins file -# sh "redis-cli -u ${env.THL_REDIS} FLUSHDB" -# sh "redis-cli -u ${env.GR_REDIS} FLUSHDB" - -# script { -# env.GR_REDIS_DB = new Random().nextInt(1024).toString() -# env.GR_REDIS = "redis://${env.REDIS}:6379/${env.GR_REDIS_DB}" -# echo "Using GR Redis: ${env.GR_REDIS}" -# if (sh(script: "redis-cli -u ${env.GR_REDIS} SET jenkins_lock 1 NX EX 3600", returnStdout: true).trim() != 'OK') -# error('Redis already locked... aborting.') -# } +if TYPE_CHECKING: + from generalresearch.config import GRLBaseSettings + from generalresearch.pg_helper import PostgresConfig @pytest.fixture(scope="session") diff --git a/test_utils/grliq/conftest.py b/test_utils/grliq/conftest.py index 891b73c..bb1a167 100644 --- a/test_utils/grliq/conftest.py +++ b/test_utils/grliq/conftest.py @@ -2,19 +2,15 @@ from __future__ import annotations from collections.abc import Callable from datetime import UTC, datetime, timedelta -from typing import Any +from typing import TYPE_CHECKING, Any from uuid import uuid4 import pytest from pydantic import PostgresDsn -from generalresearch.config import GRLBaseSettings from generalresearch.grliq.managers.forensic_data import ( GrlIqDataManager, ) -from generalresearch.grliq.managers.forensic_events import ( - GrlIqEventManager, -) from generalresearch.grliq.managers.forensic_results import ( GrlIqCategoryResultsReader, ) @@ -25,6 +21,12 @@ from generalresearch.grliq.models.forensic_result import ( ) from generalresearch.pg_helper import PostgresConfig +if TYPE_CHECKING: + from generalresearch.config import GRLBaseSettings + from generalresearch.grliq.managers.forensic_events import ( + GrlIqEventManager, + ) + # === Miscellaneous === diff --git a/test_utils/incite/collections/conftest.py b/test_utils/incite/collections/conftest.py index 631bb7b..f490e14 100644 --- a/test_utils/incite/collections/conftest.py +++ b/test_utils/incite/collections/conftest.py @@ -6,12 +6,11 @@ from typing import TYPE_CHECKING import pytest -from generalresearch.pg_helper import PostgresConfig from test_utils.conftest import clear_directory if TYPE_CHECKING: from generalresearch.incite.base import DFCollectionType, GRLDatasets - from generalresearch.incite.collections import DFCollection + from generalresearch.incite.collections.base import DFCollection from generalresearch.incite.collections.thl_web import ( AuditLogDFCollection, LedgerDFCollection, @@ -20,6 +19,7 @@ if TYPE_CHECKING: UserDFCollection, WallDFCollection, ) + from generalresearch.pg_helper import PostgresConfig @pytest.fixture diff --git a/test_utils/incite/mergers/conftest.py b/test_utils/incite/mergers/conftest.py index 1f88804..4eb3f2d 100644 --- a/test_utils/incite/mergers/conftest.py +++ b/test_utils/incite/mergers/conftest.py @@ -2,37 +2,40 @@ from __future__ import annotations from collections.abc import Callable from datetime import datetime, timedelta +from typing import TYPE_CHECKING import pytest -from generalresearch.incite.base import GRLDatasets -from generalresearch.incite.mergers.base import MergeType -from generalresearch.incite.mergers.foundations.enriched_session import ( - EnrichedSessionMerge, -) -from generalresearch.incite.mergers.foundations.enriched_task_adjust import ( - EnrichedTaskAdjustMerge, -) -from generalresearch.incite.mergers.foundations.enriched_wall import ( - EnrichedWallMerge, -) -from generalresearch.incite.mergers.foundations.user_id_product import ( - UserIdProductMerge, -) -from generalresearch.incite.mergers.pop_ledger import ( - PopLedgerMerge, - PopLedgerMergeItem, -) -from generalresearch.incite.mergers.ym_survey_wall import ( - YMSurveyWallMerge, - YMSurveyWallMergeCollectionItem, -) -from generalresearch.incite.mergers.ym_wall_summary import ( - YMWallSummaryMerge, - YMWallSummaryMergeItem, -) from test_utils.conftest import clear_directory +if TYPE_CHECKING: + from generalresearch.incite.base import GRLDatasets + from generalresearch.incite.mergers.base import MergeType + from generalresearch.incite.mergers.foundations.enriched_session import ( + EnrichedSessionMerge, + ) + from generalresearch.incite.mergers.foundations.enriched_task_adjust import ( + EnrichedTaskAdjustMerge, + ) + from generalresearch.incite.mergers.foundations.enriched_wall import ( + EnrichedWallMerge, + ) + from generalresearch.incite.mergers.foundations.user_id_product import ( + UserIdProductMerge, + ) + from generalresearch.incite.mergers.pop_ledger import ( + PopLedgerMerge, + PopLedgerMergeItem, + ) + from generalresearch.incite.mergers.ym_survey_wall import ( + YMSurveyWallMerge, + YMSurveyWallMergeCollectionItem, + ) + from generalresearch.incite.mergers.ym_wall_summary import ( + YMWallSummaryMerge, + YMWallSummaryMergeItem, + ) + # -------------------------- # Merges # -------------------------- diff --git a/test_utils/managers/cashout_methods.py b/test_utils/managers/cashout_methods.py index 238cdda..adf82f4 100644 --- a/test_utils/managers/cashout_methods.py +++ b/test_utils/managers/cashout_methods.py @@ -6,11 +6,11 @@ from uuid import uuid4 import pytest -from generalresearch.models.thl.wallet import Currency, PayoutType from generalresearch.models.thl.wallet.cashout_method import ( CashoutMethod, TangoCashoutMethodData, ) +from generalresearch.models.thl.wallet.definitions import Currency, PayoutType @pytest.fixture(scope="session") diff --git a/test_utils/managers/conftest.py b/test_utils/managers/conftest.py index 9c6a1a7..ed771c7 100644 --- a/test_utils/managers/conftest.py +++ b/test_utils/managers/conftest.py @@ -1,41 +1,44 @@ from __future__ import annotations from collections.abc import Callable +from typing import TYPE_CHECKING import pytest -from generalresearch.managers.gr.business import ( - BusinessAddressManager, - BusinessBankAccountManager, - BusinessManager, -) -from generalresearch.managers.gr.team import ( - MembershipManager, - TeamManager, -) -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, - UserIpHistoryManager, -) -from generalresearch.models import Source -from generalresearch.models.thl.wallet.cashout_method import CashoutMethod -from generalresearch.pg_helper import PostgresConfig -from generalresearch.redis_helper import RedisConfig -from generalresearch.sql_helper import SqlHelper +from generalresearch.models.definitions import Source + +if TYPE_CHECKING: + from generalresearch.managers.gr.business import ( + BusinessAddressManager, + BusinessBankAccountManager, + BusinessManager, + ) + from generalresearch.managers.gr.team import ( + MembershipManager, + TeamManager, + ) + from generalresearch.managers.spectrum.survey import SpectrumSurveyManager + from generalresearch.managers.thl.buyer import BuyerManager + from generalresearch.managers.thl.ipinfo import ( + GeoIpInfoManager, + IPGeonameManager, + IPInformationManager, + ) + from generalresearch.managers.thl.userhealth import ( + AuditLogManager, + IPRecordManager, + UserIpHistoryManager, + ) + from generalresearch.models.thl.wallet.cashout_method import CashoutMethod + from generalresearch.pg_helper import PostgresConfig + from generalresearch.redis_helper import RedisConfig + from generalresearch.sql_helper import SqlHelper # === THL === diff --git a/test_utils/managers/contest/conftest.py b/test_utils/managers/contest/conftest.py index a9375f6..b29cf18 100644 --- a/test_utils/managers/contest/conftest.py +++ b/test_utils/managers/contest/conftest.py @@ -1,10 +1,14 @@ from __future__ import annotations +from typing import TYPE_CHECKING + import pytest from generalresearch.managers.base import Permission from generalresearch.managers.thl.contest_manager import ContestManager -from generalresearch.pg_helper import PostgresConfig + +if TYPE_CHECKING: + from generalresearch.pg_helper import PostgresConfig @pytest.fixture(scope="session") diff --git a/test_utils/managers/gr/conftest.py b/test_utils/managers/gr/conftest.py index 40bd7b3..5392c69 100644 --- a/test_utils/managers/gr/conftest.py +++ b/test_utils/managers/gr/conftest.py @@ -3,6 +3,7 @@ from __future__ import annotations import subprocess from collections.abc import Callable, Generator from random import randint +from typing import TYPE_CHECKING import pytest import redis @@ -10,8 +11,6 @@ import redis.asyncio as redis_async from pydantic import PostgresDsn from redis import Redis -from generalresearch.config import GRLBaseSettings -from generalresearch.managers.gr.authentication import GRTokenManager, GRUserManager from generalresearch.managers.gr.business import ( BusinessAddressManager, BusinessBankAccountManager, @@ -20,6 +19,10 @@ from generalresearch.managers.gr.business import ( from generalresearch.pg_helper import PostgresConfig from generalresearch.redis_helper import RedisConfig +if TYPE_CHECKING: + from generalresearch.config import GRLBaseSettings + from generalresearch.managers.gr.authentication import GRTokenManager, GRUserManager + # === Msc === @pytest.fixture(scope="session") diff --git a/test_utils/managers/ledger/conftest.py b/test_utils/managers/ledger/conftest.py index ce8348e..c60ee1b 100644 --- a/test_utils/managers/ledger/conftest.py +++ b/test_utils/managers/ledger/conftest.py @@ -1,18 +1,24 @@ from __future__ import annotations +from typing import TYPE_CHECKING + import pytest from generalresearch.managers.base import Permission from generalresearch.managers.thl.ledger_manager.ledger import ( - LedgerAccountManager, LedgerManager, - LedgerTransactionManager, ) from generalresearch.managers.thl.ledger_manager.thl_ledger import ( ThlLedgerManager, ) -from generalresearch.pg_helper import PostgresConfig -from generalresearch.redis_helper import RedisConfig + +if TYPE_CHECKING: + from generalresearch.managers.thl.ledger_manager.ledger import ( + LedgerAccountManager, + LedgerTransactionManager, + ) + from generalresearch.pg_helper import PostgresConfig + from generalresearch.redis_helper import RedisConfig # --- Ledger --- diff --git a/test_utils/managers/thl/conftest.py b/test_utils/managers/thl/conftest.py index d40b7d2..af3fd23 100644 --- a/test_utils/managers/thl/conftest.py +++ b/test_utils/managers/thl/conftest.py @@ -1,44 +1,47 @@ from __future__ import annotations from collections.abc import Callable +from typing import TYPE_CHECKING import pytest from pydantic import PostgresDsn -from generalresearch.config import GRLBaseSettings from generalresearch.managers.base import Permission -from generalresearch.managers.thl.buyer import BuyerManager -from generalresearch.managers.thl.category import CategoryManager -from generalresearch.managers.thl.payout import ( - BrokerageProductPayoutEventManager, - BusinessPayoutEventManager, - PayoutEventManager, - UserPayoutEventManager, -) -from generalresearch.managers.thl.product import ProductManager -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, -) -from generalresearch.managers.thl.user_manager.user_metadata_manager import ( - UserMetadataManager, -) -from generalresearch.managers.thl.wall import ( - WallCacheManager, - WallManager, -) from generalresearch.pg_helper import PostgresConfig from generalresearch.redis_helper import RedisConfig +if TYPE_CHECKING: + from generalresearch.config import GRLBaseSettings + from generalresearch.managers.thl.buyer import BuyerManager + from generalresearch.managers.thl.category import CategoryManager + from generalresearch.managers.thl.payout import ( + BrokerageProductPayoutEventManager, + BusinessPayoutEventManager, + PayoutEventManager, + UserPayoutEventManager, + ) + from generalresearch.managers.thl.product import ProductManager + from generalresearch.managers.thl.session import SessionManager + from generalresearch.managers.thl.task_adjustment import ( + TaskAdjustmentManager, + ) + from generalresearch.managers.thl.user_manager.user_manager import ( + UserManager, + ) + from generalresearch.managers.thl.user_manager.user_metadata_manager import ( + UserMetadataManager, + ) + from generalresearch.managers.thl.wall import ( + WallCacheManager, + WallManager, + ) + @pytest.fixture(scope="session") def thl_web_rr(django_db_factory: Callable[..., PostgresDsn]) -> PostgresConfig: diff --git a/test_utils/managers/upk/conftest.py b/test_utils/managers/upk/conftest.py index 7eabee1..f581278 100644 --- a/test_utils/managers/upk/conftest.py +++ b/test_utils/managers/upk/conftest.py @@ -1,4 +1,5 @@ from collections.abc import Callable, Generator +from typing import TYPE_CHECKING import pytest @@ -12,9 +13,11 @@ from generalresearch.managers.thl.profiling.uqa import UQAManager from generalresearch.managers.thl.profiling.user_upk import ( UserUpkManager, ) -from generalresearch.models.thl.user import User -from generalresearch.pg_helper import PostgresConfig -from generalresearch.redis_helper import RedisConfig + +if TYPE_CHECKING: + from generalresearch.models.thl.user import User + from generalresearch.pg_helper import PostgresConfig + from generalresearch.redis_helper import RedisConfig @pytest.fixture(scope="session") diff --git a/test_utils/models/conftest.py b/test_utils/models/conftest.py index 3a10ea3..089f2e6 100644 --- a/test_utils/models/conftest.py +++ b/test_utils/models/conftest.py @@ -12,13 +12,12 @@ import pytest from pydantic import AwareDatetime, PositiveInt from pytest import FixtureRequest as Request -from generalresearch.models import Source +from generalresearch.models.definitions import Source from generalresearch.models.thl.definitions import ( WALL_ALLOWED_STATUS_STATUS_CODE, Status, ) from generalresearch.models.thl.survey.model import Buyer, Survey -from generalresearch.pg_helper import PostgresConfig if TYPE_CHECKING: from generalresearch.currency import USDCent @@ -53,6 +52,7 @@ if TYPE_CHECKING: from generalresearch.models.thl.user import User from generalresearch.models.thl.user_iphistory import IPRecord from generalresearch.models.thl.userhealth import AuditLog, AuditLogLevel + from generalresearch.pg_helper import PostgresConfig # === THL === diff --git a/test_utils/models/contest/conftest.py b/test_utils/models/contest/conftest.py index 84930b8..91425dc 100644 --- a/test_utils/models/contest/conftest.py +++ b/test_utils/models/contest/conftest.py @@ -3,36 +3,41 @@ from __future__ import annotations from collections.abc import Callable from datetime import UTC, datetime from decimal import Decimal +from typing import TYPE_CHECKING from uuid import uuid4 import pytest from pytest import FixtureRequest as Request from generalresearch.currency import USDCent -from generalresearch.managers.thl.contest_manager import ContestManager -from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager from generalresearch.models.thl.contest import ( ContestEndCondition, ContestPrize, ) -from generalresearch.models.thl.contest.contest import Contest from generalresearch.models.thl.contest.definitions import ( ContestPrizeKind, ContestType, ) -from generalresearch.models.thl.contest.leaderboard import ( - LeaderboardContestCreate, -) -from generalresearch.models.thl.contest.milestone import ( - MilestoneContestCreate, -) from generalresearch.models.thl.contest.raffle import ( ContestEntryType, - RaffleContest, RaffleContestCreate, ) -from generalresearch.models.thl.product import Product -from generalresearch.models.thl.user import User + +if TYPE_CHECKING: + from generalresearch.managers.thl.contest_manager import ContestManager + from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager + from generalresearch.models.thl.contest.contest import Contest + from generalresearch.models.thl.contest.leaderboard import ( + LeaderboardContestCreate, + ) + from generalresearch.models.thl.contest.milestone import ( + MilestoneContestCreate, + ) + from generalresearch.models.thl.contest.raffle import ( + RaffleContest, + ) + from generalresearch.models.thl.product import Product + from generalresearch.models.thl.user import User # === Miscellaneous === diff --git a/test_utils/models/gr/conftest.py b/test_utils/models/gr/conftest.py index b623255..6c1877a 100644 --- a/test_utils/models/gr/conftest.py +++ b/test_utils/models/gr/conftest.py @@ -2,30 +2,32 @@ from __future__ import annotations from collections.abc import Callable from random import randint +from typing import TYPE_CHECKING from uuid import uuid4 import pytest from pydantic import PositiveInt from pydantic_extra_types.phone_numbers import PhoneNumber -from generalresearch.managers.gr.authentication import GRTokenManager, GRUserManager -from generalresearch.managers.gr.business import ( - BusinessAddressManager, - BusinessBankAccountManager, - BusinessManager, -) -from generalresearch.managers.gr.team import MembershipManager, TeamManager -from generalresearch.models.custom_types import UUIDStr -from generalresearch.models.gr.authentication import GRToken, GRUser -from generalresearch.models.gr.business import ( - Business, - BusinessAddress, - BusinessBankAccount, - TransferMethod, -) -from generalresearch.models.gr.team import Membership, Team -from generalresearch.pg_helper import PostgresConfig -from generalresearch.redis_helper import RedisConfig +if TYPE_CHECKING: + from generalresearch.managers.gr.authentication import GRTokenManager, GRUserManager + from generalresearch.managers.gr.business import ( + BusinessAddressManager, + BusinessBankAccountManager, + BusinessManager, + ) + from generalresearch.managers.gr.team import MembershipManager, TeamManager + from generalresearch.models.custom_types import UUIDStr + from generalresearch.models.gr.authentication import GRToken, GRUser + from generalresearch.models.gr.business import ( + Business, + BusinessAddress, + BusinessBankAccount, + TransferMethod, + ) + from generalresearch.models.gr.team import Membership, Team + from generalresearch.pg_helper import PostgresConfig + from generalresearch.redis_helper import RedisConfig # --- Static --- diff --git a/test_utils/models/ledger/conftest.py b/test_utils/models/ledger/conftest.py index 1c1027c..8437c7f 100644 --- a/test_utils/models/ledger/conftest.py +++ b/test_utils/models/ledger/conftest.py @@ -11,7 +11,6 @@ import pytest from pytest import FixtureRequest as Request from generalresearch.currency import USDCent -from generalresearch.managers.base import PostgresManager from test_utils.models.conftest import ( payout_config, product_amt_true, @@ -24,6 +23,9 @@ from test_utils.models.conftest import ( wall_factory, ) +if TYPE_CHECKING: + from generalresearch.managers.base import PostgresManager + _ = ( user_factory, product_user_wallet_no, diff --git a/test_utils/models/network/conftest.py b/test_utils/models/network/conftest.py index 6ba37a3..4ff59ee 100644 --- a/test_utils/models/network/conftest.py +++ b/test_utils/models/network/conftest.py @@ -1,5 +1,6 @@ import os from datetime import UTC, datetime, timedelta +from typing import TYPE_CHECKING from uuid import uuid4 import pytest @@ -23,7 +24,9 @@ from generalresearch.models.network.tool_run_command import ( RDNSRunCommand, RDNSRunCommandOptions, ) -from generalresearch.pg_helper import PostgresConfig + +if TYPE_CHECKING: + from generalresearch.pg_helper import PostgresConfig @pytest.fixture(scope="session") diff --git a/test_utils/models/thl/conftest.py b/test_utils/models/thl/conftest.py index fc57c73..3c77e27 100644 --- a/test_utils/models/thl/conftest.py +++ b/test_utils/models/thl/conftest.py @@ -5,7 +5,7 @@ from datetime import UTC, datetime from decimal import ROUND_DOWN, Decimal from random import choice as rand_choice from random import randint, random -from typing import Any +from typing import TYPE_CHECKING, Any from uuid import uuid4 import faker @@ -13,47 +13,53 @@ import pytest from grip_client.enums import AccessType from pydantic import PositiveInt -from generalresearch.managers.thl.ipinfo import IPGeonameManager, IPInformationManager -from generalresearch.managers.thl.payout import UserPayoutEventManager -from generalresearch.managers.thl.product import ProductManager -from generalresearch.managers.thl.session import SessionManager -from generalresearch.managers.thl.user_manager.user_manager import UserManager -from generalresearch.managers.thl.userhealth import AuditLogManager, IPRecordManager -from generalresearch.managers.thl.wall import WallManager -from generalresearch.models import DeviceType -from generalresearch.models.custom_types import ( - AwareDatetimeISO, - IPvAnyAddressStr, - UUIDStr, -) -from generalresearch.models.legacy.bucket import Bucket -from generalresearch.models.thl.definitions import ( - PayoutStatus, -) -from generalresearch.models.thl.ipinfo import IPGeoname, IPInformation -from generalresearch.models.thl.payout import UserPayoutEvent -from generalresearch.models.thl.product import ( - PayoutConfig, - Product, - ProfilingConfig, - SessionConfig, - SourcesConfig, - SupplyConfig, - UserCreateConfig, - UserHealthConfig, - UserWalletConfig, -) +from generalresearch.models.thl.definitions import PayoutStatus from generalresearch.models.thl.session import ( - Session, Source, Status, - Wall, ) from generalresearch.models.thl.user import User -from generalresearch.models.thl.user_iphistory import IPRecord -from generalresearch.models.thl.userhealth import AuditLog, AuditLogLevel -from generalresearch.models.thl.wallet import PayoutType -from generalresearch.models.thl.wallet.cashout_method import CashMailOrderData +from generalresearch.models.thl.userhealth import AuditLogLevel +from generalresearch.models.thl.wallet.definitions import PayoutType + +if TYPE_CHECKING: + from generalresearch.managers.thl.ipinfo import ( + IPGeonameManager, + IPInformationManager, + ) + from generalresearch.managers.thl.payout import UserPayoutEventManager + from generalresearch.managers.thl.product import ProductManager + from generalresearch.managers.thl.session import SessionManager + from generalresearch.managers.thl.user_manager.user_manager import UserManager + from generalresearch.managers.thl.userhealth import AuditLogManager, IPRecordManager + from generalresearch.managers.thl.wall import WallManager + from generalresearch.models.custom_types import ( + AwareDatetimeISO, + IPvAnyAddressStr, + UUIDStr, + ) + from generalresearch.models.definitions import DeviceType + from generalresearch.models.legacy.bucket import Bucket + from generalresearch.models.thl.ipinfo import IPGeoname, IPInformation + from generalresearch.models.thl.payout import UserPayoutEvent + from generalresearch.models.thl.product import ( + PayoutConfig, + Product, + ProfilingConfig, + SessionConfig, + SourcesConfig, + SupplyConfig, + UserCreateConfig, + UserHealthConfig, + UserWalletConfig, + ) + from generalresearch.models.thl.session import ( + Session, + Wall, + ) + from generalresearch.models.thl.user_iphistory import IPRecord + from generalresearch.models.thl.userhealth import AuditLog + from generalresearch.models.thl.wallet.cashout_method import CashMailOrderData fake = faker.Faker() diff --git a/test_utils/models/upk/conftest.py b/test_utils/models/upk/conftest.py index ef77dd6..59266b2 100644 --- a/test_utils/models/upk/conftest.py +++ b/test_utils/models/upk/conftest.py @@ -9,10 +9,9 @@ from uuid import UUID import pandas as pd import pytest -from generalresearch.pg_helper import PostgresConfig - if TYPE_CHECKING: from generalresearch.managers.thl.category import CategoryManager + from generalresearch.pg_helper import PostgresConfig def insert_data_from_csv( diff --git a/test_utils/spectrum/conftest.py b/test_utils/spectrum/conftest.py index a8ce9d9..cc91cff 100644 --- a/test_utils/spectrum/conftest.py +++ b/test_utils/spectrum/conftest.py @@ -3,16 +3,15 @@ from __future__ import annotations import time from datetime import UTC, datetime from decimal import Decimal -from typing import Any +from typing import TYPE_CHECKING, Any import pytest -from generalresearch.config import GRLBaseSettings from generalresearch.managers.spectrum.survey import ( SpectrumCriteriaManager, SpectrumSurveyManager, ) -from generalresearch.models import ( +from generalresearch.models.definitions import ( LogicalOperator, ) from generalresearch.models.spectrum.survey import ( @@ -22,6 +21,9 @@ from generalresearch.models.spectrum.survey import ( from generalresearch.models.thl.survey.condition import ConditionValueType from generalresearch.sql_helper import SqlHelper +if TYPE_CHECKING: + from generalresearch.config import GRLBaseSettings + @pytest.fixture(scope="session") def spectrum_rw(settings: GRLBaseSettings) -> SqlHelper: diff --git a/tests/incite/collections/test_df_collection_base.py b/tests/incite/collections/test_df_collection_base.py index c236700..e20b44b 100644 --- a/tests/incite/collections/test_df_collection_base.py +++ b/tests/incite/collections/test_df_collection_base.py @@ -9,10 +9,10 @@ from generalresearch.incite.collections import ( DFCollection, DFCollectionType, ) -from generalresearch.pg_helper import PostgresConfig if TYPE_CHECKING: from generalresearch.incite.base import GRLDatasets + from generalresearch.pg_helper import PostgresConfig df_collection_types = [e for e in DFCollectionType if e is not DFCollectionType.TEST] diff --git a/tests/incite/collections/test_df_collection_item_base.py b/tests/incite/collections/test_df_collection_item_base.py index e0171c2..fd70bf0 100644 --- a/tests/incite/collections/test_df_collection_item_base.py +++ b/tests/incite/collections/test_df_collection_item_base.py @@ -5,15 +5,16 @@ from typing import TYPE_CHECKING import pytest -from generalresearch.incite.collections import ( +from generalresearch.incite.collections.base import ( DFCollection, DFCollectionItem, DFCollectionType, ) -from generalresearch.pg_helper import PostgresConfig if TYPE_CHECKING: from generalresearch.incite.base import GRLDatasets + from generalresearch.pg_helper import PostgresConfig + df_collection_types = [e for e in DFCollectionType if e is not DFCollectionType.TEST] diff --git a/tests/incite/collections/test_df_collection_item_thl_web.py b/tests/incite/collections/test_df_collection_item_thl_web.py index 5f9a3f6..061c576 100644 --- a/tests/incite/collections/test_df_collection_item_thl_web.py +++ b/tests/incite/collections/test_df_collection_item_thl_web.py @@ -5,6 +5,7 @@ from datetime import UTC, datetime, timedelta from itertools import product as iter_product from os.path import join as pjoin from pathlib import Path, PurePath +from typing import TYPE_CHECKING from uuid import uuid4 import dask.dataframe as dd @@ -21,19 +22,24 @@ from faker import Faker from pandera.pandas import DataFrameSchema from pydantic import FilePath -from generalresearch.incite.base import CollectionItemBase, GRLDatasets -from generalresearch.incite.collections import ( - DFCollection, - DFCollectionItem, +from generalresearch.incite.base import CollectionItemBase +from generalresearch.incite.collections.base import ( DFCollectionType, ) from generalresearch.incite.schemas import ARCHIVE_AFTER -from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager -from generalresearch.models.thl.product import Product -from generalresearch.models.thl.user import User from generalresearch.pg_helper import PostgresConfig from generalresearch.sql_helper import PostgresDsn +if TYPE_CHECKING: + from generalresearch.incite.base import GRLDatasets + from generalresearch.incite.collections.base import ( + DFCollection, + DFCollectionItem, + ) + from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager + from generalresearch.models.thl.product import Product + from generalresearch.models.thl.user import User + fake = Faker() df_collections = [ diff --git a/tests/incite/collections/test_df_collection_thl_marketplaces.py b/tests/incite/collections/test_df_collection_thl_marketplaces.py index 6a7e5c9..6ad0cb4 100644 --- a/tests/incite/collections/test_df_collection_thl_marketplaces.py +++ b/tests/incite/collections/test_df_collection_thl_marketplaces.py @@ -1,18 +1,21 @@ from datetime import UTC, datetime from itertools import product +from typing import TYPE_CHECKING import pytest from pandera.pandas import Column, DataFrameSchema, Index -from generalresearch.incite.base import GRLDatasets -from generalresearch.incite.collections import DFCollection, DFCollectionType +from generalresearch.incite.collections.base import DFCollection, DFCollectionType from generalresearch.incite.collections.thl_marketplaces import ( InnovateSurveyHistoryCollection, MorningSurveyTimeseriesCollection, SagoSurveyHistoryCollection, SpectrumSurveyTimeseriesCollection, ) -from generalresearch.pg_helper import PostgresConfig + +if TYPE_CHECKING: + from generalresearch.incite.base import GRLDatasets + from generalresearch.pg_helper import PostgresConfig def combo_object(): diff --git a/tests/incite/collections/test_df_collection_thl_web.py b/tests/incite/collections/test_df_collection_thl_web.py index 6d509bc..20d7187 100644 --- a/tests/incite/collections/test_df_collection_thl_web.py +++ b/tests/incite/collections/test_df_collection_thl_web.py @@ -9,7 +9,7 @@ import pandas as pd import pytest from pandera.pandas import DataFrameSchema -from generalresearch.incite.collections import ( +from generalresearch.incite.collections.base import ( DFCollection, DFCollectionType, ) diff --git a/tests/incite/mergers/foundations/test_enriched_session.py b/tests/incite/mergers/foundations/test_enriched_session.py index 2a161e4..71b2442 100644 --- a/tests/incite/mergers/foundations/test_enriched_session.py +++ b/tests/incite/mergers/foundations/test_enriched_session.py @@ -4,29 +4,32 @@ from collections.abc import Callable from datetime import UTC, datetime, timedelta from decimal import Decimal from itertools import product +from typing import TYPE_CHECKING import dask.dataframe as dd import pandas as pd import pytest from dask.distributed import Client as DaskClient -from generalresearch.incite.collections.thl_web import ( - SessionDFCollection, - WallDFCollection, -) -from generalresearch.incite.mergers.foundations.enriched_session import ( - EnrichedSessionMerge, -) from generalresearch.incite.schemas.admin_responses import ( AdminPOPSessionSchema, ) -from generalresearch.models.admin.request import ( - ReportRequest, -) -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 + +if TYPE_CHECKING: + from generalresearch.incite.collections.thl_web import ( + SessionDFCollection, + WallDFCollection, + ) + from generalresearch.incite.mergers.foundations.enriched_session import ( + EnrichedSessionMerge, + ) + from generalresearch.models.admin.request import ( + ReportRequest, + ) + 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 @pytest.mark.parametrize( diff --git a/tests/incite/mergers/foundations/test_enriched_task_adjust.py b/tests/incite/mergers/foundations/test_enriched_task_adjust.py index 0606b6f..877d22f 100644 --- a/tests/incite/mergers/foundations/test_enriched_task_adjust.py +++ b/tests/incite/mergers/foundations/test_enriched_task_adjust.py @@ -3,26 +3,28 @@ from __future__ import annotations from collections.abc import Callable from datetime import timedelta from itertools import product as iter_product +from typing import TYPE_CHECKING import dask.dataframe as dd import pandas as pd import pytest from dask.distributed import Client as DaskClient -from generalresearch.incite.collections.thl_web import ( - SessionDFCollection, - TaskAdjustmentDFCollection, - WallDFCollection, -) -from generalresearch.incite.mergers.foundations.enriched_task_adjust import ( - EnrichedTaskAdjustMerge, -) -from generalresearch.incite.mergers.foundations.enriched_wall import ( - EnrichedWallMerge, -) -from generalresearch.models.thl.product import Product -from generalresearch.models.thl.user import User -from generalresearch.pg_helper import PostgresConfig +if TYPE_CHECKING: + from generalresearch.incite.collections.thl_web import ( + SessionDFCollection, + TaskAdjustmentDFCollection, + WallDFCollection, + ) + from generalresearch.incite.mergers.foundations.enriched_task_adjust import ( + EnrichedTaskAdjustMerge, + ) + from generalresearch.incite.mergers.foundations.enriched_wall import ( + EnrichedWallMerge, + ) + from generalresearch.models.thl.product import Product + from generalresearch.models.thl.user import User + from generalresearch.pg_helper import PostgresConfig @pytest.mark.parametrize( diff --git a/tests/incite/mergers/foundations/test_enriched_wall.py b/tests/incite/mergers/foundations/test_enriched_wall.py index 0cb8f60..2b9afb8 100644 --- a/tests/incite/mergers/foundations/test_enriched_wall.py +++ b/tests/incite/mergers/foundations/test_enriched_wall.py @@ -2,27 +2,32 @@ from collections.abc import Callable from datetime import UTC, datetime, timedelta from decimal import Decimal from itertools import product as iter_product +from typing import TYPE_CHECKING import dask.dataframe as dd import pandas as pd import pytest from dask.distributed import Client as DaskClient -from generalresearch.incite.collections.thl_web import ( - SessionDFCollection, - WallDFCollection, -) - -# noinspection PyUnresolvedReferences from generalresearch.incite.mergers.foundations.enriched_wall import ( - EnrichedWallMerge, EnrichedWallMergeItem, ) -from generalresearch.models.admin.request import ReportRequest -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 + +if TYPE_CHECKING: + from generalresearch.incite.collections.thl_web import ( + SessionDFCollection, + WallDFCollection, + ) + + # noinspection PyUnresolvedReferences + from generalresearch.incite.mergers.foundations.enriched_wall import ( + EnrichedWallMerge, + ) + from generalresearch.models.admin.request import ReportRequest + 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 @pytest.mark.parametrize( diff --git a/tests/incite/mergers/foundations/test_user_id_product.py b/tests/incite/mergers/foundations/test_user_id_product.py index 7367056..8c4b2f7 100644 --- a/tests/incite/mergers/foundations/test_user_id_product.py +++ b/tests/incite/mergers/foundations/test_user_id_product.py @@ -2,17 +2,22 @@ from __future__ import annotations from datetime import UTC, datetime, timedelta from itertools import product +from typing import TYPE_CHECKING import pandas as pd import pytest from dask.distributed import Client as DaskClient -# noinspection PyUnresolvedReferences from generalresearch.incite.mergers.foundations.user_id_product import ( - UserIdProductMerge, UserIdProductMergeItem, ) +if TYPE_CHECKING: + # noinspection PyUnresolvedReferences + from generalresearch.incite.mergers.foundations.user_id_product import ( + UserIdProductMerge, + ) + @pytest.mark.parametrize( argnames="offset, duration, start", diff --git a/tests/incite/mergers/test_merge_collection.py b/tests/incite/mergers/test_merge_collection.py index cf8315f..3f478bd 100644 --- a/tests/incite/mergers/test_merge_collection.py +++ b/tests/incite/mergers/test_merge_collection.py @@ -2,17 +2,20 @@ from __future__ import annotations from datetime import UTC, datetime, timedelta from itertools import product +from typing import TYPE_CHECKING import pandas as pd import pytest from pandera.pandas import DataFrameSchema -from generalresearch.incite.base import GRLDatasets -from generalresearch.incite.mergers import ( +from generalresearch.incite.mergers.base import ( MergeCollection, MergeType, ) +if TYPE_CHECKING: + from generalresearch.incite.base import GRLDatasets + merge_types = [e for e in MergeType if e != MergeType.TEST] diff --git a/tests/incite/mergers/test_merge_collection_item.py b/tests/incite/mergers/test_merge_collection_item.py index 5ca2f6b..baf1bc4 100644 --- a/tests/incite/mergers/test_merge_collection_item.py +++ b/tests/incite/mergers/test_merge_collection_item.py @@ -3,14 +3,17 @@ from __future__ import annotations from datetime import timedelta from itertools import product from pathlib import PurePath +from typing import TYPE_CHECKING import pytest -from generalresearch.incite.mergers import ( - MergeCollection, - MergeCollectionItem, - MergeType, -) +from generalresearch.incite.mergers.base import MergeType + +if TYPE_CHECKING: + from generalresearch.incite.mergers.base import ( + MergeCollection, + MergeCollectionItem, + ) @pytest.mark.parametrize( diff --git a/tests/incite/mergers/test_pop_ledger.py b/tests/incite/mergers/test_pop_ledger.py index 2146344..9ec188b 100644 --- a/tests/incite/mergers/test_pop_ledger.py +++ b/tests/incite/mergers/test_pop_ledger.py @@ -3,23 +3,26 @@ from __future__ import annotations from collections.abc import Callable from datetime import UTC, datetime, timedelta from itertools import product as iter_product +from typing import TYPE_CHECKING import pandas as pd import pytest from dask.distributed import Client as DaskClient -from generalresearch.incite.base import GRLDatasets -from generalresearch.incite.collections.thl_web import ( - LedgerDFCollection, - SessionDFCollection, -) -from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge from generalresearch.incite.schemas.mergers.pop_ledger import ( numerical_col_names, ) -from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager -from generalresearch.models.thl.product import Product -from generalresearch.models.thl.user import User + +if TYPE_CHECKING: + from generalresearch.incite.base import GRLDatasets + from generalresearch.incite.collections.thl_web import ( + LedgerDFCollection, + SessionDFCollection, + ) + from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge + from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager + from generalresearch.models.thl.product import Product + from generalresearch.models.thl.user import User @pytest.mark.parametrize( diff --git a/tests/incite/mergers/test_ym_survey_merge.py b/tests/incite/mergers/test_ym_survey_merge.py index 8a4897b..d83a98c 100644 --- a/tests/incite/mergers/test_ym_survey_merge.py +++ b/tests/incite/mergers/test_ym_survey_merge.py @@ -3,22 +3,24 @@ from __future__ import annotations from collections.abc import Callable from datetime import UTC, datetime, timedelta from itertools import product +from typing import TYPE_CHECKING import pandas as pd import pytest from dask.distributed import Client as DaskClient -from generalresearch.incite.collections.thl_web import ( - SessionDFCollection, - WallDFCollection, -) -from generalresearch.incite.mergers.foundations.enriched_session import ( - EnrichedSessionMerge, -) -from generalresearch.incite.mergers.ym_survey_wall import YMSurveyWallMerge -from generalresearch.models.thl.product import Product -from generalresearch.models.thl.user import User -from generalresearch.pg_helper import PostgresConfig +if TYPE_CHECKING: + from generalresearch.incite.collections.thl_web import ( + SessionDFCollection, + WallDFCollection, + ) + from generalresearch.incite.mergers.foundations.enriched_session import ( + EnrichedSessionMerge, + ) + from generalresearch.incite.mergers.ym_survey_wall import YMSurveyWallMerge + from generalresearch.models.thl.product import Product + from generalresearch.models.thl.user import User + from generalresearch.pg_helper import PostgresConfig # noinspection PyUnresolvedReferences diff --git a/tests/incite/test_collection_base.py b/tests/incite/test_collection_base.py index 577eda9..1a664a2 100644 --- a/tests/incite/test_collection_base.py +++ b/tests/incite/test_collection_base.py @@ -4,6 +4,7 @@ from datetime import UTC, datetime, timedelta, timezone from os.path import exists as pexists from os.path import join as pjoin from pathlib import Path +from typing import TYPE_CHECKING from uuid import uuid4 import numpy as np @@ -11,7 +12,10 @@ import pandas as pd import pytest from _pytest._code.code import ExceptionInfo -from generalresearch.incite.base import CollectionBase, GRLDatasets +from generalresearch.incite.base import CollectionBase + +if TYPE_CHECKING: + from generalresearch.incite.base import GRLDatasets AGO_15min = (datetime.now(tz=UTC) - timedelta(minutes=15)).replace(microsecond=0) AGO_1HR = (datetime.now(tz=UTC) - timedelta(hours=1)).replace(microsecond=0) diff --git a/tests/incite/test_collection_base_item.py b/tests/incite/test_collection_base_item.py index e09f54a..b9f1c26 100644 --- a/tests/incite/test_collection_base_item.py +++ b/tests/incite/test_collection_base_item.py @@ -3,6 +3,7 @@ from __future__ import annotations from datetime import UTC, datetime from os.path import join as pjoin from pathlib import Path +from typing import TYPE_CHECKING from uuid import uuid4 import dask.dataframe as dd @@ -10,7 +11,10 @@ import pandas as pd import pytest from pydantic import ValidationError -from generalresearch.incite.base import CollectionItemBase, GRLDatasets +from generalresearch.incite.base import CollectionItemBase + +if TYPE_CHECKING: + from generalresearch.incite.base import GRLDatasets class TestCollectionItemBase: diff --git a/tests/managers/gr/test_business.py b/tests/managers/gr/test_business.py index ed141b1..1a5d4fa 100644 --- a/tests/managers/gr/test_business.py +++ b/tests/managers/gr/test_business.py @@ -1,21 +1,24 @@ +from typing import TYPE_CHECKING from uuid import uuid4 import pytest -from generalresearch.managers.gr.business import ( - BusinessAddressManager, - BusinessBankAccountManager, - BusinessManager, -) -from generalresearch.managers.gr.team import MembershipManager, TeamManager -from generalresearch.models.gr.authentication import GRUser from generalresearch.models.gr.business import ( Business, BusinessAddress, BusinessBankAccount, TransferMethod, ) -from generalresearch.pg_helper import PostgresConfig + +if TYPE_CHECKING: + from generalresearch.managers.gr.business import ( + BusinessAddressManager, + BusinessBankAccountManager, + BusinessManager, + ) + from generalresearch.managers.gr.team import MembershipManager, TeamManager + from generalresearch.models.gr.authentication import GRUser + from generalresearch.pg_helper import PostgresConfig class TestBusinessBankAccountManager: diff --git a/tests/managers/gr/test_team.py b/tests/managers/gr/test_team.py index ae3e1bb..17e0470 100644 --- a/tests/managers/gr/test_team.py +++ b/tests/managers/gr/test_team.py @@ -1,15 +1,18 @@ from __future__ import annotations from collections.abc import Callable +from typing import TYPE_CHECKING from uuid import uuid4 -from generalresearch.managers.gr.authentication import GRUserManager -from generalresearch.managers.gr.team import MembershipManager, TeamManager -from generalresearch.models.gr.authentication import GRUser 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 + +if TYPE_CHECKING: + from generalresearch.managers.gr.authentication import GRUserManager + from generalresearch.managers.gr.team import MembershipManager, TeamManager + from generalresearch.models.gr.authentication import GRUser + from generalresearch.models.thl.product import Product + from generalresearch.pg_helper import PostgresConfig + from generalresearch.redis_helper import RedisConfig class TestMembershipManager: diff --git a/tests/managers/leaderboard.py b/tests/managers/leaderboard.py index d97714d..197477b 100644 --- a/tests/managers/leaderboard.py +++ b/tests/managers/leaderboard.py @@ -6,6 +6,7 @@ import zoneinfo from collections.abc import Callable from datetime import UTC, datetime from decimal import Decimal +from typing import TYPE_CHECKING from uuid import uuid4 import pytest @@ -26,7 +27,9 @@ from generalresearch.models.thl.product import ( ) from generalresearch.models.thl.session import Session from generalresearch.models.thl.user import User -from generalresearch.redis_helper import RedisConfig + +if TYPE_CHECKING: + from generalresearch.redis_helper import RedisConfig # random uuid for leaderboard tests product_id = uuid4().hex diff --git a/tests/managers/network/test_label.py b/tests/managers/network/test_label.py index 71efa95..abdd28f 100644 --- a/tests/managers/network/test_label.py +++ b/tests/managers/network/test_label.py @@ -1,11 +1,12 @@ import ipaddress +from datetime import datetime +from typing import TYPE_CHECKING import faker import pytest from psycopg.errors import UniqueViolation from pydantic import ValidationError -from generalresearch.managers.network.label import IPLabelManager from generalresearch.models.network.label import ( IPLabel, IPLabelKind, @@ -14,11 +15,14 @@ from generalresearch.models.network.label import ( ) from generalresearch.models.thl.ipinfo import normalize_ip +if TYPE_CHECKING: + from generalresearch.managers.network.label import IPLabelManager + fake = faker.Faker() @pytest.fixture -def ip_label(utc_now) -> IPLabel: +def ip_label(utc_now: datetime) -> IPLabel: ip = ipaddress.IPv6Network((fake.ipv6(), 64), strict=False) return IPLabel( label_kind=IPLabelKind.VPN, @@ -31,7 +35,7 @@ def ip_label(utc_now) -> IPLabel: ) -def test_model(utc_now): +def test_model(utc_now: datetime): ip = fake.ipv4_public() lbl = IPLabel( label_kind=IPLabelKind.VPN, @@ -142,7 +146,7 @@ def test_filter_network( assert len(res) == 2 -def test_network(iplabel_manager: IPLabelManager, utc_now): +def test_network(iplabel_manager: IPLabelManager, utc_now: datetime): # This is a fully-specific /128 ipv6 address. # e.g. '51b7:b38d:8717:6c5b:cd3e:f5c3:3aba:17d' ip = fake.ipv6() @@ -174,7 +178,10 @@ def test_network(iplabel_manager: IPLabelManager, utc_now): def test_label_cidr_and_ipinfo( - iplabel_manager: IPLabelManager, ip_information_factory, ip_geoname, utc_now + iplabel_manager: IPLabelManager, + ip_information_factory, + ip_geoname, + utc_now: datetime, ): # We have network_iplabel.ip as a cidr col and # thl_ipinformation.ip as a inet col. Make sure we can join appropriately diff --git a/tests/managers/test_events.py b/tests/managers/test_events.py index cb32275..8745126 100644 --- a/tests/managers/test_events.py +++ b/tests/managers/test_events.py @@ -8,13 +8,13 @@ from datetime import UTC, datetime, timedelta from decimal import Decimal from functools import partial from math import floor +from typing import TYPE_CHECKING from uuid import uuid4 import pytest -from generalresearch.managers.events import EventManager, EventSubscriber -from generalresearch.managers.thl.product import ProductManager -from generalresearch.models import Source +from generalresearch.managers.events import EventSubscriber +from generalresearch.models.definitions import Source from generalresearch.models.events import ( AggregateBySource, EventType, @@ -25,7 +25,11 @@ 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 + +if TYPE_CHECKING: + from generalresearch.managers.events import EventManager + from generalresearch.managers.thl.product import ProductManager + from generalresearch.redis_helper import RedisConfig # We don't need anything in the db, so not using the db fixtures diff --git a/tests/managers/test_lucid.py b/tests/managers/test_lucid.py index 20dca22..6771a0c 100644 --- a/tests/managers/test_lucid.py +++ b/tests/managers/test_lucid.py @@ -1,9 +1,13 @@ from __future__ import annotations +from typing import TYPE_CHECKING + import pytest from generalresearch.managers.lucid.profiling import get_profiling_library -from generalresearch.pg_helper import PostgresConfig + +if TYPE_CHECKING: + from generalresearch.pg_helper import PostgresConfig qids = ["42", "43", "45", "97", "120", "639", "15297"] diff --git a/tests/managers/thl/test_buyer.py b/tests/managers/thl/test_buyer.py index 6776ab3..0ab2d52 100644 --- a/tests/managers/thl/test_buyer.py +++ b/tests/managers/thl/test_buyer.py @@ -1,9 +1,12 @@ from __future__ import annotations from collections.abc import Callable +from typing import TYPE_CHECKING -from generalresearch.managers.thl.buyer import BuyerManager -from generalresearch.models import Source +from generalresearch.models.definitions import Source + +if TYPE_CHECKING: + from generalresearch.managers.thl.buyer import BuyerManager class TestBuyer: diff --git a/tests/managers/thl/test_cashout_method.py b/tests/managers/thl/test_cashout_method.py index ca85c6b..877d7b2 100644 --- a/tests/managers/thl/test_cashout_method.py +++ b/tests/managers/thl/test_cashout_method.py @@ -1,21 +1,26 @@ from __future__ import annotations from collections.abc import Callable +from typing import TYPE_CHECKING 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, - CashoutMethod, PaypalCashoutMethodData, USDeliveryAddress, ) +from generalresearch.models.thl.wallet.definitions import PayoutType + +if TYPE_CHECKING: + 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.cashout_method import ( + CashoutMethod, + ) class TestTangoCashoutMethods: diff --git a/tests/managers/thl/test_category.py b/tests/managers/thl/test_category.py index ec52aae..4d00643 100644 --- a/tests/managers/thl/test_category.py +++ b/tests/managers/thl/test_category.py @@ -1,12 +1,15 @@ from __future__ import annotations from collections.abc import Callable +from typing import TYPE_CHECKING import pytest -from generalresearch.managers.thl.category import CategoryManager from generalresearch.models.thl.category import Category -from generalresearch.pg_helper import PostgresConfig + +if TYPE_CHECKING: + from generalresearch.managers.thl.category import CategoryManager + from generalresearch.pg_helper import PostgresConfig class TestCategory: diff --git a/tests/managers/thl/test_contest/test_leaderboard.py b/tests/managers/thl/test_contest/test_leaderboard.py index 3a63075..d80d512 100644 --- a/tests/managers/thl/test_contest/test_leaderboard.py +++ b/tests/managers/thl/test_contest/test_leaderboard.py @@ -1,23 +1,28 @@ from __future__ import annotations from datetime import UTC, datetime, timedelta +from typing import TYPE_CHECKING from zoneinfo import ZoneInfo from generalresearch.currency import USDCent -from generalresearch.managers.thl.contest_manager import ContestManager -from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager -from generalresearch.managers.thl.user_manager.user_manager import UserManager from generalresearch.models.thl.contest.definitions import ( ContestEndReason, ContestStatus, ) from generalresearch.models.thl.contest.leaderboard import ( LeaderboardContest, - LeaderboardContestCreate, ) -from generalresearch.models.thl.product import Product -from generalresearch.models.thl.user import User -from generalresearch.redis_helper import RedisConfig + +if TYPE_CHECKING: + from generalresearch.managers.thl.contest_manager import ContestManager + from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager + from generalresearch.managers.thl.user_manager.user_manager import UserManager + from generalresearch.models.thl.contest.leaderboard import ( + LeaderboardContestCreate, + ) + from generalresearch.models.thl.product import Product + from generalresearch.models.thl.user import User + from generalresearch.redis_helper import RedisConfig class TestLeaderboardContestCRUD: diff --git a/tests/managers/thl/test_contest/test_milestone.py b/tests/managers/thl/test_contest/test_milestone.py index e29ba4c..dbb2016 100644 --- a/tests/managers/thl/test_contest/test_milestone.py +++ b/tests/managers/thl/test_contest/test_milestone.py @@ -2,9 +2,8 @@ from __future__ import annotations from collections.abc import Callable from datetime import UTC, datetime +from typing import TYPE_CHECKING -from generalresearch.managers.thl.contest_manager import ContestManager -from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager from generalresearch.models.thl.contest.definitions import ( ContestEndReason, ContestStatus, @@ -12,12 +11,18 @@ from generalresearch.models.thl.contest.definitions import ( from generalresearch.models.thl.contest.milestone import ( ContestEntryTrigger, MilestoneContest, - MilestoneContestCreate, MilestoneUserView, ) -from generalresearch.models.thl.contest.raffle import RaffleContest -from generalresearch.models.thl.product import Product -from generalresearch.models.thl.user import User + +if TYPE_CHECKING: + from generalresearch.managers.thl.contest_manager import ContestManager + from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager + from generalresearch.models.thl.contest.milestone import ( + MilestoneContestCreate, + ) + from generalresearch.models.thl.contest.raffle import RaffleContest + from generalresearch.models.thl.product import Product + from generalresearch.models.thl.user import User class TestMilestoneContest: @@ -25,8 +30,6 @@ class TestMilestoneContest: def test_should_end( self, contest: MilestoneContest, - thl_ledger_manager: ThlLedgerManager, - contest_manager: ContestManager, ): # contest is active and has no entries should, msg = contest.should_end() @@ -53,7 +56,6 @@ class TestMilestoneContestCRUD: self, contest_create: MilestoneContestCreate, product_user_wallet_yes: Product, - thl_ledger_manager: ThlLedgerManager, contest_manager: ContestManager, ): c = contest_manager.create( diff --git a/tests/managers/thl/test_contest/test_raffle.py b/tests/managers/thl/test_contest/test_raffle.py index 06d4676..7803952 100644 --- a/tests/managers/thl/test_contest/test_raffle.py +++ b/tests/managers/thl/test_contest/test_raffle.py @@ -2,19 +2,17 @@ from __future__ import annotations from collections.abc import Callable from datetime import UTC, datetime +from typing import TYPE_CHECKING import pytest from pydantic import ValidationError from pytest import approx from generalresearch.currency import USDCent -from generalresearch.managers.thl.contest_manager import ContestManager from generalresearch.managers.thl.ledger_manager.exceptions import ( LedgerTransactionConditionFailedError, ) -from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager from generalresearch.models.thl.contest import ( - Contest, ContestEndCondition, ContestEntryRule, ContestPrize, @@ -29,11 +27,20 @@ from generalresearch.models.thl.contest.raffle import ( ContestEntry, ContestEntryType, RaffleContest, - RaffleContestCreate, - RaffleUserView, ) -from generalresearch.models.thl.product import Product -from generalresearch.models.thl.user import User + +if TYPE_CHECKING: + from generalresearch.managers.thl.contest_manager import ContestManager + from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager + from generalresearch.models.thl.contest import ( + Contest, + ) + from generalresearch.models.thl.contest.raffle import ( + RaffleContestCreate, + RaffleUserView, + ) + from generalresearch.models.thl.product import Product + from generalresearch.models.thl.user import User class TestRaffleContest: diff --git a/tests/managers/thl/test_harmonized_uqa.py b/tests/managers/thl/test_harmonized_uqa.py index 84eeb56..2fc0ff0 100644 --- a/tests/managers/thl/test_harmonized_uqa.py +++ b/tests/managers/thl/test_harmonized_uqa.py @@ -1,15 +1,18 @@ from __future__ import annotations from datetime import UTC, datetime +from typing import TYPE_CHECKING import pytest -from generalresearch.managers.thl.profiling.uqa import UQAManager from generalresearch.models.thl.profiling.user_question_answer import ( DUMMY_UQA, UserQuestionAnswer, ) -from generalresearch.models.thl.user import User + +if TYPE_CHECKING: + from generalresearch.managers.thl.profiling.uqa import UQAManager + from generalresearch.models.thl.user import User @pytest.mark.usefixtures("uqa_db_index", "upk_data", "uqa_manager_clear_cache") diff --git a/tests/managers/thl/test_ipinfo.py b/tests/managers/thl/test_ipinfo.py index 48b9efd..6954163 100644 --- a/tests/managers/thl/test_ipinfo.py +++ b/tests/managers/thl/test_ipinfo.py @@ -1,4 +1,5 @@ from collections.abc import Callable +from typing import TYPE_CHECKING import faker @@ -12,8 +13,10 @@ from generalresearch.models.thl.ipinfo import ( IPGeoname, IPInformation, ) -from generalresearch.pg_helper import PostgresConfig -from generalresearch.redis_helper import RedisConfig + +if TYPE_CHECKING: + from generalresearch.pg_helper import PostgresConfig + from generalresearch.redis_helper import RedisConfig fake = faker.Faker() diff --git a/tests/managers/thl/test_ledger/test_lm_accounts.py b/tests/managers/thl/test_ledger/test_lm_accounts.py index 7b65b2d..f5ed883 100644 --- a/tests/managers/thl/test_ledger/test_lm_accounts.py +++ b/tests/managers/thl/test_ledger/test_lm_accounts.py @@ -2,6 +2,7 @@ from __future__ import annotations from itertools import product as iproduct from random import randint +from typing import TYPE_CHECKING from uuid import uuid4 import pytest @@ -13,13 +14,18 @@ from generalresearch.managers.thl.ledger_manager.exceptions import ( LedgerAccountDoesntExistError, ) from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager -from generalresearch.models.custom_types import AccountType, Direction, UUIDStr +from generalresearch.models.custom_types import AccountType, Direction from generalresearch.models.thl.ledger import ( LedgerAccount, LedgerEntry, - LedgerTransaction, ) +if TYPE_CHECKING: + from generalresearch.models.custom_types import UUIDStr + from generalresearch.models.thl.ledger import ( + LedgerTransaction, + ) + @pytest.mark.parametrize( argnames="currency, kind, acct_id", diff --git a/tests/managers/thl/test_ledger/test_lm_tx.py b/tests/managers/thl/test_ledger/test_lm_tx.py index ce609d6..445405e 100644 --- a/tests/managers/thl/test_ledger/test_lm_tx.py +++ b/tests/managers/thl/test_ledger/test_lm_tx.py @@ -2,6 +2,7 @@ from __future__ import annotations from decimal import Decimal from random import randint +from typing import TYPE_CHECKING from uuid import uuid4 import pytest @@ -12,11 +13,15 @@ from generalresearch.managers.thl.ledger_manager.ledger import ( ) from generalresearch.models.thl.ledger import ( Direction, - LedgerAccount, LedgerEntry, LedgerTransaction, ) +if TYPE_CHECKING: + from generalresearch.models.thl.ledger import ( + LedgerAccount, + ) + class TestLedgerManagerCreateTx: diff --git a/tests/managers/thl/test_ledger/test_lm_tx_entries.py b/tests/managers/thl/test_ledger/test_lm_tx_entries.py index 9925b87..03c6e02 100644 --- a/tests/managers/thl/test_ledger/test_lm_tx_entries.py +++ b/tests/managers/thl/test_ledger/test_lm_tx_entries.py @@ -1,11 +1,17 @@ from __future__ import annotations -from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager +from typing import TYPE_CHECKING + from generalresearch.models.thl.ledger import ( LedgerEntry, - LedgerTransaction, ) +if TYPE_CHECKING: + from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager + from generalresearch.models.thl.ledger import ( + LedgerTransaction, + ) + class TestLedgerEntryManager: 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 e603632..9ecc1bc 100644 --- a/tests/managers/thl/test_ledger/test_lm_tx_locks.py +++ b/tests/managers/thl/test_ledger/test_lm_tx_locks.py @@ -5,10 +5,10 @@ from collections.abc import Callable, Generator from datetime import UTC, datetime, timedelta from decimal import Decimal from logging import LogCaptureFixture +from typing import TYPE_CHECKING import pytest -from generalresearch.currency import LedgerCurrency from generalresearch.managers.thl.ledger_manager.conditions import ( generate_condition_mp_payment, ) @@ -17,11 +17,8 @@ from generalresearch.managers.thl.ledger_manager.exceptions import ( LedgerTransactionCreateLockError, LedgerTransactionFlagAlreadyExistsError, ) -from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager -from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager -from generalresearch.models import Source +from generalresearch.models.definitions import Source from generalresearch.models.thl.ledger import LedgerTransaction -from generalresearch.models.thl.product import Product from generalresearch.models.thl.session import ( Session, Status, @@ -29,7 +26,13 @@ from generalresearch.models.thl.session import ( Wall, WallAdjustedStatus, ) -from generalresearch.models.thl.user import User + +if TYPE_CHECKING: + from generalresearch.currency import LedgerCurrency + from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager + from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager + from generalresearch.models.thl.product import Product + from generalresearch.models.thl.user import User logger = logging.getLogger("LedgerManager") diff --git a/tests/managers/thl/test_ledger/test_lm_tx_metadata.py b/tests/managers/thl/test_ledger/test_lm_tx_metadata.py index f63efa4..3d8cf89 100644 --- a/tests/managers/thl/test_ledger/test_lm_tx_metadata.py +++ b/tests/managers/thl/test_ledger/test_lm_tx_metadata.py @@ -1,9 +1,12 @@ from __future__ import annotations -from generalresearch.managers.thl.ledger_manager.ledger import ( - LedgerManager, - LedgerTransaction, -) +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from generalresearch.managers.thl.ledger_manager.ledger import ( + LedgerManager, + LedgerTransaction, + ) class TestLedgerMetadataManager: diff --git a/tests/managers/thl/test_ledger/test_thl_lm_accounts.py b/tests/managers/thl/test_ledger/test_thl_lm_accounts.py index 60eb71c..adff446 100644 --- a/tests/managers/thl/test_ledger/test_thl_lm_accounts.py +++ b/tests/managers/thl/test_ledger/test_thl_lm_accounts.py @@ -1,6 +1,7 @@ from __future__ import annotations from collections.abc import Callable +from typing import TYPE_CHECKING from uuid import uuid4 import pytest @@ -9,18 +10,20 @@ from generalresearch.currency import LedgerCurrency from generalresearch.managers.thl.ledger_manager.exceptions import ( LedgerAccountDoesntExistError, ) -from generalresearch.managers.thl.ledger_manager.ledger import ( - LedgerAccountManager, - LedgerManager, -) -from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager from generalresearch.models.thl.ledger import ( AccountType, Direction, LedgerAccount, ) from generalresearch.models.thl.product import Product -from generalresearch.models.thl.user import User + +if TYPE_CHECKING: + from generalresearch.managers.thl.ledger_manager.ledger import ( + LedgerAccountManager, + LedgerManager, + ) + from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager + from generalresearch.models.thl.user import User class TestThlLedgerManagerAccounts: diff --git a/tests/managers/thl/test_ledger/test_thl_lm_bp_payout.py b/tests/managers/thl/test_ledger/test_thl_lm_bp_payout.py index b518453..14c5270 100644 --- a/tests/managers/thl/test_ledger/test_thl_lm_bp_payout.py +++ b/tests/managers/thl/test_ledger/test_thl_lm_bp_payout.py @@ -5,6 +5,7 @@ from collections.abc import Callable from datetime import UTC, datetime, timedelta from decimal import Decimal from random import randint +from typing import TYPE_CHECKING from uuid import uuid4 import pytest @@ -12,7 +13,7 @@ import redis from pydantic import RedisDsn from redis.lock import Lock -from generalresearch.currency import LedgerCurrency, USDCent +from generalresearch.currency import USDCent from generalresearch.managers.base import Permission from generalresearch.managers.thl.ledger_manager.exceptions import ( LedgerTransactionConditionFailedError, @@ -22,24 +23,27 @@ from generalresearch.managers.thl.ledger_manager.exceptions import ( ) from generalresearch.managers.thl.ledger_manager.ledger import LedgerTransaction from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager -from generalresearch.managers.thl.payout import ( - BrokerageProductPayoutEventManager, -) -from generalresearch.models import Source +from generalresearch.models.definitions import Source from generalresearch.models.thl.definitions import PayoutStatus from generalresearch.models.thl.ledger import Direction, TransactionType -from generalresearch.models.thl.product import Product from generalresearch.models.thl.session import ( Session, Status, StatusCode1, Wall, ) -from generalresearch.models.thl.user import User -from generalresearch.models.thl.wallet import PayoutType -from generalresearch.pg_helper import PostgresConfig +from generalresearch.models.thl.wallet.definitions import PayoutType from generalresearch.redis_helper import RedisConfig +if TYPE_CHECKING: + from generalresearch.currency import LedgerCurrency + from generalresearch.managers.thl.payout import ( + BrokerageProductPayoutEventManager, + ) + from generalresearch.models.thl.product import Product + from generalresearch.models.thl.user import User + from generalresearch.pg_helper import PostgresConfig + def broken_acquire(self, *args, **kwargs): raise redis.exceptions.TimeoutError("Simulated timeout during acquire") diff --git a/tests/managers/thl/test_ledger/test_thl_lm_tx.py b/tests/managers/thl/test_ledger/test_thl_lm_tx.py index 1860d6d..2e4ab5e 100644 --- a/tests/managers/thl/test_ledger/test_thl_lm_tx.py +++ b/tests/managers/thl/test_ledger/test_thl_lm_tx.py @@ -5,26 +5,21 @@ from collections.abc import Callable from datetime import UTC, datetime, timedelta from decimal import Decimal from random import randint +from typing import TYPE_CHECKING from uuid import uuid4 import pytest -from generalresearch.currency import LedgerCurrency, USDCent +from generalresearch.currency import USDCent from generalresearch.managers.thl.ledger_manager.ledger import ( - LedgerManager, LedgerTransaction, ) -from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager -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.definitions import Source from generalresearch.models.thl.definitions import ( WALL_ALLOWED_STATUS_STATUS_CODE, ) from generalresearch.models.thl.ledger import ( Direction, - LedgerAccount, TransactionType, ) from generalresearch.models.thl.payout import UserPayoutEvent @@ -41,8 +36,21 @@ from generalresearch.models.thl.session import ( Wall, WallAdjustedStatus, ) -from generalresearch.models.thl.user import User -from generalresearch.models.thl.wallet import PayoutType +from generalresearch.models.thl.wallet.definitions import PayoutType + +if TYPE_CHECKING: + from generalresearch.currency import LedgerCurrency + 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.managers.thl.session import SessionManager + from generalresearch.managers.thl.wall import WallManager + from generalresearch.models.thl.ledger import ( + LedgerAccount, + ) + from generalresearch.models.thl.user import User logger = logging.getLogger("LedgerManager") diff --git a/tests/managers/thl/test_ledger/test_thl_lm_tx__user_payouts.py b/tests/managers/thl/test_ledger/test_thl_lm_tx__user_payouts.py index 82dc143..3fd21dc 100644 --- a/tests/managers/thl/test_ledger/test_thl_lm_tx__user_payouts.py +++ b/tests/managers/thl/test_ledger/test_thl_lm_tx__user_payouts.py @@ -3,6 +3,7 @@ from __future__ import annotations import logging from collections.abc import Callable from decimal import Decimal +from typing import TYPE_CHECKING from uuid import uuid4 import pytest @@ -11,12 +12,14 @@ from generalresearch.managers.thl.ledger_manager.exceptions import ( LedgerTransactionConditionFailedError, LedgerTransactionFlagAlreadyExistsError, ) -from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager -from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager from generalresearch.models.thl.payout import UserPayoutEvent -from generalresearch.models.thl.product import Product -from generalresearch.models.thl.user import User -from generalresearch.models.thl.wallet import PayoutType +from generalresearch.models.thl.wallet.definitions import PayoutType + +if TYPE_CHECKING: + from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager + from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager + from generalresearch.models.thl.product import Product + from generalresearch.models.thl.user import User class TestLedgerManagerAMT: diff --git a/tests/managers/thl/test_ledger/test_thl_pem.py b/tests/managers/thl/test_ledger/test_thl_pem.py index 29341cf..9dbec48 100644 --- a/tests/managers/thl/test_ledger/test_thl_pem.py +++ b/tests/managers/thl/test_ledger/test_thl_pem.py @@ -2,28 +2,31 @@ from __future__ import annotations from collections.abc import Callable from random import randint +from typing import TYPE_CHECKING from uuid import UUID, uuid4 import pytest from generalresearch.currency import USDCent -from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager -from generalresearch.managers.thl.ledger_manager.thl_ledger import ( - ThlLedgerManager, -) -from generalresearch.managers.thl.payout import ( - BrokerageProductPayoutEventManager, - UserPayoutEventManager, -) from generalresearch.models.thl.definitions import PayoutStatus from generalresearch.models.thl.payout import ( BrokerageProductPayoutEvent, ) -from generalresearch.models.thl.product import Product from generalresearch.models.thl.wallet.cashout_method import ( CashoutRequestInfo, ) +if TYPE_CHECKING: + from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager + from generalresearch.managers.thl.ledger_manager.thl_ledger import ( + ThlLedgerManager, + ) + from generalresearch.managers.thl.payout import ( + BrokerageProductPayoutEventManager, + UserPayoutEventManager, + ) + from generalresearch.models.thl.product import Product + class TestThlPayoutEventManager: diff --git a/tests/managers/thl/test_ledger/test_user_txs.py b/tests/managers/thl/test_ledger/test_user_txs.py index 56dc485..1c08498 100644 --- a/tests/managers/thl/test_ledger/test_user_txs.py +++ b/tests/managers/thl/test_ledger/test_user_txs.py @@ -3,12 +3,9 @@ from __future__ import annotations from collections.abc import Callable from datetime import UTC, datetime from decimal import Decimal +from typing import TYPE_CHECKING from uuid import uuid4 -from generalresearch.config import GRLBaseSettings -from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager -from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager -from generalresearch.managers.thl.payout import UserPayoutEventManager from generalresearch.managers.thl.user_compensate import user_compensate from generalresearch.models.thl.definitions import ( Status, @@ -18,10 +15,16 @@ from generalresearch.models.thl.ledger import ( UserLedgerTransactionTypesSummary, UserLedgerTransactionTypeSummary, ) -from generalresearch.models.thl.product import Product -from generalresearch.models.thl.session import Session -from generalresearch.models.thl.user import User -from generalresearch.models.thl.wallet import PayoutType +from generalresearch.models.thl.wallet.definitions import PayoutType + +if TYPE_CHECKING: + from generalresearch.config import GRLBaseSettings + from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager + from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager + from generalresearch.managers.thl.payout import UserPayoutEventManager + from generalresearch.models.thl.product import Product + from generalresearch.models.thl.session import Session + from generalresearch.models.thl.user import User def test_user_txs( diff --git a/tests/managers/thl/test_ledger/test_wallet.py b/tests/managers/thl/test_ledger/test_wallet.py index cad3ea4..1ee9bf9 100644 --- a/tests/managers/thl/test_ledger/test_wallet.py +++ b/tests/managers/thl/test_ledger/test_wallet.py @@ -2,12 +2,11 @@ from __future__ import annotations from collections.abc import Callable from decimal import Decimal +from typing import TYPE_CHECKING from uuid import uuid4 import pytest -from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager -from generalresearch.managers.thl.product import ProductManager from generalresearch.models.thl.product import ( PayoutConfig, PayoutTransformation, @@ -15,7 +14,11 @@ from generalresearch.models.thl.product import ( Product, UserWalletConfig, ) -from generalresearch.models.thl.user import User + +if TYPE_CHECKING: + from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager + from generalresearch.managers.thl.product import ProductManager + from generalresearch.models.thl.user import User @pytest.fixture() diff --git a/tests/managers/thl/test_payout.py b/tests/managers/thl/test_payout.py index 0f3f103..2494de8 100644 --- a/tests/managers/thl/test_payout.py +++ b/tests/managers/thl/test_payout.py @@ -6,6 +6,7 @@ from datetime import UTC, datetime, timedelta from decimal import Decimal from random import choice as rand_choice from random import randint +from typing import TYPE_CHECKING from uuid import uuid4 import pandas as pd @@ -13,35 +14,39 @@ import pytest from dask.distributed import Client as DaskClient from generalresearch.currency import USDCent -from generalresearch.incite.base import GRLDatasets -from generalresearch.incite.collections.thl_web import ( - LedgerDFCollection, -) -from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge -from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager -from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager -from generalresearch.managers.thl.payout import ( - BrokerageProductPayoutEventManager, - BusinessPayoutEventManager, - PayoutEventManager, - UserPayoutEventManager, -) -from generalresearch.managers.thl.product import ProductManager -from generalresearch.models.gr.business import Business from generalresearch.models.thl.definitions import PayoutStatus from generalresearch.models.thl.finance import BusinessBalances -from generalresearch.models.thl.ledger import LedgerAccount from generalresearch.models.thl.payout import ( BrokerageProductPayoutEvent, BusinessPayoutEvent, - UserPayoutEvent, ) -from generalresearch.models.thl.product import Product -from generalresearch.models.thl.session import Session -from generalresearch.models.thl.user import User -from generalresearch.models.thl.wallet import PayoutType -from generalresearch.pg_helper import PostgresConfig -from generalresearch.redis_helper import RedisConfig +from generalresearch.models.thl.wallet.definitions import PayoutType + +if TYPE_CHECKING: + from generalresearch.incite.base import GRLDatasets + from generalresearch.incite.collections.thl_web import ( + LedgerDFCollection, + ) + from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge + from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager + from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager + from generalresearch.managers.thl.payout import ( + BrokerageProductPayoutEventManager, + BusinessPayoutEventManager, + PayoutEventManager, + UserPayoutEventManager, + ) + from generalresearch.managers.thl.product import ProductManager + from generalresearch.models.gr.business import Business + from generalresearch.models.thl.ledger import LedgerAccount + from generalresearch.models.thl.payout import ( + UserPayoutEvent, + ) + 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 + from generalresearch.redis_helper import RedisConfig logger = logging.getLogger() @@ -189,7 +194,7 @@ class TestPayout: utc_now: datetime, ): from generalresearch.models.thl.definitions import PayoutStatus - from generalresearch.models.thl.wallet import PayoutType + from generalresearch.models.thl.wallet.definitions import PayoutType user_account = thl_ledger_manager.get_account_or_create_user_wallet(user=user) bp_account = thl_ledger_manager.get_account_or_create_bp_wallet(product=product) diff --git a/tests/managers/thl/test_product.py b/tests/managers/thl/test_product.py index f93ac36..644dc90 100644 --- a/tests/managers/thl/test_product.py +++ b/tests/managers/thl/test_product.py @@ -1,13 +1,12 @@ from __future__ import annotations from collections.abc import Callable +from typing import TYPE_CHECKING 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.definitions import Source from generalresearch.models.thl.product import ( Product, ProfilingConfig, @@ -19,6 +18,10 @@ from generalresearch.models.thl.product import ( UserHealthConfig, ) +if TYPE_CHECKING: + from generalresearch.managers.thl.product import ProductManager + from generalresearch.models.gr.team import Team + class TestProductManagerGetMethods: def test_get_by_uuid(self, product_manager: ProductManager): diff --git a/tests/managers/thl/test_product_prod.py b/tests/managers/thl/test_product_prod.py index 8734210..d584527 100644 --- a/tests/managers/thl/test_product_prod.py +++ b/tests/managers/thl/test_product_prod.py @@ -2,13 +2,16 @@ from __future__ import annotations import logging from collections.abc import Callable +from typing import TYPE_CHECKING from uuid import uuid4 import pytest -from generalresearch.managers.thl.product import ProductManager from generalresearch.models.thl.product import Product +if TYPE_CHECKING: + from generalresearch.managers.thl.product import ProductManager + logger = logging.getLogger() diff --git a/tests/managers/thl/test_profiling/test_question.py b/tests/managers/thl/test_profiling/test_question.py index 97e7365..e4afb87 100644 --- a/tests/managers/thl/test_profiling/test_question.py +++ b/tests/managers/thl/test_profiling/test_question.py @@ -1,8 +1,11 @@ from collections.abc import Callable +from typing import TYPE_CHECKING from uuid import uuid4 -from generalresearch.managers.thl.profiling.question import QuestionManager -from generalresearch.models import Source +from generalresearch.models.definitions import Source + +if TYPE_CHECKING: + from generalresearch.managers.thl.profiling.question import QuestionManager class TestQuestionManager: diff --git a/tests/managers/thl/test_profiling/test_schema.py b/tests/managers/thl/test_profiling/test_schema.py index b0eae31..feab902 100644 --- a/tests/managers/thl/test_profiling/test_schema.py +++ b/tests/managers/thl/test_profiling/test_schema.py @@ -1,10 +1,13 @@ from collections.abc import Callable +from typing import TYPE_CHECKING -from generalresearch.managers.thl.profiling.schema import ( - UpkSchemaManager, -) from generalresearch.models.thl.profiling.upk_property import PropertyType +if TYPE_CHECKING: + from generalresearch.managers.thl.profiling.schema import ( + UpkSchemaManager, + ) + class TestUpkSchemaManager: diff --git a/tests/managers/thl/test_profiling/test_user_upk.py b/tests/managers/thl/test_profiling/test_user_upk.py index fa10b67..0f3140c 100644 --- a/tests/managers/thl/test_profiling/test_user_upk.py +++ b/tests/managers/thl/test_profiling/test_user_upk.py @@ -1,8 +1,10 @@ from collections.abc import Callable from datetime import UTC, datetime +from typing import TYPE_CHECKING -from generalresearch.managers.thl.profiling.user_upk import UserUpkManager -from generalresearch.models.thl.user import User +if TYPE_CHECKING: + from generalresearch.managers.thl.profiling.user_upk import UserUpkManager + from generalresearch.models.thl.user import User now = datetime.now(tz=UTC) base = { diff --git a/tests/managers/thl/test_session_manager.py b/tests/managers/thl/test_session_manager.py index 05a49c1..30fd9ec 100644 --- a/tests/managers/thl/test_session_manager.py +++ b/tests/managers/thl/test_session_manager.py @@ -3,24 +3,27 @@ from __future__ import annotations from collections.abc import Callable from datetime import datetime, timedelta from decimal import Decimal +from typing import TYPE_CHECKING 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.definitions import DeviceType 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 + +if TYPE_CHECKING: + from generalresearch.managers.thl.session import SessionManager + from generalresearch.models.gr.business import Business + from generalresearch.models.gr.team import Team + from generalresearch.models.thl.product import Product + from generalresearch.pg_helper import PostgresConfig fake = Faker() diff --git a/tests/managers/thl/test_survey.py b/tests/managers/thl/test_survey.py index c3ab162..e114b70 100644 --- a/tests/managers/thl/test_survey.py +++ b/tests/managers/thl/test_survey.py @@ -4,16 +4,11 @@ import uuid from collections.abc import Callable from datetime import UTC, datetime from decimal import Decimal +from typing import TYPE_CHECKING import pytest -from generalresearch.managers.thl.buyer import BuyerManager -from generalresearch.managers.thl.profiling.question import ( - QuestionManager, -) -from generalresearch.managers.thl.profiling.uqa import UQAManager -from generalresearch.managers.thl.survey import SurveyManager, SurveyStatManager -from generalresearch.models import Source +from generalresearch.models.definitions import Source from generalresearch.models.legacy.bucket import ( DurationSummary, PayoutSummary, @@ -30,6 +25,14 @@ from generalresearch.models.thl.survey.model import ( SurveyStat, ) +if TYPE_CHECKING: + from generalresearch.managers.thl.buyer import BuyerManager + from generalresearch.managers.thl.profiling.question import ( + QuestionManager, + ) + from generalresearch.managers.thl.profiling.uqa import UQAManager + from generalresearch.managers.thl.survey import SurveyManager, SurveyStatManager + @pytest.fixture(scope="session") def surveys_fixture() -> list[Survey]: @@ -223,7 +226,6 @@ class TestSurvey: class TestSurveyStat: def test( self, - delete_buyers_surveys: Callable[..., None], surveystat_manager, survey_manager: SurveyManager, surveys_fixture: list[Survey], diff --git a/tests/managers/thl/test_survey_penalty.py b/tests/managers/thl/test_survey_penalty.py index 9c29a0a..04f69d2 100644 --- a/tests/managers/thl/test_survey_penalty.py +++ b/tests/managers/thl/test_survey_penalty.py @@ -1,16 +1,19 @@ from __future__ import annotations import uuid +from typing import TYPE_CHECKING import pytest -from generalresearch.managers.thl.survey_penalty import SurveyPenaltyManager -from generalresearch.models import Source +from generalresearch.models.definitions import Source from generalresearch.models.thl.survey.penalty import ( BPSurveyPenalty, TeamSurveyPenalty, ) +if TYPE_CHECKING: + from generalresearch.managers.thl.survey_penalty import SurveyPenaltyManager + @pytest.fixture def product_uuid() -> str: diff --git a/tests/managers/thl/test_task_adjustment.py b/tests/managers/thl/test_task_adjustment.py index a7324c3..a14401e 100644 --- a/tests/managers/thl/test_task_adjustment.py +++ b/tests/managers/thl/test_task_adjustment.py @@ -5,23 +5,26 @@ from collections.abc import Callable from datetime import UTC, datetime, timedelta from decimal import Decimal from random import randint +from typing import TYPE_CHECKING 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.definitions 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 + +if TYPE_CHECKING: + 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.thl.session import Session + from generalresearch.models.thl.user import User @pytest.fixture() diff --git a/tests/managers/thl/test_task_status.py b/tests/managers/thl/test_task_status.py index b47f650..9846ce0 100644 --- a/tests/managers/thl/test_task_status.py +++ b/tests/managers/thl/test_task_status.py @@ -3,13 +3,11 @@ from __future__ import annotations from collections.abc import Callable from datetime import UTC, datetime, timedelta from decimal import Decimal +from typing import TYPE_CHECKING 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.definitions import Source from generalresearch.models.thl.definitions import ( Status, StatusCode1, @@ -19,12 +17,18 @@ from generalresearch.models.thl.product import ( PayoutConfig, PayoutTransformation, PayoutTransformationPercentArgs, - Product, UserWalletConfig, ) -from generalresearch.models.thl.session import Session, WallOut +from generalresearch.models.thl.session import WallOut from generalresearch.models.thl.task_status import TaskStatusResponse -from generalresearch.models.thl.user import User + +if TYPE_CHECKING: + from generalresearch.managers.thl.product import ProductManager + from generalresearch.managers.thl.session import SessionManager + from generalresearch.managers.thl.wall import WallManager + from generalresearch.models.thl.product import Product + from generalresearch.models.thl.session import Session + from generalresearch.models.thl.user import User start1 = datetime(2023, 2, 1, tzinfo=UTC) finish1 = start1 + timedelta(minutes=5) diff --git a/tests/managers/thl/test_user_manager/test_base.py b/tests/managers/thl/test_user_manager/test_base.py index 8cd83ad..4a9750e 100644 --- a/tests/managers/thl/test_user_manager/test_base.py +++ b/tests/managers/thl/test_user_manager/test_base.py @@ -1,11 +1,11 @@ import logging from datetime import UTC, datetime from random import randint +from typing import TYPE_CHECKING from uuid import uuid4 import pytest -from generalresearch.managers.thl.product import ProductManager from generalresearch.managers.thl.user_manager import ( UserCreateNotAllowedError, get_bp_user_create_limit_hourly, @@ -17,13 +17,17 @@ 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, UserCreateConfig +from generalresearch.models.thl.product import UserCreateConfig from generalresearch.models.thl.user import User -from generalresearch.pg_helper import PostgresConfig + +if TYPE_CHECKING: + from generalresearch.managers.thl.product import ProductManager + from generalresearch.managers.thl.user_manager.user_manager import ( + UserManager, + ) + from generalresearch.managers.thl.userhealth import AuditLogManager + from generalresearch.models.thl.product import Product + from generalresearch.pg_helper import PostgresConfig logger = logging.getLogger() diff --git a/tests/managers/thl/test_user_manager/test_mysql.py b/tests/managers/thl/test_user_manager/test_mysql.py index e6f43ef..ed7d458 100644 --- a/tests/managers/thl/test_user_manager/test_mysql.py +++ b/tests/managers/thl/test_user_manager/test_mysql.py @@ -1,9 +1,12 @@ from __future__ import annotations -from generalresearch.managers.thl.user_manager.mysql_user_manager import ( - MysqlUserManager, -) -from generalresearch.models.thl.user import User +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from generalresearch.managers.thl.user_manager.mysql_user_manager import ( + MysqlUserManager, + ) + from generalresearch.models.thl.user import User class TestUserManagerMysqlNew: diff --git a/tests/managers/thl/test_user_manager/test_redis.py b/tests/managers/thl/test_user_manager/test_redis.py index 04071ee..e51aae9 100644 --- a/tests/managers/thl/test_user_manager/test_redis.py +++ b/tests/managers/thl/test_user_manager/test_redis.py @@ -1,14 +1,18 @@ from __future__ import annotations +from typing import TYPE_CHECKING + 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 + +if TYPE_CHECKING: + from generalresearch.config import GRLBaseSettings + from generalresearch.models.thl.user import User + from generalresearch.pg_helper import PostgresConfig class TestUserManagerRedis: 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 87d010a..9a279ed 100644 --- a/tests/managers/thl/test_user_manager/test_user_fetch.py +++ b/tests/managers/thl/test_user_manager/test_user_fetch.py @@ -1,13 +1,15 @@ from __future__ import annotations from collections.abc import Callable +from typing import TYPE_CHECKING 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 +if TYPE_CHECKING: + 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: 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 670e38a..eb6a272 100644 --- a/tests/managers/thl/test_user_manager/test_user_metadata.py +++ b/tests/managers/thl/test_user_manager/test_user_metadata.py @@ -1,17 +1,20 @@ from __future__ import annotations from collections.abc import Callable +from typing import TYPE_CHECKING 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 +if TYPE_CHECKING: + 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 + class TestUserMetadataManager: diff --git a/tests/managers/thl/test_user_streak.py b/tests/managers/thl/test_user_streak.py index 61e2947..564a142 100644 --- a/tests/managers/thl/test_user_streak.py +++ b/tests/managers/thl/test_user_streak.py @@ -3,17 +3,15 @@ from __future__ import annotations import copy from datetime import UTC, date, datetime, timedelta from decimal import Decimal +from typing import TYPE_CHECKING from zoneinfo import ZoneInfo import pytest -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, @@ -21,6 +19,13 @@ from generalresearch.models.thl.user_streak import ( UserStreak, ) +if TYPE_CHECKING: + from generalresearch.managers.thl.session import SessionManager + from generalresearch.managers.thl.user_streak import ( + UserStreakManager, + ) + from generalresearch.models.thl.user import User + def test_compute_streaks_from_days(): days = [ diff --git a/tests/managers/thl/test_userhealth.py b/tests/managers/thl/test_userhealth.py index ea54359..ce6c221 100644 --- a/tests/managers/thl/test_userhealth.py +++ b/tests/managers/thl/test_userhealth.py @@ -2,6 +2,7 @@ from __future__ import annotations from collections.abc import Callable from datetime import UTC, datetime +from typing import TYPE_CHECKING from uuid import uuid4 import faker @@ -12,16 +13,24 @@ from generalresearch.managers.thl.userhealth import ( IPRecordManager, UserIpHistoryManager, ) -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.ipinfo import ( + GeoIPInformation, +) 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 + +if TYPE_CHECKING: + from generalresearch.models.thl.ipinfo import ( + IPGeoname, + IPInformation, + ) + from generalresearch.models.thl.product import Product + from generalresearch.models.thl.user import User + from generalresearch.pg_helper import PostgresConfig + from generalresearch.redis_helper import RedisConfig fake = faker.Faker() diff --git a/tests/managers/thl/test_wall_manager.py b/tests/managers/thl/test_wall_manager.py index b8a636f..3215de8 100644 --- a/tests/managers/thl/test_wall_manager.py +++ b/tests/managers/thl/test_wall_manager.py @@ -3,21 +3,24 @@ from __future__ import annotations from collections.abc import Callable from datetime import UTC, datetime, timedelta from decimal import Decimal +from typing import TYPE_CHECKING 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.definitions import Source from generalresearch.models.thl.session import ( ReportValue, - Session, Status, StatusCode1, ) -from generalresearch.models.thl.user import User + +if TYPE_CHECKING: + from generalresearch.managers.thl.session import SessionManager + from generalresearch.managers.thl.wall import WallCacheManager, WallManager + from generalresearch.models.thl.session import Session + from generalresearch.models.thl.user import User class TestWallManager: diff --git a/tests/models/custom_types/test_aware_datetime.py b/tests/models/custom_types/test_aware_datetime.py index e8a5aa3..54a5d9b 100644 --- a/tests/models/custom_types/test_aware_datetime.py +++ b/tests/models/custom_types/test_aware_datetime.py @@ -2,12 +2,14 @@ from __future__ import annotations import logging from datetime import UTC, datetime +from typing import TYPE_CHECKING import pytest import pytz from pydantic import BaseModel, Field, ValidationError -from generalresearch.models.custom_types import AwareDatetimeISO +if TYPE_CHECKING: + from generalresearch.models.custom_types import AwareDatetimeISO logger = logging.getLogger() diff --git a/tests/models/custom_types/test_dsn.py b/tests/models/custom_types/test_dsn.py index d8c7c53..2aae579 100644 --- a/tests/models/custom_types/test_dsn.py +++ b/tests/models/custom_types/test_dsn.py @@ -1,12 +1,14 @@ from __future__ import annotations +from typing import TYPE_CHECKING from uuid import uuid4 import pytest from pydantic import BaseModel, Field, MySQLDsn, ValidationError from pydantic_core import Url -from generalresearch.models.custom_types import DaskDsn, SentryDsn +if TYPE_CHECKING: + from generalresearch.models.custom_types import DaskDsn, SentryDsn # --- Test Pydantic Models --- diff --git a/tests/models/custom_types/test_uuid_str.py b/tests/models/custom_types/test_uuid_str.py index 02e6a8b..92489a0 100644 --- a/tests/models/custom_types/test_uuid_str.py +++ b/tests/models/custom_types/test_uuid_str.py @@ -1,11 +1,13 @@ from __future__ import annotations +from typing import TYPE_CHECKING from uuid import uuid4 import pytest from pydantic import BaseModel, Field, ValidationError -from generalresearch.models.custom_types import UUIDStr +if TYPE_CHECKING: + from generalresearch.models.custom_types import UUIDStr class UUIDStrModel(BaseModel): diff --git a/tests/models/dynata/test_survey.py b/tests/models/dynata/test_survey.py index ad953a3..3e33897 100644 --- a/tests/models/dynata/test_survey.py +++ b/tests/models/dynata/test_survey.py @@ -1,3 +1,6 @@ +from __future__ import annotations + + class TestDynataCondition: def test_condition_create(self): diff --git a/tests/models/gr/test_authentication.py b/tests/models/gr/test_authentication.py index d2a7054..ac1298f 100644 --- a/tests/models/gr/test_authentication.py +++ b/tests/models/gr/test_authentication.py @@ -6,17 +6,21 @@ import os from collections.abc import Callable from datetime import UTC, datetime from random import randint +from typing import TYPE_CHECKING from uuid import uuid4 import pytest from redis import Redis 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 +from generalresearch.models.gr.team import Team + +if TYPE_CHECKING: + from generalresearch.models.gr.business import Business + from generalresearch.models.gr.team import Membership + from generalresearch.models.thl.product import Product + from generalresearch.pg_helper import PostgresConfig + from generalresearch.redis_helper import RedisConfig SSO_ISSUER = "" diff --git a/tests/models/gr/test_base.py b/tests/models/gr/test_base.py index f019fc0..fba0960 100644 --- a/tests/models/gr/test_base.py +++ b/tests/models/gr/test_base.py @@ -3,11 +3,13 @@ from __future__ import annotations import subprocess from collections.abc import Callable from pathlib import Path +from typing import TYPE_CHECKING import pytest from pydantic import PostgresDsn -from generalresearch.pg_helper import PostgresConfig +if TYPE_CHECKING: + from generalresearch.pg_helper import PostgresConfig class TestGRPostgresDjangoCreation: diff --git a/tests/models/gr/test_business.py b/tests/models/gr/test_business.py index 9310d2c..2c12da1 100644 --- a/tests/models/gr/test_business.py +++ b/tests/models/gr/test_business.py @@ -19,40 +19,42 @@ from distributed.utils_test import ( 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 -from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager -from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager -from generalresearch.managers.thl.payout import ( - BusinessPayoutEventManager, - PayoutEventManager, -) from generalresearch.models.gr.business import ( Business, BusinessAddress, - BusinessBankAccount, BusinessContact, ) -from generalresearch.models.gr.team import Team from generalresearch.models.thl.finance import ( BusinessBalances, ProductBalances, ) from generalresearch.models.thl.product import Product -from generalresearch.models.thl.user import User -from generalresearch.pg_helper import PostgresConfig -from generalresearch.redis_helper import RedisConfig if TYPE_CHECKING: + 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 + from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager + from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager + from generalresearch.managers.thl.payout import ( + BusinessPayoutEventManager, + PayoutEventManager, + ) from generalresearch.managers.thl.product import ProductManager + from generalresearch.models.gr.business import ( + BusinessBankAccount, + ) + from generalresearch.models.gr.team import Team from generalresearch.models.thl.product import BrokerageProductPayoutEvent from generalresearch.models.thl.session import Session + from generalresearch.models.thl.user import User + from generalresearch.pg_helper import PostgresConfig + from generalresearch.redis_helper import RedisConfig class TestBusinessBankAccount: diff --git a/tests/models/gr/test_team.py b/tests/models/gr/test_team.py index dc7d4b9..c1ae6d6 100644 --- a/tests/models/gr/test_team.py +++ b/tests/models/gr/test_team.py @@ -5,6 +5,7 @@ from collections.abc import Callable from datetime import datetime, timedelta from decimal import Decimal from pathlib import Path +from typing import TYPE_CHECKING import pandas as pd from dask.distributed import Client as DaskClient @@ -12,24 +13,29 @@ from distributed.utils_test import ( client_no_amm, ) -from generalresearch.incite.base import GRLDatasets -from generalresearch.incite.collections.thl_web import ( - SessionDFCollection, - WallDFCollection, -) -from generalresearch.incite.mergers.foundations.enriched_session import ( - EnrichedSessionMerge, -) -from generalresearch.incite.mergers.foundations.enriched_wall import EnrichedWallMerge -from generalresearch.managers.gr.team import MembershipManager, TeamManager -from generalresearch.models.gr.authentication import GRUser from generalresearch.models.gr.business import Business -from generalresearch.models.gr.team import Membership, Team +from generalresearch.models.gr.team import Team 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 -from generalresearch.redis_helper import RedisConfig + +if TYPE_CHECKING: + from generalresearch.incite.base import GRLDatasets + from generalresearch.incite.collections.thl_web import ( + SessionDFCollection, + WallDFCollection, + ) + from generalresearch.incite.mergers.foundations.enriched_session import ( + EnrichedSessionMerge, + ) + from generalresearch.incite.mergers.foundations.enriched_wall import ( + EnrichedWallMerge, + ) + from generalresearch.managers.gr.team import MembershipManager, TeamManager + from generalresearch.models.gr.authentication import GRUser + from generalresearch.models.gr.team import Membership + from generalresearch.models.thl.session import Session + from generalresearch.models.thl.user import User + from generalresearch.pg_helper import PostgresConfig + from generalresearch.redis_helper import RedisConfig class TestTeam: diff --git a/tests/models/innovate/test_question.py b/tests/models/innovate/test_question.py index b206177..ea2fc8c 100644 --- a/tests/models/innovate/test_question.py +++ b/tests/models/innovate/test_question.py @@ -1,6 +1,6 @@ from __future__ import annotations -from generalresearch.models import Source +from generalresearch.models.definitions import Source from generalresearch.models.innovate.question import ( InnovateQuestion, InnovateQuestionOption, diff --git a/tests/models/legacy/test_offerwall_parse_response.py b/tests/models/legacy/test_offerwall_parse_response.py index 56ba077..93f5c26 100644 --- a/tests/models/legacy/test_offerwall_parse_response.py +++ b/tests/models/legacy/test_offerwall_parse_response.py @@ -2,7 +2,7 @@ from __future__ import annotations import json -from generalresearch.models import Source +from generalresearch.models.definitions import Source from generalresearch.models.legacy.bucket import ( BucketTask, DurationSummary, diff --git a/tests/models/legacy/test_user_question_answer_in.py b/tests/models/legacy/test_user_question_answer_in.py index 3fdaa05..f14c1a7 100644 --- a/tests/models/legacy/test_user_question_answer_in.py +++ b/tests/models/legacy/test_user_question_answer_in.py @@ -4,19 +4,22 @@ import json from collections.abc import Callable from datetime import datetime from decimal import Decimal +from typing import TYPE_CHECKING from uuid import uuid4 import pytest -from generalresearch.managers.thl.user_manager.user_manager import UserManager -from generalresearch.models import Source +from generalresearch.models.definitions import Source from generalresearch.models.legacy.questions import ( UserQuestionAnswers, ) -from generalresearch.models.thl.product import Product from generalresearch.models.thl.session import Session, Wall from generalresearch.models.thl.user import User +if TYPE_CHECKING: + from generalresearch.managers.thl.user_manager.user_manager import UserManager + from generalresearch.models.thl.product import Product + class TestUserQuestionAnswers: """This is for the GRS POST submission that may contain multiple diff --git a/tests/models/network/test_mtr.py b/tests/models/network/test_mtr.py index 7f8a736..5d136c4 100644 --- a/tests/models/network/test_mtr.py +++ b/tests/models/network/test_mtr.py @@ -1,11 +1,15 @@ from __future__ import annotations +from typing import TYPE_CHECKING + import faker -from generalresearch.managers.network.tool_run import ToolRunManager from generalresearch.models.network.mtr.execute import execute_mtr from generalresearch.models.network.tool_run import ToolClass, ToolName +if TYPE_CHECKING: + from generalresearch.managers.network.tool_run import ToolRunManager + fake = faker.Faker() diff --git a/tests/models/network/test_nmap.py b/tests/models/network/test_nmap.py index db39997..6adc9e4 100644 --- a/tests/models/network/test_nmap.py +++ b/tests/models/network/test_nmap.py @@ -1,14 +1,18 @@ from __future__ import annotations import subprocess +from typing import TYPE_CHECKING import faker -from generalresearch.managers.network.tool_run import ToolRunManager from generalresearch.models.network.definitions import IPProtocol from generalresearch.models.network.nmap.execute import execute_nmap from generalresearch.models.network.nmap.result import NmapResult, PortState -from generalresearch.models.network.tool_run import NmapRun, ToolClass, ToolName +from generalresearch.models.network.tool_run import ToolClass, ToolName + +if TYPE_CHECKING: + from generalresearch.managers.network.tool_run import ToolRunManager + from generalresearch.models.network.tool_run import NmapRun fake = faker.Faker() diff --git a/tests/models/network/test_nmap_parser.py b/tests/models/network/test_nmap_parser.py index 473a63f..fc9884b 100644 --- a/tests/models/network/test_nmap_parser.py +++ b/tests/models/network/test_nmap_parser.py @@ -1,14 +1,15 @@ from __future__ import annotations import os +from typing import TYPE_CHECKING import pytest from generalresearch.models.network.nmap.parser import parse_nmap_xml -from generalresearch.models.network.nmap.result import ( - NmapResult, - NmapTrace, -) +from generalresearch.models.network.nmap.result import NmapTrace + +if TYPE_CHECKING: + from generalresearch.models.network.nmap.result import NmapResult @pytest.fixture diff --git a/tests/models/network/test_rdns.py b/tests/models/network/test_rdns.py index 1a15a28..82126dd 100644 --- a/tests/models/network/test_rdns.py +++ b/tests/models/network/test_rdns.py @@ -1,11 +1,15 @@ from __future__ import annotations +from typing import TYPE_CHECKING + import faker -from generalresearch.managers.network.tool_run import ToolRunManager from generalresearch.models.network.rdns.execute import execute_rdns from generalresearch.models.network.tool_run import ToolClass, ToolName +if TYPE_CHECKING: + from generalresearch.managers.network.tool_run import ToolRunManager + fake = faker.Faker() diff --git a/tests/models/spectrum/test_question.py b/tests/models/spectrum/test_question.py index a44286d..d469530 100644 --- a/tests/models/spectrum/test_question.py +++ b/tests/models/spectrum/test_question.py @@ -2,7 +2,7 @@ from __future__ import annotations from datetime import UTC, datetime -from generalresearch.models import Source +from generalresearch.models.definitions import Source from generalresearch.models.spectrum.question import ( SpectrumQuestion, SpectrumQuestionClass, diff --git a/tests/models/spectrum/test_survey.py b/tests/models/spectrum/test_survey.py index bad6857..02c5d3f 100644 --- a/tests/models/spectrum/test_survey.py +++ b/tests/models/spectrum/test_survey.py @@ -3,7 +3,7 @@ from __future__ import annotations from datetime import UTC, datetime from decimal import Decimal -from generalresearch.models import ( +from generalresearch.models.definitions import ( LogicalOperator, Source, TaskCalculationType, diff --git a/tests/models/spectrum/test_survey_manager.py b/tests/models/spectrum/test_survey_manager.py index 11dc01f..0300956 100644 --- a/tests/models/spectrum/test_survey_manager.py +++ b/tests/models/spectrum/test_survey_manager.py @@ -3,15 +3,17 @@ from __future__ import annotations import logging from datetime import UTC, datetime from decimal import Decimal -from typing import Any +from typing import TYPE_CHECKING, Any from pymysql import IntegrityError from generalresearch.config import is_debug -from generalresearch.managers.spectrum.survey import ( - SpectrumSurveyManager, -) -from generalresearch.sql_helper import SqlHelper + +if TYPE_CHECKING: + from generalresearch.managers.spectrum.survey import ( + SpectrumSurveyManager, + ) + from generalresearch.sql_helper import SqlHelper logger = logging.getLogger() diff --git a/tests/models/test_device.py b/tests/models/test_device.py index 8e1251a..fdbd906 100644 --- a/tests/models/test_device.py +++ b/tests/models/test_device.py @@ -15,7 +15,7 @@ chromebook_ua_string = ( ) -from generalresearch.models import DeviceType +from generalresearch.models.definitions import DeviceType from generalresearch.models.device import parse_device_from_useragent diff --git a/tests/models/test_finance.py b/tests/models/test_finance.py index 72f4f4d..eabc877 100644 --- a/tests/models/test_finance.py +++ b/tests/models/test_finance.py @@ -4,6 +4,7 @@ from collections.abc import Callable from datetime import UTC, datetime, timedelta from itertools import product as iter_product from random import randint +from typing import TYPE_CHECKING from uuid import uuid4 import pandas as pd @@ -16,25 +17,27 @@ from distributed.utils_test import ( ) from faker import Faker -from generalresearch.incite.collections.thl_web import LedgerDFCollection -from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge from generalresearch.incite.schemas.mergers.pop_ledger import ( numerical_col_names, ) -from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager from generalresearch.models.thl.finance import ( BusinessBalances, POPFinancial, ProductBalances, ) -from generalresearch.models.thl.ledger import LedgerAccount -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 from test_utils.incite.collections.conftest import ledger_collection from test_utils.incite.mergers.conftest import pop_ledger_merge +if TYPE_CHECKING: + from generalresearch.incite.collections.thl_web import LedgerDFCollection + from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge + from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager + from generalresearch.models.thl.ledger import LedgerAccount + 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() diff --git a/tests/models/thl/test_adjustments.py b/tests/models/thl/test_adjustments.py index c5b3f6b..cd75318 100644 --- a/tests/models/thl/test_adjustments.py +++ b/tests/models/thl/test_adjustments.py @@ -3,22 +3,27 @@ from __future__ import annotations from collections.abc import Callable from datetime import UTC, datetime, timedelta from decimal import Decimal +from typing import TYPE_CHECKING import pytest -from generalresearch.managers.thl.session import SessionManager -from generalresearch.managers.thl.wall import WallManager -from generalresearch.models import Source +from generalresearch.models.definitions import Source from generalresearch.models.thl.product import Product from generalresearch.models.thl.session import ( - Session, SessionAdjustedStatus, Status, StatusCode1, - Wall, WallAdjustedStatus, ) -from generalresearch.models.thl.user import User + +if TYPE_CHECKING: + from generalresearch.managers.thl.session import SessionManager + from generalresearch.managers.thl.wall import WallManager + from generalresearch.models.thl.session import ( + Session, + Wall, + ) + from generalresearch.models.thl.user import User started1 = datetime(2023, 1, 1, tzinfo=UTC) started2 = datetime(2023, 1, 1, 0, 10, 0, tzinfo=UTC) diff --git a/tests/models/thl/test_buyer.py b/tests/models/thl/test_buyer.py index 02093e2..ef97166 100644 --- a/tests/models/thl/test_buyer.py +++ b/tests/models/thl/test_buyer.py @@ -1,6 +1,6 @@ from __future__ import annotations -from generalresearch.models import Source +from generalresearch.models.definitions import Source from generalresearch.models.thl.survey.buyer import BuyerCountryStat diff --git a/tests/models/thl/test_contest/test_contest.py b/tests/models/thl/test_contest/test_contest.py index e1053f4..ed8477b 100644 --- a/tests/models/thl/test_contest/test_contest.py +++ b/tests/models/thl/test_contest/test_contest.py @@ -1,11 +1,13 @@ from __future__ import annotations from collections.abc import Callable +from typing import TYPE_CHECKING import pytest -from generalresearch.models.thl.product import Product -from generalresearch.models.thl.user import User +if TYPE_CHECKING: + from generalresearch.models.thl.product import Product + from generalresearch.models.thl.user import User class TestContest: diff --git a/tests/models/thl/test_contest/test_leaderboard_contest.py b/tests/models/thl/test_contest/test_leaderboard_contest.py index 99cfb37..c49776b 100644 --- a/tests/models/thl/test_contest/test_leaderboard_contest.py +++ b/tests/models/thl/test_contest/test_leaderboard_contest.py @@ -1,6 +1,7 @@ from __future__ import annotations from datetime import UTC +from typing import TYPE_CHECKING from uuid import uuid4 import pytest @@ -8,7 +9,6 @@ from redis import Redis from generalresearch.currency import USDCent from generalresearch.managers.leaderboard.manager import LeaderboardManager -from generalresearch.managers.thl.user_manager.user_manager import UserManager from generalresearch.models.thl.contest import ContestPrize from generalresearch.models.thl.contest.definitions import ( ContestPrizeKind, @@ -21,10 +21,13 @@ from generalresearch.models.thl.contest.utils import ( distribute_leaderboard_prizes, ) from generalresearch.models.thl.leaderboard import LeaderboardRow -from generalresearch.models.thl.product import Product from generalresearch.models.thl.user import User from tests.models.thl.test_contest.test_contest import TestContest +if TYPE_CHECKING: + from generalresearch.managers.thl.user_manager.user_manager import UserManager + from generalresearch.models.thl.product import Product + class TestLeaderboardContest(TestContest): diff --git a/tests/models/thl/test_contest/test_raffle_contest.py b/tests/models/thl/test_contest/test_raffle_contest.py index 8812cb3..e71851e 100644 --- a/tests/models/thl/test_contest/test_raffle_contest.py +++ b/tests/models/thl/test_contest/test_raffle_contest.py @@ -2,6 +2,7 @@ from __future__ import annotations from collections import Counter from datetime import datetime +from typing import TYPE_CHECKING from uuid import uuid4 import pytest @@ -21,10 +22,12 @@ from generalresearch.models.thl.contest.definitions import ( ContestType, ) from generalresearch.models.thl.contest.raffle import RaffleContest -from generalresearch.models.thl.product import Product -from generalresearch.models.thl.user import User from tests.models.thl.test_contest.test_contest import TestContest +if TYPE_CHECKING: + from generalresearch.models.thl.product import Product + from generalresearch.models.thl.user import User + class TestRaffleContest(TestContest): diff --git a/tests/models/thl/test_marketplace_condition.py b/tests/models/thl/test_marketplace_condition.py index 1dd25e8..6936a7c 100644 --- a/tests/models/thl/test_marketplace_condition.py +++ b/tests/models/thl/test_marketplace_condition.py @@ -3,7 +3,7 @@ from __future__ import annotations import pytest from pydantic import ValidationError -from generalresearch.models import LogicalOperator +from generalresearch.models.definitions import LogicalOperator from generalresearch.models.thl.survey.condition import ( ConditionValueType, MarketplaceCondition, @@ -130,7 +130,7 @@ class TestMarketplaceCondition: assert c.evaluate_criterion(user_qas) is None def test_list_and_negate(self): - from generalresearch.models import LogicalOperator + from generalresearch.models.definitions import LogicalOperator from generalresearch.models.thl.survey.condition import ( ConditionValueType, MarketplaceCondition, @@ -247,7 +247,7 @@ class TestMarketplaceCondition: assert ["1", "10", "11", "12", "2", "3", "4", "5"] == c.values def test_ranges_infinity(self): - from generalresearch.models import LogicalOperator + from generalresearch.models.definitions import LogicalOperator from generalresearch.models.thl.survey.condition import ( ConditionValueType, MarketplaceCondition, diff --git a/tests/models/thl/test_payout.py b/tests/models/thl/test_payout.py index daf1bd7..927687e 100644 --- a/tests/models/thl/test_payout.py +++ b/tests/models/thl/test_payout.py @@ -7,12 +7,16 @@ from pydantic import ValidationError from generalresearch.currency import USDCent from generalresearch.models.gr import Team -from generalresearch.models.gr.business import Business, BusinessAddress, BusinessType +from generalresearch.models.gr.business import ( + Business, + BusinessAddress, + BusinessType, +) from generalresearch.models.thl.payout import ( BrokerageProductPayoutEvent, BusinessPayoutEvent, ) -from generalresearch.models.thl.wallet import PayoutType +from generalresearch.models.thl.wallet.definitions import PayoutType class TestBusinessPayoutEvent: diff --git a/tests/models/thl/test_payout_format.py b/tests/models/thl/test_payout_format.py index fe7aea5..56eafe3 100644 --- a/tests/models/thl/test_payout_format.py +++ b/tests/models/thl/test_payout_format.py @@ -1,14 +1,20 @@ from __future__ import annotations +from typing import TYPE_CHECKING + import pytest from pydantic import BaseModel from generalresearch.models.thl.payout_format import ( PayoutFormatField, - PayoutFormatType, format_payout_format, ) +if TYPE_CHECKING: + from generalresearch.models.thl.payout_format import ( + PayoutFormatType, + ) + class PayoutFormatTestClass(BaseModel): payout_format: PayoutFormatType = PayoutFormatField diff --git a/tests/models/thl/test_product.py b/tests/models/thl/test_product.py index 880799a..cc0fa8e 100644 --- a/tests/models/thl/test_product.py +++ b/tests/models/thl/test_product.py @@ -5,6 +5,7 @@ import shutil from collections.abc import Callable from datetime import UTC, datetime, timedelta from decimal import Decimal +from typing import TYPE_CHECKING from uuid import uuid4 import pytest @@ -12,22 +13,9 @@ from dask.distributed import Client as DaskClient from pydantic import ValidationError from generalresearch.currency import USDCent -from generalresearch.incite.base import GRLDatasets -from generalresearch.incite.collections.thl_web import LedgerDFCollection -from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge -from generalresearch.managers.thl.ledger_manager.thl_ledger import ( - ThlLedgerManager, -) -from generalresearch.managers.thl.payout import PayoutEventManager -from generalresearch.managers.thl.product import ProductManager -from generalresearch.models import Source -from generalresearch.models.gr.business import Business +from generalresearch.models.definitions import Source from generalresearch.models.thl.finance import ProductBalances -from generalresearch.models.thl.payout import ( - BrokerageProductPayoutEvent, -) from generalresearch.models.thl.product import ( - BrokerageProductPayoutEventManager, IntegrationMode, PayoutConfig, PayoutTransformation, @@ -39,9 +27,24 @@ from generalresearch.models.thl.product import ( SupplyConfig, SupplyPolicy, ) -from generalresearch.models.thl.session import Session -from generalresearch.models.thl.user import User -from generalresearch.redis_helper import RedisConfig + +if TYPE_CHECKING: + from generalresearch.incite.base import GRLDatasets + from generalresearch.incite.collections.thl_web import LedgerDFCollection + from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge + from generalresearch.managers.thl.ledger_manager.thl_ledger import ( + ThlLedgerManager, + ) + from generalresearch.managers.thl.payout import PayoutEventManager + from generalresearch.managers.thl.product import ProductManager + from generalresearch.models.gr.business import Business + from generalresearch.models.thl.payout import ( + BrokerageProductPayoutEvent, + ) + from generalresearch.models.thl.product import BrokerageProductPayoutEventManager + from generalresearch.models.thl.session import Session + from generalresearch.models.thl.user import User + from generalresearch.redis_helper import RedisConfig class TestProduct: diff --git a/tests/models/thl/test_product_userwalletconfig.py b/tests/models/thl/test_product_userwalletconfig.py index b348981..b1bfa60 100644 --- a/tests/models/thl/test_product_userwalletconfig.py +++ b/tests/models/thl/test_product_userwalletconfig.py @@ -6,7 +6,7 @@ from random import shuffle as rshuffle from generalresearch.models.thl.product import ( UserWalletConfig, ) -from generalresearch.models.thl.wallet import PayoutType +from generalresearch.models.thl.wallet.definitions import PayoutType def all_equal(iterable: list[str]) -> bool: diff --git a/tests/models/thl/test_soft_pair.py b/tests/models/thl/test_soft_pair.py index 3cf835e..34902e2 100644 --- a/tests/models/thl/test_soft_pair.py +++ b/tests/models/thl/test_soft_pair.py @@ -1,6 +1,6 @@ from __future__ import annotations -from generalresearch.models import Source +from generalresearch.models.definitions import Source from generalresearch.models.dynata.survey import ( ConditionValueType, DynataCondition, diff --git a/tests/models/thl/test_user.py b/tests/models/thl/test_user.py index 9c4b548..bc941d4 100644 --- a/tests/models/thl/test_user.py +++ b/tests/models/thl/test_user.py @@ -6,16 +6,19 @@ from datetime import UTC, datetime, timedelta, timezone from decimal import Decimal from random import choice as rand_choice from random import randint +from typing import TYPE_CHECKING from uuid import uuid4 import pytest from pydantic import ValidationError -from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager -from generalresearch.managers.thl.userhealth import AuditLogManager -from generalresearch.models.thl.product import Product from generalresearch.models.thl.user import User +if TYPE_CHECKING: + from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager + from generalresearch.managers.thl.userhealth import AuditLogManager + from generalresearch.models.thl.product import Product + class TestUserUserID: diff --git a/tests/models/thl/test_user_metadata.py b/tests/models/thl/test_user_metadata.py index a7b479d..7e84f3e 100644 --- a/tests/models/thl/test_user_metadata.py +++ b/tests/models/thl/test_user_metadata.py @@ -2,7 +2,7 @@ from __future__ import annotations import pytest -from generalresearch.models import MAX_INT32 +from generalresearch.models.definitions import MAX_INT32 from generalresearch.models.thl.user_profile import UserMetadata diff --git a/tests/models/thl/test_wall.py b/tests/models/thl/test_wall.py index 58e9825..61ca11d 100644 --- a/tests/models/thl/test_wall.py +++ b/tests/models/thl/test_wall.py @@ -7,7 +7,7 @@ from uuid import uuid4 import pytest from pydantic import ValidationError -from generalresearch.models import Source +from generalresearch.models.definitions import Source from generalresearch.models.thl.definitions import ( Status, StatusCode1, diff --git a/tests/models/thl/test_wall_session.py b/tests/models/thl/test_wall_session.py index 48b89ea..40d3619 100644 --- a/tests/models/thl/test_wall_session.py +++ b/tests/models/thl/test_wall_session.py @@ -5,7 +5,7 @@ from decimal import Decimal import pytest -from generalresearch.models import Source +from generalresearch.models.definitions import Source from generalresearch.models.thl.definitions import Status, StatusCode1 from generalresearch.models.thl.session import Session, Wall from generalresearch.models.thl.user import User diff --git a/tests/test_postgres.py b/tests/test_postgres.py index ed5a7ae..c53f644 100644 --- a/tests/test_postgres.py +++ b/tests/test_postgres.py @@ -1,12 +1,15 @@ import socket import subprocess from collections.abc import Callable +from typing import TYPE_CHECKING from pydantic import PostgresDsn -from generalresearch.models.custom_types import InternalHostname, PostgresDict from generalresearch.pg_helper import PostgresConfig +if TYPE_CHECKING: + from generalresearch.models.custom_types import InternalHostname, PostgresDict + def is_port_open(host: InternalHostname, port: int = 5432, timeout: int = 3): try: -- cgit v1.2.3 From 6469e7e55a53cfe18bd015b3c455ecbbb550cbb9 Mon Sep 17 00:00:00 2001 From: Max Nanis Date: Tue, 1 Sep 2026 12:29:10 -0700 Subject: WIP Business tests, fixture cleanup(s) --- generalresearch/incite/base.py | 4 +- generalresearch/incite/defaults.py | 10 +- generalresearch/managers/__init__.py | 16 -- generalresearch/managers/gr/business.py | 3 +- generalresearch/managers/pollfish/user_pid.py | 2 +- generalresearch/managers/thl/cashout_method.py | 9 +- generalresearch/models/__init__.py | 114 ------------- generalresearch/models/gr/business.py | 13 +- generalresearch/models/gr/definitions.py | 13 ++ generalresearch/models/thl/__init__.py | 18 +-- generalresearch/models/thl/session.py | 8 +- generalresearch/models/thl/task_status.py | 2 +- generalresearch/models/thl/utils.py | 11 ++ generalresearch/models/thl/wallet/__init__.py | 87 ---------- test_utils/conftest.py | 2 +- test_utils/incite/collections/conftest.py | 2 +- test_utils/incite/conftest.py | 10 +- test_utils/incite/mergers/conftest.py | 16 +- test_utils/managers/gr/conftest.py | 28 ---- test_utils/managers/thl/conftest.py | 37 ++++- test_utils/models/conftest.py | 4 +- test_utils/models/contest/conftest.py | 12 +- test_utils/models/gr/conftest.py | 2 +- test_utils/models/ledger/conftest.py | 108 +++++++------ .../incite/collections/test_df_collection_base.py | 6 +- .../collections/test_df_collection_item_base.py | 6 +- tests/incite/test_interval_idx.py | 2 +- tests/managers/gr/test_business.py | 32 ++-- tests/managers/thl/test_ledger/test_lm_accounts.py | 96 ++++++----- tests/managers/thl/test_ledger/test_thl_lm_tx.py | 5 +- tests/managers/thl/test_payout.py | 176 ++++++++++----------- tests/managers/thl/test_session_manager.py | 10 +- tests/models/gr/test_authentication.py | 55 +++---- tests/models/gr/test_business.py | 86 +++++----- tests/models/gr/test_team.py | 6 +- tests/models/test_finance.py | 14 +- tests/models/thl/test_payout.py | 2 +- tests/models/thl/test_product.py | 99 ++++++++---- 38 files changed, 484 insertions(+), 642 deletions(-) create mode 100644 generalresearch/models/gr/definitions.py create mode 100644 generalresearch/models/thl/utils.py (limited to 'tests/incite/collections') diff --git a/generalresearch/incite/base.py b/generalresearch/incite/base.py index 473a124..a06aac9 100644 --- a/generalresearch/incite/base.py +++ b/generalresearch/incite/base.py @@ -95,7 +95,7 @@ class GRLDatasets(BaseModel): from generalresearch.incite.collections.thl_marketplaces import ( DFCollectionType, ) - from generalresearch.incite.mergers import MergeType + from generalresearch.incite.mergers.base import MergeType assert self.data_src, "data src must be defined" @@ -128,7 +128,7 @@ class GRLDatasets(BaseModel): type.. """ - from generalresearch.incite.mergers import MergeType + from generalresearch.incite.mergers.base import MergeType folder = "mergers" if isinstance(enum_type, MergeType) else "raw/df-collections" assert self.incite is not None diff --git a/generalresearch/incite/defaults.py b/generalresearch/incite/defaults.py index 368b74a..5ee305b 100644 --- a/generalresearch/incite/defaults.py +++ b/generalresearch/incite/defaults.py @@ -3,7 +3,7 @@ from __future__ import annotations from datetime import UTC, datetime from generalresearch.incite.base import GRLDatasets -from generalresearch.incite.collections import DFCollectionType +from generalresearch.incite.collections.base import DFCollectionType from generalresearch.incite.collections.thl_marketplaces import ( InnovateSurveyHistoryCollection, MorningSurveyTimeseriesCollection, @@ -82,7 +82,7 @@ def ledger_df_collection( ds: GRLDatasets, pg_config: PostgresConfig ) -> LedgerDFCollection: return LedgerDFCollection( - offset="12d", + offset="12D", pg_config=pg_config, # thl_web:ledger_transaction - 1st record is 2018-03-14 20:22:17.408232 start=datetime(year=2018, month=3, day=14, hour=0, tzinfo=UTC), @@ -153,7 +153,7 @@ def user_id_product(ds: GRLDatasets) -> UserIdProductMerge: def enriched_session(ds: GRLDatasets) -> EnrichedSessionMerge: return EnrichedSessionMerge( start=datetime(year=2023, month=5, day=1, tzinfo=UTC), - offset="14d", + offset="14D", archive_path=ds.archive_path(enum_type=MergeType.ENRICHED_SESSION), ) @@ -162,7 +162,7 @@ def enriched_wall(ds: GRLDatasets) -> EnrichedWallMerge: return EnrichedWallMerge( # start=datetime(year=2022, month=5, day=1, tzinfo=timezone.utc), start=datetime(year=2023, month=7, day=23, tzinfo=UTC), - offset="14d", + offset="14D", archive_path=ds.archive_path(enum_type=MergeType.ENRICHED_WALL), ) @@ -180,7 +180,7 @@ def pop_ledger(ds: GRLDatasets) -> PopLedgerMerge: return PopLedgerMerge( # thl_web:ledger_transaction - 1st record is 2018-03-14 20:22:17.408232 start=datetime(year=2018, month=3, day=14, hour=0, tzinfo=UTC), - offset="30d", + offset="30D", archive_path=ds.archive_path(enum_type=MergeType.POP_LEDGER), ) diff --git a/generalresearch/managers/__init__.py b/generalresearch/managers/__init__.py index bc745fd..e69de29 100644 --- a/generalresearch/managers/__init__.py +++ b/generalresearch/managers/__init__.py @@ -1,16 +0,0 @@ -def parse_order_by(order_by_str: str) -> str: - """ - Converts django-rest-framework ordering str to mysql clause - :param order_by_str: e.g. 'created,-name' - :return: mysql clause e.g. ORDER BY created ASC, name DESC - """ - fields = order_by_str.split(",") - - order_clause = [] - for field in fields: - if field.startswith("-"): - order_clause.append(f"{field[1:]} DESC") - else: - order_clause.append(f"{field} ASC") - - return "ORDER BY " + ", ".join(order_clause) diff --git a/generalresearch/managers/gr/business.py b/generalresearch/managers/gr/business.py index ef26f30..9bf6ef2 100644 --- a/generalresearch/managers/gr/business.py +++ b/generalresearch/managers/gr/business.py @@ -14,14 +14,13 @@ from generalresearch.managers.base import ( from generalresearch.models.gr.business import ( Business, BusinessBankAccount, - BusinessType, ) +from generalresearch.models.gr.definitions import BusinessType, TransferMethod if TYPE_CHECKING: from generalresearch.models.custom_types import UUIDStr from generalresearch.models.gr.business import ( BusinessAddress, - TransferMethod, ) from generalresearch.models.gr.team import Team diff --git a/generalresearch/managers/pollfish/user_pid.py b/generalresearch/managers/pollfish/user_pid.py index 1068405..f3983cf 100644 --- a/generalresearch/managers/pollfish/user_pid.py +++ b/generalresearch/managers/pollfish/user_pid.py @@ -1,5 +1,5 @@ from generalresearch.managers.marketplace.user_pid import UserPidManager -from generalresearch.models import Source +from generalresearch.models.definitions import Source class PollfishUserPidManager(UserPidManager): diff --git a/generalresearch/managers/thl/cashout_method.py b/generalresearch/managers/thl/cashout_method.py index c12c920..ee86bec 100644 --- a/generalresearch/managers/thl/cashout_method.py +++ b/generalresearch/managers/thl/cashout_method.py @@ -9,15 +9,13 @@ from uuid import UUID, uuid4 from pydantic import NonNegativeInt from generalresearch.managers.base import PostgresManager -from generalresearch.models.thl.wallet.cashout_method import ( - CashoutMethod, -) from generalresearch.models.thl.wallet.definitions import PayoutType if TYPE_CHECKING: from generalresearch.models.thl.user import User from generalresearch.models.thl.wallet.cashout_method import ( CashMailCashoutMethodData, + CashoutMethod, PaypalCashoutMethodData, ) @@ -82,6 +80,7 @@ class CashoutMethodManager(PostgresManager): :return: the uuid of the created cashout method """ # todo: validate shipping address? + from generalresearch.models.thl.wallet.cashout_method import CashoutMethod cm = CashoutMethod( name="Cash in Mail", @@ -126,6 +125,8 @@ class CashoutMethodManager(PostgresManager): :param user: :return: the uuid of the created cashout method """ + from generalresearch.models.thl.wallet.cashout_method import CashoutMethod + cm = CashoutMethod( name="PayPal", description="Cashout via PayPal", @@ -290,6 +291,8 @@ class CashoutMethodManager(PostgresManager): # The data column here is inconsistent. Pulling keys from the mysql 'data' col # and putting them into the base level. Renamed so that we don't overwrite # a col called "data" within the "_data_" field. + from generalresearch.models.thl.wallet.cashout_method import CashoutMethod + for k in list(x["_data_"].keys()): if k in CashoutMethod.model_fields: x[k] = x["_data_"].pop(k) diff --git a/generalresearch/models/__init__.py b/generalresearch/models/__init__.py index c0348d7..e69de29 100644 --- a/generalresearch/models/__init__.py +++ b/generalresearch/models/__init__.py @@ -1,114 +0,0 @@ -from __future__ import annotations - -from enum import IntEnum, StrEnum - -from generalresearch.utils.enum import ReprEnumMeta - - -class Source(StrEnum, metaclass=ReprEnumMeta): - # The external marketplace, or the source of the survey / work. - # Max length of the value is 2. - GRS = "g" - CINT = "c" - DALIA = "a" # deprecated - DYNATA = "d" - ETX = "et" - FULL_CIRCLE = "f" - INNOVATE = "i" - LUCID = "l" - MORNING_CONSULT = "m" - OPEN_LABS = "n" - POLLFISH = "o" - PRECISION = "e" - PRODEGE_USER = "r" # deprecated - PRODEGE = "pr" # using 'r' for vendor_wall - PULLEY = "p" # deprecated - REPDATA = "rd" # using 'q' for vendor_wall - SAGO = "h" - SPECTRUM = "s" - TESTING = "t" # Used internally for testing - TESTING2 = "u" # Used internally for testing - WXET = "w" - - -class DebitKey(IntEnum, metaclass=ReprEnumMeta): - # The debit key for marketplaces - CINT = 8 - DALIA = 9 - DYNATA = 6 - # ETX = None - FULL_CIRCLE = 15 - INNOVATE = 7 - LUCID = 0 - MORNING_CONSULT = 12 - # OPEN_LABS = None - POLLFISH = 13 - PRECISION = 14 - PRODEGE = 11 - SAGO = 10 - SPECTRUM = 5 - # WXET = None - - -class DeviceType(IntEnum, metaclass=ReprEnumMeta): - UNKNOWN = 0 - MOBILE = 1 - DESKTOP = 2 - TABLET = 3 - - -class LogicalOperator(StrEnum, metaclass=ReprEnumMeta): - OR = "OR" - AND = "AND" - # There is currently no use case for NOT. See MarketplaceCondition.explain_not - NOT = "NOT" - - -class TaskStatus(StrEnum, metaclass=ReprEnumMeta): - # A survey is live if it is open and, given all conditions are met, is - # possible to send in traffic. All other statuses are just variants of - # NOT Live (not accepting traffic) - LIVE = "LIVE" - - # This is a generic NOT Live status. A marketplace may use other more - # specific statuses but in practice they don't matter because all we care - # about is if the task is LIVE. - NOT_LIVE = "NOT_LIVE" - - # We need a status to mark if a survey we thought was live does not come - # back from the API, we'll mark it as NOT_FOUND. - NOT_FOUND = "NOT_FOUND" - - -class TaskCalculationType(StrEnum): - COMPLETES = "COMPLETES" - STARTS = "STARTS" - - @classmethod - def from_api(cls, v: str) -> TaskCalculationType: - return { - "complete": cls.COMPLETES, - "completes": cls.COMPLETES, - "survey start": cls.STARTS, - "survey starts": cls.STARTS, - "start": cls.STARTS, - "prescreens": cls.STARTS, - "prescreen": cls.STARTS, - }[v.lower()] - - @classmethod - def prodege_from_api(cls, v: int) -> TaskCalculationType: - return {1: cls.COMPLETES, 2: cls.STARTS}[v] - - @classmethod - def innovate_from_api(cls, v: int) -> TaskCalculationType: - return {0: cls.COMPLETES, 1: cls.STARTS}[v] - - -class URLQueryKey(StrEnum, metaclass=ReprEnumMeta): - PRODUCT_ID = "39057c8b" - PRODUCT_USER_ID = "c184efc0" - SESSION_ID = "0bb50182" - - -MAX_INT32 = 2**31 diff --git a/generalresearch/models/gr/business.py b/generalresearch/models/gr/business.py index e11c54d..c6d3468 100644 --- a/generalresearch/models/gr/business.py +++ b/generalresearch/models/gr/business.py @@ -4,7 +4,6 @@ import json import logging import os from datetime import UTC, datetime -from enum import Enum, StrEnum from pathlib import Path from typing import TYPE_CHECKING from uuid import uuid4 @@ -29,11 +28,11 @@ from generalresearch.models.custom_types import ( UUIDStr, UUIDStrCoerce, ) +from generalresearch.models.gr.definitions import BusinessType, TransferMethod from generalresearch.models.gr.team import Team from generalresearch.models.thl.finance import BusinessBalances, POPFinancial from generalresearch.models.thl.ledger import OrderBy from generalresearch.utils.aggregation import group_by_year -from generalresearch.utils.enum import ReprEnumMeta if TYPE_CHECKING: from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge @@ -69,16 +68,6 @@ if TYPE_CHECKING: from generalresearch.models.thl.product import Product -class TransferMethod(Enum, metaclass=ReprEnumMeta): - ACH = 0 - WIRE = 1 - - -class BusinessType(StrEnum, metaclass=ReprEnumMeta): - INDIVIDUAL = "i" - COMPANY = "c" - - class BusinessBankAccount(BaseModel): model_config = ConfigDict( use_enum_values=True, diff --git a/generalresearch/models/gr/definitions.py b/generalresearch/models/gr/definitions.py new file mode 100644 index 0000000..2e06c03 --- /dev/null +++ b/generalresearch/models/gr/definitions.py @@ -0,0 +1,13 @@ +from enum import Enum, StrEnum + +from generalresearch.utils.enum import ReprEnumMeta + + +class TransferMethod(Enum, metaclass=ReprEnumMeta): + ACH = 0 + WIRE = 1 + + +class BusinessType(StrEnum, metaclass=ReprEnumMeta): + INDIVIDUAL = "i" + COMPANY = "c" diff --git a/generalresearch/models/thl/__init__.py b/generalresearch/models/thl/__init__.py index 7f2b8a9..45278f8 100644 --- a/generalresearch/models/thl/__init__.py +++ b/generalresearch/models/thl/__init__.py @@ -1,14 +1,12 @@ -from decimal import Decimal - # from generalresearch.models.thl.finance import ( # POPFinancial, # ProductBalances, # ) # from generalresearch.models.thl.payout import ( -# BrokerageProductPayoutEvent, +# # BrokerageProductPayoutEvent, # PayoutEvent, # ) -from generalresearch.models.thl.product import Product +# from generalresearch.models.thl.product import Product # _ = ( # Product, @@ -18,16 +16,6 @@ from generalresearch.models.thl.product import Product # POPFinancial, # ) -Product.model_rebuild() +# Product.model_rebuild() # PayoutEvent.model_rebuild() # BrokerageProductPayoutEvent.model_rebuild() - - -def decimal_to_int_cents(usd: Decimal | None) -> int | None: - return round(usd * 100) if usd is not None else None - - -def int_cents_to_decimal(value: int | None, decimals: int = 2) -> Decimal | None: - if value is None: - return None - return (Decimal(value) / Decimal(100)).quantize(Decimal(10) ** -decimals) diff --git a/generalresearch/models/thl/session.py b/generalresearch/models/thl/session.py index 404cff7..65b885e 100644 --- a/generalresearch/models/thl/session.py +++ b/generalresearch/models/thl/session.py @@ -19,10 +19,6 @@ from pydantic import ( ) from generalresearch.models.definitions import Source -from generalresearch.models.thl import ( - decimal_to_int_cents, - int_cents_to_decimal, -) from generalresearch.models.thl.definitions import ( WALL_ALLOWED_STATUS_CODE_1_2, WALL_ALLOWED_STATUS_STATUS_CODE, @@ -32,6 +28,10 @@ from generalresearch.models.thl.definitions import ( WallAdjustedStatus, WallStatusCode2, ) +from generalresearch.models.thl.utils import ( + decimal_to_int_cents, + int_cents_to_decimal, +) if TYPE_CHECKING: from generalresearch.managers.thl.ledger_manager.thl_ledger import ( diff --git a/generalresearch/models/thl/task_status.py b/generalresearch/models/thl/task_status.py index 817f4c5..6cff884 100644 --- a/generalresearch/models/thl/task_status.py +++ b/generalresearch/models/thl/task_status.py @@ -13,7 +13,6 @@ from pydantic import ( model_validator, ) -from generalresearch.models.thl import decimal_to_int_cents from generalresearch.models.thl.definitions import ( SessionAdjustedStatus, SessionStatusCode2, @@ -25,6 +24,7 @@ from generalresearch.models.thl.payout_format import ( PayoutFormatOptionalField, ) from generalresearch.models.thl.session import WallOut +from generalresearch.models.thl.utils import decimal_to_int_cents if TYPE_CHECKING: from generalresearch.models.custom_types import ( diff --git a/generalresearch/models/thl/utils.py b/generalresearch/models/thl/utils.py new file mode 100644 index 0000000..3e14065 --- /dev/null +++ b/generalresearch/models/thl/utils.py @@ -0,0 +1,11 @@ +from decimal import Decimal + + +def decimal_to_int_cents(usd: Decimal | None) -> int | None: + return round(usd * 100) if usd is not None else None + + +def int_cents_to_decimal(value: int | None, decimals: int = 2) -> Decimal | None: + if value is None: + return None + return (Decimal(value) / Decimal(100)).quantize(Decimal(10) ** -decimals) diff --git a/generalresearch/models/thl/wallet/__init__.py b/generalresearch/models/thl/wallet/__init__.py index 2d1eb8d..e69de29 100644 --- a/generalresearch/models/thl/wallet/__init__.py +++ b/generalresearch/models/thl/wallet/__init__.py @@ -1,87 +0,0 @@ -from enum import StrEnum - -from generalresearch.utils.enum import ReprEnumMeta - - -class PayoutType(StrEnum, metaclass=ReprEnumMeta): - """ - The method in which the requested payout is delivered. - """ - - # The max size of the db field that holds this value is 14, so please - # don't add new values longer than that! - - # User is paid out to their personal PayPal email address - PAYPAL = "PAYPAL" - # User is paid out via a Tango Gift Card - TANGO = "TANGO" - # DWOLLA - DWOLLA = "DWOLLA" - # A payment is made to a bank account using ACH - ACH = "ACH" - # A payment is made to a bank account using ACH - WIRE = "WIRE" - # A payment is made in cash and mailed to the user. - CASH_IN_MAIL = "CASH_IN_MAIL" - # A payment is made as a prize with some monetary value - PRIZE = "PRIZE" - - # This is used to designate either AMT_BONUS or AMT_HIT - AMT = "AMT" - # Amazon Mechanical Turk as a Bonus - AMT_BONUS = "AMT_BONUS" - # Amazon Mechanical Turk for a HIT - AMT_HIT = "AMT_ASSIGNMENT" - AMT_ASSIGNMENT = "AMT_ASSIGNMENT" - - -class Currency(StrEnum): - # United States Dollar - USD = "USD" - # Canadian Dollar - CAD = "CAD" - # British Pound Sterling - GBP = "GBP" - # Euro - EUR = "EUR" - # Indian Rupee - INR = "INR" - # Australian Dollar - AUD = "AUD" - # Polish Zloty - PLN = "PLN" - # Swedish Krona - SEK = "SEK" - # Singapore Dollar - SGD = "SGD" - # Mexican Peso - MXN = "MXN" - - -CURRENCY_FORMATTER = { - "USD": lambda x: f"${x / 100:,.2f}", - "CAD": lambda x: f"${x / 100:,.2f} CAD", - "GBP": lambda x: f"{x / 100:,.2f} £", - "EUR": lambda x: f"€{x / 100:,.2f}", - "INR": lambda x: f"₹{x / 100:,.2f}", - "AUD": lambda x: f"${x / 100:,.2f} AUD", - "PLN": lambda x: f"{x / 100:,.2f} zł", - "SEK": lambda x: f"{x / 100:,.2f} kr", - "SGD": lambda x: f"${x / 100:,.2f} SGD", - "MXN": lambda x: f"${x / 100:,.2f} MXN", -} - -# The max value user can redeem in one go in foreign currencies. should be < $250 -# in order to avoid exchange rate issues -CURRENCY_MAX_VALUE = { - "USD": 250, - "CAD": 200, - "GBP": 100, - "EUR": 100, - "INR": 10000, - "AUD": 200, - "PLN": 500, - "SEK": 1000, - "SGD": 200, - "MXN": 4000, -} diff --git a/test_utils/conftest.py b/test_utils/conftest.py index 397d98f..daf6b43 100644 --- a/test_utils/conftest.py +++ b/test_utils/conftest.py @@ -342,7 +342,7 @@ def delete_df_collection( thl_web_rw: PostgresConfig, create_main_accounts: Callable[..., None] ) -> Callable[..., None]: - from generalresearch.incite.collections import ( + from generalresearch.incite.collections.base import ( DFCollection, DFCollectionType, ) diff --git a/test_utils/incite/collections/conftest.py b/test_utils/incite/collections/conftest.py index f490e14..499f90b 100644 --- a/test_utils/incite/collections/conftest.py +++ b/test_utils/incite/collections/conftest.py @@ -197,7 +197,7 @@ def df_collection( utc_90days_ago: datetime, thl_web_rr: PostgresConfig, ) -> DFCollection: - from generalresearch.incite.collections import DFCollection + from generalresearch.incite.collections.base import DFCollection start = utc_90days_ago.replace(microsecond=0) diff --git a/test_utils/incite/conftest.py b/test_utils/incite/conftest.py index 2968d18..bcf0511 100644 --- a/test_utils/incite/conftest.py +++ b/test_utils/incite/conftest.py @@ -16,11 +16,11 @@ from faker import Faker if TYPE_CHECKING: from generalresearch.config import GRLBaseSettings from generalresearch.incite.base import GRLDatasets - from generalresearch.incite.collections import ( + from generalresearch.incite.collections.base import ( DFCollectionItem, DFCollectionType, ) - from generalresearch.incite.mergers import MergeType + from generalresearch.incite.mergers.base import MergeType from generalresearch.models.admin.request import ( ReportRequest, ) @@ -131,14 +131,14 @@ def duration() -> timedelta | None: @pytest.fixture def df_collection_data_type() -> DFCollectionType: - from generalresearch.incite.collections import DFCollectionType + from generalresearch.incite.collections.base import DFCollectionType return DFCollectionType.TEST @pytest.fixture def merge_type() -> MergeType: - from generalresearch.incite.mergers import MergeType + from generalresearch.incite.mergers.base import MergeType return MergeType.TEST @@ -156,7 +156,7 @@ def incite_item_factory( observations: int = 3, user: User | None = None, ): - from generalresearch.incite.collections import ( + from generalresearch.incite.collections.base import ( DFCollection, DFCollectionType, ) diff --git a/test_utils/incite/mergers/conftest.py b/test_utils/incite/mergers/conftest.py index 4eb3f2d..fb95c81 100644 --- a/test_utils/incite/mergers/conftest.py +++ b/test_utils/incite/mergers/conftest.py @@ -58,7 +58,7 @@ def pop_ledger_merge( duration: timedelta, ) -> PopLedgerMerge: - from generalresearch.incite.mergers import MergeType + from generalresearch.incite.mergers.base import MergeType from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge return PopLedgerMerge( @@ -88,7 +88,7 @@ def ym_survey_wall_merge( mnt_filepath: GRLDatasets, start: datetime, ) -> YMSurveyWallMerge: - from generalresearch.incite.mergers import MergeType + from generalresearch.incite.mergers.base import MergeType from generalresearch.incite.mergers.ym_survey_wall import YMSurveyWallMerge return YMSurveyWallMerge( @@ -119,7 +119,7 @@ def ym_wall_summary_merge( duration: timedelta, start: datetime, ) -> YMWallSummaryMerge: - from generalresearch.incite.mergers import MergeType + from generalresearch.incite.mergers.base import MergeType from generalresearch.incite.mergers.ym_wall_summary import YMWallSummaryMerge return YMWallSummaryMerge( @@ -155,7 +155,7 @@ def enriched_session_merge( duration: timedelta, start: datetime, ) -> EnrichedSessionMerge: - from generalresearch.incite.mergers import MergeType + from generalresearch.incite.mergers.base import MergeType from generalresearch.incite.mergers.foundations.enriched_session import ( EnrichedSessionMerge, ) @@ -175,7 +175,7 @@ def enriched_task_adjust_merge( duration: timedelta, start: datetime, ) -> EnrichedTaskAdjustMerge: - from generalresearch.incite.mergers import MergeType + from generalresearch.incite.mergers.base import MergeType from generalresearch.incite.mergers.foundations.enriched_task_adjust import ( EnrichedTaskAdjustMerge, ) @@ -197,7 +197,7 @@ def enriched_wall_merge( duration: timedelta, start: datetime, ) -> EnrichedWallMerge: - from generalresearch.incite.mergers import MergeType + from generalresearch.incite.mergers.base import MergeType from generalresearch.incite.mergers.foundations.enriched_wall import ( EnrichedWallMerge, ) @@ -217,7 +217,7 @@ def user_id_product_merge( offset: str, start: datetime, ) -> UserIdProductMerge: - from generalresearch.incite.mergers import MergeType + from generalresearch.incite.mergers.base import MergeType from generalresearch.incite.mergers.foundations.user_id_product import ( UserIdProductMerge, ) @@ -243,7 +243,7 @@ def merge_collection( duration: timedelta, start: datetime, ): - from generalresearch.incite.mergers import MergeCollection + from generalresearch.incite.mergers.base import MergeCollection return MergeCollection( merge_type=merge_type, diff --git a/test_utils/managers/gr/conftest.py b/test_utils/managers/gr/conftest.py index 5392c69..a7fa9e9 100644 --- a/test_utils/managers/gr/conftest.py +++ b/test_utils/managers/gr/conftest.py @@ -9,7 +9,6 @@ import pytest import redis import redis.asyncio as redis_async from pydantic import PostgresDsn -from redis import Redis from generalresearch.managers.gr.business import ( BusinessAddressManager, @@ -30,33 +29,6 @@ def gr_redis_config_db() -> str: return str(randint(99, 1_023)) -@pytest.fixture(scope="session") -def gr_redis(settings: GRLBaseSettings) -> Redis: - assert "unittest" in str(settings.testing_redis) or "127.0.0.1" in str( - settings.testing_redis - ) - return Redis.from_url( - url=str(settings.gr_redis), - decode_responses=True, - socket_timeout=settings.redis_timeout, - socket_connect_timeout=settings.redis_timeout, - ) - - -@pytest.fixture -def gr_redis_async(settings: GRLBaseSettings) -> redis_async.Redis: - assert "unittest" in str(settings.testing_redis) or "127.0.0.1" in str( - settings.testing_redis - ) - - return redis_async.Redis.from_url( - str(settings.testing_redis), - decode_responses=True, - socket_timeout=0.20, - socket_connect_timeout=0.20, - ) - - @pytest.fixture(scope="session") def gr_redis_config( settings: GRLBaseSettings, gr_redis_config_db: str diff --git a/test_utils/managers/thl/conftest.py b/test_utils/managers/thl/conftest.py index af3fd23..391b74c 100644 --- a/test_utils/managers/thl/conftest.py +++ b/test_utils/managers/thl/conftest.py @@ -1,9 +1,12 @@ from __future__ import annotations -from collections.abc import Callable +import subprocess +from collections.abc import Callable, Generator +from random import randint from typing import TYPE_CHECKING import pytest +import redis from pydantic import PostgresDsn from generalresearch.managers.base import Permission @@ -59,14 +62,40 @@ def thl_web_rw(thl_web_rr: PostgresConfig) -> PostgresConfig: @pytest.fixture(scope="session") -def thl_redis_config(settings: GRLBaseSettings) -> RedisConfig: - return RedisConfig( - dsn=settings.thl_redis, +def thl_redis_config_db() -> str: + return str(randint(99, 1_023)) + + +@pytest.fixture(scope="session") +def thl_redis_config( + settings: GRLBaseSettings, thl_redis_config_db: str +) -> Generator[RedisConfig]: + assert "unittest" in str(settings.testing_redis) or "127.0.0.1" in str( + settings.testing_redis + ) + + uri = f"redis://{settings.testing_redis}/{thl_redis_config_db}" + + res = subprocess.run( + ["redis-cli", "-u", uri, "SET", "jenkins_lock", "1", "NX", "EX", "3600"], + check=True, + text=True, + capture_output=True, + ) + + if res.stdout.strip() != "OK": + raise ValueError("Redis already locked... aborting.") + + yield RedisConfig( + dsn=uri, decode_responses=True, socket_timeout=settings.redis_timeout, socket_connect_timeout=settings.redis_timeout, ) + r = redis.from_url(uri) + r.flushdb() + @pytest.fixture(scope="session") def payout_event_manager( diff --git a/test_utils/models/conftest.py b/test_utils/models/conftest.py index 089f2e6..ed4da08 100644 --- a/test_utils/models/conftest.py +++ b/test_utils/models/conftest.py @@ -370,7 +370,7 @@ def product_amt_true( @pytest.fixture def bp_payout_factory( - thl_lm: ThlLedgerManager, + thl_ledger_manager: ThlLedgerManager, product_manager: ProductManager, business_payout_event_manager: BusinessPayoutEventManager, ) -> Callable[..., BrokerageProductPayoutEvent]: @@ -389,7 +389,7 @@ def bp_payout_factory( amount = amount or USDCent(randint(1, 99_99)) return business_payout_event_manager.create_bp_payout_event( - thl_ledger_manager=thl_lm, + thl_ledger_manager=thl_ledger_manager, product=product, amount=amount, ext_ref_id=ext_ref_id or uuid4().hex, diff --git a/test_utils/models/contest/conftest.py b/test_utils/models/contest/conftest.py index 91425dc..18a8e5f 100644 --- a/test_utils/models/contest/conftest.py +++ b/test_utils/models/contest/conftest.py @@ -275,24 +275,26 @@ def user_with_money( request: Request, user_factory: Callable[..., User], product_user_wallet_yes: Product, - thl_lm: ThlLedgerManager, + thl_ledger_manager: ThlLedgerManager, ) -> User: params = getattr(request, "param", {}) or {} min_balance = int(params.get("min_balance", USDCent(1_00))) user: User = user_factory(product=product_user_wallet_yes) - wallet = thl_lm.get_account_or_create_user_wallet(user) - balance = thl_lm.get_account_balance(wallet) + wallet = thl_ledger_manager.get_account_or_create_user_wallet(user) + balance = thl_ledger_manager.get_account_balance(wallet) todo = min_balance - balance if todo > 0: # # Put money in user's wallet - thl_lm.create_tx_user_bonus( + thl_ledger_manager.create_tx_user_bonus( user=user, ref_uuid=uuid4().hex, description="bonus", amount=Decimal(todo) / 100, ) - print(f"wallet balance: {thl_lm.get_user_wallet_balance(user=user)}") + print( + f"wallet balance: {thl_ledger_manager.get_user_wallet_balance(user=user)}" + ) return user diff --git a/test_utils/models/gr/conftest.py b/test_utils/models/gr/conftest.py index 6c1877a..e493f20 100644 --- a/test_utils/models/gr/conftest.py +++ b/test_utils/models/gr/conftest.py @@ -23,8 +23,8 @@ if TYPE_CHECKING: Business, BusinessAddress, BusinessBankAccount, - TransferMethod, ) + from generalresearch.models.gr.definitions import TransferMethod from generalresearch.models.gr.team import Membership, Team from generalresearch.pg_helper import PostgresConfig from generalresearch.redis_helper import RedisConfig diff --git a/test_utils/models/ledger/conftest.py b/test_utils/models/ledger/conftest.py index 8437c7f..31e5eb4 100644 --- a/test_utils/models/ledger/conftest.py +++ b/test_utils/models/ledger/conftest.py @@ -65,7 +65,7 @@ if TYPE_CHECKING: @pytest.fixture def ledger_account( - request: Request, lm: LedgerManager, currency: LedgerCurrency + request: Request, ledger_manager: LedgerManager, currency: LedgerCurrency ) -> LedgerAccount: from generalresearch.models.thl.ledger import ( AccountType, @@ -87,14 +87,14 @@ def ledger_account( account_type=account_type, normal_balance=direction, ) - return lm.create_account(account=acct_model) + return ledger_manager.create_account(account=acct_model) @pytest.fixture def ledger_account_factory( request: Request, - thl_lm: ThlLedgerManager, - lm: LedgerManager, + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, currency: LedgerCurrency, ) -> Callable[..., LedgerAccount]: @@ -109,7 +109,7 @@ def ledger_account_factory( account_type: AccountType = AccountType.CASH, direction: Direction = Direction.CREDIT, ) -> LedgerAccount: - thl_lm.get_account_or_create_bp_wallet(product=product) + thl_ledger_manager.get_account_or_create_bp_wallet(product=product) acct_uuid = uuid4().hex qn = f"{currency}:{account_type}:{acct_uuid}" @@ -121,14 +121,14 @@ def ledger_account_factory( account_type=account_type, normal_balance=direction, ) - return lm.create_account(account=acct_model) + return ledger_manager.create_account(account=acct_model) return _inner @pytest.fixture def ledger_account_credit( - request: Request, lm: LedgerManager, currency: LedgerCurrency + request: Request, ledger_manager: LedgerManager, currency: LedgerCurrency ) -> LedgerAccount: from generalresearch.models.thl.ledger import AccountType, Direction @@ -146,12 +146,12 @@ def ledger_account_credit( account_type=account_type, normal_balance=Direction.CREDIT, ) - return lm.create_account(account=acct_model) + return ledger_manager.create_account(account=acct_model) @pytest.fixture def ledger_account_debit( - request: Request, lm: LedgerManager, currency: LedgerCurrency + request: Request, ledger_manager: LedgerManager, currency: LedgerCurrency ) -> LedgerAccount: from generalresearch.models.thl.ledger import AccountType, Direction @@ -169,11 +169,11 @@ def ledger_account_debit( account_type=account_type, normal_balance=Direction.DEBIT, ) - return lm.create_account(account=acct_model) + return ledger_manager.create_account(account=acct_model) @pytest.fixture -def tag(request: Request, lm: LedgerManager) -> str: +def tag(request: Request) -> str: from generalresearch.currency import LedgerCurrency return ( @@ -194,11 +194,11 @@ def bp_payout_event( product: Product, usd_cent: USDCent, business_payout_event_manager: BusinessPayoutEventManager, - thl_lm: ThlLedgerManager, + thl_ledger_manager: ThlLedgerManager, ) -> BrokerageProductPayoutEvent: return business_payout_event_manager.create_bp_payout_event( - thl_ledger_manager=thl_lm, + thl_ledger_manager=thl_ledger_manager, product=product, amount=usd_cent, skip_wallet_balance_check=True, @@ -209,7 +209,7 @@ def bp_payout_event( @pytest.fixture def bp_payout_event_factory( brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager, - thl_lm: ThlLedgerManager, + thl_ledger_manager: ThlLedgerManager, ) -> Callable[..., BrokerageProductPayoutEvent]: def _inner( @@ -217,7 +217,7 @@ def bp_payout_event_factory( ) -> BrokerageProductPayoutEvent: return brokerage_product_payout_event_manager.create_bp_payout_event( - thl_ledger_manager=thl_lm, + thl_ledger_manager=thl_ledger_manager, product=product, amount=usd_cent, ext_ref_id=ext_ref_id, @@ -229,10 +229,12 @@ def bp_payout_event_factory( @pytest.fixture -def currency(lm: LedgerManager) -> LedgerCurrency: +def currency(ledger_manager: LedgerManager) -> LedgerCurrency: # return request.param if hasattr(request, "currency") else LedgerCurrency.TEST - assert lm.currency, "LedgerManager must have a currency specified for these tests" - return lm.currency + assert ( + ledger_manager.currency + ), "LedgerManager must have a currency specified for these tests" + return ledger_manager.currency @pytest.fixture @@ -252,7 +254,7 @@ def ledger_tx( tag: str, currency: LedgerCurrency, tx_metadata: dict[str, str] | None, - lm: LedgerManager, + ledger_manager: LedgerManager, ) -> LedgerTransaction: from generalresearch.models.thl.ledger import Direction, LedgerEntry @@ -271,12 +273,12 @@ def ledger_tx( ), ] - return lm.create_tx(entries=entries, tag=tag, metadata=tx_metadata) + return ledger_manager.create_tx(entries=entries, tag=tag, metadata=tx_metadata) @pytest.fixture def create_main_accounts( - lm: LedgerManager, currency: LedgerCurrency + ledger_manager: LedgerManager, currency: LedgerCurrency ) -> Callable[..., None]: def _inner() -> None: @@ -291,9 +293,9 @@ def create_main_accounts( qualified_name=f"{currency.value}:revenue:task_complete", normal_balance=Direction.CREDIT, account_type=AccountType.REVENUE, - currency=lm.currency, + currency=ledger_manager.currency, ) - lm.get_account_or_create(account=account) + ledger_manager.get_account_or_create(account=account) account = LedgerAccount( display_name="Operating Cash Account", @@ -303,7 +305,7 @@ def create_main_accounts( currency=currency, ) - lm.get_account_or_create(account=account) + ledger_manager.get_account_or_create(account=account) return _inner @@ -327,7 +329,7 @@ def delete_ledger_db(thl_web_rw: PostgresManager) -> Callable[..., None]: @pytest.fixture def wipe_main_accounts( - thl_web_rw: PostgresManager, lm: LedgerManager, currency: LedgerCurrency + thl_web_rw: PostgresManager, ledger_manager: LedgerManager, currency: LedgerCurrency ) -> Callable[..., None]: def _inner() -> None: @@ -397,7 +399,9 @@ def wipe_main_accounts( @pytest.fixture -def account_cash(lm: LedgerManager, currency: LedgerCurrency) -> LedgerAccount: +def account_cash( + ledger_manager: LedgerManager, currency: LedgerCurrency +) -> LedgerAccount: from generalresearch.models.thl.ledger import ( AccountType, Direction, @@ -411,12 +415,12 @@ def account_cash(lm: LedgerManager, currency: LedgerCurrency) -> LedgerAccount: account_type=AccountType.CASH, currency=currency, ) - return lm.get_account_or_create(account=account) + return ledger_manager.get_account_or_create(account=account) @pytest.fixture def account_revenue_task_complete( - lm: LedgerManager, currency: LedgerCurrency + ledger_manager: LedgerManager, currency: LedgerCurrency ) -> LedgerAccount: from generalresearch.models.thl.ledger import ( AccountType, @@ -431,11 +435,13 @@ def account_revenue_task_complete( account_type=AccountType.REVENUE, currency=currency, ) - return lm.get_account_or_create(account=account) + return ledger_manager.get_account_or_create(account=account) @pytest.fixture -def account_expense_tango(lm: LedgerManager, currency: LedgerCurrency) -> LedgerAccount: +def account_expense_tango( + ledger_manager: LedgerManager, currency: LedgerCurrency +) -> LedgerAccount: from generalresearch.models.thl.ledger import ( AccountType, Direction, @@ -449,12 +455,12 @@ def account_expense_tango(lm: LedgerManager, currency: LedgerCurrency) -> Ledger account_type=AccountType.EXPENSE, currency=currency, ) - return lm.get_account_or_create(account=account) + return ledger_manager.get_account_or_create(account=account) @pytest.fixture def user_account_user_wallet( - lm: LedgerManager, user: User, currency: LedgerCurrency + ledger_manager: LedgerManager, user: User, currency: LedgerCurrency ) -> LedgerAccount: from generalresearch.models.thl.ledger import ( AccountType, @@ -471,12 +477,12 @@ def user_account_user_wallet( reference_uuid=user.uuid, currency=currency, ) - return lm.get_account_or_create(account=account) + return ledger_manager.get_account_or_create(account=account) @pytest.fixture def product_account_bp_wallet( - lm: LedgerManager, product: Product, currency: LedgerCurrency + ledger_manager: LedgerManager, product: Product, currency: LedgerCurrency ) -> LedgerAccount: from generalresearch.models.thl.ledger import ( AccountType, @@ -495,13 +501,13 @@ def product_account_bp_wallet( "currency": currency, } ) - return lm.get_account_or_create(account=account) + return ledger_manager.get_account_or_create(account=account) @pytest.fixture def setup_accounts( product_factory: Callable[..., Product], - lm: LedgerManager, + ledger_manager: LedgerManager, user: User, currency: LedgerCurrency, ) -> Callable[..., None]: @@ -524,7 +530,7 @@ def setup_accounts( reference_uuid=p1.uuid, currency=currency, ) - lm.get_account_or_create(account=account) + ledger_manager.get_account_or_create(account=account) account = LedgerAccount.model_validate( { @@ -537,7 +543,7 @@ def setup_accounts( "currency": currency, } ) - lm.get_account_or_create(account=account) + ledger_manager.get_account_or_create(account=account) # BP's wallet, user's wallet, and a revenue from their commissions account. p2 = product_factory() @@ -550,7 +556,7 @@ def setup_accounts( reference_uuid=p2.uuid, currency=currency, ) - lm.get_account_or_create(account) + ledger_manager.get_account_or_create(account) account = LedgerAccount( display_name=f"{p2.name} Wallet", @@ -561,7 +567,7 @@ def setup_accounts( reference_uuid=p2.uuid, currency=currency, ) - lm.get_account_or_create(account) + ledger_manager.get_account_or_create(account) account = LedgerAccount( display_name=f"{user.uuid} Wallet", @@ -572,7 +578,7 @@ def setup_accounts( reference_uuid=user.uuid, currency="test", ) - lm.get_account_or_create(account=account) + ledger_manager.get_account_or_create(account=account) return _inner @@ -583,7 +589,7 @@ def session_with_tx_factory( session_manager: SessionManager, wall_manager: WallManager, utc_hour_ago: datetime, - thl_lm: ThlLedgerManager, + thl_ledger_manager: ThlLedgerManager, ) -> Callable[..., Session]: from generalresearch.models.thl.session import ( @@ -624,14 +630,16 @@ def session_with_tx_factory( status_code_1=status_code_1, ) - thl_lm.create_tx_task_complete( + thl_ledger_manager.create_tx_task_complete( wall=last_wall, user=user, created=last_wall.finished, force=True, ) - thl_lm.create_tx_bp_payment(session=s, created=last_wall.finished, force=True) + thl_ledger_manager.create_tx_bp_payment( + session=s, created=last_wall.finished, force=True + ) return s @@ -642,7 +650,7 @@ def session_with_tx_factory( def adj_to_fail_with_tx_factory( session_manager: SessionManager, wall_manager: WallManager, - thl_lm: ThlLedgerManager, + thl_ledger_manager: ThlLedgerManager, ) -> Callable[..., None]: from datetime import timedelta @@ -675,7 +683,7 @@ def adj_to_fail_with_tx_factory( adjusted_timestamp=created, ) - thl_lm.create_tx_task_adjustment( + thl_ledger_manager.create_tx_task_adjustment( wall=w1, user=session.user, created=created + timedelta(milliseconds=1), @@ -684,7 +692,7 @@ def adj_to_fail_with_tx_factory( session.wall_events = wall_manager.get_wall_events(session_id=session.id) session_manager.adjust_status(session=session) - thl_lm.create_tx_bp_adjustment( + thl_ledger_manager.create_tx_bp_adjustment( session=session, created=created + timedelta(milliseconds=2) ) @@ -695,7 +703,7 @@ def adj_to_fail_with_tx_factory( def adj_to_complete_with_tx_factory( session_manager: SessionManager, wall_manager: WallManager, - thl_lm: ThlLedgerManager, + thl_ledger_manager: ThlLedgerManager, ) -> Callable[..., None]: from datetime import timedelta @@ -714,7 +722,7 @@ def adj_to_complete_with_tx_factory( adjusted_timestamp=created, ) - thl_lm.create_tx_task_adjustment( + thl_ledger_manager.create_tx_task_adjustment( wall=w1, user=session.user, created=created + timedelta(milliseconds=1), @@ -723,7 +731,7 @@ def adj_to_complete_with_tx_factory( session.wall_events = wall_manager.get_wall_events(session_id=session.id) session_manager.adjust_status(session=session) - thl_lm.create_tx_bp_adjustment( + thl_ledger_manager.create_tx_bp_adjustment( session=session, created=created + timedelta(milliseconds=2) ) diff --git a/tests/incite/collections/test_df_collection_base.py b/tests/incite/collections/test_df_collection_base.py index e20b44b..6d715fa 100644 --- a/tests/incite/collections/test_df_collection_base.py +++ b/tests/incite/collections/test_df_collection_base.py @@ -5,7 +5,7 @@ import pandas as pd import pytest from pandera.pandas import DataFrameSchema -from generalresearch.incite.collections import ( +from generalresearch.incite.collections.base import ( DFCollection, DFCollectionType, ) @@ -53,7 +53,7 @@ class TestDFCollectionBaseProperties: data_type=df_coll_type, start=datetime(year=1800, month=1, day=1, tzinfo=UTC), finished=datetime(year=1900, month=1, day=1, tzinfo=UTC), - offset="100d", + offset="100D", archive_path=mnt_filepath.archive_path(enum_type=df_coll_type), ) @@ -67,7 +67,7 @@ class TestDFCollectionBaseProperties: data_type=df_coll_type, start=datetime(year=1800, month=1, day=1, tzinfo=UTC), finished=datetime(year=1900, month=1, day=1, tzinfo=UTC), - offset="100d", + offset="100D", archive_path=mnt_filepath.archive_path(enum_type=df_coll_type), ) diff --git a/tests/incite/collections/test_df_collection_item_base.py b/tests/incite/collections/test_df_collection_item_base.py index fd70bf0..83d4973 100644 --- a/tests/incite/collections/test_df_collection_item_base.py +++ b/tests/incite/collections/test_df_collection_item_base.py @@ -25,7 +25,7 @@ class TestDFCollectionItemBase: def test_init(self, mnt_filepath: GRLDatasets, df_coll_type: DFCollectionType): collection = DFCollection( data_type=df_coll_type, - offset="100d", + offset="100D", start=datetime(year=1800, month=1, day=1, tzinfo=UTC), finished=datetime(year=1900, month=1, day=1, tzinfo=UTC), archive_path=mnt_filepath.archive_path(enum_type=df_coll_type), @@ -53,7 +53,7 @@ class TestDFCollectionItemMethods: ): collection = DFCollection( data_type=df_coll_type, - offset="100d", + offset="100D", start=datetime(year=1800, month=1, day=1, tzinfo=UTC), finished=datetime(year=1900, month=1, day=1, tzinfo=UTC), archive_path=mnt_filepath.archive_path(enum_type=df_coll_type), @@ -70,7 +70,7 @@ class TestDFCollectionItemMethods: ): collection = DFCollection( data_type=df_coll_type, - offset="100d", + offset="100D", start=datetime(year=1800, month=1, day=1, tzinfo=UTC), finished=datetime(year=1900, month=1, day=1, tzinfo=UTC), archive_path=mnt_filepath.archive_path(enum_type=df_coll_type), diff --git a/tests/incite/test_interval_idx.py b/tests/incite/test_interval_idx.py index 03d29ea..04d0bb2 100644 --- a/tests/incite/test_interval_idx.py +++ b/tests/incite/test_interval_idx.py @@ -18,7 +18,7 @@ class TestIntervalIndex: # If the offset is longer than the end - start it will not # error. It will simply have 0 rows. iv_r: pd.IntervalIndex = pd.interval_range( - start=start, end=end, freq="30d", closed="left" + start=start, end=end, freq="30D", closed="left" ) assert isinstance(iv_r, pd.IntervalIndex) assert len(iv_r.to_list()) == 0 diff --git a/tests/managers/gr/test_business.py b/tests/managers/gr/test_business.py index 1a5d4fa..35c471e 100644 --- a/tests/managers/gr/test_business.py +++ b/tests/managers/gr/test_business.py @@ -7,8 +7,8 @@ from generalresearch.models.gr.business import ( Business, BusinessAddress, BusinessBankAccount, - TransferMethod, ) +from generalresearch.models.gr.definitions import TransferMethod if TYPE_CHECKING: from generalresearch.managers.gr.business import ( @@ -32,12 +32,12 @@ class TestBusinessBankAccountManager: def test_create( self, - business: Business, + gr_business: Business, business_bank_account_manager: BusinessBankAccountManager, ): instance = business_bank_account_manager.create( - business_id=business.id, + business_id=gr_business.id, uuid=uuid4().hex, transfer_method=TransferMethod.ACH, ) @@ -56,10 +56,12 @@ class TestBusinessBankAccountManager: class TestBusinessAddressManager: def test_create( - self, business: Business, business_address_manager: BusinessAddressManager + self, gr_business: Business, business_address_manager: BusinessAddressManager ): - res = business_address_manager.create(uuid=uuid4().hex, business_id=business.id) + res = business_address_manager.create( + uuid=uuid4().hex, business_id=gr_business.id + ) assert isinstance(res, BusinessAddress) assert isinstance(res.id, int) @@ -140,18 +142,20 @@ class TestBusinessManager: def test_get_uuids_by_user_id(self): pass - def test_get_by_uuid(self, business: Business, business_manager: BusinessManager): - instance = business_manager.get_by_uuid(business_uuid=business.uuid) + def test_get_by_uuid( + self, gr_business: Business, business_manager: BusinessManager + ): + instance = business_manager.get_by_uuid(business_uuid=gr_business.uuid) assert isinstance(instance, Business) - assert business.id == instance.id + assert gr_business.id == instance.id - def test_get_by_id(self, business: Business, business_manager: BusinessManager): - instance = business_manager.get_by_id(business_id=business.id) + def test_get_by_id(self, gr_business: Business, business_manager: BusinessManager): + instance = business_manager.get_by_id(business_id=gr_business.id) assert isinstance(instance, Business) - assert business.uuid == instance.uuid + assert gr_business.uuid == instance.uuid - def test_cache_key(self, business: Business): - assert "business:" in business.cache_key + def test_cache_key(self, gr_business: Business): + assert "business:" in gr_business.cache_key # def test_create_raise_on_duplicate(self): # b_uuid = uuid4().hex @@ -160,7 +164,7 @@ class TestBusinessManager: # business = BusinessManager.create( # uuid=b_uuid, # name=f"test-{b_uuid[:6]}") - # assert isinstance(business: Business, Business) + # assert isinstance(gr_business: Business, Business) # # # Try to make it again # with pytest.raises(expected_exception=psycopg.errors.UniqueViolation): diff --git a/tests/managers/thl/test_ledger/test_lm_accounts.py b/tests/managers/thl/test_ledger/test_lm_accounts.py index f5ed883..cdef99a 100644 --- a/tests/managers/thl/test_ledger/test_lm_accounts.py +++ b/tests/managers/thl/test_ledger/test_lm_accounts.py @@ -44,7 +44,7 @@ class TestLedgerAccountManagerNoResults: currency: LedgerCurrency, kind: str, acct_id: UUIDStr, - lm: LedgerManager, + ledger_manager: LedgerManager, ): """Try to query for accounts that we know don't exist and confirm that we either get the expected None result or it raises the correct @@ -54,40 +54,50 @@ class TestLedgerAccountManagerNoResults: # (1) .get_account is just a wrapper for .get_account_many_ but # call it either way - assert lm.get_account(qualified_name=qn, raise_on_error=False) is None + assert ( + ledger_manager.get_account(qualified_name=qn, raise_on_error=False) is None + ) with pytest.raises(expected_exception=LedgerAccountDoesntExistError): - lm.get_account(qualified_name=qn, raise_on_error=True) + ledger_manager.get_account(qualified_name=qn, raise_on_error=True) # (2) .get_account_if_exists is another wrapper - assert lm.get_account(qualified_name=qn, raise_on_error=False) is None + assert ( + ledger_manager.get_account(qualified_name=qn, raise_on_error=False) is None + ) def test_get_account_no_results_many( self, currency: LedgerCurrency, kind: str, acct_id: UUIDStr, - lm: LedgerManager, + ledger_manager: LedgerManager, ): qn = f"{currency}:{kind}:{acct_id}" # (1) .get_many_ - assert lm.get_account_many_(qualified_names=[qn], raise_on_error=False) == [] + assert ( + ledger_manager.get_account_many_(qualified_names=[qn], raise_on_error=False) + == [] + ) with pytest.raises(expected_exception=LedgerAccountDoesntExistError): - lm.get_account_many_(qualified_names=[qn], raise_on_error=True) + ledger_manager.get_account_many_(qualified_names=[qn], raise_on_error=True) # (2) .get_many - assert lm.get_account_many(qualified_names=[qn], raise_on_error=False) == [] + assert ( + ledger_manager.get_account_many(qualified_names=[qn], raise_on_error=False) + == [] + ) with pytest.raises(expected_exception=LedgerAccountDoesntExistError): - lm.get_account_many(qualified_names=[qn], raise_on_error=True) + ledger_manager.get_account_many(qualified_names=[qn], raise_on_error=True) # (3) .get_accounts(..) - assert lm.get_accounts_if_exists(qualified_names=[qn]) == [] + assert ledger_manager.get_accounts_if_exists(qualified_names=[qn]) == [] with pytest.raises(expected_exception=LedgerAccountDoesntExistError): - lm.get_accounts(qualified_names=[qn]) + ledger_manager.get_accounts(qualified_names=[qn]) @pytest.mark.parametrize( @@ -107,7 +117,7 @@ class TestLedgerAccountManagerCreate: currency: LedgerCurrency, account_type: AccountType, direction: Direction, - lm: LedgerManager, + ledger_manager: LedgerManager, ): """Confirm that the Permission values that are set on the Ledger Manger allow the Creation action to occur. @@ -124,11 +134,11 @@ class TestLedgerAccountManagerCreate: # (1) With no Permissions defined test_lm = LedgerManager( - pg_config=lm.pg_config, + pg_config=ledger_manager.pg_config, permissions=[], - redis_config=lm.redis_config, - cache_prefix=lm.cache_prefix, - testing=lm.testing, + redis_config=ledger_manager.redis_config, + cache_prefix=ledger_manager.cache_prefix, + testing=ledger_manager.testing, ) with pytest.raises(expected_exception=AssertionError) as excinfo: @@ -139,11 +149,11 @@ class TestLedgerAccountManagerCreate: # (2) With Permissions defined, but not CREATE test_lm = LedgerManager( - pg_config=lm.pg_config, + pg_config=ledger_manager.pg_config, permissions=[Permission.READ, Permission.UPDATE, Permission.DELETE], - redis_config=lm.redis_config, - cache_prefix=lm.cache_prefix, - testing=lm.testing, + redis_config=ledger_manager.redis_config, + cache_prefix=ledger_manager.cache_prefix, + testing=ledger_manager.testing, ) with pytest.raises(expected_exception=AssertionError) as excinfo: @@ -157,7 +167,7 @@ class TestLedgerAccountManagerCreate: currency: LedgerCurrency, account_type: AccountType, direction: Direction, - lm: LedgerManager, + ledger_manager: LedgerManager, ): """Confirm that the Permission values that are set on the Ledger Manger allow the Creation action to occur. @@ -174,11 +184,11 @@ class TestLedgerAccountManagerCreate: account_type=account_type, normal_balance=direction, ) - account = lm.create_account(account=acct_model) + account = ledger_manager.create_account(account=acct_model) assert isinstance(account, LedgerAccount) # Query for, and make sure the Account was saved in the DB - res = lm.get_account(qualified_name=qn, raise_on_error=True) + res = ledger_manager.get_account(qualified_name=qn, raise_on_error=True) assert res is not None assert account.uuid == res.uuid @@ -187,7 +197,7 @@ class TestLedgerAccountManagerCreate: currency: LedgerCurrency, account_type: AccountType, direction: Direction, - lm: LedgerManager, + ledger_manager: LedgerManager, ): """Confirm that the Permission values that are set on the Ledger Manger allow the Creation action to occur. @@ -204,27 +214,31 @@ class TestLedgerAccountManagerCreate: account_type=account_type, normal_balance=direction, ) - account = lm.get_account_or_create(account=acct_model) + account = ledger_manager.get_account_or_create(account=acct_model) assert isinstance(account, LedgerAccount) # Query for, and make sure the Account was saved in the DB - res = lm.get_account(qualified_name=qn, raise_on_error=True) + res = ledger_manager.get_account(qualified_name=qn, raise_on_error=True) assert res is not None assert account.uuid == res.uuid class TestLedgerAccountManagerGet: - def test_get(self, ledger_account: LedgerAccount, lm: LedgerManager): - res = lm.get_account(qualified_name=ledger_account.qualified_name) + def test_get(self, ledger_account: LedgerAccount, ledger_manager: LedgerManager): + res = ledger_manager.get_account(qualified_name=ledger_account.qualified_name) assert res is not None assert res.uuid == ledger_account.uuid - res = lm.get_account_many(qualified_names=[ledger_account.qualified_name]) + res = ledger_manager.get_account_many( + qualified_names=[ledger_account.qualified_name] + ) assert len(res) == 1 assert res[0].uuid == ledger_account.uuid - res = lm.get_accounts(qualified_names=[ledger_account.qualified_name]) + res = ledger_manager.get_accounts( + qualified_names=[ledger_account.qualified_name] + ) assert len(res) == 1 assert res[0].uuid == ledger_account.uuid @@ -237,15 +251,15 @@ class TestLedgerAccountManagerGet: ledger_account_credit: LedgerAccount, ledger_account_debit: LedgerAccount, ledger_tx: LedgerTransaction, - lm: LedgerManager, + ledger_manager: LedgerManager, ): - res = lm.get_account_balance(account=ledger_account) + res = ledger_manager.get_account_balance(account=ledger_account) assert res == 0 - res = lm.get_account_balance(account=ledger_account_credit) + res = ledger_manager.get_account_balance(account=ledger_account_credit) assert res == 100 - res = lm.get_account_balance(account=ledger_account_debit) + res = ledger_manager.get_account_balance(account=ledger_account_debit) assert res == 100 @pytest.mark.parametrize("n_times", range(5)) @@ -256,7 +270,7 @@ class TestLedgerAccountManagerGet: ledger_account_debit: LedgerAccount, ledger_tx: LedgerTransaction, n_times: PositiveInt, - lm: LedgerManager, + ledger_manager: LedgerManager, ): """Try searching for random metadata and confirm it's always 0 because Tx can be found. @@ -265,7 +279,7 @@ class TestLedgerAccountManagerGet: rand_value = uuid4().hex assert ( - lm.get_account_filtered_balance( + ledger_manager.get_account_filtered_balance( account=ledger_account, metadata_key=rand_key, metadata_value=rand_value ) == 0 @@ -275,7 +289,7 @@ class TestLedgerAccountManagerGet: # and that we can filter it back rand_amount = randint(10, 1_000) - lm.create_tx( + ledger_manager.create_tx( entries=[ LedgerEntry( direction=Direction.CREDIT, @@ -292,7 +306,7 @@ class TestLedgerAccountManagerGet: ) assert ( - lm.get_account_filtered_balance( + ledger_manager.get_account_filtered_balance( account=ledger_account_credit, metadata_key=rand_key, metadata_value=rand_value, @@ -301,7 +315,7 @@ class TestLedgerAccountManagerGet: ) assert ( - lm.get_account_filtered_balance( + ledger_manager.get_account_filtered_balance( account=ledger_account_debit, metadata_key=rand_key, metadata_value=rand_value, @@ -310,7 +324,7 @@ class TestLedgerAccountManagerGet: ) def test_get_balance_timerange_empty( - self, ledger_account: LedgerAccount, lm: LedgerManager + self, ledger_account: LedgerAccount, ledger_manager: LedgerManager ): - res = lm.get_account_balance_timerange(account=ledger_account) + res = ledger_manager.get_account_balance_timerange(account=ledger_account) assert res == 0 diff --git a/tests/managers/thl/test_ledger/test_thl_lm_tx.py b/tests/managers/thl/test_ledger/test_thl_lm_tx.py index 2e4ab5e..b0484ae 100644 --- a/tests/managers/thl/test_ledger/test_thl_lm_tx.py +++ b/tests/managers/thl/test_ledger/test_thl_lm_tx.py @@ -311,15 +311,14 @@ class TestThlLedgerTxManager: def test_create_tx_bp_payout_( self, product: Product, - thl_lm: ThlLedgerManager, - ledger_manager: LedgerManager, + thl_ledger_manager: ThlLedgerManager, currency: LedgerCurrency, ): rand_amount: USDCent = USDCent(randint(100, 1_000)) payoutevent_uuid = uuid4().hex # Create a BP Payout for a Product without any activity. - tx = thl_lm.create_tx_bp_payout_( + tx = thl_ledger_manager.create_tx_bp_payout_( product=product, amount=rand_amount, payoutevent_uuid=payoutevent_uuid, diff --git a/tests/managers/thl/test_payout.py b/tests/managers/thl/test_payout.py index 2494de8..ad101a4 100644 --- a/tests/managers/thl/test_payout.py +++ b/tests/managers/thl/test_payout.py @@ -86,7 +86,6 @@ class TestPayout: self, user: User, user_payout_event_manager: UserPayoutEventManager, - ledger_manager: LedgerManager, thl_ledger_manager: ThlLedgerManager, utc_now: datetime, ): @@ -128,11 +127,11 @@ class TestPayout: self, thl_web_rw: PostgresConfig, product: Product, - thl_lm: ThlLedgerManager, + thl_ledger_manager: ThlLedgerManager, brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager, utc_now: datetime, ) -> BrokerageProductPayoutEvent: - account = thl_lm.get_account_or_create_bp_wallet(product=product) + account = thl_ledger_manager.get_account_or_create_bp_wallet(product=product) bp_pe = BrokerageProductPayoutEvent( product_id=product.uuid, amount=USDCent(100), @@ -161,15 +160,14 @@ class TestPayout: self, product: Product, brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager, - thl_lm: ThlLedgerManager, - ledger_manager: LedgerManager, + thl_ledger_manager: ThlLedgerManager, utc_now: datetime, pending_bp_pe: BrokerageProductPayoutEvent, ): - thl_lm.get_account_or_create_bp_wallet(product=product) + thl_ledger_manager.get_account_or_create_bp_wallet(product=product) brokerage_product_payout_event_manager.create_tx_bp_payout_from_payout_event( - thl_ledger_manager=thl_lm, + thl_ledger_manager=thl_ledger_manager, bp_pe=pending_bp_pe, product=product, created=utc_now, @@ -177,7 +175,7 @@ class TestPayout: with pytest.raises(ValueError) as cm: brokerage_product_payout_event_manager.create_tx_bp_payout_from_payout_event( - thl_ledger_manager=thl_lm, + thl_ledger_manager=thl_ledger_manager, product=product, bp_pe=pending_bp_pe, created=utc_now, @@ -187,7 +185,6 @@ class TestPayout: def test_filter( self, thl_ledger_manager: ThlLedgerManager, - ledger_manager: LedgerManager, product: Product, user: User, user_payout_event_manager: UserPayoutEventManager, @@ -280,19 +277,18 @@ class TestBusinessPayoutEventManager: def test_base( self, - brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager, business_payout_event_manager: BusinessPayoutEventManager, delete_ledger_db: Callable[..., None], create_main_accounts: Callable[..., None], thl_ledger_manager: ThlLedgerManager, product_factory: Callable[..., Product], bp_payout_factory: Callable[..., BrokerageProductPayoutEvent], - business: Business, + gr_business: Business, ): delete_ledger_db() create_main_accounts() - p1: Product = product_factory(business=business) + p1: Product = product_factory(business=gr_business) thl_ledger_manager.get_account_or_create_bp_wallet(product=p1) ach_id1 = uuid4().hex @@ -310,23 +306,25 @@ class TestBusinessPayoutEventManager: bp_payout_factory(product=p1, amount=USDCent(50), ext_ref_id=ach_id2) - business.prebuild_payouts( + gr_business.prebuild_payouts( bpem=business_payout_event_manager, ) - assert isinstance(business.payouts, list) - assert len(business.payouts) == 3 - assert business.payouts_total == sum([pe.amount for pe in business.payouts]) - assert business.payouts[0].created > business.payouts[1].created - assert len(business.payouts[0].bp_payouts) == 1 + assert isinstance(gr_business.payouts, list) + assert len(gr_business.payouts) == 3 + assert gr_business.payouts_total == sum( + [pe.amount for pe in gr_business.payouts] + ) + assert gr_business.payouts[0].created > gr_business.payouts[1].created + assert len(gr_business.payouts[0].bp_payouts) == 1 # Cannot pay out the same product twice in the same business payout # assert len(business.payouts[1].bp_payouts) == 2 - assert len(business.payouts[1].bp_payouts) == 1 + assert len(gr_business.payouts[1].bp_payouts) == 1 - assert business.payouts[0].ext_ref_id == ach_id2 - assert business.payouts[1].ext_ref_id == ach_id1 - assert business.payouts[2].ext_ref_id == "none" + assert gr_business.payouts[0].ext_ref_id == ach_id2 + assert gr_business.payouts[1].ext_ref_id == ach_id1 + assert gr_business.payouts[2].ext_ref_id == "none" def test_update_ext_reference_ids( self, @@ -345,13 +343,13 @@ class TestBusinessPayoutEventManager: mnt_filepath: GRLDatasets, product_manager: ProductManager, start: datetime, - business: Business, + gr_business: Business, ): delete_ledger_db() create_main_accounts() delete_df_collection(coll=ledger_collection) - p1: Product = product_factory(business=business) + p1: Product = product_factory(business=gr_business) u1: User = user_factory(product=p1) thl_ledger_manager.get_account_or_create_bp_wallet(product=p1) @@ -377,7 +375,7 @@ class TestBusinessPayoutEventManager: # We must build the balance to issue ACH/Wire ledger_collection.initial_load(client=None, sync=True) pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) - business.prebuild_balance( + gr_business.prebuild_balance( thl_pg_config=thl_web_rr, lm=thl_ledger_manager, ds=mnt_filepath, @@ -386,7 +384,7 @@ class TestBusinessPayoutEventManager: ) res = business_payout_event_manager.create_from_ach_or_wire( - business=business, + business=gr_business, amount=USDCent(100_01), pm=product_manager, thl_lm=thl_ledger_manager, @@ -558,7 +556,7 @@ class TestBusinessPayoutEventManager: create_main_accounts: Callable[..., None], delete_df_collection: Callable[..., None], ledger_collection: LedgerDFCollection, - business: Business, + gr_business: Business, user_factory: Callable[..., User], product_factory: Callable[..., Product], session_with_tx_factory: Callable[..., Session], @@ -581,7 +579,7 @@ class TestBusinessPayoutEventManager: create_main_accounts() delete_df_collection(coll=ledger_collection) - p1: Product = product_factory(business=business) + p1: Product = product_factory(business=gr_business) u1: User = user_factory(product=p1) thl_ledger_manager.get_account_or_create_bp_wallet(product=p1) @@ -603,7 +601,7 @@ class TestBusinessPayoutEventManager: ledger_collection.initial_load(client=None, sync=True) pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) - business.prebuild_balance( + gr_business.prebuild_balance( thl_pg_config=thl_web_rr, lm=ledger_manager, ds=mnt_filepath, @@ -613,7 +611,7 @@ class TestBusinessPayoutEventManager: with pytest.raises(expected_exception=AssertionError) as cm: business_payout_event_manager.create_from_ach_or_wire( - business=business, + business=gr_business, amount=USDCent(500), pm=product_manager, thl_lm=thl_ledger_manager, @@ -631,7 +629,7 @@ class TestBusinessPayoutEventManager: create_main_accounts: Callable[..., None], delete_df_collection: Callable[..., None], ledger_collection: LedgerDFCollection, - business: Business, + gr_business: Business, user_factory: Callable[..., User], product_factory: Callable[..., Product], session_with_tx_factory: Callable[..., None], @@ -648,9 +646,9 @@ class TestBusinessPayoutEventManager: create_main_accounts() delete_df_collection(coll=ledger_collection) - p1: Product = product_factory(business=business) - p2: Product = product_factory(business=business) - p3: Product = product_factory(business=business) + p1: Product = product_factory(business=gr_business) + p2: Product = product_factory(business=gr_business) + p3: Product = product_factory(business=gr_business) _: User = user_factory(product=p1) u2: User = user_factory(product=p2) u3: User = user_factory(product=p3) @@ -679,7 +677,7 @@ class TestBusinessPayoutEventManager: ledger_collection.initial_load(client=None, sync=True) pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) - business.prebuild_balance( + gr_business.prebuild_balance( thl_pg_config=thl_web_rr, lm=ledger_manager, ds=mnt_filepath, @@ -687,13 +685,13 @@ class TestBusinessPayoutEventManager: pop_ledger=pop_ledger_merge, ) - bb = business.balance + bb = gr_business.balance assert isinstance(bb, BusinessBalances) assert bb.payout == 475_00 # $500 * .95% = $475 assert bb.net == 475_00 bp1 = business_payout_event_manager.create_from_ach_or_wire( - business=business, + business=gr_business, amount=USDCent(100_00), pm=product_manager, thl_lm=thl_ledger_manager, @@ -705,7 +703,7 @@ class TestBusinessPayoutEventManager: assert len(bp1.bp_payouts) == 2 bp2 = business_payout_event_manager.create_from_ach_or_wire( - business=business, + business=gr_business, amount=USDCent(bb.available_balance), pm=product_manager, thl_lm=thl_ledger_manager, @@ -743,7 +741,7 @@ class TestBusinessPayoutEventManager: create_main_accounts: Callable[..., None], delete_df_collection: Callable[..., None], ledger_collection: LedgerDFCollection, - business: Business, + gr_business: Business, user_factory: Callable[..., User], product_factory: Callable[..., Product], session_with_tx_factory: Callable[..., None], @@ -768,9 +766,9 @@ class TestBusinessPayoutEventManager: create_main_accounts() delete_df_collection(coll=ledger_collection) - p1: Product = product_factory(business=business) - p2: Product = product_factory(business=business) - p3: Product = product_factory(business=business) + p1: Product = product_factory(business=gr_business) + p2: Product = product_factory(business=gr_business) + p3: Product = product_factory(business=gr_business) u1: User = user_factory(product=p1) u2: User = user_factory(product=p2) u3: User = user_factory(product=p3) @@ -813,10 +811,10 @@ class TestBusinessPayoutEventManager: started=start + timedelta(days=1, hours=3, minutes=1 + idx), ) - # Now that we paid out the business: Business, let's confirm the updated balances + # Now that we paid out the gr_business: Business, let's confirm the updated balances ledger_collection.initial_load(client=None, sync=True) pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) - business.prebuild_balance( + gr_business.prebuild_balance( thl_pg_config=thl_web_rr, lm=ledger_manager, ds=mnt_filepath, @@ -824,7 +822,7 @@ class TestBusinessPayoutEventManager: pop_ledger=pop_ledger_merge, ) - bb1 = business.balance + bb1 = gr_business.balance assert isinstance(bb1, BusinessBalances) pb1 = bb1.product_balances[0] pb2 = bb1.product_balances[1] @@ -848,18 +846,18 @@ class TestBusinessPayoutEventManager: assert pb2.recoup_usd_str == "$0.00" assert pb3.recoup_usd_str == "$0.00" - assert business.payouts is None - business.prebuild_payouts( + assert gr_business.payouts is None + gr_business.prebuild_payouts( thl_pg_config=thl_web_rr, thl_lm=thl_ledger_manager, bpem=business_payout_event_manager, ) - assert isinstance(business.payouts, list) - assert len(business.payouts) == 1 - assert business.payouts[0].ext_ref_id == ach_id1 + assert isinstance(gr_business.payouts, list) + assert len(gr_business.payouts) == 1 + assert gr_business.payouts[0].ext_ref_id == ach_id1 bp1 = business_payout_event_manager.create_from_ach_or_wire( - business=business, + business=gr_business, amount=USDCent(bb1.available_balance), pm=product_manager, thl_lm=thl_ledger_manager, @@ -937,7 +935,7 @@ class TestBusinessPayoutEventManager: create_main_accounts: Callable[..., None], delete_df_collection: Callable[..., None], ledger_collection: LedgerDFCollection, - business: Business, + gr_business: Business, user_factory: Callable[..., User], product_factory: Callable[..., Product], session_with_tx_factory: Callable[..., None], @@ -950,7 +948,7 @@ class TestBusinessPayoutEventManager: rm_pop_ledger_merge: Callable[..., None], ): """There are valid instances when we want issue a ACH or Wire to a - business: Business, but not for the full Available Balance amount in their + gr_business: Business, but not for the full Available Balance amount in their account. To test this, we'll create a Business with multiple Products, and @@ -965,9 +963,9 @@ class TestBusinessPayoutEventManager: create_main_accounts() delete_df_collection(coll=ledger_collection) - p1: Product = product_factory(business=business) - p2: Product = product_factory(business=business) - p3: Product = product_factory(business=business) + p1: Product = product_factory(business=gr_business) + p2: Product = product_factory(business=gr_business) + p3: Product = product_factory(business=gr_business) u1: User = user_factory(product=p1) u2: User = user_factory(product=p2) u3: User = user_factory(product=p3) @@ -988,20 +986,20 @@ class TestBusinessPayoutEventManager: # Now that we paid out the business: Business, let's confirm the updated balances ledger_collection.initial_load(client=None, sync=True) pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) - business.prebuild_balance( + gr_business.prebuild_balance( thl_pg_config=thl_web_rr, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, pop_ledger=pop_ledger_merge, ) - business.prebuild_payouts( + gr_business.prebuild_payouts( bpem=business_payout_event_manager, ) # Confirm the initial amounts. - assert len(business.payouts) == 0 - bb1 = business.balance + assert len(gr_business.payouts) == 0 + bb1 = gr_business.balance assert isinstance(bb1, BusinessBalances) assert bb1.payout == 3 * 5 * 4750 @@ -1015,16 +1013,16 @@ class TestBusinessPayoutEventManager: assert bb1.product_balances[x].balance == 5 * 4750 assert bb1.product_balances[x].available_balance_usd_str == "$178.13" - assert business.payouts_total_str == "$0.00" - assert isinstance(business.balance, BusinessBalances) - assert business.balance.payment_usd_str == "$0.00" - assert business.balance.available_balance_usd_str == "$534.39" + assert gr_business.payouts_total_str == "$0.00" + assert isinstance(gr_business.balance, BusinessBalances) + assert gr_business.balance.payment_usd_str == "$0.00" + assert gr_business.balance.available_balance_usd_str == "$534.39" # This is the important part, even those the Business has $534.39 # available to it, we are only trying to issue out a $250.00 ACH or # Wire to the Business bp1 = business_payout_event_manager.create_from_ach_or_wire( - business=business, + business=gr_business, amount=USDCent(250_00), pm=product_manager, thl_lm=thl_ledger_manager, @@ -1033,7 +1031,7 @@ class TestBusinessPayoutEventManager: assert isinstance(bp1, BusinessPayoutEvent) assert len(bp1.bp_payouts) == 3 - # Now that we paid out the business: Business, let's confirm the updated + # Now that we paid out the gr_business: Business, let's confirm the updated # balances. Clear and rebuild the parquet files. rm_ledger_collection() rm_pop_ledger_merge() @@ -1043,25 +1041,23 @@ class TestBusinessPayoutEventManager: # Now rebuild and confirm the payouts, balance.payment, and the # balance.available_balance are reflective of having a $250 ACH/Wire # sent to the Business - business.prebuild_balance( + gr_business.prebuild_balance( thl_pg_config=thl_web_rr, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, pop_ledger=pop_ledger_merge, ) - business.prebuild_payouts( - thl_pg_config=thl_web_rr, - thl_lm=thl_ledger_manager, + gr_business.prebuild_payouts( bpem=business_payout_event_manager, ) - assert isinstance(business.payouts, list) - assert len(business.payouts) == 1 - assert len(business.payouts[0].bp_payouts) == 3 - assert business.payouts_total_str == "$250.00" - assert isinstance(business.balance, BusinessBalances) - assert business.balance.payment_usd_str == "$250.00" - assert business.balance.available_balance_usd_str == "$346.88" + assert isinstance(gr_business.payouts, list) + assert len(gr_business.payouts) == 1 + assert len(gr_business.payouts[0].bp_payouts) == 3 + assert gr_business.payouts_total_str == "$250.00" + assert isinstance(gr_business.balance, BusinessBalances) + assert gr_business.balance.payment_usd_str == "$250.00" + assert gr_business.balance.available_balance_usd_str == "$346.88" def test_ach_tx_id_reference( self, @@ -1074,7 +1070,7 @@ class TestBusinessPayoutEventManager: create_main_accounts: Callable[..., None], delete_df_collection: Callable[..., None], ledger_collection: LedgerDFCollection, - business: Business, + gr_business: Business, user_factory: Callable[..., User], product_factory: Callable[..., Product], session_with_tx_factory: Callable[..., Session], @@ -1092,9 +1088,9 @@ class TestBusinessPayoutEventManager: create_main_accounts() delete_df_collection(coll=ledger_collection) - p1: Product = product_factory(business=business) - p2: Product = product_factory(business=business) - p3: Product = product_factory(business=business) + p1: Product = product_factory(business=gr_business) + p2: Product = product_factory(business=gr_business) + p3: Product = product_factory(business=gr_business) u1: User = user_factory(product=p1) u2: User = user_factory(product=p2) u3: User = user_factory(product=p3) @@ -1118,7 +1114,7 @@ class TestBusinessPayoutEventManager: rm_pop_ledger_merge() ledger_collection.initial_load(client=None, sync=True) pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) - business.prebuild_balance( + gr_business.prebuild_balance( thl_pg_config=thl_web_rr, lm=ledger_manager, ds=mnt_filepath, @@ -1127,7 +1123,7 @@ class TestBusinessPayoutEventManager: ) bp1 = business_payout_event_manager.create_from_ach_or_wire( - business=business, + business=gr_business, amount=USDCent(100_01), transaction_id=ach_id1, pm=product_manager, @@ -1139,7 +1135,7 @@ class TestBusinessPayoutEventManager: rm_pop_ledger_merge() ledger_collection.initial_load(client=None, sync=True) pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) - business.prebuild_balance( + gr_business.prebuild_balance( thl_pg_config=thl_web_rr, lm=ledger_manager, ds=mnt_filepath, @@ -1148,7 +1144,7 @@ class TestBusinessPayoutEventManager: ) bp2 = business_payout_event_manager.create_from_ach_or_wire( - business=business, + business=gr_business, amount=USDCent(100_02), transaction_id=ach_id2, pm=product_manager, @@ -1163,18 +1159,18 @@ class TestBusinessPayoutEventManager: rm_pop_ledger_merge() ledger_collection.initial_load(client=None, sync=True) pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) - business.prebuild_payouts( + gr_business.prebuild_payouts( thl_pg_config=thl_web_rr, thl_lm=thl_ledger_manager, bpem=business_payout_event_manager, ) - business.prebuild_balance( + gr_business.prebuild_balance( thl_pg_config=thl_web_rr, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, pop_ledger=pop_ledger_merge, ) - assert isinstance(business.payouts, list) - assert business.payouts[0].ext_ref_id == ach_id2 - assert business.payouts[1].ext_ref_id == ach_id1 + assert isinstance(gr_business.payouts, list) + assert gr_business.payouts[0].ext_ref_id == ach_id2 + assert gr_business.payouts[1].ext_ref_id == ach_id1 diff --git a/tests/managers/thl/test_session_manager.py b/tests/managers/thl/test_session_manager.py index 30fd9ec..67a802e 100644 --- a/tests/managers/thl/test_session_manager.py +++ b/tests/managers/thl/test_session_manager.py @@ -137,19 +137,19 @@ class TestSessionManagerFilter: def test_business( self, product_factory: Callable[..., Product], - business: Business, + gr_business: Business, user_factory: Callable[..., User], session_manager: SessionManager, utc_hour_ago: datetime, thl_web_rr: PostgresConfig, ): - p1 = product_factory(business=business) + p1 = product_factory(business=gr_business) for _ in range(5): u = user_factory(product=p1) session_manager.create(started=utc_hour_ago, user=u, uuid_id=uuid4().hex) - business.prefetch_products(thl_pg_config=thl_web_rr) - assert len(business.product_uuids) == 1 - res = session_manager.filter(product_uuids=business.product_uuids) + gr_business.prefetch_products(thl_pg_config=thl_web_rr) + assert len(gr_business.product_uuids) == 1 + res = session_manager.filter(product_uuids=gr_business.product_uuids) assert len(res) == 5 diff --git a/tests/models/gr/test_authentication.py b/tests/models/gr/test_authentication.py index ac1298f..059a0a4 100644 --- a/tests/models/gr/test_authentication.py +++ b/tests/models/gr/test_authentication.py @@ -116,7 +116,7 @@ class TestGRUser: class TestGRUserMethods: - def test_cache_key(self, gr_user: GRUser, gr_redis: RedisConfig): + def test_cache_key(self, gr_user: GRUser): assert isinstance(gr_user.cache_key, str) assert ":" in gr_user.cache_key assert str(gr_user.id) in gr_user.cache_key @@ -124,13 +124,12 @@ class TestGRUserMethods: def test_to_redis( self, gr_user: GRUser, - gr_redis: Redis, team: Team, - business: Business, + gr_business: Business, product_factory: Callable[..., Product], membership_factory: Callable[..., Membership], ): - product_factory(team=team, business=business) + product_factory(team=team, business=gr_business) membership_factory(team=team, gr_user=gr_user) res = gr_user.to_redis() @@ -144,31 +143,30 @@ class TestGRUserMethods: def test_set_cache( self, gr_user: GRUser, - gr_user_token: GRToken, - gr_redis: Redis, gr_db: PostgresConfig, thl_web_rr: PostgresConfig, 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 - assert gr_redis.get(name=f"{gr_user.cache_key}:business_uuids") is None - assert gr_redis.get(name=f"{gr_user.cache_key}:product_uuids") is None + + client = gr_redis_config.create_redis_client() + + assert client.get(name=gr_user.cache_key) is None + assert client.get(name=f"{gr_user.cache_key}:team_uuids") is None + assert client.get(name=f"{gr_user.cache_key}:business_uuids") is None + assert client.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, redis_config=gr_redis_config ) - assert gr_redis.get(name=gr_user.cache_key) is not None - assert gr_redis.get(name=f"{gr_user.cache_key}:team_uuids") is not None - assert gr_redis.get(name=f"{gr_user.cache_key}:business_uuids") is not None - assert gr_redis.get(name=f"{gr_user.cache_key}:product_uuids") is not None + assert client.get(name=gr_user.cache_key) is not None + assert client.get(name=f"{gr_user.cache_key}:team_uuids") is not None + assert client.get(name=f"{gr_user.cache_key}:business_uuids") is not None + assert client.get(name=f"{gr_user.cache_key}:product_uuids") is not None def test_set_cache_gr_user( self, gr_user: GRUser, - gr_user_token: GRToken, - gr_redis: RedisConfig, gr_redis_config: RedisConfig, gr_db: PostgresConfig, thl_web_rr: PostgresConfig, @@ -179,6 +177,8 @@ class TestGRUserMethods: ): from generalresearch.models.gr.authentication import GRUser + client = gr_redis_config.create_redis_client() + p1 = product_factory(team=team) membership_factory(team=team, gr_user=gr_user) @@ -186,7 +186,7 @@ class TestGRUserMethods: 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) + res: str = client.get(name=gr_user.cache_key) gru2 = GRUser.from_redis(res) assert gr_user.model_dump_json( @@ -203,9 +203,6 @@ class TestGRUserMethods: def test_set_cache_team_uuids( self, gr_user: GRUser, - membership: Membership, - gr_user_token: GRToken, - gr_redis: Redis, gr_db: PostgresConfig, thl_web_rr: PostgresConfig, product_factory: Callable[..., Product], @@ -213,11 +210,12 @@ class TestGRUserMethods: gr_redis_config: RedisConfig, ): product_factory(team=team) + client = gr_redis_config.create_redis_client() gr_user.set_cache( 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")) + res = json.loads(client.get(name=f"{gr_user.cache_key}:team_uuids")) assert len(res) == 1 assert gr_user.team_uuids == res @@ -225,29 +223,27 @@ class TestGRUserMethods: def test_set_cache_business_uuids( self, gr_user: GRUser, - gr_redis: Redis, gr_db: PostgresConfig, thl_web_rr: PostgresConfig, product_factory: Callable[..., Product], - business: Business, + gr_business: Business, team: Team, gr_redis_config: RedisConfig, ): - product_factory(team=team, business=business) + product_factory(team=team, business=gr_business) gr_user.set_cache( 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")) + + client = gr_redis_config.create_redis_client() + res = json.loads(client.get(name=f"{gr_user.cache_key}:business_uuids")) assert len(res) == 1 assert gr_user.business_uuids == res def test_set_cache_product_uuids( self, gr_user: GRUser, - membership: Membership, - gr_user_token: GRToken, - gr_redis: Redis, gr_db: PostgresConfig, thl_web_rr: PostgresConfig, product_factory: Callable[..., Product], @@ -259,7 +255,8 @@ class TestGRUserMethods: gr_user.set_cache( 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")) + client = gr_redis_config.create_redis_client() + res = json.loads(client.get(name=f"{gr_user.cache_key}:product_uuids")) assert len(res) == 1 assert gr_user.product_uuids == res diff --git a/tests/models/gr/test_business.py b/tests/models/gr/test_business.py index 2c12da1..90e69db 100644 --- a/tests/models/gr/test_business.py +++ b/tests/models/gr/test_business.py @@ -64,10 +64,8 @@ class TestBusinessBankAccount: gr_business: Business, business_bank_account_manager: BusinessBankAccountManager, ): - from generalresearch.models.gr.business import ( - BusinessBankAccount, - TransferMethod, - ) + from generalresearch.models.gr.business import BusinessBankAccount + from generalresearch.models.gr.definitions import TransferMethod instance = business_bank_account_manager.create( business_id=gr_business.id, @@ -115,7 +113,7 @@ class TestBusiness: @pytest.fixture def offset(self) -> str: - return "30d" + return "30D" @pytest.fixture def duration(self) -> timedelta | None: @@ -222,46 +220,46 @@ class TestBusiness: def test_teams( self, - business: Business, + gr_business: Business, team: Team, team_manager: TeamManager, gr_db: PostgresConfig, ): - assert business.teams is None + assert gr_business.teams is None - business.prefetch_teams(pg_config=gr_db) - assert isinstance(business.teams, list) - assert len(business.teams) == 0 + gr_business.prefetch_teams(pg_config=gr_db) + assert isinstance(gr_business.teams, list) + assert len(gr_business.teams) == 0 - team_manager.add_business(team=team, business=business) - assert len(business.teams) == 0 - business.prefetch_teams(pg_config=gr_db) - assert len(business.teams) == 1 + team_manager.add_business(team=team, business=gr_business) + assert len(gr_business.teams) == 0 + gr_business.prefetch_teams(pg_config=gr_db) + assert len(gr_business.teams) == 1 def test_products( self, - business: Business, + gr_business: Business, product_factory: Callable[..., Product], product_manager: ProductManager, ): - p1 = product_factory(business=business) - assert business.products is None + p1 = product_factory(business=gr_business) + assert gr_business.products is None - business.prefetch_products(product_manager=product_manager) - assert isinstance(business.products, list) - assert len(business.products) == 1 - assert isinstance(business.products[0], Product) + gr_business.prefetch_products(product_manager=product_manager) + assert isinstance(gr_business.products, list) + assert len(gr_business.products) == 1 + assert isinstance(gr_business.products[0], Product) - assert business.products[0].uuid == p1.uuid + assert gr_business.products[0].uuid == p1.uuid # Add two more, but list is still one until we prefetch - product_factory(business=business) - product_factory(business=business) - assert len(business.products) == 1 + product_factory(business=gr_business) + product_factory(business=gr_business) + assert len(gr_business.products) == 1 - business.prefetch_products(product_manager=product_manager) - assert len(business.products) == 3 + gr_business.prefetch_products(product_manager=product_manager) + assert len(gr_business.products) == 3 def test_bank_accounts( self, @@ -306,7 +304,6 @@ class TestBusiness: self, gr_business: Business, product_factory: Callable[..., Product], - thl_web_rr: PostgresConfig, thl_ledger_manager: ThlLedgerManager, business_payout_event_manager: BusinessPayoutEventManager, ): @@ -322,8 +319,6 @@ class TestBusiness: thl_ledger_manager.get_account_or_create_bp_wallet(product=p) gr_business.prebuild_payouts( - thl_pg_config=thl_web_rr, - thl_lm=thl_ledger_manager, bpem=business_payout_event_manager, ) assert isinstance(gr_business.payouts, list) @@ -335,7 +330,6 @@ class TestBusiness: product_factory: Callable[..., Product], bp_payout_factory: Callable[..., BrokerageProductPayoutEvent], thl_ledger_manager: ThlLedgerManager, - thl_web_rr: PostgresConfig, business_payout_event_manager: BusinessPayoutEventManager, create_main_accounts: Callable[..., None], ): @@ -351,8 +345,6 @@ class TestBusiness: ) gr_business.prebuild_payouts( - thl_pg_config=thl_web_rr, - thl_lm=thl_ledger_manager, bpem=business_payout_event_manager, ) assert len(gr_business.payouts) == 1 @@ -478,7 +470,7 @@ class TestBusinessBalance: @pytest.fixture def offset(self) -> str: - return "30d" + return "30D" @pytest.fixture def duration(self) -> timedelta | None: @@ -1190,15 +1182,14 @@ class TestBusinessMethods: ) -> timedelta | None: return None - def test_cache_key(self, business: Business): - assert isinstance(business.cache_key, str) - assert ":" in business.cache_key - assert str(business.uuid) in business.cache_key + def test_cache_key(self, gr_business: Business): + assert isinstance(gr_business.cache_key, str) + assert ":" in gr_business.cache_key + assert str(gr_business.uuid) in gr_business.cache_key def test_set_cache( self, gr_business: Business, - gr_redis: RedisConfig, gr_db: PostgresConfig, thl_web_rr: PostgresConfig, client_no_amm: DaskClient, @@ -1218,7 +1209,8 @@ class TestBusinessMethods: gr_redis_config: RedisConfig, mnt_gr_api_dir: Path, ): - assert gr_redis.get(name=gr_business.cache_key) is None + client = gr_redis_config.create_redis_client() + assert client.get(name=gr_business.cache_key) is None p1 = product_factory(team=team, business=gr_business) u1 = user_factory(product=p1) @@ -1244,7 +1236,7 @@ class TestBusinessMethods: mnt_gr_api=mnt_gr_api_dir, ) - assert gr_redis.hgetall(name=gr_business.cache_key) is not None + assert client.hgetall(name=gr_business.cache_key) is not None from generalresearch.models.gr.business import Business # We're going to pull only a specific year, but make sure that @@ -1367,7 +1359,7 @@ class TestBusinessMethods: session_factory: Callable[..., Session], product_factory: Callable[..., Product], delete_df_collection: Callable[..., None], - business: Business, + gr_business: Business, mnt_filepath: GRLDatasets, mnt_gr_api_dir: Path, ): @@ -1375,8 +1367,8 @@ class TestBusinessMethods: delete_df_collection(coll=wall_collection) delete_df_collection(coll=session_collection) - p1 = product_factory(business=business) - p2 = product_factory(business=business) + p1 = product_factory(business=gr_business) + p2 = product_factory(business=gr_business) for p in [p1, p2]: u = user_factory(product=p) @@ -1397,7 +1389,7 @@ class TestBusinessMethods: pg_config=thl_web_rr, ) - business.prebuild_enriched_session_parquet( + gr_business.prebuild_enriched_session_parquet( thl_pg_config=thl_web_rr, ds=mnt_filepath, client=client_no_amm, @@ -1407,7 +1399,9 @@ class TestBusinessMethods: # Now try to read from path df = pd.read_parquet( - os.path.join(mnt_gr_api_dir, "pop_session", f"{business.file_key}.parquet") + os.path.join( + mnt_gr_api_dir, "pop_session", f"{gr_business.file_key}.parquet" + ) ) assert isinstance(df, pd.DataFrame) diff --git a/tests/models/gr/test_team.py b/tests/models/gr/test_team.py index c1ae6d6..aa2de45 100644 --- a/tests/models/gr/test_team.py +++ b/tests/models/gr/test_team.py @@ -152,7 +152,6 @@ class TestTeamMethods: def test_set_cache( self, team: Team, - gr_redis: RedisConfig, gr_db: PostgresConfig, thl_web_rr: PostgresConfig, gr_redis_config: RedisConfig, @@ -162,7 +161,8 @@ class TestTeamMethods: enriched_wall_merge: EnrichedWallMerge, enriched_session_merge: EnrichedSessionMerge, ): - assert gr_redis.get(name=team.cache_key) is None + client = gr_redis_config.create_redis_client() + assert client.get(name=team.cache_key) is None team.set_cache( pg_config=gr_db, @@ -175,7 +175,7 @@ class TestTeamMethods: enriched_session=enriched_session_merge, ) - assert gr_redis.hgetall(name=team.cache_key) is not None + assert client.hgetall(name=team.cache_key) is not None def test_set_cache_team( self, diff --git a/tests/models/test_finance.py b/tests/models/test_finance.py index eabc877..c579d78 100644 --- a/tests/models/test_finance.py +++ b/tests/models/test_finance.py @@ -760,7 +760,7 @@ class TestPOPFinancialData: duration: timedelta, create_main_accounts: Callable[..., None], session_with_tx_factory: Callable[..., Session], - thl_lm: ThlLedgerManager, + thl_ledger_manager: ThlLedgerManager, delete_df_collection: Callable[..., None], delete_ledger_db: Callable[..., None], ): @@ -798,8 +798,10 @@ class TestPOPFinancialData: last_item_finish = item_finishes[0] accounts = [] - for _ in users: - account = thl_lm.get_account_or_create_bp_wallet(product=u.product) + for _u in users: + account = thl_ledger_manager.get_account_or_create_bp_wallet( + product=_u.product + ) accounts.append(account) account_ids = [a.uuid for a in accounts] @@ -856,7 +858,7 @@ class TestBusinessBalanceData: user_factory: Callable[..., User], product: Product, create_main_accounts: Callable[..., None], - thl_lm: ThlLedgerManager, + thl_ledger_manager: ThlLedgerManager, thl_web_rr: PostgresConfig, delete_df_collection: Callable[..., None], delete_ledger_db: Callable[..., None], @@ -886,7 +888,9 @@ class TestBusinessBalanceData: pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) # assert pop_ledger_merge.progress.has_archive.eq(True).all() - account: LedgerAccount = thl_lm.get_account_or_create_bp_wallet(product=product) + account: LedgerAccount = thl_ledger_manager.get_account_or_create_bp_wallet( + product=product + ) ddf = pop_ledger_merge.ddf( force_rr_latest=False, diff --git a/tests/models/thl/test_payout.py b/tests/models/thl/test_payout.py index 927687e..cc00f33 100644 --- a/tests/models/thl/test_payout.py +++ b/tests/models/thl/test_payout.py @@ -10,8 +10,8 @@ from generalresearch.models.gr import Team from generalresearch.models.gr.business import ( Business, BusinessAddress, - BusinessType, ) +from generalresearch.models.gr.definitions import BusinessType from generalresearch.models.thl.payout import ( BrokerageProductPayoutEvent, BusinessPayoutEvent, diff --git a/tests/models/thl/test_product.py b/tests/models/thl/test_product.py index cc0fa8e..a1b3688 100644 --- a/tests/models/thl/test_product.py +++ b/tests/models/thl/test_product.py @@ -603,7 +603,7 @@ class TestProductFinancials: @pytest.fixture def offset(self) -> str: - return "30d" + return "30D" @pytest.fixture def duration(self) -> timedelta | None: @@ -611,12 +611,12 @@ class TestProductFinancials: def test_balance( self, - business: Business, + gr_business: Business, product_factory: Callable[..., Product], user_factory: Callable[..., User], mnt_filepath: GRLDatasets, bp_payout_factory: Callable[..., BrokerageProductPayoutEvent], - thl_lm: ThlLedgerManager, + thl_ledger_manager: ThlLedgerManager, start: datetime, brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager, session_with_tx_factory: Callable[..., Session], @@ -633,33 +633,54 @@ class TestProductFinancials: from generalresearch.currency import USDCent - p1: Product = product_factory(business=business) + p1: Product = product_factory(business=gr_business) u1: User = user_factory(product=p1) - bp_wallet = thl_lm.get_account_or_create_bp_wallet(product=p1) - thl_lm.get_account_or_create_user_wallet(user=u1) + bp_wallet = thl_ledger_manager.get_account_or_create_bp_wallet(product=p1) + thl_ledger_manager.get_account_or_create_user_wallet(user=u1) brokerage_product_payout_event_manager.set_account_lookup_table(thl_lm=thl_lm) - assert len(thl_lm.get_tx_filtered_by_account(account_uuid=bp_wallet.uuid)) == 0 + assert ( + len( + thl_ledger_manager.get_tx_filtered_by_account( + account_uuid=bp_wallet.uuid + ) + ) + == 0 + ) session_with_tx_factory( user=u1, wall_req_cpi=Decimal(".50"), started=start + timedelta(days=1), ) - assert thl_lm.get_account_balance(account=bp_wallet) == 48 - assert len(thl_lm.get_tx_filtered_by_account(account_uuid=bp_wallet.uuid)) == 1 + assert thl_ledger_manager.get_account_balance(account=bp_wallet) == 48 + assert ( + len( + thl_ledger_manager.get_tx_filtered_by_account( + account_uuid=bp_wallet.uuid + ) + ) + == 1 + ) session_with_tx_factory( user=u1, wall_req_cpi=Decimal("1.00"), started=start + timedelta(days=2), ) - assert thl_lm.get_account_balance(account=bp_wallet) == 143 - assert len(thl_lm.get_tx_filtered_by_account(account_uuid=bp_wallet.uuid)) == 2 + assert thl_ledger_manager.get_account_balance(account=bp_wallet) == 143 + assert ( + len( + thl_ledger_manager.get_tx_filtered_by_account( + account_uuid=bp_wallet.uuid + ) + ) + == 2 + ) with pytest.raises(expected_exception=AssertionError) as cm: p1.prebuild_balance( - thl_lm=thl_lm, + thl_lm=thl_ledger_manager, ds=mnt_filepath, client=client_no_amm, ) @@ -669,7 +690,7 @@ class TestProductFinancials: pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) p1.prebuild_balance( - thl_lm=thl_lm, + thl_lm=thl_ledger_manager, ds=mnt_filepath, client=client_no_amm, ) @@ -683,7 +704,7 @@ class TestProductFinancials: assert p1.balance.available_balance == 108 p1.prebuild_payouts( - thl_lm=thl_lm, + thl_lm=thl_ledger_manager, bp_pem=brokerage_product_payout_event_manager, ) assert p1.payouts is not None @@ -700,7 +721,14 @@ class TestProductFinancials: skip_wallet_balance_check=True, skip_one_per_day_check=True, ) - assert len(thl_lm.get_tx_filtered_by_account(account_uuid=bp_wallet.uuid)) == 3 + assert ( + len( + thl_ledger_manager.get_tx_filtered_by_account( + account_uuid=bp_wallet.uuid + ) + ) + == 3 + ) # RM the entire directories shutil.rmtree(ledger_collection.archive_path) @@ -712,7 +740,7 @@ class TestProductFinancials: pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) p1.prebuild_balance( - thl_lm=thl_lm, + thl_lm=thl_ledger_manager, ds=mnt_filepath, client=client_no_amm, ) @@ -726,7 +754,7 @@ class TestProductFinancials: assert p1.balance.available_balance == 70 p1.prebuild_payouts( - thl_lm=thl_lm, + thl_lm=thl_ledger_manager, bp_pem=brokerage_product_payout_event_manager, ) assert p1.payouts is not None @@ -743,7 +771,14 @@ class TestProductFinancials: skip_wallet_balance_check=True, skip_one_per_day_check=True, ) - assert len(thl_lm.get_tx_filtered_by_account(account_uuid=bp_wallet.uuid)) == 4 + assert ( + len( + thl_ledger_manager.get_tx_filtered_by_account( + account_uuid=bp_wallet.uuid + ) + ) + == 4 + ) # RM the entire directories shutil.rmtree(ledger_collection.archive_path) @@ -755,7 +790,7 @@ class TestProductFinancials: pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) p1.prebuild_balance( - thl_lm=thl_lm, + thl_lm=thl_ledger_manager, ds=mnt_filepath, client=client_no_amm, ) @@ -769,7 +804,7 @@ class TestProductFinancials: assert p1.balance.available_balance == 66 p1.prebuild_payouts( - thl_lm=thl_lm, + thl_lm=thl_ledger_manager, bp_pem=brokerage_product_payout_event_manager, ) assert p1.payouts is not None @@ -786,7 +821,7 @@ class TestProductBalance: @pytest.fixture def offset(self) -> str: - return "30d" + return "30D" @pytest.fixture def duration(self) -> timedelta | None: @@ -796,7 +831,7 @@ class TestProductBalance: self, product: Product, mnt_filepath: GRLDatasets, - thl_lm: ThlLedgerManager, + thl_ledger_manager: ThlLedgerManager, client_no_amm: DaskClient, delete_ledger_db: Callable[..., None], create_main_accounts: Callable[..., None], @@ -826,7 +861,7 @@ class TestProductBalance: pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) # 2. Payout and build Parquets 2nd time - payout_event_manager.set_account_lookup_table(thl_lm=thl_lm) + payout_event_manager.set_account_lookup_table(thl_lm=thl_ledger_manager) bp_payout_factory( product=product, amount=USDCent(71), @@ -840,7 +875,7 @@ class TestProductBalance: with pytest.raises(expected_exception=AssertionError) as cm: product.prebuild_balance( - thl_lm=thl_lm, ds=mnt_filepath, client=client_no_amm + thl_lm=thl_ledger_manager, ds=mnt_filepath, client=client_no_amm ) assert "Sql and Parquet Balance inconsistent" in str(cm) @@ -848,7 +883,7 @@ class TestProductBalance: self, product: Product, mnt_filepath: GRLDatasets, - thl_lm: ThlLedgerManager, + thl_ledger_manager: ThlLedgerManager, client_no_amm: DaskClient, delete_ledger_db: Callable[..., None], create_main_accounts: Callable[..., None], @@ -885,7 +920,7 @@ class TestProductBalance: # 2. Payout and build Parquets 2nd time but this payout is "now" # so it hasn't already been archived - payout_event_manager.set_account_lookup_table(thl_lm=thl_lm) + payout_event_manager.set_account_lookup_table(thl_lm=thl_ledger_manager) bp_payout_factory( product=product, amount=USDCent(71), @@ -898,7 +933,9 @@ class TestProductBalance: pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) # We just want to call this to confirm it doesn't raise. - product.prebuild_balance(thl_lm=thl_lm, ds=mnt_filepath, client=client_no_amm) + product.prebuild_balance( + thl_lm=thl_ledger_manager, ds=mnt_filepath, client=client_no_amm + ) class TestProductPOPFinancial: @@ -909,7 +946,7 @@ class TestProductPOPFinancial: @pytest.fixture def offset(self) -> str: - return "30d" + return "30D" @pytest.fixture def duration(self) -> timedelta | None: @@ -919,7 +956,7 @@ class TestProductPOPFinancial: self, product: Product, mnt_filepath: GRLDatasets, - thl_lm: ThlLedgerManager, + thl_ledger_manager: ThlLedgerManager, client_no_amm: DaskClient, delete_ledger_db: Callable[..., None], create_main_accounts: Callable[..., None], @@ -955,7 +992,7 @@ class TestProductPOPFinancial: # --- test --- assert product.pop_financial is None product.prebuild_pop_financial( - thl_lm=thl_lm, + thl_lm=thl_ledger_manager, ds=mnt_filepath, client=client_no_amm, pop_ledger=pop_ledger_merge, @@ -982,7 +1019,7 @@ class TestProductCache: @pytest.fixture def offset(self) -> str: - return "30d" + return "30D" @pytest.fixture def duration(self) -> timedelta | None: -- cgit v1.2.3 From b2450fa8be202941bb1b3770267fe6b9215d4c27 Mon Sep 17 00:00:00 2001 From: Max Nanis Date: Thu, 3 Sep 2026 11:20:32 -0700 Subject: pytest tests -vs RUNS, but obv doesn't all PASS --- generalresearch/incite/collections/base.py | 8 +- generalresearch/managers/marketplace/__init__.py | 23 ---- generalresearch/managers/marketplace/managers.py | 23 ++++ generalresearch/models/cint/question.py | 5 +- generalresearch/models/precision/question.py | 5 +- generalresearch/models/sago/question.py | 7 +- generalresearch/models/spectrum/question.py | 9 +- generalresearch/thl_django/app/test_settings.py | 2 +- .../collections/test_df_collection_item_thl_web.py | 126 ++++++++++++--------- .../test_df_collection_thl_marketplaces.py | 8 +- .../collections/test_df_collection_thl_web.py | 7 +- tests/incite/mergers/test_merge_collection.py | 8 +- tests/models/test_finance.py | 43 ++----- 13 files changed, 143 insertions(+), 131 deletions(-) create mode 100644 generalresearch/managers/marketplace/managers.py (limited to 'tests/incite/collections') diff --git a/generalresearch/incite/collections/base.py b/generalresearch/incite/collections/base.py index bf6d9ce..47bb70a 100644 --- a/generalresearch/incite/collections/base.py +++ b/generalresearch/incite/collections/base.py @@ -3,6 +3,7 @@ from __future__ import annotations import os import subprocess import time +import warnings from datetime import datetime from enum import StrEnum from sys import platform @@ -644,9 +645,14 @@ class DFCollection(CollectionBase): return items @property - def _schema(self) -> DataFrameSchema: + def type_schema(self) -> DataFrameSchema: return DFCollectionTypeSchemas[self.data_type] + @property + def _schema(self) -> DataFrameSchema: + warnings.deprecated("The _schema attribute on DFCollection is Deprecated") + return self.type_schema + # --- Methods --- def initial_load( diff --git a/generalresearch/managers/marketplace/__init__.py b/generalresearch/managers/marketplace/__init__.py index 3349434..e69de29 100644 --- a/generalresearch/managers/marketplace/__init__.py +++ b/generalresearch/managers/marketplace/__init__.py @@ -1,23 +0,0 @@ -from generalresearch.managers.cint.user_pid import CintUserPidManager -from generalresearch.managers.dynata.user_pid import DynataUserPidManager -from generalresearch.managers.innovate.user_pid import InnovateUserPidManager -from generalresearch.managers.morning.user_pid import MorningUserPidManager -from generalresearch.managers.precision.user_pid import PrecisionUserPidManager -from generalresearch.managers.prodege.user_pid import ProdegeUserPidManager -from generalresearch.managers.repdata.user_pid import RepdataUserPidManager -from generalresearch.managers.sago.user_pid import SagoUserPidManager -from generalresearch.managers.spectrum.user_pid import SpectrumUserPidManager - -_managers = [ - CintUserPidManager, - DynataUserPidManager, - InnovateUserPidManager, - MorningUserPidManager, - PrecisionUserPidManager, - ProdegeUserPidManager, - RepdataUserPidManager, - SagoUserPidManager, - SpectrumUserPidManager, -] - -USER_PID_MANAGERS = {x.SOURCE: x for x in _managers} diff --git a/generalresearch/managers/marketplace/managers.py b/generalresearch/managers/marketplace/managers.py new file mode 100644 index 0000000..3349434 --- /dev/null +++ b/generalresearch/managers/marketplace/managers.py @@ -0,0 +1,23 @@ +from generalresearch.managers.cint.user_pid import CintUserPidManager +from generalresearch.managers.dynata.user_pid import DynataUserPidManager +from generalresearch.managers.innovate.user_pid import InnovateUserPidManager +from generalresearch.managers.morning.user_pid import MorningUserPidManager +from generalresearch.managers.precision.user_pid import PrecisionUserPidManager +from generalresearch.managers.prodege.user_pid import ProdegeUserPidManager +from generalresearch.managers.repdata.user_pid import RepdataUserPidManager +from generalresearch.managers.sago.user_pid import SagoUserPidManager +from generalresearch.managers.spectrum.user_pid import SpectrumUserPidManager + +_managers = [ + CintUserPidManager, + DynataUserPidManager, + InnovateUserPidManager, + MorningUserPidManager, + PrecisionUserPidManager, + ProdegeUserPidManager, + RepdataUserPidManager, + SagoUserPidManager, + SpectrumUserPidManager, +] + +USER_PID_MANAGERS = {x.SOURCE: x for x in _managers} diff --git a/generalresearch/models/cint/question.py b/generalresearch/models/cint/question.py index 5f703ee..0f5453b 100644 --- a/generalresearch/models/cint/question.py +++ b/generalresearch/models/cint/question.py @@ -9,7 +9,8 @@ from uuid import UUID from pydantic import BaseModel, Field, field_validator, model_validator from generalresearch.models.custom_types import AwareDatetimeISO -from generalresearch.models.definitions import Source, string_utils +from generalresearch.models.definitions import Source +from generalresearch.models.string_utils import remove_nbsp from generalresearch.models.thl.profiling.marketplace import ( MarketplaceQuestion, MarketplaceUserQuestionAnswer, @@ -106,7 +107,7 @@ class CintQuestion(MarketplaceQuestion): @field_validator("question_name", "question_text", mode="after") def remove_nbsp(cls, s: str | None) -> str | None: - return string_utils.remove_nbsp(s) + return remove_nbsp(s) @model_validator(mode="after") def check_type_options_agreement(self) -> Self: diff --git a/generalresearch/models/precision/question.py b/generalresearch/models/precision/question.py index 6ed6bbd..3e39124 100644 --- a/generalresearch/models/precision/question.py +++ b/generalresearch/models/precision/question.py @@ -8,7 +8,8 @@ from typing import TYPE_CHECKING, Any, Literal from pydantic import BaseModel, Field, ValidationError, field_validator, model_validator -from generalresearch.models.definitions import Source, string_utils +from generalresearch.models.definitions import Source +from generalresearch.models.string_utils import remove_nbsp from generalresearch.models.thl.profiling.marketplace import ( MarketplaceQuestion, MarketplaceUserQuestionAnswer, @@ -94,7 +95,7 @@ class PrecisionQuestion(MarketplaceQuestion): @field_validator("question_text", mode="after") def remove_nbsp(cls, s: str | None): - return string_utils.remove_nbsp(s) + return remove_nbsp(s) @model_validator(mode="after") def check_type_options_agreement(self): diff --git a/generalresearch/models/sago/question.py b/generalresearch/models/sago/question.py index bb51d31..216b278 100644 --- a/generalresearch/models/sago/question.py +++ b/generalresearch/models/sago/question.py @@ -18,7 +18,8 @@ from pydantic import ( model_validator, ) -from generalresearch.models.definitions import MAX_INT32, Source, string_utils +from generalresearch.models.definitions import MAX_INT32, Source +from generalresearch.models.string_utils import remove_nbsp from generalresearch.models.thl.profiling.marketplace import MarketplaceQuestion if TYPE_CHECKING: @@ -55,7 +56,7 @@ class SagoQuestionOption(BaseModel): @field_validator("text", mode="after") def remove_nbsp(cls, s: str): - return string_utils.remove_nbsp(s) + return remove_nbsp(s) class SagoQuestionType(StrEnum): @@ -169,7 +170,7 @@ class SagoQuestion(MarketplaceQuestion): @field_validator("question_name", "question_text", "tags", mode="after") def remove_nbsp(cls, s: str | None): - return string_utils.remove_nbsp(s) + return remove_nbsp(s) @classmethod def from_api( diff --git a/generalresearch/models/spectrum/question.py b/generalresearch/models/spectrum/question.py index 9c9bfa0..4b854fb 100644 --- a/generalresearch/models/spectrum/question.py +++ b/generalresearch/models/spectrum/question.py @@ -18,7 +18,8 @@ from pydantic import ( model_validator, ) -from generalresearch.models.definitions import MAX_INT32, Source, string_utils +from generalresearch.models.definitions import MAX_INT32, Source +from generalresearch.models.string_utils import remove_nbsp from generalresearch.models.thl.profiling.marketplace import ( MarketplaceQuestion, ) @@ -91,7 +92,9 @@ class SpectrumQuestionOption(BaseModel): @field_validator("text", mode="after") def remove_nbsp(cls, s: str) -> str: - return string_utils.remove_nbsp(s) + res = remove_nbsp(s) + assert isinstance(res, str), "Spectrum Question Option text must be str" + return res class SpectrumQuestionType(StrEnum): @@ -206,7 +209,7 @@ class SpectrumQuestion(MarketplaceQuestion): @field_validator("question_name", "question_text", "tags", mode="after") def remove_nbsp(cls, s: str | None): - return string_utils.remove_nbsp(s) + return remove_nbsp(s) @model_validator(mode="before") @classmethod diff --git a/generalresearch/thl_django/app/test_settings.py b/generalresearch/thl_django/app/test_settings.py index 57cb9b9..4168513 100644 --- a/generalresearch/thl_django/app/test_settings.py +++ b/generalresearch/thl_django/app/test_settings.py @@ -1,7 +1,7 @@ DATABASES = { "default": { "ENGINE": "django.db.backends.postgresql", - "NAME": 'unittest-2026-09-03-a0a584', + "NAME": 'unittest-2026-09-03-38dbfc', "USER": 'jenkins', "PASSWORD": '123456789', "HOST": 'unittest-postgresql.fmt2.grl.internal', diff --git a/tests/incite/collections/test_df_collection_item_thl_web.py b/tests/incite/collections/test_df_collection_item_thl_web.py index 061c576..b8b955f 100644 --- a/tests/incite/collections/test_df_collection_item_thl_web.py +++ b/tests/incite/collections/test_df_collection_item_thl_web.py @@ -12,7 +12,8 @@ import dask.dataframe as dd import pandas as pd import pytest from dask.distributed import Client as DaskClient -from distributed import Client, Scheduler, Worker +from dask.distributed import Scheduler as DaskScheduler +from dask.distributed import Worker as DaskWorker # noinspection PyUnresolvedReferences from distributed.utils_test import ( @@ -36,7 +37,6 @@ if TYPE_CHECKING: DFCollection, DFCollectionItem, ) - from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager from generalresearch.models.thl.product import Product from generalresearch.models.thl.user import User @@ -57,12 +57,11 @@ unsupported_mock_types = { } -def combo_object() -> Generator[str]: - for x in iter_product( +def combo_object() -> Generator[tuple[DFCollectionType, str]]: + yield from iter_product( df_collections, ["15min", "45min", "1H"], - ): - yield from x + ) class TestDFCollectionItemBase: @@ -170,9 +169,12 @@ class TestDFCollectionItemMethod: def test_has_mysql( self, + df_collection_data_type: DFCollectionType, + offset: str, + duration: timedelta, + delete_df_collection: Callable[..., None], df_collection: DFCollection, thl_web_rr: PostgresConfig, - delete_df_collection: Callable[..., None], ): delete_df_collection(coll=df_collection) @@ -197,6 +199,9 @@ class TestDFCollectionItemMethod: @pytest.mark.skip def test_update_partial_archive( self, + df_collection_data_type: DFCollectionType, + offset: str, + duration: timedelta, ): # for i in collection.items: # assert i.update_partial_archive() @@ -206,8 +211,11 @@ class TestDFCollectionItemMethod: @pytest.mark.skip def test_create_partial_archive( self, + df_collection_data_type: DFCollectionType, + offset: str, + duration: timedelta, ): - assert 1 + 1 == 2 + pass def test_dict( self, @@ -262,14 +270,14 @@ class TestDFCollectionItemMethod: else: assert isinstance(df, pd.DataFrame) assert df.empty - assert set(df.columns) == set(df_collection._schema.columns.keys()) + assert set(df.columns) == set(df_collection.type_schema.columns.keys()) incite_item_factory(user=u1, item=item) df = item.from_db() assert isinstance(df, pd.DataFrame) assert not df.empty - assert set(df.columns) == set(df_collection._schema.columns.keys()) + assert set(df.columns) == set(df_collection.type_schema.columns.keys()) if df_collection.data_type == DFCollectionType.LEDGER: # The number of rows in this dataframe will change depending # on the mocking of data. It's because if the account has @@ -315,32 +323,23 @@ class TestDFCollectionItemMethod: df = item.from_mysql_standard() assert isinstance(df, pd.DataFrame) assert df.empty - assert set(df.columns) == set(df_collection._schema.columns.keys()) + assert set(df.columns) == set(df_collection.type_schema.columns.keys()) incite_item_factory(user=u1, item=item) df = item.from_mysql_standard() assert isinstance(df, pd.DataFrame) assert not df.empty - assert set(df.columns) == set(df_collection._schema.columns.keys()) + assert set(df.columns) == set(df_collection.type_schema.columns.keys()) assert df.shape[0] > 0 def test_from_mysql_ledger( self, df_collection: DFCollection, - user: User, - create_main_accounts: Callable[..., None], - offset: str, - duration: timedelta, - thl_web_rw: PostgresConfig, - thl_ledger_manager: ThlLedgerManager, - df_collection_data_type: DFCollectionType, user_factory: Callable[..., User], product: Product, - client_no_amm: DaskClient, incite_item_factory: Callable[..., None], delete_df_collection: Callable[..., None], - mnt_filepath: GRLDatasets, ): if df_collection.data_type != DFCollectionType.LEDGER: @@ -382,16 +381,10 @@ class TestDFCollectionItemMethod: def test_to_archive( self, df_collection: DFCollection, - user: User, - offset: str, - duration: timedelta, - df_collection_data_type: DFCollectionType, user_factory: Callable[..., User], product: Product, - client_no_amm: DaskClient, incite_item_factory: Callable[..., None], delete_df_collection: Callable[..., None], - mnt_filepath: GRLDatasets, ): if df_collection.data_type in unsupported_mock_types: @@ -418,14 +411,9 @@ class TestDFCollectionItemMethod: def test__to_archive( self, - df_collection_data_type: DFCollectionType, df_collection: DFCollection, user_factory: Callable[..., User], product: Product, - offset: str, - duration: timedelta, - client_no_amm: DaskClient, - user: User, incite_item_factory: Callable[..., None], delete_df_collection: Callable[..., None], mnt_filepath: GRLDatasets, @@ -492,18 +480,27 @@ class TestDFCollectionItemMethod: @pytest.mark.skip def test_to_archive_numbered_partial( self, + df_collection_data_type: DFCollectionType, + offset: str, + duration: timedelta, ): pass @pytest.mark.skip def test_initial_load( self, + df_collection_data_type: DFCollectionType, + offset: str, + duration: timedelta, ): pass @pytest.mark.skip def test_clear_corrupt_archive( self, + df_collection_data_type: DFCollectionType, + offset: str, + duration: timedelta, ): pass @@ -514,18 +511,6 @@ class TestDFCollectionItemMethod: ) class TestDFCollectionItemMethodBase: - @pytest.mark.skip - def test_path_exists( - self, - ): - pass - - @pytest.mark.skip - def test_next_numbered_path( - self, - ): - pass - @pytest.mark.skip def test_search_highest_numbered_path( self, @@ -538,12 +523,18 @@ class TestDFCollectionItemMethodBase: @pytest.mark.skip def test_tmp_filename( self, + df_collection_data_type: DFCollectionType, + offset: str, + duration: timedelta, ): pass @pytest.mark.skip def test_tmp_path( self, + df_collection_data_type: DFCollectionType, + offset: str, + duration: timedelta, ): pass @@ -640,7 +631,8 @@ class TestDFCollectionItemMethodBase: self, df_collection: DFCollection, ): - schema: DataFrameSchema = df_collection._schema + schema: DataFrameSchema = df_collection.type_schema + assert schema.metadata aa = schema.metadata[ARCHIVE_AFTER] # It shouldn't be None, it can be timedelta(seconds=0) @@ -657,6 +649,9 @@ class TestDFCollectionItemMethodBase: @pytest.mark.skip def test_set_empty( self, + df_collection_data_type: DFCollectionType, + offset: str, + duration: timedelta, ): pass @@ -687,17 +682,26 @@ class TestDFCollectionItemMethodBase: @pytest.mark.skip def test_validate_df( self, + df_collection_data_type: DFCollectionType, + offset: str, + duration: timedelta, ): pass @pytest.mark.skip def test_from_archive( self, + df_collection_data_type: DFCollectionType, + offset: str, + duration: timedelta, ): pass def test__to_dict( self, + df_collection_data_type: DFCollectionType, + offset: str, + duration: timedelta, df_collection: DFCollection, ): @@ -718,29 +722,38 @@ class TestDFCollectionItemMethodBase: @pytest.mark.skip def test_delete_partial( self, + df_collection_data_type: DFCollectionType, + offset: str, + duration: timedelta, ): pass @pytest.mark.skip def test_cleanup_partials( self, + df_collection_data_type: DFCollectionType, + offset: str, + duration: timedelta, ): pass @pytest.mark.skip def test_delete_dangling_partials( self, + df_collection_data_type: DFCollectionType, + offset: str, + duration: timedelta, ): pass @gen_cluster(client=True, nthreads=[("127.0.0.1", 1)]) -async def test_client(client, s, worker): +async def test_client(client: DaskClient, s: DaskScheduler, worker: DaskWorker): """c,s,a are all required - the secondary Worker (b) is not required""" - assert isinstance(client, Client) - assert isinstance(s, Scheduler) - assert isinstance(worker, Worker) + assert isinstance(client, DaskClient) + assert isinstance(s, DaskScheduler) + assert isinstance(worker, DaskWorker) @pytest.mark.parametrize( @@ -750,13 +763,17 @@ async def test_client(client, s, worker): @gen_cluster(client=True, nthreads=[("127.0.0.1", 1)]) @pytest.mark.anyio async def test_client_parametrize( - c, s, w, df_collection_data_type: DFCollectionType, offset: str + c: DaskClient, + s: DaskScheduler, + w: DaskWorker, + df_collection_data_type: DFCollectionType, + offset: str, ): """c,s,a are all required - the secondary Worker (b) is not required""" - assert isinstance(c, Client), f"c is not Client, it's {type(c)}" - assert isinstance(s, Scheduler), f"s is not Scheduler, it's {type(s)}" - assert isinstance(w, Worker), f"w is not Worker, it's {type(w)}" + assert isinstance(c, DaskClient), f"c is not Client, it's {type(c)}" + assert isinstance(s, DaskScheduler), f"s is not Scheduler, it's {type(s)}" + assert isinstance(w, DaskWorker), f"w is not Worker, it's {type(w)}" assert df_collection_data_type is not None assert isinstance(offset, str) @@ -923,7 +940,8 @@ class TestDFCollectionItemFunctionalTest: return u1: User = user_factory(product=product) - schema: DataFrameSchema = df_collection._schema + schema: DataFrameSchema = df_collection.type_schema + assert schema.metadata aa = schema.metadata[ARCHIVE_AFTER] assert isinstance(aa, timedelta) diff --git a/tests/incite/collections/test_df_collection_thl_marketplaces.py b/tests/incite/collections/test_df_collection_thl_marketplaces.py index 6ad0cb4..0f79b81 100644 --- a/tests/incite/collections/test_df_collection_thl_marketplaces.py +++ b/tests/incite/collections/test_df_collection_thl_marketplaces.py @@ -1,3 +1,4 @@ +from collections.abc import Generator from datetime import UTC, datetime from itertools import product from typing import TYPE_CHECKING @@ -18,8 +19,8 @@ if TYPE_CHECKING: from generalresearch.pg_helper import PostgresConfig -def combo_object(): - for x in product( +def combo_object() -> Generator[tuple[type, str]]: + yield from product( [ InnovateSurveyHistoryCollection, MorningSurveyTimeseriesCollection, @@ -27,8 +28,7 @@ def combo_object(): SpectrumSurveyTimeseriesCollection, ], ["5min", "6H", "30D"], - ): - yield from x + ) @pytest.mark.parametrize("df_coll, offset", combo_object()) diff --git a/tests/incite/collections/test_df_collection_thl_web.py b/tests/incite/collections/test_df_collection_thl_web.py index 20d7187..3c5e41b 100644 --- a/tests/incite/collections/test_df_collection_thl_web.py +++ b/tests/incite/collections/test_df_collection_thl_web.py @@ -15,8 +15,8 @@ from generalresearch.incite.collections.base import ( ) -def combo_object() -> Generator[tuple]: - for x in product( +def combo_object() -> Generator[tuple[DFCollectionType, str]]: + yield from product( [ DFCollectionType.USER, DFCollectionType.WALL, @@ -26,8 +26,7 @@ def combo_object() -> Generator[tuple]: DFCollectionType.LEDGER, ], ["30min", "1H"], - ): - yield from x + ) @pytest.mark.parametrize( diff --git a/tests/incite/mergers/test_merge_collection.py b/tests/incite/mergers/test_merge_collection.py index 3f478bd..7ed3996 100644 --- a/tests/incite/mergers/test_merge_collection.py +++ b/tests/incite/mergers/test_merge_collection.py @@ -34,8 +34,11 @@ class TestMergeCollection: def test_init( self, - mnt_filepath: GRLDatasets, merge_type: MergeType, + offset: str, + duration: timedelta, + start: datetime, + mnt_filepath: GRLDatasets, ): with pytest.raises(expected_exception=ValueError) as cm: MergeCollection(archive_path=mnt_filepath.data_src) @@ -88,6 +91,9 @@ class TestMergeCollection: def test_schema( self, + offset: str, + duration: timedelta, + start: datetime, mnt_filepath: GRLDatasets, merge_type: MergeType, ): diff --git a/tests/models/test_finance.py b/tests/models/test_finance.py index c579d78..502c596 100644 --- a/tests/models/test_finance.py +++ b/tests/models/test_finance.py @@ -7,6 +7,7 @@ from random import randint from typing import TYPE_CHECKING from uuid import uuid4 +import dask.dataframe as dd import pandas as pd import pytest from dask.distributed import Client as DaskClient @@ -32,11 +33,11 @@ if TYPE_CHECKING: from generalresearch.incite.collections.thl_web import LedgerDFCollection from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager + from generalresearch.managers.thl.product import ProductManager from generalresearch.models.thl.ledger import LedgerAccount 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() @@ -659,10 +660,9 @@ class TestBusinessBalanceInitialize: @pytest.mark.parametrize( - argnames="offset, duration", + argnames="duration", argvalues=list( iter_product( - ["12h", "2D"], [timedelta(days=2), timedelta(days=5)], ) ), @@ -671,17 +671,15 @@ class TestProductFinanceData: def test_base( self, + duration: timedelta, product: Product, user_factory: Callable[..., User], start: datetime, - duration: timedelta, thl_ledger_manager: ThlLedgerManager, session_with_tx_factory: Callable[..., None], ): # -- Build & Setup - # assert ledger_collection.start is None - # assert ledger_collection.offset is None u: User = user_factory(product=product, created=ledger_collection.start) for item in ledger_collection.items: @@ -738,15 +736,6 @@ class TestProductFinanceData: assert len(res) == len({i.time for i in res}) -@pytest.mark.parametrize( - argnames="offset, duration", - argvalues=list( - iter_product( - ["12h", "2D"], - [timedelta(days=2), timedelta(days=5)], - ) - ), -) class TestPOPFinancialData: def test_base( @@ -757,7 +746,6 @@ class TestPOPFinancialData: user_factory: Callable[..., User], product: Product, start: datetime, - duration: timedelta, create_main_accounts: Callable[..., None], session_with_tx_factory: Callable[..., Session], thl_ledger_manager: ThlLedgerManager, @@ -768,8 +756,6 @@ class TestPOPFinancialData: delete_ledger_db() create_main_accounts() delete_df_collection(coll=ledger_collection) - # assert ledger_collection.start is None - # assert ledger_collection.offset is None users = [] for _ in range(5): @@ -833,22 +819,10 @@ class TestPOPFinancialData: # 1 product: Product, multiple Users assert len(users) == len(accounts) - # We group on days, and duration is a parameter to parametrize - assert isinstance(duration, timedelta) - # -- Teardown delete_df_collection(ledger_collection) -@pytest.mark.parametrize( - argnames="offset, duration", - argvalues=list( - iter_product( - ["12h", "1D"], - [timedelta(days=2), timedelta(days=3)], - ) - ), -) class TestBusinessBalanceData: def test_from_pandas( self, @@ -859,7 +833,7 @@ class TestBusinessBalanceData: product: Product, create_main_accounts: Callable[..., None], thl_ledger_manager: ThlLedgerManager, - thl_web_rr: PostgresConfig, + product_manager: ProductManager, delete_df_collection: Callable[..., None], delete_ledger_db: Callable[..., None], session_with_tx_factory: Callable[..., Session], @@ -898,15 +872,18 @@ class TestBusinessBalanceData: columns=numerical_col_names + ["account_id"], filters=[("account_id", "in", [account.uuid])], ) + assert isinstance(ddf, dd.DataFrame) ddf = ddf.groupby("account_id").sum() df: pd.DataFrame = client_no_amm.compute(collections=ddf, sync=True) assert isinstance(df, pd.DataFrame) instance = BusinessBalances.from_pandas( - input_data=df, accounts=[account], thl_pg_config=thl_web_rr + product_manager=product_manager, + input_data=df, + accounts=[account], ) - balance: int = thl_lm.get_account_balance(account=account) + balance: int = thl_ledger_manager.get_account_balance(account=account) assert instance.balance == balance assert instance.net == balance -- cgit v1.2.3 From 044ba3a4f5f10e22d41a8fc5b810e33e7aabbf26 Mon Sep 17 00:00:00 2001 From: stuppie Date: Fri, 4 Sep 2026 11:08:53 -0600 Subject: from db fix in tests. '1H' -> '1h'. remove mysql from non-mysql items --- .../collections/test_df_collection_item_base.py | 13 ++----- .../collections/test_df_collection_item_thl_web.py | 44 ++++++++++------------ .../collections/test_df_collection_thl_web.py | 2 +- 3 files changed, 24 insertions(+), 35 deletions(-) (limited to 'tests/incite/collections') diff --git a/tests/incite/collections/test_df_collection_item_base.py b/tests/incite/collections/test_df_collection_item_base.py index 83d4973..7a8793d 100644 --- a/tests/incite/collections/test_df_collection_item_base.py +++ b/tests/incite/collections/test_df_collection_item_base.py @@ -6,6 +6,7 @@ from typing import TYPE_CHECKING import pytest from generalresearch.incite.collections.base import ( + MYSQL_ALLOWED_COLL_TYPES, DFCollection, DFCollectionItem, DFCollectionType, @@ -16,12 +17,8 @@ if TYPE_CHECKING: from generalresearch.pg_helper import PostgresConfig -df_collection_types = [e for e in DFCollectionType if e is not DFCollectionType.TEST] - - -@pytest.mark.parametrize("df_coll_type", df_collection_types) +@pytest.mark.parametrize("df_coll_type", MYSQL_ALLOWED_COLL_TYPES) class TestDFCollectionItemBase: - def test_init(self, mnt_filepath: GRLDatasets, df_coll_type: DFCollectionType): collection = DFCollection( data_type=df_coll_type, @@ -37,17 +34,15 @@ class TestDFCollectionItemBase: assert isinstance(item, DFCollectionItem) -@pytest.mark.parametrize("df_coll_type", df_collection_types) +@pytest.mark.parametrize("df_coll_type", MYSQL_ALLOWED_COLL_TYPES) class TestDFCollectionItemProperties: - @pytest.mark.skip def test_filename(self, df_coll_type: DFCollectionType): pass -@pytest.mark.parametrize("df_coll_type", df_collection_types) +@pytest.mark.parametrize("df_coll_type", MYSQL_ALLOWED_COLL_TYPES) class TestDFCollectionItemMethods: - def test_has_mysql_false( self, mnt_filepath: GRLDatasets, df_coll_type: DFCollectionType ): diff --git a/tests/incite/collections/test_df_collection_item_thl_web.py b/tests/incite/collections/test_df_collection_item_thl_web.py index b8b955f..9e00bca 100644 --- a/tests/incite/collections/test_df_collection_item_thl_web.py +++ b/tests/incite/collections/test_df_collection_item_thl_web.py @@ -60,7 +60,7 @@ unsupported_mock_types = { def combo_object() -> Generator[tuple[DFCollectionType, str]]: yield from iter_product( df_collections, - ["15min", "45min", "1H"], + ["15min", "45min", "1h"], ) @@ -167,7 +167,7 @@ class TestDFCollectionItemPropertiesBase: ) class TestDFCollectionItemMethod: - def test_has_mysql( + def test_has_postgres( self, df_collection_data_type: DFCollectionType, offset: str, @@ -180,12 +180,12 @@ class TestDFCollectionItemMethod: df_collection.pg_config = None for i in df_collection.items: - assert not i.has_mysql() + assert not i.has_postgres() # Confirm that the regular connection should work as expected df_collection.pg_config = thl_web_rr for i in df_collection.items: - assert i.has_mysql() + assert i.has_postgres() # Make a fake connection and confirm it does NOT work df_collection.pg_config = PostgresConfig( @@ -194,7 +194,7 @@ class TestDFCollectionItemMethod: statement_timeout=1, ) for i in df_collection.items: - assert not i.has_mysql() + assert not i.has_postgres() @pytest.mark.skip def test_update_partial_archive( @@ -263,29 +263,23 @@ class TestDFCollectionItemMethod: for item in df_collection.items: # Unlike .from_mysql_ledger(), .from_mysql_standard() will return # back and empty df with the correct columns in place - delete_df_collection(coll=df_collection) - df = item.from_db() if df_collection.data_type == DFCollectionType.LEDGER: - assert df is None - else: - assert isinstance(df, pd.DataFrame) - assert df.empty - assert set(df.columns) == set(df_collection.type_schema.columns.keys()) + continue + delete_df_collection(coll=df_collection) + df = item.from_postgres_standard() + assert isinstance(df, pd.DataFrame) + assert df.empty + assert set(df.columns) == set(df_collection.type_schema.columns.keys()) incite_item_factory(user=u1, item=item) - df = item.from_db() + df = item.from_postgres_standard() assert isinstance(df, pd.DataFrame) assert not df.empty assert set(df.columns) == set(df_collection.type_schema.columns.keys()) - if df_collection.data_type == DFCollectionType.LEDGER: - # The number of rows in this dataframe will change depending - # on the mocking of data. It's because if the account has - # user wallet on, then there will be more transactions for - # example. - assert df.shape[0] > 0 - def test_from_mysql_standard( + + def test_from_postgres_standard( self, df_collection_data_type: DFCollectionType, df_collection: DFCollection, @@ -333,7 +327,7 @@ class TestDFCollectionItemMethod: assert set(df.columns) == set(df_collection.type_schema.columns.keys()) assert df.shape[0] > 0 - def test_from_mysql_ledger( + def test_from_postgres_ledger( self, df_collection: DFCollection, user_factory: Callable[..., User], @@ -355,14 +349,14 @@ class TestDFCollectionItemMethod: # Okay, now continue with the actual Ledger Item tests... we need # to ensure that this item.start - item.finish range hasn't had # any prior transactions created within that range. - assert item.from_mysql_ledger() is None + assert item.from_postgres_ledger() is None # Create main accounts doesn't matter because it doesn't # add any transactions to the db - assert item.from_mysql_ledger() is None + assert item.from_postgres_ledger() is None incite_item_factory(user=u1, item=item) - df = item.from_mysql_ledger() + df = item.from_postgres_ledger() assert isinstance(df, pd.DataFrame) # Not only is this a np.int64 to int comparison, but I also know it @@ -400,7 +394,7 @@ class TestDFCollectionItemMethod: # Load up the data that we'll be using for various to_archive # methods. - df = item.from_db() + df = item.from_postgres_standard() ddf = dd.from_pandas(df, npartitions=1) # (1) Write the basic archive, the issue is that because it's diff --git a/tests/incite/collections/test_df_collection_thl_web.py b/tests/incite/collections/test_df_collection_thl_web.py index 3c5e41b..7253dd0 100644 --- a/tests/incite/collections/test_df_collection_thl_web.py +++ b/tests/incite/collections/test_df_collection_thl_web.py @@ -25,7 +25,7 @@ def combo_object() -> Generator[tuple[DFCollectionType, str]]: DFCollectionType.AUDIT_LOG, DFCollectionType.LEDGER, ], - ["30min", "1H"], + ["30min", "1h"], ) -- cgit v1.2.3 From 493f131d6ffe1629ffb3610165271c3c87b32d3b Mon Sep 17 00:00:00 2001 From: stuppie Date: Fri, 4 Sep 2026 11:31:02 -0600 Subject: incite collection tests: i think most of these work? they just take forever to run. moving on for now --- .../collections/test_df_collection_item_thl_web.py | 19 ++----------------- 1 file changed, 2 insertions(+), 17 deletions(-) (limited to 'tests/incite/collections') diff --git a/tests/incite/collections/test_df_collection_item_thl_web.py b/tests/incite/collections/test_df_collection_item_thl_web.py index 9e00bca..eeabb41 100644 --- a/tests/incite/collections/test_df_collection_item_thl_web.py +++ b/tests/incite/collections/test_df_collection_item_thl_web.py @@ -75,7 +75,6 @@ class TestDFCollectionItemBase: argnames="df_collection_data_type, offset", argvalues=combo_object() ) class TestDFCollectionItemProperties: - def test_filename( self, df_collection_data_type: DFCollectionType, @@ -97,7 +96,6 @@ class TestDFCollectionItemProperties: argnames="df_collection_data_type, offset", argvalues=combo_object() ) class TestDFCollectionItemPropertiesBase: - def test_name( self, df_collection: DFCollection, @@ -166,7 +164,6 @@ class TestDFCollectionItemPropertiesBase: ), ) class TestDFCollectionItemMethod: - def test_has_postgres( self, df_collection_data_type: DFCollectionType, @@ -278,7 +275,6 @@ class TestDFCollectionItemMethod: assert not df.empty assert set(df.columns) == set(df_collection.type_schema.columns.keys()) - def test_from_postgres_standard( self, df_collection_data_type: DFCollectionType, @@ -301,27 +297,18 @@ class TestDFCollectionItemMethod: item: DFCollectionItem if df_collection.data_type == DFCollectionType.LEDGER: - # We're using parametrize, so this If statement is just to - # confirm other Item Types will always raise an assertion - with pytest.raises(expected_exception=AssertionError) as cm: - _ = item.from_mysql_standard() - assert ( - "Can't call from_mysql_standard for Ledger DFCollectionItem" - in str(cm.value) - ) - continue # Unlike .from_mysql_ledger(), .from_mysql_standard() will return # back and empty df with the correct columns in place - df = item.from_mysql_standard() + df = item.from_postgres_standard() assert isinstance(df, pd.DataFrame) assert df.empty assert set(df.columns) == set(df_collection.type_schema.columns.keys()) incite_item_factory(user=u1, item=item) - df = item.from_mysql_standard() + df = item.from_postgres_standard() assert isinstance(df, pd.DataFrame) assert not df.empty assert set(df.columns) == set(df_collection.type_schema.columns.keys()) @@ -504,7 +491,6 @@ class TestDFCollectionItemMethod: argvalues=list(iter_product(df_collections, ["12h", "10D"], [timedelta(days=15)])), ) class TestDFCollectionItemMethodBase: - @pytest.mark.skip def test_search_highest_numbered_path( self, @@ -783,7 +769,6 @@ async def test_client_parametrize( argvalues=list(iter_product(df_collections, ["12h", "10D"], [timedelta(days=15)])), ) class TestDFCollectionItemFunctionalTest: - def test_to_archive_and_ddf( self, client_no_amm: DaskClient, -- cgit v1.2.3