From cff96d2c7ad155f3c3b5e51c76f8d46696457a09 Mon Sep 17 00:00:00 2001 From: stuppie Date: Fri, 21 Aug 2026 14:55:42 -0600 Subject: fix import (fastapi.Request -> pytest.FixtureRequest) --- test_utils/models/ledger/conftest.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) (limited to 'test_utils/models/ledger/conftest.py') diff --git a/test_utils/models/ledger/conftest.py b/test_utils/models/ledger/conftest.py index 5bef113..14a7465 100644 --- a/test_utils/models/ledger/conftest.py +++ b/test_utils/models/ledger/conftest.py @@ -7,7 +7,7 @@ from typing import TYPE_CHECKING, Callable from uuid import uuid4 import pytest -from fastapi import Request +from pytest import FixtureRequest as Request from generalresearch.currency import USDCent from generalresearch.managers.base import PostgresManager -- cgit v1.2.3 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 --- generalresearch/__init__.py | 5 +- generalresearch/config.py | 4 +- generalresearch/grliq/managers/event_plotter.py | 6 +- generalresearch/grliq/managers/forensic_data.py | 3 +- generalresearch/grliq/managers/forensic_events.py | 29 +++---- generalresearch/grliq/managers/forensic_results.py | 23 +++--- generalresearch/grliq/managers/forensic_summary.py | 14 ++-- generalresearch/grliq/models/custom_types.py | 2 +- generalresearch/grliq/models/decider.py | 4 +- generalresearch/grliq/models/events.py | 32 ++++---- generalresearch/grliq/models/forensic_data.py | 6 +- generalresearch/grliq/models/forensic_summary.py | 2 +- generalresearch/grliq/models/useragents.py | 2 +- generalresearch/grliq/utils.py | 4 +- generalresearch/grpc.py | 8 +- generalresearch/incite/base.py | 25 +++--- generalresearch/incite/defaults.py | 60 +++++++------- generalresearch/incite/mergers/__init__.py | 24 +++--- generalresearch/incite/mergers/ym_wall_summary.py | 2 +- generalresearch/incite/schemas/__init__.py | 4 +- .../mergers/foundations/enriched_task_adjust.py | 2 +- generalresearch/incite/schemas/thl_web.py | 10 +-- generalresearch/locales/__init__.py | 4 +- generalresearch/locales/timezone.py | 4 +- generalresearch/managers/cint/profiling.py | 2 +- generalresearch/managers/cint/survey.py | 4 +- generalresearch/managers/criteria.py | 4 +- generalresearch/managers/dynata/profiling.py | 2 +- generalresearch/managers/dynata/survey.py | 6 +- generalresearch/managers/events.py | 6 +- generalresearch/managers/gr/authentication.py | 12 +-- generalresearch/managers/gr/team.py | 4 +- generalresearch/managers/innovate/survey.py | 6 +- generalresearch/managers/leaderboard/manager.py | 6 +- generalresearch/managers/morning/survey.py | 8 +- generalresearch/managers/network/label.py | 6 +- generalresearch/managers/precision/survey.py | 6 +- generalresearch/managers/prodege/survey.py | 8 +- generalresearch/managers/repdata/survey.py | 6 +- generalresearch/managers/sago/survey.py | 8 +- generalresearch/managers/spectrum/survey.py | 8 +- generalresearch/managers/thl/buyer.py | 4 +- generalresearch/managers/thl/cashout_method.py | 4 +- generalresearch/managers/thl/contest_manager.py | 8 +- .../managers/thl/ledger_manager/conditions.py | 7 +- .../managers/thl/ledger_manager/ledger.py | 13 ++-- .../managers/thl/ledger_manager/thl_ledger.py | 17 ++-- generalresearch/managers/thl/product.py | 4 +- generalresearch/managers/thl/profiling/uqa.py | 4 +- generalresearch/managers/thl/profiling/user_upk.py | 4 +- generalresearch/managers/thl/session.py | 20 ++--- generalresearch/managers/thl/survey.py | 4 +- generalresearch/managers/thl/task_adjustment.py | 6 +- generalresearch/managers/thl/user_compensate.py | 4 +- .../thl/user_manager/mysql_user_manager.py | 6 +- generalresearch/managers/thl/userhealth.py | 8 +- generalresearch/managers/thl/wall.py | 12 ++- generalresearch/managers/thl/wallet/__init__.py | 4 +- generalresearch/managers/thl/wallet/tango.py | 4 +- generalresearch/models/admin/__init__.py | 14 ++-- generalresearch/models/admin/request.py | 12 +-- generalresearch/models/cint/__init__.py | 3 +- generalresearch/models/cint/question.py | 9 +-- generalresearch/models/cint/survey.py | 13 ++-- generalresearch/models/cint/task_collection.py | 4 +- generalresearch/models/custom_types.py | 15 ++-- generalresearch/models/dynata/survey.py | 11 ++- generalresearch/models/events.py | 70 +++++++---------- generalresearch/models/gr/authentication.py | 11 ++- generalresearch/models/gr/business.py | 11 ++- generalresearch/models/gr/team.py | 7 +- generalresearch/models/innovate/__init__.py | 2 +- generalresearch/models/innovate/survey.py | 14 ++-- generalresearch/models/legacy/bucket.py | 3 +- generalresearch/models/legacy/questions.py | 3 +- generalresearch/models/lucid/__init__.py | 3 +- generalresearch/models/lucid/question.py | 3 +- generalresearch/models/marketplace/summary.py | 3 +- generalresearch/models/morning/__init__.py | 2 +- generalresearch/models/morning/question.py | 13 ++-- generalresearch/models/morning/survey.py | 84 ++++++++++---------- generalresearch/models/network/mtr/execute.py | 6 +- generalresearch/models/network/nmap/parser.py | 6 +- generalresearch/models/network/nmap/result.py | 4 +- generalresearch/models/network/rdns/execute.py | 6 +- generalresearch/models/precision/__init__.py | 2 +- generalresearch/models/precision/survey.py | 69 ++++++++-------- .../models/precision/task_collection.py | 8 +- generalresearch/models/prodege/__init__.py | 3 +- generalresearch/models/prodege/question.py | 6 +- generalresearch/models/prodege/survey.py | 12 +-- generalresearch/models/prodege/task_collection.py | 4 +- generalresearch/models/repdata/survey.py | 13 ++-- generalresearch/models/sago/__init__.py | 2 +- generalresearch/models/sago/survey.py | 13 ++-- generalresearch/models/spectrum/__init__.py | 2 +- generalresearch/models/spectrum/question.py | 15 ++-- generalresearch/models/spectrum/survey.py | 19 ++--- generalresearch/models/string_utils.py | 2 +- generalresearch/models/thl/category.py | 3 +- generalresearch/models/thl/contest/__init__.py | 6 +- generalresearch/models/thl/contest/contest.py | 19 +++-- .../models/thl/contest/contest_entry.py | 10 +-- generalresearch/models/thl/contest/io.py | 4 +- generalresearch/models/thl/contest/leaderboard.py | 9 +-- generalresearch/models/thl/contest/milestone.py | 2 +- generalresearch/models/thl/contest/raffle.py | 7 +- generalresearch/models/thl/finance.py | 8 +- generalresearch/models/thl/ipinfo.py | 13 ++-- generalresearch/models/thl/leaderboard.py | 8 +- generalresearch/models/thl/ledger.py | 21 ++--- generalresearch/models/thl/ledger_example.py | 10 +-- generalresearch/models/thl/offerwall/__init__.py | 3 +- generalresearch/models/thl/offerwall/base.py | 3 +- generalresearch/models/thl/offerwall/cache.py | 10 +-- generalresearch/models/thl/payout_format.py | 2 +- generalresearch/models/thl/product.py | 8 +- .../models/thl/profiling/marketplace.py | 6 +- .../models/thl/profiling/other_option.py | 2 +- .../models/thl/profiling/upk_question.py | 13 ++-- .../models/thl/profiling/upk_question_answer.py | 9 +-- .../models/thl/profiling/user_question_answer.py | 14 ++-- generalresearch/models/thl/session.py | 27 +++---- generalresearch/models/thl/survey/__init__.py | 2 +- generalresearch/models/thl/survey/buyer.py | 4 +- generalresearch/models/thl/survey/condition.py | 3 +- generalresearch/models/thl/survey/model.py | 17 ++-- generalresearch/models/thl/survey/penalty.py | 9 +-- generalresearch/models/thl/task_adjustment.py | 6 +- generalresearch/models/thl/task_status.py | 3 +- generalresearch/models/thl/user.py | 13 ++-- generalresearch/models/thl/user_iphistory.py | 8 +- generalresearch/models/thl/user_profile.py | 3 +- generalresearch/models/thl/user_quality_event.py | 6 +- generalresearch/models/thl/userhealth.py | 21 +++-- .../models/thl/wallet/cashout_method.py | 6 +- generalresearch/models/thl/wallet/payout.py | 8 +- generalresearch/pg_helper.py | 4 +- generalresearch/sql_helper.py | 4 +- generalresearch/utils/aggregation.py | 2 +- generalresearch/utils/copying_cache.py | 2 +- generalresearch/utils/enum.py | 2 +- generalresearch/wall_status_codes/__init__.py | 6 +- generalresearch/wall_status_codes/cint.py | 6 +- generalresearch/wall_status_codes/dynata.py | 16 ++-- generalresearch/wall_status_codes/fullcircle.py | 14 ++-- generalresearch/wall_status_codes/innovate.py | 12 +-- generalresearch/wall_status_codes/lucid.py | 16 ++-- generalresearch/wall_status_codes/morning.py | 14 ++-- generalresearch/wall_status_codes/pollfish.py | 14 ++-- generalresearch/wall_status_codes/precision.py | 12 +-- generalresearch/wall_status_codes/prodege.py | 10 +-- generalresearch/wall_status_codes/repdata.py | 14 ++-- generalresearch/wall_status_codes/sago.py | 16 ++-- generalresearch/wall_status_codes/spectrum.py | 12 +-- generalresearch/wall_status_codes/wxet.py | 14 ++-- generalresearch/wxet/models/definitions.py | 34 ++++---- generalresearch/wxet/models/finish_type.py | 8 +- test_utils/conftest.py | 91 ++++++++++++++-------- test_utils/grliq/conftest.py | 8 +- test_utils/incite/collections/conftest.py | 3 +- test_utils/incite/conftest.py | 7 +- test_utils/incite/mergers/conftest.py | 2 +- test_utils/managers/conftest.py | 2 +- test_utils/managers/gr/conftest.py | 2 +- test_utils/managers/thl/conftest.py | 2 +- test_utils/managers/upk/conftest.py | 2 +- test_utils/models/conftest.py | 7 +- test_utils/models/contest/conftest.py | 7 +- test_utils/models/gr/conftest.py | 2 +- test_utils/models/ledger/conftest.py | 3 +- test_utils/models/network/conftest.py | 6 +- test_utils/models/thl/conftest.py | 17 ++-- test_utils/spectrum/conftest.py | 16 ++-- tests/grliq/models/test_forensic_data.py | 8 +- .../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 +- .../mergers/foundations/test_enriched_session.py | 8 +- .../mergers/foundations/test_enriched_wall.py | 16 ++-- .../mergers/foundations/test_user_id_product.py | 16 ++-- tests/incite/mergers/test_merge_collection.py | 8 +- tests/incite/mergers/test_pop_ledger.py | 12 +-- tests/incite/mergers/test_ym_survey_merge.py | 18 ++--- tests/incite/schemas/test_admin_responses.py | 29 ++++--- tests/incite/test_collection_base.py | 22 +++--- tests/incite/test_collection_base_item.py | 6 +- tests/managers/leaderboard.py | 14 ++-- tests/managers/test_events.py | 16 ++-- .../managers/thl/test_contest/test_leaderboard.py | 10 ++- tests/managers/thl/test_contest/test_milestone.py | 16 ++-- tests/managers/thl/test_contest/test_raffle.py | 22 +++--- tests/managers/thl/test_harmonized_uqa.py | 12 +-- tests/managers/thl/test_ledger/test_lm_accounts.py | 65 ++++++++-------- tests/managers/thl/test_ledger/test_lm_tx_locks.py | 22 +++--- .../thl/test_ledger/test_thl_lm_bp_payout.py | 48 ++++++------ tests/managers/thl/test_ledger/test_thl_lm_tx.py | 55 +++++++------ .../test_ledger/test_thl_lm_tx__user_payouts.py | 10 +-- tests/managers/thl/test_ledger/test_user_txs.py | 31 ++++---- tests/managers/thl/test_maxmind.py | 7 +- tests/managers/thl/test_profiling/test_user_upk.py | 4 +- tests/managers/thl/test_survey.py | 12 +-- tests/managers/thl/test_task_adjustment.py | 12 ++- tests/managers/thl/test_task_status.py | 16 ++-- tests/managers/thl/test_user_manager/test_base.py | 4 +- tests/managers/thl/test_user_streak.py | 16 ++-- tests/managers/thl/test_userhealth.py | 8 +- tests/managers/thl/test_wall_manager.py | 12 +-- tests/models/admin/test_report_request.py | 20 ++--- tests/models/custom_types/test_aware_datetime.py | 6 +- tests/models/custom_types/test_dsn.py | 6 +- tests/models/dynata/test_eligbility.py | 6 +- tests/models/gr/test_authentication.py | 6 +- tests/models/gr/test_base.py | 24 +++--- tests/models/gr/test_business.py | 20 ++--- tests/models/morning/test.py | 6 +- tests/models/prodege/test_survey_participation.py | 6 +- tests/models/spectrum/test_question.py | 14 ++-- tests/models/spectrum/test_survey.py | 26 +++---- tests/models/spectrum/test_survey_manager.py | 13 ++-- tests/models/test_finance.py | 10 +-- tests/models/thl/test_adjustments.py | 14 ++-- tests/models/thl/test_contest/test_contest.py | 2 +- .../thl/test_contest/test_leaderboard_contest.py | 6 +- tests/models/thl/test_ledger.py | 4 +- tests/models/thl/test_product.py | 14 ++-- tests/models/thl/test_user.py | 40 +++++----- tests/models/thl/test_user_iphistory.py | 4 +- tests/models/thl/test_wall.py | 38 ++++----- tests/models/thl/test_wall_session.py | 20 ++--- tests/test_postgres.py | 4 +- 233 files changed, 1298 insertions(+), 1379 deletions(-) (limited to 'test_utils/models/ledger/conftest.py') diff --git a/generalresearch/__init__.py b/generalresearch/__init__.py index 2100d41..604b7e2 100644 --- a/generalresearch/__init__.py +++ b/generalresearch/__init__.py @@ -1,7 +1,8 @@ import threading import time +from collections.abc import Callable from functools import wraps -from typing import Any, Callable, Optional +from typing import Any, Optional from decorator import decorator from wrapt import FunctionWrapper, ObjectProxy @@ -12,7 +13,7 @@ def retry( tries: int = 4, delay: float = 0.5, backoff: int = 2, - logger: Optional[Any] = None, + logger: Any | None = None, ) -> Callable: """ https://www.calazan.com/retry-decorator-for-python-3/ diff --git a/generalresearch/config.py b/generalresearch/config.py index 44f3db7..76e3995 100644 --- a/generalresearch/config.py +++ b/generalresearch/config.py @@ -1,7 +1,7 @@ from __future__ import annotations import os -from datetime import datetime, timezone +from datetime import UTC, datetime, timezone from pathlib import Path from pydantic import DirectoryPath, Field, MariaDBDsn, PostgresDsn, RedisDsn @@ -125,4 +125,4 @@ EXAMPLE_PRODUCT_ID = "1108d053e4fa47c5b0dbdcd03a7981e7" # AMT accounting was changed many times and txs before this date # are either missing AMT bonuses, or not accounting for hit rewards. JAMES_BILLINGS_BPID = "888dbc589987425fa846d6e2a8daed04" -JAMES_BILLINGS_TX_CUTOFF = datetime(2026, 1, 1, tzinfo=timezone.utc) +JAMES_BILLINGS_TX_CUTOFF = datetime(2026, 1, 1, tzinfo=UTC) diff --git a/generalresearch/grliq/managers/event_plotter.py b/generalresearch/grliq/managers/event_plotter.py index 54105ce..ed01d2e 100644 --- a/generalresearch/grliq/managers/event_plotter.py +++ b/generalresearch/grliq/managers/event_plotter.py @@ -11,7 +11,7 @@ from generalresearch.grliq.models.events import KeyboardEvent, MouseEvent def make_events_svg( - mouse_events: List[MouseEvent], keyboard_events: List[KeyboardEvent] + mouse_events: list[MouseEvent], keyboard_events: list[KeyboardEvent] ) -> str: if len(mouse_events) + len(keyboard_events) == 0: return f'\n' + "\n" @@ -119,8 +119,8 @@ def svg_multiline_text( def group_input_events_by_xy( - mouse_events: List[MouseEvent], keyboard_events: List[KeyboardEvent] -) -> List[tuple[tuple[float, float], List[str]]]: + mouse_events: list[MouseEvent], keyboard_events: list[KeyboardEvent] +) -> list[tuple[tuple[float, float], list[str]]]: """ Each keypress is its own event. For plotting, we want to group together all keypresses that were made when the mouse was at the same position, diff --git a/generalresearch/grliq/managers/forensic_data.py b/generalresearch/grliq/managers/forensic_data.py index 739c520..7567552 100644 --- a/generalresearch/grliq/managers/forensic_data.py +++ b/generalresearch/grliq/managers/forensic_data.py @@ -1,7 +1,8 @@ from __future__ import annotations from datetime import datetime -from typing import Any, Collection +from typing import Any +from collections.abc import Collection from psycopg import sql from pydantic import NonNegativeInt, PositiveInt diff --git a/generalresearch/grliq/managers/forensic_events.py b/generalresearch/grliq/managers/forensic_events.py index bbc6b6d..85e9620 100644 --- a/generalresearch/grliq/managers/forensic_events.py +++ b/generalresearch/grliq/managers/forensic_events.py @@ -1,6 +1,7 @@ import json from datetime import datetime -from typing import Any, Collection, Dict, List, Optional +from typing import Any, Dict, List, Optional +from collections.abc import Collection from uuid import uuid4 from psycopg import sql @@ -25,7 +26,7 @@ class GrlIqEventManager: def update_or_create_timing( self, session_uuid: UUIDStr, - timing_data: Optional[TimingData] = None, + timing_data: TimingData | None = None, ) -> PositiveInt: data = { "session_uuid": session_uuid, @@ -77,8 +78,8 @@ class GrlIqEventManager: session_uuid: UUIDStr, event_start: datetime, event_end: datetime, - events: Optional[List[Dict]] = None, - mouse_events: Optional[List[Dict]] = None, + events: list[dict] | None = None, + mouse_events: list[dict] | None = None, ) -> PositiveInt: data = { "uuid": uuid4().hex, @@ -135,14 +136,14 @@ class GrlIqEventManager: def filter( self, - select_str: Optional[str] = None, - session_uuid: Optional[str] = None, - session_uuids: Optional[Collection[str]] = None, - uuids: Optional[Collection[str]] = None, - started_since: Optional[datetime] = None, - limit: Optional[int] = None, + select_str: str | None = None, + session_uuid: str | None = None, + session_uuids: Collection[str] | None = None, + uuids: Collection[str] | None = None, + started_since: datetime | None = None, + limit: int | None = None, order_by: str = "event_start DESC", - ) -> List[Dict[str, Any]]: + ) -> list[dict[str, Any]]: if not limit: limit = 100 @@ -199,7 +200,7 @@ class GrlIqEventManager: def filter_distinct_timing( self, session_uuids: Collection[str], - ) -> List[Dict[str, Any]]: + ) -> list[dict[str, Any]]: params = {"session_uuids": list(session_uuids)} query = sql.SQL( """ @@ -229,7 +230,7 @@ class GrlIqEventManager: return res @staticmethod - def process_mouse_events(pointer_moves: List[PointerMove], events: List[Dict]): + def process_mouse_events(pointer_moves: list[PointerMove], events: list[dict]): """ In the db column 'mouse_events' we put all 'pointermove' events. Pull those out, and then any 'pointerdown' and 'pointerup' events from the @@ -274,7 +275,7 @@ class GrlIqEventManager: return mouse_events @staticmethod - def process_keyboard_events(events: List[Dict]): + def process_keyboard_events(events: list[dict]): res = [ KeyboardEvent( type=x["type"], diff --git a/generalresearch/grliq/managers/forensic_results.py b/generalresearch/grliq/managers/forensic_results.py index 30db53d..52bde99 100644 --- a/generalresearch/grliq/managers/forensic_results.py +++ b/generalresearch/grliq/managers/forensic_results.py @@ -1,5 +1,6 @@ from datetime import datetime -from typing import Any, Collection, Dict, List, Optional, Tuple +from typing import Any, Dict, List, Optional, Tuple +from collections.abc import Collection from generalresearch.grliq.models.forensic_result import ( GrlIqForensicCategoryResult, @@ -16,16 +17,16 @@ class GrlIqCategoryResultsReader: def filter_category_results( self, - session_uuid: Optional[str] = None, - fingerprint: Optional[str] = None, - phase: Optional[Phase] = None, - uuids: Optional[Collection[str]] = None, - product_ids: Optional[Collection[str]] = None, - created_since: Optional[datetime] = None, - created_between: Optional[Tuple[datetime, datetime]] = None, - user: Optional[User] = None, - limit: Optional[int] = None, - ) -> List[Dict[str, Any]]: + session_uuid: str | None = None, + fingerprint: str | None = None, + phase: Phase | None = None, + uuids: Collection[str] | None = None, + product_ids: Collection[str] | None = None, + created_since: datetime | None = None, + created_between: tuple[datetime, datetime] | None = None, + user: User | None = None, + limit: int | None = None, + ) -> list[dict[str, Any]]: """ For retrieving GrlIqForensicCategoryResult objects from db. diff --git a/generalresearch/grliq/managers/forensic_summary.py b/generalresearch/grliq/managers/forensic_summary.py index 21b7e4b..5039a38 100644 --- a/generalresearch/grliq/managers/forensic_summary.py +++ b/generalresearch/grliq/managers/forensic_summary.py @@ -2,7 +2,7 @@ from __future__ import annotations import statistics from collections import defaultdict -from datetime import datetime, timedelta, timezone +from datetime import datetime, timedelta, timezone, UTC from typing import Any, Dict, List import numpy as np @@ -27,7 +27,7 @@ from generalresearch.redis_helper import RedisConfig def calculate_category_summary( - res: List[GrlIqForensicCategoryResult], + res: list[GrlIqForensicCategoryResult], ) -> GrlIqForensicCategorySummary: totals = defaultdict(int) is_complete_count = 0 @@ -55,7 +55,7 @@ def calculate_category_summary( def calculate_checker_summary( - res: List[GrlIqCheckerResults], + res: list[GrlIqCheckerResults], ) -> GrlIqCheckerResultsSummary: totals = defaultdict(list) none_totals = defaultdict(int) @@ -85,8 +85,8 @@ def calculate_checker_summary( def calculate_timing_summary( - redis_config: RedisConfig, timing_res: List[Dict[str, Any]] -) -> Dict[str, TimingDataCountrySummary]: + redis_config: RedisConfig, timing_res: list[dict[str, Any]] +) -> dict[str, TimingDataCountrySummary]: country_median_rtts = defaultdict(list) for x in timing_res: @@ -137,7 +137,7 @@ def run_user_forensic_summary( user: User, ) -> UserForensicSummary: - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) created_between = (now - timedelta(days=90), now) select_str = "id, session_uuid, product_id, product_user_id, created_at, result_data, category_result" res = iq_dm.filter( @@ -158,7 +158,7 @@ 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] = iq_em.filter_distinct_timing(session_uuids=session_uuids) country_timing_data_summary = ( calculate_timing_summary(redis_config=redis_config, timing_res=timing_res) diff --git a/generalresearch/grliq/models/custom_types.py b/generalresearch/grliq/models/custom_types.py index c5eb93c..5c7f155 100644 --- a/generalresearch/grliq/models/custom_types.py +++ b/generalresearch/grliq/models/custom_types.py @@ -1,5 +1,5 @@ import annotated_types -from typing_extensions import Annotated +from typing import Annotated GrlIqScore = Annotated[int, annotated_types.Ge(0), annotated_types.Le(100)] GrlIqAvgScore = Annotated[float, annotated_types.Ge(0), annotated_types.Le(100)] diff --git a/generalresearch/grliq/models/decider.py b/generalresearch/grliq/models/decider.py index 4464a7f..d24e150 100644 --- a/generalresearch/grliq/models/decider.py +++ b/generalresearch/grliq/models/decider.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import datetime, timezone +from datetime import datetime, timezone, UTC from enum import Enum from pydantic import BaseModel, ConfigDict, Field @@ -35,7 +35,7 @@ class GrlIqAttemptResult(BaseModel): timestamp: AwareDatetimeISO = Field( description="When this decision was made", - default_factory=lambda: datetime.now(tz=timezone.utc), + default_factory=lambda: datetime.now(tz=UTC), ) decider: Decider = Field(description="Where this decision was made") decision: AttemptDecision = Field( diff --git a/generalresearch/grliq/models/events.py b/generalresearch/grliq/models/events.py index 7e9ebee..5daa032 100644 --- a/generalresearch/grliq/models/events.py +++ b/generalresearch/grliq/models/events.py @@ -14,7 +14,7 @@ from pydantic import ( NonNegativeInt, PositiveFloat, ) -from typing_extensions import Self +from typing import Self from generalresearch.models.custom_types import AwareDatetimeISO, IPvAnyAddressStr @@ -37,14 +37,14 @@ class Event: # in microseconds, since page load (?) timeStamp: float # optional ID of the event target (e.g.: where the mouse is hovering) - _elementId: Optional[str] = None + _elementId: str | None = None # optional tag name of the event target - _elementTagName: Optional[str] = None + _elementTagName: str | None = None # extracted coordinates for the element being interacted with - _elementBounds: Optional[Bounds] = None + _elementBounds: Bounds | None = None @classmethod - def from_dict(cls, data: Dict[str, Any]) -> Self: + def from_dict(cls, data: dict[str, Any]) -> Self: data = {k: v for k, v in data.items() if k in cls.__dataclass_fields__} bounds = data.get("_elementBounds") if bounds is not None and not isinstance(bounds, Bounds): @@ -104,13 +104,13 @@ class KeyboardEvent(Event): # "insertText", "insertCompositionText", "deleteCompositionText", # "insertFromComposition", "deleteContentBackward" - inputType: Optional[str] + inputType: str | None # e.g., 'Enter', 'a', 'Backspace' - key: Optional[str] = None + key: str | None = None # This is the actual text, if applicable - data: Optional[str] = None + data: str | None = None @property def key_text(self): @@ -159,18 +159,18 @@ class TimingData(BaseModel): """ model_config = ConfigDict(extra="forbid", validate_assignment=True) - client_rtts: List[float] = Field() - server_rtts: List[float] = Field() + client_rtts: list[float] = Field() + server_rtts: list[float] = Field() # Have to be optional for backwards-compatibility, but should always be set. - started_at: Optional[AwareDatetimeISO] = Field(default=None) - ended_at: Optional[AwareDatetimeISO] = Field(default=None) - client_ip: Optional[IPvAnyAddressStr] = Field( + started_at: AwareDatetimeISO | None = Field(default=None) + ended_at: AwareDatetimeISO | None = Field(default=None) + client_ip: IPvAnyAddressStr | None = Field( description="This comes from the websocket request's headers", examples=["72.39.217.116"], default=None, ) - server_hostname: Optional[str] = Field( + server_hostname: str | None = Field( description="The hostname of the server that handled this request", examples=["grliq-web-0"], default=None, @@ -189,7 +189,7 @@ class TimingData(BaseModel): def has_data(self): return len(self.client_rtts) > 0 and len(self.server_rtts) > 0 - def filter_rtts(self, rtts: List[float]) -> List[float]: + def filter_rtts(self, rtts: list[float]) -> list[float]: # Skip the first 5 pings, unless we have <10 pings, then get the last # 5 instead. # The first couple pings are usually outliers as they are running @@ -234,7 +234,7 @@ class TimingData(BaseModel): return rtts @property - def summarize(self) -> Optional[TimingDataSummary]: + def summarize(self) -> TimingDataSummary | None: if len(self.filtered_rtts) < 5: return None diff --git a/generalresearch/grliq/models/forensic_data.py b/generalresearch/grliq/models/forensic_data.py index f8bdd98..eda7186 100644 --- a/generalresearch/grliq/models/forensic_data.py +++ b/generalresearch/grliq/models/forensic_data.py @@ -3,7 +3,7 @@ from __future__ import annotations import hashlib import re from collections import Counter -from datetime import datetime, timedelta, timezone +from datetime import datetime, timedelta, timezone, UTC from enum import Enum from functools import cached_property from typing import Any, Literal @@ -23,7 +23,7 @@ from pydantic import ( ) from pydantic.json_schema import SkipJsonSchema from pydantic_extra_types.timezone_name import TimeZoneName -from typing_extensions import Annotated, Self +from typing import Annotated, Self from generalresearch.grliq.models import ( AUDIO_CODEC_NAMES, @@ -776,7 +776,7 @@ class GrlIqData(BaseModel): ), "product_user_id mismatch" # validate the Session's mid is "recent" - assert (datetime.now(tz=timezone.utc) - session.started) < timedelta( + assert (datetime.now(tz=UTC) - session.started) < timedelta( minutes=90 ), "expired session" diff --git a/generalresearch/grliq/models/forensic_summary.py b/generalresearch/grliq/models/forensic_summary.py index d6f46f8..5ecf1a4 100644 --- a/generalresearch/grliq/models/forensic_summary.py +++ b/generalresearch/grliq/models/forensic_summary.py @@ -228,7 +228,7 @@ class CountryRTTDistribution(BaseModel): rtt_mean: float = Field(gt=0, examples=[179.302]) rtt_max: float = Field(gt=0, examples=[890.006]) rtt_std: float = Field(gt=0, examples=[46.831]) - rtt_percentiles: List[float] = Field( + rtt_percentiles: list[float] = Field( min_length=101, max_length=101, examples=[example_rtt_percentiles] ) diff --git a/generalresearch/grliq/models/useragents.py b/generalresearch/grliq/models/useragents.py index 1953f6d..4bb340e 100644 --- a/generalresearch/grliq/models/useragents.py +++ b/generalresearch/grliq/models/useragents.py @@ -4,7 +4,7 @@ import hashlib from enum import Enum from pydantic import BaseModel, ConfigDict, Field, field_validator -from typing_extensions import Self +from typing import Self from user_agents import parse as ua_parse from user_agents.parsers import UserAgent diff --git a/generalresearch/grliq/utils.py b/generalresearch/grliq/utils.py index 95390a8..ca8c6a1 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 +from datetime import datetime, timezone, UTC from pathlib import Path from uuid import UUID @@ -16,7 +16,7 @@ def get_screenshot_fp( grliq_ss_dir_name: str = "canvas2html", create_dir_if_not_exists: bool = True, ) -> Path | None: - assert created_at.tzinfo == timezone.utc + assert created_at.tzinfo == UTC if isinstance(forensic_uuid, UUID): forensic_uuid = forensic_uuid.hex diff --git a/generalresearch/grpc.py b/generalresearch/grpc.py index 040fd26..178521e 100644 --- a/generalresearch/grpc.py +++ b/generalresearch/grpc.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta, timezone from google.protobuf.duration_pb2 import Duration from google.protobuf.timestamp_pb2 import Timestamp @@ -20,13 +20,13 @@ def timestamp_from_datetime_nullable(dt: datetime | None) -> Timestamp: def timestamp_to_datetime(ts: Timestamp) -> datetime: - return datetime.fromtimestamp(ts.seconds + ts.nanos / 1e9, tz=timezone.utc) + return datetime.fromtimestamp(ts.seconds + ts.nanos / 1e9, tz=UTC) def timestamp_to_datetime_nullable(ts: Timestamp) -> datetime | None: # grpc has no None. If a google.protobuf.Timestamp field is not set, it gets interpreted as timestamp 0 - default = datetime.fromtimestamp(0, tz=timezone.utc) - d = datetime.fromtimestamp(ts.seconds + ts.nanos / 1e9, tz=timezone.utc) + default = datetime.fromtimestamp(0, tz=UTC) + d = datetime.fromtimestamp(ts.seconds + ts.nanos / 1e9, tz=UTC) return None if d == default else d diff --git a/generalresearch/incite/base.py b/generalresearch/incite/base.py index a8088ac..6888ca2 100644 --- a/generalresearch/incite/base.py +++ b/generalresearch/incite/base.py @@ -8,7 +8,7 @@ import shutil import subprocess import warnings from concurrent.futures import Future -from datetime import datetime, timedelta, timezone +from datetime import datetime, timedelta, timezone, UTC from os import R_OK, access, listdir from os.path import isdir from os.path import join as pjoin @@ -17,9 +17,8 @@ from sys import platform from typing import ( TYPE_CHECKING, Any, - Callable, - Sequence, ) +from collections.abc import Callable, Sequence from uuid import uuid4 import dask @@ -43,7 +42,7 @@ from pydantic import ( ) from pydantic.json_schema import SkipJsonSchema from sentry_sdk import capture_exception -from typing_extensions import Self +from typing import Self from generalresearch.config import is_debug from generalresearch.incite.schemas import ( @@ -166,7 +165,7 @@ class CollectionBase(BaseModel): offset: str = Field(default="72h", max_length=5) start: AwareDatetimeISO = Field( - default=datetime(year=2018, month=1, day=1, tzinfo=timezone.utc), + default=datetime(year=2018, month=1, day=1, tzinfo=UTC), description="This is the starting point in which data will be retrieved" "in chunks from.", frozen=True, @@ -208,7 +207,7 @@ class CollectionBase(BaseModel): return self offset_total_sec = pd.Timedelta(self.offset).total_seconds() - start_total_sec = (datetime.now(tz=timezone.utc) - self.start).total_seconds() + start_total_sec = (datetime.now(tz=UTC) - self.start).total_seconds() if offset_total_sec > start_total_sec: raise ValueError("Offset must be equal to, or smaller the start timestamp") @@ -294,14 +293,14 @@ class CollectionBase(BaseModel): @property def interval_range(self) -> list[tuple[datetime, datetime]]: """closed='left', so 0 <= x < 5""" - end = self.finished or datetime.now(tz=timezone.utc).replace(microsecond=0) + end = self.finished or datetime.now(tz=UTC).replace(microsecond=0) iv_r = self._interval_range(end) return [(iv.left.to_pydatetime(), iv.right.to_pydatetime()) for iv in iv_r] @property def progress(self) -> pd.DataFrame: records = [i.to_dict() for i in self.items] - end = self.finished if self.finished else datetime.now(tz=timezone.utc) + end = self.finished if self.finished else datetime.now(tz=UTC) return pd.DataFrame.from_records(records, index=self._interval_range(end)) @property @@ -626,15 +625,15 @@ class CollectionBase(BaseModel): return res def get_items_from_year(self, year: int) -> Items: - ts = datetime(year=year, month=1, day=1, tzinfo=timezone.utc) + ts = datetime(year=year, month=1, day=1, tzinfo=UTC) return self.get_items(since=ts) def get_items_last90(self) -> Items: - ts = datetime.now(tz=timezone.utc) - timedelta(days=90) + ts = datetime.now(tz=UTC) - timedelta(days=90) return self.get_items(since=ts) def get_items_last365(self) -> Items: - ts = datetime.now(tz=timezone.utc) - timedelta(days=365) + ts = datetime.now(tz=UTC) - timedelta(days=365) return self.get_items(since=ts) @@ -642,7 +641,7 @@ class CollectionItemBase(BaseModel): # I want to intentionally keep these as native python types, and not # pandas specific types. start: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(tz=timezone.utc).replace(microsecond=0) + default_factory=lambda: datetime.now(tz=UTC).replace(microsecond=0) ) # --- Private attrs --- @@ -839,7 +838,7 @@ class CollectionItemBase(BaseModel): if archive_after is None: return False - return datetime.now(tz=timezone.utc) > self.finish + archive_after + return datetime.now(tz=UTC) > self.finish + archive_after def set_empty(self): assert ( diff --git a/generalresearch/incite/defaults.py b/generalresearch/incite/defaults.py index 421710e..5a95607 100644 --- a/generalresearch/incite/defaults.py +++ b/generalresearch/incite/defaults.py @@ -1,4 +1,4 @@ -from datetime import datetime, timezone +from datetime import datetime, timezone, UTC from generalresearch.incite.base import GRLDatasets from generalresearch.incite.collections import DFCollectionType @@ -37,69 +37,69 @@ from generalresearch.sql_helper import SqlHelper def session_df_collection( - ds: "GRLDatasets", pg_config: PostgresConfig + ds: GRLDatasets, pg_config: PostgresConfig ) -> SessionDFCollection: return SessionDFCollection( offset="37h", pg_config=pg_config, - start=datetime(year=2022, month=5, day=3, hour=12, tzinfo=timezone.utc), + start=datetime(year=2022, month=5, day=3, hour=12, tzinfo=UTC), archive_path=ds.archive_path(enum_type=DFCollectionType.SESSION), ) def wall_df_collection( - ds: "GRLDatasets", pg_config: PostgresConfig + ds: GRLDatasets, pg_config: PostgresConfig ) -> WallDFCollection: return WallDFCollection( offset="49h", pg_config=pg_config, - start=datetime(year=2022, month=5, day=3, hour=12, tzinfo=timezone.utc), + start=datetime(year=2022, month=5, day=3, hour=12, tzinfo=UTC), archive_path=ds.archive_path(enum_type=DFCollectionType.WALL), ) def user_df_collection( - ds: "GRLDatasets", pg_config: PostgresConfig + ds: GRLDatasets, pg_config: PostgresConfig ) -> UserDFCollection: return UserDFCollection( offset="73h", pg_config=pg_config, - start=datetime(year=2016, month=7, day=13, hour=1, tzinfo=timezone.utc), + start=datetime(year=2016, month=7, day=13, hour=1, tzinfo=UTC), archive_path=ds.archive_path(enum_type=DFCollectionType.USER), ) def task_df_collection( - ds: "GRLDatasets", pg_config: PostgresConfig + ds: GRLDatasets, pg_config: PostgresConfig ) -> TaskAdjustmentDFCollection: return TaskAdjustmentDFCollection( offset="48h", pg_config=pg_config, - start=datetime(year=2022, month=7, day=16, hour=0, tzinfo=timezone.utc), + start=datetime(year=2022, month=7, day=16, hour=0, tzinfo=UTC), archive_path=ds.archive_path(enum_type=DFCollectionType.TASK_ADJUSTMENT), ) def ledger_df_collection( - ds: "GRLDatasets", pg_config: PostgresConfig + ds: GRLDatasets, pg_config: PostgresConfig ) -> LedgerDFCollection: return LedgerDFCollection( 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=timezone.utc), + start=datetime(year=2018, month=3, day=14, hour=0, tzinfo=UTC), archive_path=ds.archive_path(enum_type=DFCollectionType.LEDGER), ) # --- Marketplace Specifics --- # def innovate_survey_history_collection( - ds: "GRLDatasets", sql_helper: SqlHelper + ds: GRLDatasets, sql_helper: SqlHelper ) -> InnovateSurveyHistoryCollection: return InnovateSurveyHistoryCollection( offset="12h", sql_helper=sql_helper, - start=datetime(year=2024, month=3, day=1, hour=0, tzinfo=timezone.utc), + start=datetime(year=2024, month=3, day=1, hour=0, tzinfo=UTC), archive_path=ds.archive_path( enum_type=DFCollectionType.INNOVATE_SURVEY_HISTORY ), @@ -107,12 +107,12 @@ def innovate_survey_history_collection( def morning_survey_ts_collection( - ds: "GRLDatasets", sql_helper: SqlHelper + ds: GRLDatasets, sql_helper: SqlHelper ) -> MorningSurveyTimeseriesCollection: return MorningSurveyTimeseriesCollection( offset="12h", sql_helper=sql_helper, - start=datetime(year=2024, month=3, day=1, hour=0, tzinfo=timezone.utc), + start=datetime(year=2024, month=3, day=1, hour=0, tzinfo=UTC), archive_path=ds.archive_path( enum_type=DFCollectionType.MORNING_SURVEY_TIMESERIES ), @@ -120,23 +120,23 @@ def morning_survey_ts_collection( def sago_survey_history_collection( - ds: "GRLDatasets", sql_helper: SqlHelper + ds: GRLDatasets, sql_helper: SqlHelper ) -> SagoSurveyHistoryCollection: return SagoSurveyHistoryCollection( offset="12h", sql_helper=sql_helper, - start=datetime(year=2024, month=3, day=1, hour=0, tzinfo=timezone.utc), + start=datetime(year=2024, month=3, day=1, hour=0, tzinfo=UTC), archive_path=ds.archive_path(enum_type=DFCollectionType.SAGO_SURVEY_HISTORY), ) def spectrum_survey_ts_collection( - ds: "GRLDatasets", sql_helper: SqlHelper + ds: GRLDatasets, sql_helper: SqlHelper ) -> SpectrumSurveyTimeseriesCollection: return SpectrumSurveyTimeseriesCollection( offset="12h", sql_helper=sql_helper, - start=datetime(year=2024, month=3, day=1, hour=0, tzinfo=timezone.utc), + start=datetime(year=2024, month=3, day=1, hour=0, tzinfo=UTC), archive_path=ds.archive_path( enum_type=DFCollectionType.SPECTRUM_SURVEY_TIMESERIES ), @@ -144,50 +144,50 @@ def spectrum_survey_ts_collection( # --- Mergers: Foundations --- # -def user_id_product(ds: "GRLDatasets") -> UserIdProductMerge: +def user_id_product(ds: GRLDatasets) -> UserIdProductMerge: return UserIdProductMerge( - start=datetime(year=2010, month=1, day=1, tzinfo=timezone.utc), + start=datetime(year=2010, month=1, day=1, tzinfo=UTC), offset=None, archive_path=ds.archive_path(enum_type=MergeType.USER_ID_PRODUCT), ) -def enriched_session(ds: "GRLDatasets") -> EnrichedSessionMerge: +def enriched_session(ds: GRLDatasets) -> EnrichedSessionMerge: return EnrichedSessionMerge( - start=datetime(year=2023, month=5, day=1, tzinfo=timezone.utc), + start=datetime(year=2023, month=5, day=1, tzinfo=UTC), offset="14d", archive_path=ds.archive_path(enum_type=MergeType.ENRICHED_SESSION), ) -def enriched_wall(ds: "GRLDatasets") -> EnrichedWallMerge: +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=timezone.utc), + start=datetime(year=2023, month=7, day=23, tzinfo=UTC), offset="14d", archive_path=ds.archive_path(enum_type=MergeType.ENRICHED_WALL), ) -def enriched_task_adjust(ds: "GRLDatasets") -> EnrichedTaskAdjustMerge: +def enriched_task_adjust(ds: GRLDatasets) -> EnrichedTaskAdjustMerge: return EnrichedTaskAdjustMerge( - start=datetime(year=2010, month=1, day=1, tzinfo=timezone.utc), + start=datetime(year=2010, month=1, day=1, tzinfo=UTC), offset=None, archive_path=ds.archive_path(enum_type=MergeType.ENRICHED_TASK_ADJUST), ) # --- Mergers: Others --- # -def pop_ledger(ds: "GRLDatasets") -> PopLedgerMerge: +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=timezone.utc), + start=datetime(year=2018, month=3, day=14, hour=0, tzinfo=UTC), offset="30d", archive_path=ds.archive_path(enum_type=MergeType.POP_LEDGER), ) -def ym_survey_wall(ds: "GRLDatasets") -> YMSurveyWallMerge: +def ym_survey_wall(ds: GRLDatasets) -> YMSurveyWallMerge: return YMSurveyWallMerge( start=None, offset="10D", diff --git a/generalresearch/incite/mergers/__init__.py b/generalresearch/incite/mergers/__init__.py index b9c3789..9c5276c 100644 --- a/generalresearch/incite/mergers/__init__.py +++ b/generalresearch/incite/mergers/__init__.py @@ -1,7 +1,7 @@ import logging import os.path import subprocess -from datetime import datetime, timezone +from datetime import datetime, timezone, UTC from enum import Enum from sys import platform from typing import List, Optional, Type @@ -11,7 +11,7 @@ import pandas as pd from dask.distributed import Client from pandera.pandas import DataFrameSchema from pydantic import Field, ValidationInfo, field_validator, model_validator -from typing_extensions import Self +from typing import Self from generalresearch.incite.base import CollectionBase, CollectionItemBase from generalresearch.incite.schemas import PARTITION_ON @@ -80,7 +80,7 @@ class MergeCollectionItem(CollectionItemBase): pd.Timestamp(self.start) + pd.Timedelta(self._collection.offset) ).to_pydatetime() else: - return datetime.now(tz=timezone.utc).replace(microsecond=0) + return datetime.now(tz=UTC).replace(microsecond=0) @property def filename(self) -> str: @@ -236,20 +236,20 @@ class MergeCollection(CollectionBase): # In a merge, we can set offset = None which indicates that there is only 1 # period/item where the range is 'start' until now. - offset: Optional[str] = Field(default="72h") + offset: str | None = Field(default="72h") # In a merge, we can set start = None which indicates that there is only 1 # period/item where the range is (now - offset) until now. - start: Optional[AwareDatetimeISO] = Field( + start: AwareDatetimeISO | None = Field( default=None, description="This is the starting point in which data will" " be retrieved in chunks from.", frozen=True, ) - merge_type: Optional[MergeType] = Field(default=None) - group_by: Optional[str] = Field(default=None) - grouped_key: Optional[str] = Field(default=None) - collection_item_class: Type[MergeCollectionItem] = MergeCollectionItem + merge_type: MergeType | None = Field(default=None) + group_by: str | None = Field(default=None) + grouped_key: str | None = Field(default=None) + collection_item_class: type[MergeCollectionItem] = MergeCollectionItem @model_validator(mode="after") def check_start_and_offset_nullable(self) -> Self: @@ -269,16 +269,16 @@ class MergeCollection(CollectionBase): # --- Properties --- @property - def interval_start(self) -> Optional[datetime]: + def interval_start(self) -> datetime | None: # if self.start is None and self.offset is set, the inferred start is (now - offset) if self.start is None: - return datetime.now(tz=timezone.utc).replace(microsecond=0) - pd.Timedelta( + return datetime.now(tz=UTC).replace(microsecond=0) - pd.Timedelta( self.offset ) return self.start @property - def items(self) -> List[MergeCollectionItem]: + def items(self) -> list[MergeCollectionItem]: items = [] for iv in self.interval_range: cm = self.collection_item_class(start=iv[0]) diff --git a/generalresearch/incite/mergers/ym_wall_summary.py b/generalresearch/incite/mergers/ym_wall_summary.py index 2f5995f..4816c05 100644 --- a/generalresearch/incite/mergers/ym_wall_summary.py +++ b/generalresearch/incite/mergers/ym_wall_summary.py @@ -82,7 +82,7 @@ class YMWallSummaryMergeItem(MergeCollectionItem): class YMWallSummaryMerge(MergeCollection): merge_type: Literal[MergeType.YM_WALL_SUMMARY] = MergeType.YM_WALL_SUMMARY _schema = YMWallSummarySchema - collection_item_class: Type[YMWallSummaryMergeItem] = YMWallSummaryMergeItem + collection_item_class: type[YMWallSummaryMergeItem] = YMWallSummaryMergeItem items: list[YMWallSummaryMergeItem] = Field(default_factory=list) @field_validator("offset") diff --git a/generalresearch/incite/schemas/__init__.py b/generalresearch/incite/schemas/__init__.py index c0000d1..6fc83b0 100644 --- a/generalresearch/incite/schemas/__init__.py +++ b/generalresearch/incite/schemas/__init__.py @@ -11,8 +11,8 @@ ARCHIVE_AFTER = "archive_after" PARTITION_ON = "partition_on" -def empty_dataframe_from_schema(schema: pa.DataFrameSchema) -> "pd.DataFrame": - index_names: List[str] = schema.index.names +def empty_dataframe_from_schema(schema: pa.DataFrameSchema) -> pd.DataFrame: + index_names: list[str] = schema.index.names columns = set(schema.dtypes.keys()) if len(index_names) > 1: diff --git a/generalresearch/incite/schemas/mergers/foundations/enriched_task_adjust.py b/generalresearch/incite/schemas/mergers/foundations/enriched_task_adjust.py index cc909f6..97e73a3 100644 --- a/generalresearch/incite/schemas/mergers/foundations/enriched_task_adjust.py +++ b/generalresearch/incite/schemas/mergers/foundations/enriched_task_adjust.py @@ -13,7 +13,7 @@ from generalresearch.models.thl.definitions import ( thl_task_adj_columns = THLTaskAdjustmentSchema.columns.copy() -COUNTRY_ISOS: Set[str] = Localelator().get_all_countries() +COUNTRY_ISOS: set[str] = Localelator().get_all_countries() kosovo = "xk" COUNTRY_ISOS.add(kosovo) BIGINT = 9223372036854775807 diff --git a/generalresearch/incite/schemas/thl_web.py b/generalresearch/incite/schemas/thl_web.py index 5073a18..b831b9a 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 +from datetime import datetime, timedelta, timezone, UTC import pandas as pd from pandera.pandas import Check, Column, DataFrameSchema, Index, MultiIndex @@ -105,7 +105,7 @@ THLWallSchema = DataFrameSchema( ), "started": Column( dtype=pd.DatetimeTZDtype(tz="UTC"), - checks=[Check(lambda x: x < datetime.now(tz=timezone.utc))], + checks=[Check(lambda x: x < datetime.now(tz=UTC))], nullable=False, ), "session_id": Column( @@ -205,12 +205,12 @@ THLSessionSchema = DataFrameSchema( ), "started": Column( dtype=pd.DatetimeTZDtype(tz="UTC"), - checks=[Check(lambda x: x < datetime.now(tz=timezone.utc))], + checks=[Check(lambda x: x < datetime.now(tz=UTC))], nullable=True, ), "finished": Column( dtype=pd.DatetimeTZDtype(tz="UTC"), - checks=[Check(lambda x: x < datetime.now(tz=timezone.utc))], + checks=[Check(lambda x: x < datetime.now(tz=UTC))], nullable=True, ), "loi_min": Column(dtype="Int64", nullable=True), @@ -450,7 +450,7 @@ THLTaskAdjustmentSchema = DataFrameSchema( ), "started": Column( dtype=pd.DatetimeTZDtype(tz="UTC"), - checks=[Check(lambda x: x < datetime.now(tz=timezone.utc))], + checks=[Check(lambda x: x < datetime.now(tz=UTC))], ), "source": Column( dtype=str, diff --git a/generalresearch/locales/__init__.py b/generalresearch/locales/__init__.py index 88b72e6..813966e 100644 --- a/generalresearch/locales/__init__.py +++ b/generalresearch/locales/__init__.py @@ -43,11 +43,11 @@ class Localelator: pkgutil.get_data(__name__, "country_default_lang.json") ) - def get_all_languages(self) -> Set[str]: + def get_all_languages(self) -> set[str]: # returns only the ISO 639-2/B (three-letter codes) return set(self.lang_alpha2_to_alpha3b.values()) - def get_all_countries(self) -> Set[str]: + def get_all_countries(self) -> set[str]: # returns only the ISO 3166-1 alpha-2 (two-letter codes) return set(self.country_alpha3_to_alpha2.values()) diff --git a/generalresearch/locales/timezone.py b/generalresearch/locales/timezone.py index 50d539d..fce6e0e 100644 --- a/generalresearch/locales/timezone.py +++ b/generalresearch/locales/timezone.py @@ -3,7 +3,7 @@ from typing import Optional from pytz import country_timezones -def get_default_timezone(country_iso: str) -> Optional[str]: +def get_default_timezone(country_iso: str) -> str | None: # to list all: # from pytz import country_names, country_timezones # [country_timezones.get(country) for country in country_names] @@ -72,6 +72,6 @@ country_default_locale = { } -def get_default_locale(country_iso: str) -> Optional[str]: +def get_default_locale(country_iso: str) -> str | None: # todo: "https://cdn.simplelocalize.io/public/v1/locales" to fill in the rest? return country_default_locale.get(country_iso, None) diff --git a/generalresearch/managers/cint/profiling.py b/generalresearch/managers/cint/profiling.py index d549e94..9216aa5 100644 --- a/generalresearch/managers/cint/profiling.py +++ b/generalresearch/managers/cint/profiling.py @@ -1,7 +1,7 @@ from __future__ import annotations import json -from typing import Collection +from collections.abc import Collection from generalresearch.models.cint.question import CintQuestion from generalresearch.sql_helper import SqlHelper diff --git a/generalresearch/managers/cint/survey.py b/generalresearch/managers/cint/survey.py index f80542e..da1ecd9 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 +from datetime import datetime, timezone, UTC import pymysql from pymysql import IntegrityError @@ -107,7 +107,7 @@ class CintSurveyManager(SurveyManager): return True def update(self, surveys: list[CintSurvey]) -> bool: - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) for survey in surveys: survey.last_updated = now diff --git a/generalresearch/managers/criteria.py b/generalresearch/managers/criteria.py index fe70732..b5d9830 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 datetime, timezone +from datetime import UTC, datetime, timezone from more_itertools import chunked @@ -65,7 +65,7 @@ class CriteriaManager(SqlManager, ABC): new_hashes = this_hashes - known_hashes if new_hashes: - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) values = [ condition.to_mysql() for condition in conditions diff --git a/generalresearch/managers/dynata/profiling.py b/generalresearch/managers/dynata/profiling.py index 25afbe2..661bdad 100644 --- a/generalresearch/managers/dynata/profiling.py +++ b/generalresearch/managers/dynata/profiling.py @@ -1,7 +1,7 @@ from __future__ import annotations import json -from typing import Collection +from collections.abc import Collection from generalresearch.models.dynata.question import DynataQuestion from generalresearch.sql_helper import SqlHelper diff --git a/generalresearch/managers/dynata/survey.py b/generalresearch/managers/dynata/survey.py index 372a57d..3a15c4d 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 +from datetime import datetime, timezone, UTC import pymysql from pymysql import IntegrityError @@ -102,7 +102,7 @@ class DynataSurveyManager(SurveyManager): return surveys def create(self, survey: DynataSurvey) -> bool: - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) d = survey.to_mysql() conn: pymysql.Connection = self.sql_helper.make_connection() conn.autocommit(True) @@ -123,7 +123,7 @@ class DynataSurveyManager(SurveyManager): return True def update(self, surveys: list[DynataSurvey]) -> bool: - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) update_fields = self.SURVEY_FIELDS + ["last_updated"] data = [survey.to_mysql() for survey in surveys] diff --git a/generalresearch/managers/events.py b/generalresearch/managers/events.py index 0be2bb9..f3c6a04 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 datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta, timezone from decimal import Decimal from typing import TYPE_CHECKING @@ -141,7 +141,7 @@ class UserStatsManager(RedisManager): pipe.execute() def mark_user_active(self, user: User) -> None: - now = datetime.now(tz=timezone.utc).isoformat() + now = datetime.now(tz=UTC).isoformat() r = self.redis_client pipe = r.pipeline(transaction=False) @@ -175,7 +175,7 @@ class UserStatsManager(RedisManager): # This call is idempotent; it can be called multiple times (for the # same user) and won't falsely increase a counter; it will just # reset the expiration for this user (times out after 60 min) - now = datetime.now(tz=timezone.utc).isoformat() + now = datetime.now(tz=UTC).isoformat() r = self.redis_client pipe = r.pipeline(transaction=False) diff --git a/generalresearch/managers/gr/authentication.py b/generalresearch/managers/gr/authentication.py index a402693..409cb10 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 +from datetime import datetime, timezone, UTC from typing import TYPE_CHECKING, Any from psycopg import sql @@ -29,7 +29,7 @@ class GRUserManager(PostgresManagerWithRedis): ) -> GRUser: from generalresearch.models.gr.authentication import GRUser - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) instance = GRUser.model_validate( { @@ -147,7 +147,7 @@ class GRUserManager(PostgresManagerWithRedis): for item in res: for k, v in item.items(): if isinstance(item[k], datetime): - item[k] = item[k].replace(tzinfo=timezone.utc) + item[k] = item[k].replace(tzinfo=UTC) return [GRUser.model_validate(item) for item in res] @@ -216,7 +216,7 @@ class GRTokenManager(PostgresManager): "key": api_key, "user_id": gr_user.id, "user": gr_user, - "created": datetime.now(tz=timezone.utc), + "created": datetime.now(tz=UTC), } ) @@ -251,7 +251,7 @@ class GRTokenManager(PostgresManager): token = GRToken.model_validate( { "key": binascii.hexlify(os.urandom(20)).decode(), - "created": datetime.now(tz=timezone.utc), + "created": datetime.now(tz=UTC), "user_id": user_id, } ) @@ -298,6 +298,6 @@ class GRTokenManager(PostgresManager): for k, _ in res.items(): if isinstance(res[k], datetime): - res[k] = res[k].replace(tzinfo=timezone.utc) + res[k] = res[k].replace(tzinfo=UTC) return GRToken.model_validate(res) diff --git a/generalresearch/managers/gr/team.py b/generalresearch/managers/gr/team.py index 6de82b0..393f446 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 +from datetime import datetime, timezone, UTC from typing import TYPE_CHECKING from uuid import uuid4 @@ -43,7 +43,7 @@ class MembershipManager(PostgresManager): owner=False, team_id=team.id, user_id=gr_user.id, - created=datetime.now(tz=timezone.utc), + created=datetime.now(tz=UTC), ) data = membership.model_dump(by_alias=True) diff --git a/generalresearch/managers/innovate/survey.py b/generalresearch/managers/innovate/survey.py index 7db2f49..c65b100 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 +from datetime import datetime, timezone, UTC import pymysql from pymysql import IntegrityError @@ -121,7 +121,7 @@ class InnovateSurveyManager(SurveyManager): return surveys def create(self, survey: InnovateSurvey) -> bool: - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) d = survey.to_mysql() conn: pymysql.Connection = self.sql_helper.make_connection() conn.autocommit(True) @@ -142,7 +142,7 @@ class InnovateSurveyManager(SurveyManager): return True def update(self, surveys: list[InnovateSurvey]) -> bool: - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) update_fields = self.SURVEY_FIELDS + ["updated"] data = [survey.to_mysql() for survey in surveys] diff --git a/generalresearch/managers/leaderboard/manager.py b/generalresearch/managers/leaderboard/manager.py index 0bf0312..71c3a73 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 +from datetime import datetime, timedelta, timezone, UTC from decimal import Decimal from functools import cached_property from typing import TYPE_CHECKING @@ -45,7 +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=timezone.utc).astimezone( + self.within_time_aware = datetime.now(tz=UTC).astimezone( self.timezone ) elif within_time.tzinfo is not None: @@ -57,7 +57,7 @@ class LeaderboardManager: @cached_property def period(self) -> Period: local_ts = self.within_time_aware - assert local_ts.tzinfo != timezone.utc and local_ts.tzinfo is not None + assert local_ts.tzinfo != UTC and local_ts.tzinfo is not None t = pd.Timestamp(local_ts).tz_localize(tz=None) freq_pd = { LeaderboardFrequency.WEEKLY: "W-SUN", diff --git a/generalresearch/managers/morning/survey.py b/generalresearch/managers/morning/survey.py index 2d86f0f..5fba70d 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 +from datetime import datetime, timezone, UTC import pymysql from pymysql import IntegrityError @@ -138,7 +138,7 @@ class MorningSurveyManager(SurveyManager): return bids def create(self, bid: MorningBid) -> bool: - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) d = bid.to_mysql() create_fields = self.BID_FIELDS + ["created", "updated"] @@ -179,14 +179,14 @@ class MorningSurveyManager(SurveyManager): return True def update(self, surveys: list[MorningBid]) -> None: - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) for survey in surveys: self.update_one(survey, now=now) def update_one(self, bid: MorningBid, now: datetime | None = None) -> bool: if now is None: - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) d = bid.to_mysql() d["updated"] = now diff --git a/generalresearch/managers/network/label.py b/generalresearch/managers/network/label.py index 1efe875..c5306a6 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 +from datetime import datetime, timedelta, timezone, UTC from psycopg import sql from pydantic import IPvAnyNetwork, TypeAdapter @@ -48,8 +48,8 @@ class IPLabelManager(PostgresManager): filters = [] params = {} if labeled_after or labeled_before: - time_end = labeled_before or datetime.now(tz=timezone.utc) - time_start = labeled_after or datetime(2017, 1, 1, tzinfo=timezone.utc) + time_end = labeled_before or datetime.now(tz=UTC) + time_start = labeled_after or datetime(2017, 1, 1, tzinfo=UTC) assert time_start.tzinfo.utcoffset(time_start) == timedelta(), "must be UTC" assert time_end.tzinfo.utcoffset(time_end) == timedelta(), "must be UTC" filters.append("labeled_at BETWEEN %(time_start)s AND %(time_end)s") diff --git a/generalresearch/managers/precision/survey.py b/generalresearch/managers/precision/survey.py index 6fb30f2..833cb28 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 +from datetime import datetime, timezone, UTC import pymysql from pymysql import IntegrityError @@ -104,7 +104,7 @@ class PrecisionSurveyManager(SurveyManager): return surveys def create(self, survey: PrecisionSurvey) -> bool: - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) d = survey.to_mysql() conn: pymysql.Connection = self.sql_helper.make_connection() conn.autocommit(False) @@ -151,7 +151,7 @@ class PrecisionSurveyManager(SurveyManager): return True def update_one(self, survey: PrecisionSurvey) -> bool: - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) d = survey.to_mysql() d["updated"] = now diff --git a/generalresearch/managers/prodege/survey.py b/generalresearch/managers/prodege/survey.py index f555290..750383f 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 +from datetime import datetime, timezone, UTC import pymysql @@ -93,7 +93,7 @@ class ProdegeSurveyManager(SurveyManager): return surveys def create(self, survey: ProdegeSurvey) -> bool: - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) d = survey.to_mysql() conn: pymysql.Connection = self.sql_helper.make_connection() conn.autocommit(True) @@ -114,7 +114,7 @@ class ProdegeSurveyManager(SurveyManager): return True def update(self, surveys: list[ProdegeSurvey]) -> None: - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) # Do to stupidity with bid/actual loi/ir values (see ProdegeSurvey.to_mysql), we now # can't do a bulk update b/c the fields may be different in different rows. Just do @@ -124,7 +124,7 @@ class ProdegeSurveyManager(SurveyManager): def update_one(self, survey: ProdegeSurvey, now=None) -> bool: if now is None: - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) d = survey.to_mysql() # We have to have special logic for bid/actual loi/ir here. The api is # stupid and only returns one set of them. If we just do the db diff --git a/generalresearch/managers/repdata/survey.py b/generalresearch/managers/repdata/survey.py index 2e2224f..1e1c3c6 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 +from datetime import datetime, timezone, UTC import pymysql @@ -122,7 +122,7 @@ class RepDataSurveyManager(SurveyManager): return list(surveys.values()) def create(self, survey: RepDataSurvey | RepDataSurveyHashed) -> bool: - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) d = survey.to_mysql() conn: pymysql.Connection = self.sql_helper.make_connection() conn.autocommit(True) @@ -160,7 +160,7 @@ class RepDataSurveyManager(SurveyManager): return True def update(self, surveys: list[RepDataSurveyHashed]) -> bool: - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) update_fields = self.SURVEY_FIELDS + ["last_updated"] data = [survey.to_mysql() for survey in surveys] diff --git a/generalresearch/managers/sago/survey.py b/generalresearch/managers/sago/survey.py index 325639f..2582902 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 +from datetime import datetime, timezone, UTC import pymysql from pymysql import IntegrityError @@ -101,7 +101,7 @@ class SagoSurveyManager(SurveyManager): return surveys def create(self, survey: SagoSurvey) -> bool: - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) d = survey.to_mysql() conn: pymysql.Connection = self.sql_helper.make_connection() conn.autocommit(True) @@ -122,7 +122,7 @@ class SagoSurveyManager(SurveyManager): return True def update(self, surveys: list[SagoSurvey]) -> bool: - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) update_fields = self.SURVEY_FIELDS + ["updated"] data = [survey.to_mysql() for survey in surveys] @@ -131,7 +131,7 @@ class SagoSurveyManager(SurveyManager): return True def update_field(self, survey: SagoSurvey, field: str) -> bool: - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) conn: pymysql.Connection = self.sql_helper.make_connection() value = survey.to_mysql()[field] c = conn.cursor() diff --git a/generalresearch/managers/spectrum/survey.py b/generalresearch/managers/spectrum/survey.py index 3ff2db8..58f8a1a 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 +from datetime import datetime, timezone, UTC import pymysql from pymysql import IntegrityError @@ -110,7 +110,7 @@ class SpectrumSurveyManager(SurveyManager): return surveys def create(self, survey: SpectrumSurvey) -> bool: - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) d = survey.to_mysql() conn: pymysql.Connection = self.sql_helper.make_connection() conn.autocommit(True) @@ -134,7 +134,7 @@ class SpectrumSurveyManager(SurveyManager): return True def update(self, surveys: list[SpectrumSurvey]) -> bool: - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) # Due to stupidity with bid/actual loi/ir values (last block nonsense), # we can't do a bulk update b/c the fields may be different in @@ -146,7 +146,7 @@ class SpectrumSurveyManager(SurveyManager): def update_one(self, survey: SpectrumSurvey, now: datetime | None = None) -> bool: if now is None: - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) d = survey.to_mysql() # We have to have special logic for bid/actual loi/ir here. The api diff --git a/generalresearch/managers/thl/buyer.py b/generalresearch/managers/thl/buyer.py index 04452cd..ae40bc8 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 +from datetime import datetime, timezone, UTC from generalresearch.managers.base import Permission, PostgresManager from generalresearch.models import Source @@ -45,7 +45,7 @@ class BuyerManager(PostgresManager): return None def bulk_get_or_create(self, source: Source, codes: Collection[str]) -> list[Buyer]: - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) buyers = [] params_seq = [] diff --git a/generalresearch/managers/thl/cashout_method.py b/generalresearch/managers/thl/cashout_method.py index 7365878..94617f7 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 +from datetime import datetime, timezone, UTC from typing import Any from uuid import UUID, uuid4 @@ -21,7 +21,7 @@ from generalresearch.models.thl.wallet.cashout_method import ( class CashoutMethodManager(PostgresManager): def create(self, cm: CashoutMethod) -> None: - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) query = """ INSERT INTO accounting_cashoutmethod ( id, last_updated, is_live, provider, diff --git a/generalresearch/managers/thl/contest_manager.py b/generalresearch/managers/thl/contest_manager.py index 517f677..286de3d 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 +from datetime import datetime, timezone, UTC from typing import Any, Literal, cast from uuid import UUID @@ -199,10 +199,10 @@ class ContestBaseManager(PostgresManager): params["contest_type"] = contest_type.value filters.append("contest_type = %(contest_type)s") if starts_at_before is True: - params["starts_at"] = datetime.now(tz=timezone.utc) + params["starts_at"] = datetime.now(tz=UTC) filters.append("starts_at < %(starts_at)s") elif starts_at_before: - assert starts_at_before.tzinfo == timezone.utc + assert starts_at_before.tzinfo == UTC params["starts_at"] = starts_at_before filters.append("starts_at < %(starts_at)s") if name is not None: @@ -822,7 +822,7 @@ class MilestoneContestManager(ContestBaseManager): if decision: contest.update( status=ContestStatus.COMPLETED, - ended_at=datetime.now(tz=timezone.utc), + ended_at=datetime.now(tz=UTC), end_reason=reason, ) self.end_milestone_contest(contest) diff --git a/generalresearch/managers/thl/ledger_manager/conditions.py b/generalresearch/managers/thl/ledger_manager/conditions.py index a457f30..f79a0f0 100644 --- a/generalresearch/managers/thl/ledger_manager/conditions.py +++ b/generalresearch/managers/thl/ledger_manager/conditions.py @@ -1,8 +1,9 @@ from __future__ import annotations import logging -from datetime import datetime, timedelta, timezone -from typing import TYPE_CHECKING, Callable +from datetime import datetime, timedelta, timezone, UTC +from typing import TYPE_CHECKING +from collections.abc import Callable from generalresearch.config import JAMES_BILLINGS_BPID, JAMES_BILLINGS_TX_CUTOFF from generalresearch.currency import USDCent @@ -73,7 +74,7 @@ def generate_condition_bp_payout( skip_one_per_day_check: bool = False, skip_wallet_balance_check: bool = False, ) -> Callable[..., tuple[bool, str]]: - created = datetime.now(tz=timezone.utc) + created = datetime.now(tz=UTC) def _condition( lm: ThlLedgerManager, diff --git a/generalresearch/managers/thl/ledger_manager/ledger.py b/generalresearch/managers/thl/ledger_manager/ledger.py index 864f1dd..00fac27 100644 --- a/generalresearch/managers/thl/ledger_manager/ledger.py +++ b/generalresearch/managers/thl/ledger_manager/ledger.py @@ -3,8 +3,9 @@ from __future__ import annotations import logging from collections import defaultdict from collections.abc import Collection -from datetime import datetime, timedelta, timezone -from typing import Any, Callable +from datetime import datetime, timedelta, timezone, UTC +from typing import Any +from collections.abc import Callable from uuid import UUID import redis @@ -104,8 +105,8 @@ class LedgerManagerBasePostgres(PostgresManager, RedisManager): filters = [] params = {} if time_start or time_end: - time_end = time_end or datetime.now(tz=timezone.utc) - time_start = time_start or datetime(2017, 1, 1, tzinfo=timezone.utc) + time_end = time_end or datetime.now(tz=UTC) + time_start = time_start or datetime(2017, 1, 1, tzinfo=UTC) assert time_start.tzinfo.utcoffset(time_start) == timedelta() assert time_end.tzinfo.utcoffset(time_end) == timedelta() filters.append("lt.created BETWEEN %(time_start)s AND %(time_end)s") @@ -152,7 +153,7 @@ class LedgerTransactionManager(LedgerManagerBasePostgres): if metadata is None: metadata = dict() if created is None: - created = datetime.now(tz=timezone.utc) + created = datetime.now(tz=UTC) t = LedgerTransaction( created=created, @@ -449,7 +450,7 @@ class LedgerTransactionManager(LedgerManagerBasePostgres): id=row["transaction_id"], entries=entries, metadata=metadata, - created=row["created"].replace(tzinfo=timezone.utc), + created=row["created"].replace(tzinfo=UTC), ext_description=row["ext_description"], tag=row["tag"], ) diff --git a/generalresearch/managers/thl/ledger_manager/thl_ledger.py b/generalresearch/managers/thl/ledger_manager/thl_ledger.py index 3119a46..c9203f6 100644 --- a/generalresearch/managers/thl/ledger_manager/thl_ledger.py +++ b/generalresearch/managers/thl/ledger_manager/thl_ledger.py @@ -2,9 +2,10 @@ from __future__ import annotations import logging from collections.abc import Collection -from datetime import datetime, timedelta, timezone +from datetime import datetime, timedelta, timezone, UTC from decimal import Decimal -from typing import TYPE_CHECKING, Callable +from typing import TYPE_CHECKING +from collections.abc import Callable from uuid import UUID import numpy as np @@ -244,10 +245,10 @@ class ThlLedgerManager(LedgerManager): time_end: datetime | None = None, ): if time_start is None: - time_start = datetime(year=2017, month=1, day=1, tzinfo=timezone.utc) + time_start = datetime(year=2017, month=1, day=1, tzinfo=UTC) if time_end is None: - time_end = datetime.now(tz=timezone.utc) + time_end = datetime.now(tz=UTC) assert all( isinstance(item, str) for item in account_uuids @@ -798,7 +799,7 @@ class ThlLedgerManager(LedgerManager): skip_flag_check = True assert ( - datetime.now(tz=timezone.utc) > created + datetime.now(tz=UTC) > created ), "created cannot be in the future" f = lambda: self.create_tx_bp_payout_( product=product, @@ -904,7 +905,7 @@ class ThlLedgerManager(LedgerManager): for retry of a failed previous call. """ assert ( - datetime.now(tz=timezone.utc) > created + datetime.now(tz=UTC) > created ), "created cannot be in the future" assert isinstance(amount, int) assert isinstance(amount, USDCent) @@ -1837,7 +1838,7 @@ class ThlLedgerManager(LedgerManager): user.product.user_wallet_config.enabled ), "Can't get wallet balance on non-managed account." - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) wallet = self.get_account_or_create_user_wallet(user) if user.product_id == JAMES_BILLINGS_BPID: assert since_days_ago is None @@ -1867,7 +1868,7 @@ class ThlLedgerManager(LedgerManager): After 3 days, about 25% of all "future" recons have happened, 7 days: 50%, 14 days: 75%, till end of next month: 100%. """ - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) # The redeemable balance can NOT ever be more than the actual user_wallet_balance # Sum up the redeemable amount for each complete diff --git a/generalresearch/managers/thl/product.py b/generalresearch/managers/thl/product.py index 46280b8..3b92361 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 +from datetime import datetime, timezone, UTC from decimal import Decimal from threading import Lock from typing import TYPE_CHECKING @@ -293,7 +293,7 @@ class ProductManager(PostgresManager): UserWalletConfig, ) - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) # TODO: Add product_id, and possibly name uniqueness validation to the # pydantic model definition itself. The create manager doesn't need diff --git a/generalresearch/managers/thl/profiling/uqa.py b/generalresearch/managers/thl/profiling/uqa.py index cbe39e7..3854333 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 +from datetime import datetime, timedelta, timezone, UTC from generalresearch.managers.base import PostgresManagerWithRedis from generalresearch.models.thl.profiling.user_question_answer import ( @@ -128,7 +128,7 @@ class UQAManager(PostgresManagerWithRedis): def get_from_db(self, user: User) -> list[UserQuestionAnswer]: logger.info(f"get_uqa_from_db: {user.user_id}") # Only store the latest row per question_id. We don't need it multiple times. - since = datetime.now(tz=timezone.utc) - timedelta(days=30) + since = datetime.now(tz=UTC) - timedelta(days=30) # We CAN use the RR, b/c either # 1) the cache expired and the user hasn't sent an answer recently diff --git a/generalresearch/managers/thl/profiling/user_upk.py b/generalresearch/managers/thl/profiling/user_upk.py index 1820103..ee36124 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 +from datetime import datetime, timedelta, timezone, UTC from typing import Any from uuid import UUID @@ -56,7 +56,7 @@ class UserUpkManager(PostgresManagerWithRedis): return res def get_user_upk_mysql(self, user_id: int) -> list[UpkQuestionAnswer]: - since = datetime.now(tz=timezone.utc) - timedelta(days=89) + since = datetime.now(tz=UTC) - timedelta(days=89) query = """ SELECT diff --git a/generalresearch/managers/thl/session.py b/generalresearch/managers/thl/session.py index 746a518..771f882 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 +from datetime import datetime, timedelta, timezone, UTC from decimal import Decimal from typing import Any from uuid import UUID, uuid4 @@ -188,7 +188,7 @@ class SessionManager(PostgresManager): # validation errors. There doesn't seem to be a clean way of doing this. # model_copy with update doesn't trigger the validators, so we # re-run model_validate after - finished = finished if finished else datetime.now(tz=timezone.utc) + finished = finished if finished else datetime.now(tz=UTC) session.update( **{ "status": status, @@ -451,13 +451,13 @@ class SessionManager(PostgresManager): params = {} if started_before or started_after: - started_after = started_after or datetime(2017, 1, 1, tzinfo=timezone.utc) - started_before = started_before or datetime.now(tz=timezone.utc) + started_after = started_after or datetime(2017, 1, 1, tzinfo=UTC) + started_before = started_before or datetime.now(tz=UTC) assert ( - started_after.tzinfo == timezone.utc + started_after.tzinfo == UTC ), "started_after must be tz-aware as UTC" assert ( - started_before.tzinfo == timezone.utc + started_before.tzinfo == UTC ), "started_before must be tz-aware as UTC" assert ( started_after < started_before @@ -467,13 +467,13 @@ class SessionManager(PostgresManager): params["started_before"] = started_before if adjusted_before or adjusted_after: - adjusted_after = adjusted_after or datetime(2017, 1, 1, tzinfo=timezone.utc) - adjusted_before = adjusted_before or datetime.now(tz=timezone.utc) + adjusted_after = adjusted_after or datetime(2017, 1, 1, tzinfo=UTC) + adjusted_before = adjusted_before or datetime.now(tz=UTC) assert ( - adjusted_after.tzinfo == timezone.utc + adjusted_after.tzinfo == UTC ), "adjusted_after must be tz-aware as UTC" assert ( - adjusted_before.tzinfo == timezone.utc + adjusted_before.tzinfo == UTC ), "adjusted_before must be tz-aware as UTC" assert ( adjusted_after < adjusted_before diff --git a/generalresearch/managers/thl/survey.py b/generalresearch/managers/thl/survey.py index c7671ce..966e96d 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 datetime, timezone +from datetime import UTC, datetime, timezone from typing import Any import pandas as pd @@ -544,7 +544,7 @@ class SurveyStatManager(PostgresManager): VALUES ({values_str}) ON CONFLICT ({unique_cols_str}) DO UPDATE SET {update_str};""" - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) params = [ss.model_dump_sql() | {"updated_at": now} for ss in survey_stats] with self.pg_config.make_connection() as conn: diff --git a/generalresearch/managers/thl/task_adjustment.py b/generalresearch/managers/thl/task_adjustment.py index e4736d4..3ec3d41 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 datetime, timezone +from datetime import UTC, datetime, timezone from decimal import Decimal from functools import cached_property @@ -120,8 +120,8 @@ class TaskAdjustmentManager(PostgresManager): CHANGES/DELTAS as just communicated by the marketplace, not what the Wall's final adjusted_* will be. """ - alert_time = alert_time or datetime.now(tz=timezone.utc) - assert alert_time.tzinfo == timezone.utc + alert_time = alert_time or datetime.now(tz=UTC) + assert alert_time.tzinfo == UTC wall = self.wall_manager.get_from_uuid(wall_uuid) session = self.session_manager.get_from_id(wall.session_id) diff --git a/generalresearch/managers/thl/user_compensate.py b/generalresearch/managers/thl/user_compensate.py index 543de87..c6c0747 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 datetime, timezone +from datetime import UTC, datetime, timezone from decimal import Decimal from uuid import uuid4 @@ -28,7 +28,7 @@ def user_compensate( pg_config = ledger_manager.pg_config redis_client = ledger_manager.redis_client - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) assert type(amount_int) is int user.prefetch_product(pg_config=pg_config) assert ( diff --git a/generalresearch/managers/thl/user_manager/mysql_user_manager.py b/generalresearch/managers/thl/user_manager/mysql_user_manager.py index dbed5de..7931ba4 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 +from datetime import datetime, timezone, UTC from functools import lru_cache from uuid import uuid4 @@ -26,7 +26,7 @@ class MysqlUserManager: def _set_last_seen(self, user: User) -> None: # Don't call this directly. Use UserManager.set_last_seen() assert not self.is_read_replica - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) self.pg_config.execute_write( """ UPDATE thl_user @@ -118,7 +118,7 @@ class MysqlUserManager: if not self.product_id_exists(product_id=product_id): raise ValueError(f"userprofile_brokerageproduct not found: {product_id}") - now = created or datetime.now(tz=timezone.utc) + now = created or datetime.now(tz=UTC) user_uuid = uuid4().hex params = { "user_uuid": user_uuid, diff --git a/generalresearch/managers/thl/userhealth.py b/generalresearch/managers/thl/userhealth.py index fe2163f..0bc60ec 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 datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta, timezone from itertools import zip_longest from typing import Any @@ -210,7 +210,7 @@ class IPRecordManager(PostgresManagerWithRedis): data = { "user_id": user_id, "ip": ipaddress.ip_address(ip).exploded, - "created": datetime.now(tz=timezone.utc), + "created": datetime.now(tz=UTC), } fips_cols = [ @@ -335,7 +335,7 @@ class AuditLogManager(PostgresManager): al = AuditLog.model_validate( { "user_id": user_id, - "created": datetime.now(tz=timezone.utc), + "created": datetime.now(tz=UTC), "level": level, "event_type": event_type, "event_msg": event_msg, @@ -495,7 +495,7 @@ class AuditLogManager(PostgresManager): ), "must pass user_id as int" if created_after is None: - created_after = datetime.now(tz=timezone.utc) - timedelta(days=7) + created_after = datetime.now(tz=UTC) - timedelta(days=7) filters = [ "user_id = ANY(%(user_ids)s)", diff --git a/generalresearch/managers/thl/wall.py b/generalresearch/managers/thl/wall.py index c2eb821..7e413d7 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 datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta, timezone from decimal import Decimal from functools import cached_property from uuid import uuid4 @@ -360,12 +360,10 @@ class WallManager(PostgresManager): params = {} filters.append("user_id = %(user_id)s") params["user_id"] = user_id - default_started = datetime.now(tz=timezone.utc) - timedelta(days=90) + default_started = datetime.now(tz=UTC) - timedelta(days=90) started_after = started_after or default_started - started_before = started_before or datetime.now(tz=timezone.utc) - assert ( - started_before.tzinfo == timezone.utc - ), "started_before must be tz-aware as UTC" + started_before = started_before or datetime.now(tz=UTC) + assert started_before.tzinfo == UTC, "started_before must be tz-aware as UTC" assert ( started_after < started_before ), "started_after must be before started_before" @@ -412,7 +410,7 @@ class WallManager(PostgresManager): started_before: datetime | None = None, order_by: str | None = "-started", ) -> list[WallAttempt]: - started_before = started_before or datetime.now(tz=timezone.utc) + started_before = started_before or datetime.now(tz=UTC) res = [] page = 1 while True: diff --git a/generalresearch/managers/thl/wallet/__init__.py b/generalresearch/managers/thl/wallet/__init__.py index b063e54..9f3ae85 100644 --- a/generalresearch/managers/thl/wallet/__init__.py +++ b/generalresearch/managers/thl/wallet/__init__.py @@ -32,8 +32,8 @@ def manage_pending_cashout( user_ip_history_manager: UserIpHistoryManager, user_manager: UserManager, ledger_manager: ThlLedgerManager, - order_data: Optional[Union[Dict[str, Any], CashMailOrderData]] = None, - tango_client: Optional[TangoClient] = None, + order_data: dict[str, Any] | CashMailOrderData | None = None, + tango_client: TangoClient | None = None, ) -> UserPayoutEvent: """ Called by a UI actions performed by Todd. This rejects/approves/cancels diff --git a/generalresearch/managers/thl/wallet/tango.py b/generalresearch/managers/thl/wallet/tango.py index 2f2dc52..445719d 100644 --- a/generalresearch/managers/thl/wallet/tango.py +++ b/generalresearch/managers/thl/wallet/tango.py @@ -65,8 +65,8 @@ def complete_tango_order( def create_tango_order( - request_data: Dict[str, Any], ref_id: str, tango_client: TangoClient -) -> Dict[str, Any]: + request_data: dict[str, Any], ref_id: str, tango_client: TangoClient +) -> dict[str, Any]: """ Create a tango gift card order. Throws exception if anything is not right. diff --git a/generalresearch/models/admin/__init__.py b/generalresearch/models/admin/__init__.py index ebe839a..ad6302b 100644 --- a/generalresearch/models/admin/__init__.py +++ b/generalresearch/models/admin/__init__.py @@ -1,14 +1,14 @@ from __future__ import annotations -from datetime import datetime, timezone +from datetime import UTC, datetime, timezone import pandas as pd from dateutil import relativedelta def get_date_list(start_datetime: datetime, end_datetime: datetime | None = None): - start_datetime = start_datetime.replace(tzinfo=timezone.utc) - end_datetime = end_datetime if end_datetime else datetime.now(tz=timezone.utc) + start_datetime = start_datetime.replace(tzinfo=UTC) + end_datetime = end_datetime if end_datetime else datetime.now(tz=UTC) return ( pd.date_range(start_datetime, end_datetime, freq="1D") .strftime("%Y-%m-%d") @@ -22,7 +22,7 @@ def year_start(periods_ago: int = 6) -> datetime: years. Goal is to provide a simple way to know when to do filters from """ - n: datetime = datetime.now(tz=timezone.utc) + n: datetime = datetime.now(tz=UTC) d: datetime = n - relativedelta.relativedelta(years=periods_ago) return d.replace(month=1, day=1, hour=0, minute=0, second=0, microsecond=0) @@ -33,7 +33,7 @@ def month_start(periods_ago: int = 6) -> datetime: months. Goal is to provide a simple way to know when to do filters from """ - n: datetime = datetime.now(tz=timezone.utc) + n: datetime = datetime.now(tz=UTC) d: datetime = n - relativedelta.relativedelta(months=periods_ago) return d.replace(day=1, hour=0, minute=0, second=0, microsecond=0) @@ -44,7 +44,7 @@ def day_start(periods_ago: int = 6) -> datetime: days. Goal is to provide a simple way to know when to do filters from """ - n: datetime = datetime.now(tz=timezone.utc) + n: datetime = datetime.now(tz=UTC) d: datetime = n - relativedelta.relativedelta(days=periods_ago) return d.replace(hour=0, minute=0, second=0, microsecond=0) @@ -55,6 +55,6 @@ def hour_start(periods_ago: int = 6) -> datetime: hours. Goal is to provide a simple way to know when to do filters from """ - n: datetime = datetime.now(tz=timezone.utc) + n: datetime = datetime.now(tz=UTC) d: datetime = n - relativedelta.relativedelta(hours=periods_ago) return d.replace(minute=0, second=0, microsecond=0) diff --git a/generalresearch/models/admin/request.py b/generalresearch/models/admin/request.py index 67bd263..2d68de1 100644 --- a/generalresearch/models/admin/request.py +++ b/generalresearch/models/admin/request.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta, timezone from enum import Enum from typing import Literal @@ -25,9 +25,9 @@ class ReportRequest(BaseModel): index1: str = Field(default="product_id") start: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(tz=timezone.utc) - timedelta(days=14) + default_factory=lambda: datetime.now(tz=UTC) - timedelta(days=14) ) - end: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=timezone.utc)) + end: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC)) interval: Literal["5min", "15min", "1h", "6h", "12h", "1d"] = "1h" include_open_bucket: bool = Field(default=True) @@ -35,7 +35,7 @@ class ReportRequest(BaseModel): @computed_field( title="Start floor", description="The datetime that this report starts from", - examples=[datetime(year=2025, month=5, day=1, tzinfo=timezone.utc)], + examples=[datetime(year=2025, month=5, day=1, tzinfo=UTC)], return_type=datetime, ) @property @@ -60,7 +60,7 @@ class ReportRequest(BaseModel): @model_validator(mode="after") def check_start_end_tz(self): - assert self.start.tzinfo == self.end.tzinfo == timezone.utc + assert self.start.tzinfo == self.end.tzinfo == UTC return self @model_validator(mode="after") @@ -150,7 +150,7 @@ class ReportRequest(BaseModel): start=self.ts_start_floor, end=self.ts_end, freq=self.interval, - tz=timezone.utc, + tz=UTC, ) def bucket_ranges(self) -> list[tuple[pd.Timestamp, pd.Timestamp]]: diff --git a/generalresearch/models/cint/__init__.py b/generalresearch/models/cint/__init__.py index 2c1be7e..d2713ab 100644 --- a/generalresearch/models/cint/__init__.py +++ b/generalresearch/models/cint/__init__.py @@ -1,5 +1,6 @@ +from typing import Annotated + from pydantic import Field -from typing_extensions import Annotated CintQuestionIdType = Annotated[ str, Field(min_length=1, max_length=16, pattern=r"^[0-9]+$") diff --git a/generalresearch/models/cint/question.py b/generalresearch/models/cint/question.py index 1ac9eea..5959141 100644 --- a/generalresearch/models/cint/question.py +++ b/generalresearch/models/cint/question.py @@ -1,13 +1,12 @@ from __future__ import annotations import json -from datetime import datetime, timezone +from datetime import UTC, datetime, timezone from enum import Enum -from typing import TYPE_CHECKING, Any, Literal +from typing import TYPE_CHECKING, Any, Literal, Self from uuid import UUID from pydantic import BaseModel, Field, field_validator, model_validator -from typing_extensions import Self from generalresearch.models import Source, string_utils from generalresearch.models.cint import CintQuestionIdType @@ -151,7 +150,7 @@ class CintQuestion(MarketplaceQuestion): options = None created_at = datetime.strptime( d["create_date"], "%Y-%m-%dT%H:%M:%S%z" - ).astimezone(timezone.utc) + ).astimezone(UTC) if d.get("question_options"): options = [ @@ -189,7 +188,7 @@ class CintQuestion(MarketplaceQuestion): ] if d.get("created_at"): - d["created_at"] = d["created_at"].replace(tzinfo=timezone.utc) + d["created_at"] = d["created_at"].replace(tzinfo=UTC) return cls( question_id=d["question_id"], diff --git a/generalresearch/models/cint/survey.py b/generalresearch/models/cint/survey.py index 56384e3..01615f6 100644 --- a/generalresearch/models/cint/survey.py +++ b/generalresearch/models/cint/survey.py @@ -2,9 +2,9 @@ from __future__ import annotations import json import logging -from datetime import datetime, timezone +from datetime import UTC, datetime, timezone from decimal import Decimal -from typing import Any, Literal, Type +from typing import Annotated, Any, Literal, Self, Type from more_itertools import flatten from pydantic import ( @@ -15,7 +15,6 @@ from pydantic import ( computed_field, model_validator, ) -from typing_extensions import Annotated, Self from generalresearch.locales import Localelator from generalresearch.models import Source, TaskCalculationType @@ -291,7 +290,7 @@ class CintSurvey(MarketplaceTask): return data @property - def condition_model(self) -> Type[MarketplaceCondition]: + def condition_model(self) -> type[MarketplaceCondition]: return CintCondition @property @@ -390,7 +389,7 @@ class CintSurvey(MarketplaceTask): d["conditions"][q.criterion_hash] = q d["quotas"] = quotas - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) d["created_at"] = now d["last_updated"] = now @@ -418,8 +417,8 @@ class CintSurvey(MarketplaceTask): @classmethod def from_mysql(cls, d: Dict[str, Any]) -> Self: - d["created_at"] = d["created_at"].replace(tzinfo=timezone.utc) - d["last_updated"] = d["last_updated"].replace(tzinfo=timezone.utc) + d["created_at"] = d["created_at"].replace(tzinfo=UTC) + d["last_updated"] = d["last_updated"].replace(tzinfo=UTC) d["qualifications"] = json.loads(d["qualifications"]) d["used_question_ids"] = json.loads(d["used_question_ids"]) d["quotas"] = json.loads(d["quotas"]) diff --git a/generalresearch/models/cint/task_collection.py b/generalresearch/models/cint/task_collection.py index 5d39090..efb6e95 100644 --- a/generalresearch/models/cint/task_collection.py +++ b/generalresearch/models/cint/task_collection.py @@ -35,8 +35,8 @@ CintTaskCollectionSchema = DataFrameSchema( "bid_ir": Column(float, Check.between(0, 1), nullable=True), "created_at": Column(dtype=pd.DatetimeTZDtype(tz="UTC")), "last_updated": Column(dtype=pd.DatetimeTZDtype(tz="UTC")), - "used_question_ids": Column(List[str]), - "all_hashes": Column(List[str]), # set >> list for column support + "used_question_ids": Column(list[str]), + "all_hashes": Column(list[str]), # set >> list for column support }, checks=[], index=Index( diff --git a/generalresearch/models/custom_types.py b/generalresearch/models/custom_types.py index 84bf8e3..9346064 100644 --- a/generalresearch/models/custom_types.py +++ b/generalresearch/models/custom_types.py @@ -3,8 +3,8 @@ from __future__ import annotations import json import re import sys as _sys -from datetime import datetime, timedelta, timezone -from typing import Any, Literal +from datetime import UTC, datetime, timedelta, timezone +from typing import Annotated, Any, Literal from uuid import UUID from pydantic import ( @@ -20,7 +20,6 @@ from pydantic.functional_serializers import PlainSerializer from pydantic.functional_validators import AfterValidator, BeforeValidator from pydantic.networks import IPvAnyNetwork, UrlConstraints from pydantic_core import MultiHostHost, Url -from typing_extensions import Annotated from generalresearch.models import DeviceType, Source @@ -57,19 +56,17 @@ def convert_str_dt(v: Any) -> AwareDatetime | None: # to parse a str that was dumped using the iso8601 format with Z suffix. if v is not None and type(v) is str: assert v.endswith("Z") and "T" in v, "invalid format" - return datetime.strptime(v, "%Y-%m-%dT%H:%M:%S.%fZ").replace( - tzinfo=timezone.utc - ) + return datetime.strptime(v, "%Y-%m-%dT%H:%M:%S.%fZ").replace(tzinfo=UTC) return v def assert_utc(v: AwareDatetime) -> AwareDatetime: if isinstance(v, datetime): # We need utcoffset b/c FastAPI parses datetimes using FixedTimezone - assert v.tzinfo == timezone.utc or v.tzinfo.utcoffset(v) == timedelta( + assert v.tzinfo == UTC or v.tzinfo.utcoffset(v) == timedelta( 0 ), "Timezone is not UTC" - v = v.astimezone(timezone.utc) + v = v.astimezone(UTC) return v @@ -309,4 +306,4 @@ PropertyCode = Annotated[ def now_utc_factory(): - return datetime.now(tz=timezone.utc) + return datetime.now(tz=UTC) diff --git a/generalresearch/models/dynata/survey.py b/generalresearch/models/dynata/survey.py index 0e1b3e5..097eea2 100644 --- a/generalresearch/models/dynata/survey.py +++ b/generalresearch/models/dynata/survey.py @@ -2,10 +2,10 @@ from __future__ import annotations import json import logging -from datetime import timezone +from datetime import UTC, timezone from decimal import Decimal from functools import cached_property -from typing import Any, Literal, Type +from typing import Any, Literal, Self, Type from more_itertools import flatten from pydantic import ( @@ -17,7 +17,6 @@ from pydantic import ( field_validator, model_validator, ) -from typing_extensions import Self from generalresearch.locales import Localelator from generalresearch.models import Source, TaskCalculationType @@ -500,7 +499,7 @@ class DynataSurvey(MarketplaceTask): return res @property - def condition_model(self) -> Type[MarketplaceCondition]: + def condition_model(self) -> type[MarketplaceCondition]: return DynataCondition @property @@ -552,8 +551,8 @@ class DynataSurvey(MarketplaceTask): @classmethod def from_db(cls, d: Dict[str, Any]) -> Self: - d["created"] = d["created"].replace(tzinfo=timezone.utc) - d["last_updated"] = d["last_updated"].replace(tzinfo=timezone.utc) + d["created"] = d["created"].replace(tzinfo=UTC) + d["last_updated"] = d["last_updated"].replace(tzinfo=UTC) d["filters"] = json.loads(d["filters"]) d["quotas"] = json.loads(d["quotas"]) d["used_question_ids"] = json.loads(d["used_question_ids"]) diff --git a/generalresearch/models/events.py b/generalresearch/models/events.py index 63ed2a1..5efd4f6 100644 --- a/generalresearch/models/events.py +++ b/generalresearch/models/events.py @@ -1,6 +1,6 @@ -from datetime import datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta, timezone from enum import StrEnum -from typing import Dict, Literal, Optional, Union +from typing import Annotated, Dict, Literal, Optional, Union from uuid import uuid4 from pydantic import ( @@ -12,7 +12,6 @@ from pydantic import ( TypeAdapter, model_validator, ) -from typing_extensions import Annotated from generalresearch.models import Source from generalresearch.models.custom_types import ( @@ -63,7 +62,7 @@ class TaskEnterPayload(BaseModel): source: Source = Field() survey_id: str = Field(min_length=1, max_length=32, examples=["127492892"]) - quota_id: Optional[str] = Field( + quota_id: str | None = Field( default=None, max_length=32, description="The marketplace's internal quota id", @@ -76,9 +75,9 @@ class TaskFinishPayload(TaskEnterPayload): duration_sec: PositiveFloat = Field() status: Status - status_code_1: Optional[StatusCode1] = None - status_code_2: Optional[WallStatusCode2] = None - cpi: Optional[NonNegativeInt] = Field(le=4000, default=None) + status_code_1: StatusCode1 | None = None + status_code_2: WallStatusCode2 | None = None + cpi: NonNegativeInt | None = Field(le=4000, default=None) class SessionEnterPayload(BaseModel): @@ -91,18 +90,13 @@ class SessionFinishPayload(SessionEnterPayload): duration_sec: PositiveFloat = Field() status: Status - status_code_1: Optional[StatusCode1] = None - status_code_2: Optional[SessionStatusCode2] = None - user_payout: Optional[NonNegativeInt] = Field(default=None, le=4000, ge=0) + status_code_1: StatusCode1 | None = None + status_code_2: SessionStatusCode2 | None = None + user_payout: NonNegativeInt | None = Field(default=None, le=4000, ge=0) EventPayload = Annotated[ - Union[ - TaskEnterPayload, - TaskFinishPayload, - SessionEnterPayload, - SessionFinishPayload, - ], + TaskEnterPayload | TaskFinishPayload | SessionEnterPayload | SessionFinishPayload, Field(discriminator="event_type"), ] @@ -110,12 +104,10 @@ EventPayload = Annotated[ class EventEnvelope(BaseModel): event_uuid: UUIDStr = Field(default_factory=lambda: uuid4().hex) event_type: EventType = Field() - timestamp: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(tz=timezone.utc) - ) + timestamp: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC)) version: int = 1 - product_user_id: Optional[str] = Field( + product_user_id: str | None = Field( min_length=3, max_length=128, examples=["app-user-9329ebd"], @@ -136,7 +128,7 @@ class EventEnvelope(BaseModel): class AggregateBySource(BaseModel): total: NonNegativeInt = Field(default=0) - by_source: Dict[Source, NonNegativeInt] = Field(default_factory=dict) + by_source: dict[Source, NonNegativeInt] = Field(default_factory=dict) @model_validator(mode="after") def remove_zero(self): @@ -145,8 +137,8 @@ class AggregateBySource(BaseModel): class MaxGaugeBySource(BaseModel): - value: Optional[NonNegativeInt] = Field(default=None) - by_source: Dict[Source, NonNegativeInt] = Field(default_factory=dict) + value: NonNegativeInt | None = Field(default=None) + by_source: dict[Source, NonNegativeInt] = Field(default_factory=dict) @model_validator(mode="after") def remove_zero(self): @@ -174,11 +166,9 @@ class StatsSnapshot(TaskStatsSnapshot): model_config = ConfigDict(ser_json_timedelta="float") # If this is set, then everything is scoped to this country. - country_iso: Optional[CountryISOLike] = Field(default=None) + country_iso: CountryISOLike | None = Field(default=None) - timestamp: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(tz=timezone.utc) - ) + timestamp: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC)) # Counts: User related active_users_last_1h: NonNegativeInt = Field( @@ -217,17 +207,17 @@ class StatsSnapshot(TaskStatsSnapshot): ) # Rolling averages - session_avg_payout_last_24h: Optional[NonNegativeInt] = Field( + session_avg_payout_last_24h: NonNegativeInt | None = Field( description="Average (actual) payout of all tasks completed in the past 24 hrs" ) - session_avg_user_payout_last_24h: Optional[NonNegativeInt] = Field( + session_avg_user_payout_last_24h: NonNegativeInt | None = Field( description="Average (actual) user payout of all tasks completed in the past 24 hrs" ) - session_fail_avg_loi_last_24h: Optional[timedelta] = Field( + session_fail_avg_loi_last_24h: timedelta | None = Field( description="Average LOI of all tasks terminated in the past 24 hrs (excludes abandons)" ) - session_complete_avg_loi_last_24h: Optional[timedelta] = Field( + session_complete_avg_loi_last_24h: timedelta | None = Field( description="Average LOI of all tasks completed in the past 24 hrs" ) @@ -246,34 +236,26 @@ class StatsSnapshot(TaskStatsSnapshot): class EventMessage(BaseModel): kind: Literal[MessageKind.EVENT] = Field(default=MessageKind.EVENT) - timestamp: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(tz=timezone.utc) - ) + timestamp: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC)) data: EventEnvelope class StatsMessage(BaseModel): kind: Literal[MessageKind.STATS] = Field(default=MessageKind.STATS) - timestamp: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(tz=timezone.utc) - ) + timestamp: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC)) # The data/StatsSnapshot can optionally be scoped to a country - country_iso: Optional[CountryISOLike] = Field(default=None) + country_iso: CountryISOLike | None = Field(default=None) data: StatsSnapshot class PingMessage(BaseModel): kind: Literal[MessageKind.PING] = Field(default=MessageKind.PING) - timestamp: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(tz=timezone.utc) - ) + timestamp: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC)) class PongMessage(BaseModel): kind: Literal[MessageKind.PONG] = Field(default=MessageKind.PONG) - timestamp: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(tz=timezone.utc) - ) + timestamp: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC)) class SubscribeMessage(BaseModel): diff --git a/generalresearch/models/gr/authentication.py b/generalresearch/models/gr/authentication.py index 4ee70f9..67a8fc2 100644 --- a/generalresearch/models/gr/authentication.py +++ b/generalresearch/models/gr/authentication.py @@ -3,8 +3,8 @@ from __future__ import annotations import binascii import json import os -from datetime import datetime, timezone -from typing import TYPE_CHECKING, Any +from datetime import UTC, datetime, timezone +from typing import TYPE_CHECKING, Any, Self from pydantic import ( AnyHttpUrl, @@ -15,7 +15,6 @@ from pydantic import ( PositiveInt, field_validator, ) -from typing_extensions import Self from generalresearch.decorators import LOG from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr @@ -199,7 +198,7 @@ class GRUser(BaseModel): @field_validator("date_joined") @classmethod def date_joined_utc(cls, v: datetime) -> datetime: - return v.replace(tzinfo=timezone.utc) + return v.replace(tzinfo=UTC) # --- Properties --- @property @@ -290,7 +289,7 @@ class GRUser(BaseModel): @classmethod def from_postgresql(cls, d: dict) -> Self: - d["date_joined"] = d["date_joined"].replace(tzinfo=timezone.utc) + d["date_joined"] = d["date_joined"].replace(tzinfo=UTC) return GRUser.model_validate(d) @classmethod @@ -354,7 +353,7 @@ class GRToken(BaseModel): @field_validator("created", mode="before") @classmethod def created_utc(cls, v: datetime) -> datetime: - return v.replace(tzinfo=timezone.utc) + return v.replace(tzinfo=UTC) # --- Properties --- diff --git a/generalresearch/models/gr/business.py b/generalresearch/models/gr/business.py index 51317da..a67cf48 100644 --- a/generalresearch/models/gr/business.py +++ b/generalresearch/models/gr/business.py @@ -3,10 +3,10 @@ from __future__ import annotations import json import logging import os -from datetime import datetime, timezone +from datetime import UTC, datetime, timezone from enum import Enum from pathlib import Path -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Self from uuid import uuid4 import pandas as pd @@ -16,7 +16,6 @@ from psycopg.rows import dict_row from pydantic import BaseModel, ConfigDict, Field, PositiveInt from pydantic.json_schema import SkipJsonSchema from pydantic_extra_types.phone_numbers import PhoneNumber -from typing_extensions import Self from generalresearch.currency import USDCent from generalresearch.decorators import LOG @@ -391,8 +390,8 @@ class Business(BaseModel): pop_ledger = plm(ds=ds) if at_timestamp is None: - at_timestamp = datetime.now(tz=timezone.utc) - assert at_timestamp.tzinfo == timezone.utc + at_timestamp = datetime.now(tz=UTC) + assert at_timestamp.tzinfo == UTC ddf = pop_ledger.ddf( force_rr_latest=False, @@ -724,7 +723,7 @@ class Business(BaseModel): if "pop_financial" in keys: # We should explicitly pass the pop_financial years we want. By default, # at least get this year. - year = datetime.now(tz=timezone.utc).year + year = datetime.now(tz=UTC).year keys = list(set(keys) | {f"pop_financial:{year}"}) rc = gr_redis_config.create_redis_client() diff --git a/generalresearch/models/gr/team.py b/generalresearch/models/gr/team.py index 8d60825..900062f 100644 --- a/generalresearch/models/gr/team.py +++ b/generalresearch/models/gr/team.py @@ -2,10 +2,10 @@ from __future__ import annotations import json import os -from datetime import datetime, timezone +from datetime import UTC, datetime, timezone from enum import Enum from pathlib import Path -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Self from uuid import uuid4 import pandas as pd @@ -18,7 +18,6 @@ from pydantic import ( field_validator, ) from pydantic.json_schema import SkipJsonSchema -from typing_extensions import Self from generalresearch.decorators import LOG from generalresearch.incite.mergers.foundations.enriched_session import ( @@ -92,7 +91,7 @@ class Membership(BaseModel): @classmethod def created_utc(cls, v: datetime | str) -> datetime | str: if isinstance(v, datetime): - return v.replace(tzinfo=timezone.utc) + return v.replace(tzinfo=UTC) return v # --- prefetch methods --- diff --git a/generalresearch/models/innovate/__init__.py b/generalresearch/models/innovate/__init__.py index 054c69d..26946e9 100644 --- a/generalresearch/models/innovate/__init__.py +++ b/generalresearch/models/innovate/__init__.py @@ -1,7 +1,7 @@ from enum import Enum +from typing import Annotated from pydantic import StringConstraints -from typing_extensions import Annotated # Note, this is called the KEY in the Question model InnovateQuestionID = Annotated[ diff --git a/generalresearch/models/innovate/survey.py b/generalresearch/models/innovate/survey.py index bcd50d3..0359bd6 100644 --- a/generalresearch/models/innovate/survey.py +++ b/generalresearch/models/innovate/survey.py @@ -2,13 +2,14 @@ from __future__ import annotations import json import logging -from datetime import date, timezone +from datetime import UTC, date, timezone from decimal import Decimal from functools import cached_property from typing import ( Annotated, Any, Literal, + Self, Type, ) @@ -20,7 +21,6 @@ from pydantic import ( computed_field, model_validator, ) -from typing_extensions import Self from generalresearch.locales import Localelator from generalresearch.models import ( @@ -290,7 +290,7 @@ class InnovateSurvey(MarketplaceTask): return cls.model_validate(d) @property - def condition_model(self) -> Type[MarketplaceCondition]: + def condition_model(self) -> type[MarketplaceCondition]: return InnovateCondition @property @@ -361,10 +361,10 @@ class InnovateSurvey(MarketplaceTask): @classmethod def from_db(cls, d: dict[str, Any]) -> Self: - d["created"] = d["created"].replace(tzinfo=timezone.utc) - d["updated"] = d["updated"].replace(tzinfo=timezone.utc) - d["modified_api"] = d["modified_api"].replace(tzinfo=timezone.utc) - d["created_api"] = d["created_api"].replace(tzinfo=timezone.utc) + d["created"] = d["created"].replace(tzinfo=UTC) + d["updated"] = d["updated"].replace(tzinfo=UTC) + d["modified_api"] = d["modified_api"].replace(tzinfo=UTC) + d["created_api"] = d["created_api"].replace(tzinfo=UTC) d["qualifications"] = json.loads(d["qualifications"]) d["used_question_ids"] = json.loads(d["used_question_ids"]) d["quotas"] = json.loads(d["quotas"]) diff --git a/generalresearch/models/legacy/bucket.py b/generalresearch/models/legacy/bucket.py index 2650b0b..812241d 100644 --- a/generalresearch/models/legacy/bucket.py +++ b/generalresearch/models/legacy/bucket.py @@ -4,7 +4,7 @@ import logging import math from datetime import timedelta from decimal import Decimal -from typing import Any, Literal +from typing import Any, Literal, Self from pydantic import ( BaseModel, @@ -14,7 +14,6 @@ from pydantic import ( field_validator, model_validator, ) -from typing_extensions import Self from generalresearch.models import Source from generalresearch.models.custom_types import ( diff --git a/generalresearch/models/legacy/questions.py b/generalresearch/models/legacy/questions.py index 81e794c..8e19e57 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, Any +from typing import TYPE_CHECKING, Annotated, Any, Self from pydantic import ( BaseModel, @@ -14,7 +14,6 @@ from pydantic import ( model_validator, ) from sentry_sdk import capture_exception -from typing_extensions import Annotated, Self from generalresearch.models.custom_types import UUIDStr from generalresearch.models.legacy.api_status import StatusResponse diff --git a/generalresearch/models/lucid/__init__.py b/generalresearch/models/lucid/__init__.py index c3365db..4653339 100644 --- a/generalresearch/models/lucid/__init__.py +++ b/generalresearch/models/lucid/__init__.py @@ -1,5 +1,6 @@ +from typing import Annotated + from pydantic import Field -from typing_extensions import Annotated LucidQuestionIdType = Annotated[ str, Field(min_length=1, max_length=16, pattern=r"^[0-9]+$") diff --git a/generalresearch/models/lucid/question.py b/generalresearch/models/lucid/question.py index 908ce70..6cbcb73 100644 --- a/generalresearch/models/lucid/question.py +++ b/generalresearch/models/lucid/question.py @@ -2,10 +2,9 @@ from __future__ import annotations import logging from enum import Enum -from typing import TYPE_CHECKING, Any, Literal +from typing import TYPE_CHECKING, Any, Literal, Self from pydantic import BaseModel, Field, field_validator, model_validator -from typing_extensions import Self from generalresearch.models import Source from generalresearch.models.lucid import LucidQuestionIdType diff --git a/generalresearch/models/marketplace/summary.py b/generalresearch/models/marketplace/summary.py index f75c530..9551417 100644 --- a/generalresearch/models/marketplace/summary.py +++ b/generalresearch/models/marketplace/summary.py @@ -2,11 +2,10 @@ from __future__ import annotations from abc import ABC from collections.abc import Collection -from typing import Literal +from typing import Literal, Self import numpy as np from pydantic import BaseModel, ConfigDict, Field, computed_field -from typing_extensions import Self from generalresearch.models.thl.stats import StatisticalSummary diff --git a/generalresearch/models/morning/__init__.py b/generalresearch/models/morning/__init__.py index 2c61c49..1bc15a7 100644 --- a/generalresearch/models/morning/__init__.py +++ b/generalresearch/models/morning/__init__.py @@ -1,7 +1,7 @@ from enum import Enum +from typing import Annotated from pydantic import StringConstraints -from typing_extensions import Annotated # This is text-based, in lowercase. e.g. 'age', 'household_income' MorningQuestionID = Annotated[ diff --git a/generalresearch/models/morning/question.py b/generalresearch/models/morning/question.py index 0ab5030..8a1f729 100644 --- a/generalresearch/models/morning/question.py +++ b/generalresearch/models/morning/question.py @@ -1,10 +1,9 @@ import json from enum import Enum -from typing import Any, Literal, Dict, List, Optional +from typing import Any, Dict, List, Literal, Optional, Self from uuid import UUID from pydantic import BaseModel, Field, field_validator, model_validator -from typing_extensions import Self from generalresearch.locales import Localelator from generalresearch.models import Source @@ -54,7 +53,7 @@ class MorningQuestionType(str, Enum): class MorningUserQuestionAnswer(MarketplaceUserQuestionAnswer): question_id: MorningQuestionID = Field() - question_type: Optional[MorningQuestionType] = Field(default=None) + question_type: MorningQuestionType | None = Field(default=None) # Did this answer come from us asking, or was it passed back from the # marketplace? Note, morning doesn't "pass back" answers, but we can # retrieve a user's profile through API, so it is possible to populate @@ -92,7 +91,7 @@ class MorningQuestion(MarketplaceQuestion): frozen=True, ) # API calls this "responses", but I think that is a confusing name - options: Optional[List[MorningQuestionOption]] = Field( + options: list[MorningQuestionOption] | None = Field( default=None, min_length=1, frozen=True ) @@ -119,7 +118,7 @@ class MorningQuestion(MarketplaceQuestion): return options @classmethod - def from_api(cls, d: Dict[str, Any], country_iso: str, language_iso: str): + def from_api(cls, d: dict[str, Any], country_iso: str, language_iso: str): options = None if d.get("responses"): options = [ @@ -138,7 +137,7 @@ class MorningQuestion(MarketplaceQuestion): ) @classmethod - def from_db(cls, d: Dict[str, Any]) -> Self: + def from_db(cls, d: dict[str, Any]) -> Self: options = None if d["options"]: options = [ @@ -162,7 +161,7 @@ class MorningQuestion(MarketplaceQuestion): ), ) - def to_mysql(self) -> Dict[str, Any]: + def to_mysql(self) -> dict[str, Any]: d = self.model_dump(mode="json", by_alias=True) d["options"] = json.dumps(d["options"]) return d diff --git a/generalresearch/models/morning/survey.py b/generalresearch/models/morning/survey.py index 255a63c..3c5a0e0 100644 --- a/generalresearch/models/morning/survey.py +++ b/generalresearch/models/morning/survey.py @@ -2,7 +2,7 @@ from __future__ import annotations import json import logging -from datetime import timezone +from datetime import UTC, timezone from decimal import Decimal from functools import cached_property from typing import ( @@ -12,6 +12,7 @@ from typing import ( List, Literal, Optional, + Self, Set, Tuple, Type, @@ -27,7 +28,6 @@ from pydantic import ( computed_field, model_validator, ) -from typing_extensions import Self from generalresearch.locales import Localelator from generalresearch.models import Source @@ -79,14 +79,14 @@ class MorningStatistics(BaseModel): # bid bid_loi: int = Field(validation_alias="estimated_length_of_interview", # le=120 * 60) # If num_completes == 0 , this gets returned as 0. it should be None - obs_median_loi: Optional[NonNegativeInt] = Field( + obs_median_loi: NonNegativeInt | None = Field( validation_alias="median_length_of_interview", default=None, le=120 * 60 ) # API returns 100 until 5 completes! Should be None. # This is calculated as the total completes divided by the total number of # finished sessions that passed the prescreener. - qualified_conversion: Optional[float] = Field( + qualified_conversion: float | None = Field( ge=0, le=1, description="conversion rate of qualified respondents" ) @@ -142,7 +142,7 @@ class MorningTaskStatistics(MorningStatistics): # relevant to quotas. # API returns 100 until 5 completes! Should be None ... - system_conversion: Optional[float] = Field( + system_conversion: float | None = Field( description="conversion rate of the system. completes divided by total number of entrants to the system", ge=0, le=1, @@ -166,8 +166,8 @@ class MorningTaskStatistics(MorningStatistics): class MorningCondition(MarketplaceCondition): model_config = ConfigDict(populate_by_name=True, frozen=False, extra="ignore") - question_id: Optional[MorningQuestionID] = Field(validation_alias="id") - values: List[Annotated[str, Field(max_length=128)]] = Field( + question_id: MorningQuestionID | None = Field(validation_alias="id") + values: list[Annotated[str, Field(max_length=128)]] = Field( validation_alias="response_ids" ) value_type: ConditionValueType = Field(default=ConditionValueType.LIST) @@ -184,11 +184,11 @@ class MorningQuota(MorningStatistics, MarketplaceTask): max_digits=5, validation_alias="cost_per_interview", ) - condition_hashes: List[str] = Field(min_length=1, default_factory=list) + condition_hashes: list[str] = Field(min_length=1, default_factory=list) # since the Quota is the MarketplaceTask, it needs these fields, copied from the Bid source: Literal[Source.MORNING_CONSULT] = Field(default=Source.MORNING_CONSULT) - used_question_ids: Set[MorningQuestionID] = Field(default_factory=set) + used_question_ids: set[MorningQuestionID] = Field(default_factory=set) country_iso: CountryISO = Field(frozen=True) country_isos: CountryISOs = Field() language_isos: LanguageISOs = Field(frozen=True) @@ -219,11 +219,11 @@ class MorningQuota(MorningStatistics, MarketplaceTask): @computed_field @cached_property - def all_hashes(self) -> Set[str]: + def all_hashes(self) -> set[str]: return set(self.condition_hashes) @property - def condition_model(self) -> Type[MarketplaceCondition]: + def condition_model(self) -> type[MarketplaceCondition]: return MorningCondition @property @@ -233,7 +233,7 @@ class MorningQuota(MorningStatistics, MarketplaceTask): @property def marketplace_genders( self, - ) -> Dict[Gender, Optional[MarketplaceCondition]]: + ) -> dict[Gender, MarketplaceCondition | None]: return { Gender.MALE: MorningCondition( question_id="gender", @@ -253,14 +253,14 @@ class MorningQuota(MorningStatistics, MarketplaceTask): # num_available includes in-progress (they're already deducted) return self.num_available >= self._min_open_spots - def passes(self, criteria_evaluation: Dict[str, Optional[bool]]) -> bool: + def passes(self, criteria_evaluation: dict[str, bool | None]) -> bool: # Passes means we 1) meet all conditions (aka "match") AND 2) the quota is open. return self.is_open and self.matches(criteria_evaluation) # TODO: I did some speed tests. This is faster than how this is implemented # in sago/spectrum/dynata/etc. We should generalize this logic instead of # copying/pasting it 7 times. (matches, matches_optional and _soft) - def matches(self, criteria_evaluation: Dict[str, Optional[bool]]) -> bool: + def matches(self, criteria_evaluation: dict[str, bool | None]) -> bool: # Matches means we meet all conditions. # In Morning, all quotas are mutually exclusive. so if it doesn't # matter if we match a closed quota, b/c that means that we won't @@ -268,8 +268,8 @@ class MorningQuota(MorningStatistics, MarketplaceTask): return self.matches_optional(criteria_evaluation) is True def matches_optional( - self, criteria_evaluation: Dict[str, Optional[bool]] - ) -> Optional[bool]: + self, criteria_evaluation: dict[str, bool | None] + ) -> bool | None: for c in self.condition_hashes: eval_value = criteria_evaluation.get(c) if eval_value is False: @@ -279,8 +279,8 @@ class MorningQuota(MorningStatistics, MarketplaceTask): return True def matches_soft( - self, criteria_evaluation: Dict[str, Optional[bool]] - ) -> Tuple[Optional[bool], List[str]]: + 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() for c in self.condition_hashes: @@ -321,22 +321,22 @@ class MorningBid(MorningTaskStatistics): timeout: PositiveInt = Field(le=24 * 60 * 60) topic_id: str = Field(min_length=1, max_length=64) - exclusions: List[MorningExclusion] = Field(default_factory=list) + exclusions: list[MorningExclusion] = Field(default_factory=list) - quotas: List[MorningQuota] = Field(default_factory=list) + quotas: list[MorningQuota] = Field(default_factory=list) source: Literal[Source.MORNING_CONSULT] = Field(default=Source.MORNING_CONSULT) - used_question_ids: Set[MorningQuestionID] = Field(default_factory=set) + used_question_ids: set[MorningQuestionID] = Field(default_factory=set) # This is a "special" key to store all conditions that are used (as # "condition_hashes") throughout this survey. In the reduced representation # of this task (nearly always, for db i/o, in global_vars) this field will # be null. - conditions: Optional[Dict[str, MorningCondition]] = Field(default=None) + conditions: dict[str, MorningCondition] | None = Field(default=None) # This doesn't get stored in the db directly - experimental_single_use_qualifications: Optional[List[MorningQuestion]] = Field( + experimental_single_use_qualifications: list[MorningQuestion] | None = Field( default=None ) @@ -345,8 +345,8 @@ class MorningBid(MorningTaskStatistics): created_api: AwareDatetimeISO = Field(validation_alias="published_at") # This does not come from the API. We set it when we update this in the db. - created: Optional[AwareDatetimeISO] = Field(default=None) - updated: Optional[AwareDatetimeISO] = Field(default=None) + created: AwareDatetimeISO | None = Field(default=None) + updated: AwareDatetimeISO | None = Field(default=None) # ignoring from API: closed_at @@ -373,7 +373,7 @@ class MorningBid(MorningTaskStatistics): @computed_field @cached_property - def all_hashes(self) -> Set[str]: + def all_hashes(self) -> set[str]: s = set() for q in self.quotas: s.update(set(q.condition_hashes)) @@ -387,7 +387,7 @@ class MorningBid(MorningTaskStatistics): @model_validator(mode="before") @classmethod - def setup_quota_fields(cls, data: Dict[str, Any]) -> Dict[str, Any]: + def setup_quota_fields(cls, data: dict[str, Any]) -> dict[str, Any]: # These fields get "inherited" by each quota from its bid. quota_fields = [ "country_iso", @@ -419,7 +419,7 @@ class MorningBid(MorningTaskStatistics): @model_validator(mode="before") @classmethod - def setup_conditions(cls, data: Dict[str, Any]) -> Dict[str, Any]: + def setup_conditions(cls, data: dict[str, Any]) -> dict[str, Any]: if "conditions" in data: return data @@ -448,7 +448,7 @@ class MorningBid(MorningTaskStatistics): @model_validator(mode="before") @classmethod - def clean_alias(cls, data: Dict[str, Any]) -> Dict[str, Any]: + def clean_alias(cls, data: dict[str, Any]) -> dict[str, Any]: # Make sure fields are named certain ways, so we don't have to check # aliases within other validators if "estimated_length_of_interview" in data: @@ -503,18 +503,16 @@ class MorningBid(MorningTaskStatistics): return d @classmethod - def from_db(cls, d: Dict[str, Any]) -> Self: - d["created"] = d["created"].replace(tzinfo=timezone.utc) - d["updated"] = d["updated"].replace(tzinfo=timezone.utc) - d["expected_end"] = d["expected_end"].replace(tzinfo=timezone.utc) - d["created_api"] = d["created_api"].replace(tzinfo=timezone.utc) + def from_db(cls, d: dict[str, Any]) -> Self: + d["created"] = d["created"].replace(tzinfo=UTC) + d["updated"] = d["updated"].replace(tzinfo=UTC) + d["expected_end"] = d["expected_end"].replace(tzinfo=UTC) + d["created_api"] = d["created_api"].replace(tzinfo=UTC) d["used_question_ids"] = json.loads(d["used_question_ids"]) d["exclusions"] = json.loads(d["exclusions"]) return cls.model_validate(d) - def passes_quotas( - self, criteria_evaluation: Dict[str, Optional[bool]] - ) -> Optional[str]: + def passes_quotas(self, criteria_evaluation: dict[str, bool | None]) -> str | None: # Quotas are mutually-exclusive. A user can only possibly match 1 quota. # Returns the passing quota ID or None (if user doesn't pass any quota) for q in self.quotas: @@ -522,8 +520,8 @@ class MorningBid(MorningTaskStatistics): return q.id def passes_quotas_soft( - self, criteria_evaluation: Dict[str, Optional[bool]] - ) -> Tuple[Optional[bool], Optional[List[str]], Optional[Set[str]]]: + self, criteria_evaluation: dict[str, bool | None] + ) -> tuple[bool | None, list[str] | None, set[str] | None]: """ Quotas are mutually-exclusive. A user can only possibly match 1 quota. As such, all unknown questions on any quota will be @@ -547,15 +545,15 @@ class MorningBid(MorningTaskStatistics): return False, None, None def determine_eligibility( - self, criteria_evaluation: dict[str, Optional[bool]] - ) -> Optional[str]: + self, criteria_evaluation: dict[str, bool | None] + ) -> str | None: if not self.is_open: return None return self.passes_quotas(criteria_evaluation) def determine_eligibility_soft( - self, criteria_evaluation: dict[str, Optional[bool]] - ) -> Tuple[Optional[bool], Optional[List[str]], Optional[Set[str]]]: + self, criteria_evaluation: dict[str, bool | None] + ) -> tuple[bool | None, list[str] | None, set[str] | None]: if not self.is_open: return False, None, None return self.passes_quotas_soft(criteria_evaluation) diff --git a/generalresearch/models/network/mtr/execute.py b/generalresearch/models/network/mtr/execute.py index d77e814..953124d 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 datetime, timezone +from datetime import UTC, datetime, timezone from uuid import uuid4 from generalresearch.models.custom_types import UUIDStr @@ -33,10 +33,10 @@ def execute_mtr( ), ) - started_at = datetime.now(tz=timezone.utc) + started_at = datetime.now(tz=UTC) tool_version = get_mtr_version() result = run_mtr(config) - finished_at = datetime.now(tz=timezone.utc) + finished_at = datetime.now(tz=UTC) return MTRRun( tool_name=ToolName.MTR, diff --git a/generalresearch/models/network/nmap/parser.py b/generalresearch/models/network/nmap/parser.py index e946e5f..6ad4ab4 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 datetime, timezone +from datetime import UTC, datetime, timezone from typing import Any from generalresearch.models.network.definitions import IPProtocol @@ -123,7 +123,7 @@ class NmapXmlParser: finished_at = None ts = finished.attrib.get("time") if ts: - finished_at = datetime.fromtimestamp(int(ts), tz=timezone.utc) + finished_at = datetime.fromtimestamp(int(ts), tz=UTC) return { "finished_at": finished_at, @@ -136,7 +136,7 @@ class NmapXmlParser: nmaprun = dict(nmaprun_el.attrib) nmap_data["command_line"] = nmaprun["args"] nmap_data["started_at"] = datetime.fromtimestamp( - float(nmaprun["start"]), tz=timezone.utc + float(nmaprun["start"]), tz=UTC ) nmap_data["version"] = nmaprun["version"] nmap_data["xmloutputversion"] = nmaprun["xmloutputversion"] diff --git a/generalresearch/models/network/nmap/result.py b/generalresearch/models/network/nmap/result.py index 3f9cae6..e6a0fd3 100644 --- a/generalresearch/models/network/nmap/result.py +++ b/generalresearch/models/network/nmap/result.py @@ -256,13 +256,13 @@ class NmapScanInfo(BaseModel): services: str = Field() @cached_property - def port_set(self) -> Set[int]: + def port_set(self) -> set[int]: """ Expand the Nmap services string into a set of port numbers. Example: "22-25,80,443" -> {22,23,24,25,80,443} """ - ports: Set[int] = set() + ports: set[int] = set() for part in self.services.split(","): if "-" in part: start, end = part.split("-", 1) diff --git a/generalresearch/models/network/rdns/execute.py b/generalresearch/models/network/rdns/execute.py index cabd13c..1d74df2 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 datetime, timezone +from datetime import UTC, datetime, timezone from uuid import uuid4 from generalresearch.models.custom_types import UUIDStr @@ -21,11 +21,11 @@ from generalresearch.models.network.tool_run_command import ( def execute_rdns(ip: str, scan_group_id: UUIDStr | None = None): - started_at = datetime.now(tz=timezone.utc) + started_at = datetime.now(tz=UTC) tool_version = get_dig_version() config = RDNSRunCommand(options=RDNSRunCommandOptions(ip=ip)) result = run_rdns(config) - finished_at = datetime.now(tz=timezone.utc) + finished_at = datetime.now(tz=UTC) run = RDNSRun( tool_name=ToolName.DIG, diff --git a/generalresearch/models/precision/__init__.py b/generalresearch/models/precision/__init__.py index 4bb2c6a..089c34a 100644 --- a/generalresearch/models/precision/__init__.py +++ b/generalresearch/models/precision/__init__.py @@ -1,7 +1,7 @@ from enum import Enum +from typing import Annotated from pydantic import StringConstraints -from typing_extensions import Annotated class PrecisionStatus(str, Enum): diff --git a/generalresearch/models/precision/survey.py b/generalresearch/models/precision/survey.py index 646d60e..be98a79 100644 --- a/generalresearch/models/precision/survey.py +++ b/generalresearch/models/precision/survey.py @@ -1,9 +1,9 @@ from __future__ import annotations import json -from datetime import timezone +from datetime import UTC, timezone from functools import cached_property -from typing import Any, Dict, List, Literal, Optional, Self, Set, Tuple, Type +from typing import Annotated, Any, Dict, List, Literal, Optional, Self, Set, Tuple, Type from more_itertools import flatten from pydantic import ( @@ -14,7 +14,6 @@ from pydantic import ( computed_field, model_validator, ) -from typing_extensions import Annotated from generalresearch.models import Source from generalresearch.models.custom_types import ( @@ -34,8 +33,8 @@ from generalresearch.models.thl.survey.condition import ( class PrecisionCondition(MarketplaceCondition): - question_id: Optional[PrecisionQuestionID] = Field() - values: List[Annotated[str, Field(max_length=128)]] = Field() + question_id: PrecisionQuestionID | None = Field() + values: list[Annotated[str, Field(max_length=128)]] = Field() value_type: ConditionValueType = Field(default=ConditionValueType.LIST) _CONVERT_LIST_TO_RANGE = ["age"] @@ -54,7 +53,7 @@ class PrecisionQuota(BaseModel): termination_count: int = Field(ge=0) overquota_count: int = Field(ge=0) - condition_hashes: List[str] = Field(min_length=1, default_factory=list) + condition_hashes: list[str] = Field(min_length=1, default_factory=list) # Min spots a quota should have open to be OPEN _min_open_spots: int = PrivateAttr(default=3) @@ -78,7 +77,7 @@ class PrecisionQuota(BaseModel): # TODO: I did some speed tests. This is faster than how this is implemented # in sago/spectrum/dynata/etc. We should generalize this logic instead of # copying/pasting it 7 times. (matches, matches_optional and _soft) - def matches(self, criteria_evaluation: Dict[str, Optional[bool]]) -> bool: + def matches(self, criteria_evaluation: dict[str, bool | None]) -> bool: # Matches means we meet all conditions. # In Morning, all quotas are mutually exclusive. so if it doesn't # matter if we match a closed quota, b/c that means that we won't @@ -86,8 +85,8 @@ class PrecisionQuota(BaseModel): return self.matches_optional(criteria_evaluation) is True def matches_optional( - self, criteria_evaluation: Dict[str, Optional[bool]] - ) -> Optional[bool]: + self, criteria_evaluation: dict[str, bool | None] + ) -> bool | None: for c in self.condition_hashes: eval_value = criteria_evaluation.get(c) if eval_value is False: @@ -97,8 +96,8 @@ class PrecisionQuota(BaseModel): return True def matches_soft( - self, criteria_evaluation: Dict[str, Optional[bool]] - ) -> Tuple[Optional[bool], List[str]]: + 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() for c in self.condition_hashes: @@ -128,7 +127,7 @@ class PrecisionSurvey(MarketplaceTask): name: str = Field(validation_alias="prj_name") survey_guid: UUIDStrCoerce = Field(validation_alias="prj_guid") - category_id: Optional[str] = Field(validation_alias="sc_id", default=None) + category_id: str | None = Field(validation_alias="sc_id", default=None) buyer_id: CoercedStr = Field(max_length=16) # This seems to always be 0 ... ? @@ -141,7 +140,7 @@ class PrecisionSurvey(MarketplaceTask): bid_ir: float = Field(ge=0, le=1, validation_alias="ir") # Be careful with this, it doesn't make any sense. See survey 452481, has 12 completes with a 100% live_ir, # but the only quotas have 0 completes and 1052 terms. .... ?? - global_conversion: Optional[float] = Field( + global_conversion: float | None = Field( ge=0, le=1, default=None, @@ -156,31 +155,31 @@ class PrecisionSurvey(MarketplaceTask): allowed_devices: DeviceTypes = Field(min_length=1) entry_link: str = Field(validation_alias="url") - excluded_surveys: Optional[AlphaNumStrSet] = Field( + excluded_surveys: AlphaNumStrSet | None = Field( description="list of excluded survey ids", default=None, validation_alias="exclusion_project_id", ) - quotas: List[PrecisionQuota] = Field(default_factory=list) + quotas: list[PrecisionQuota] = Field(default_factory=list) source: Literal[Source.PRECISION] = Field(default=Source.PRECISION) - used_question_ids: Set[PrecisionQuestionID] = Field(default_factory=set) + used_question_ids: set[PrecisionQuestionID] = Field(default_factory=set) # This is a "special" key to store all conditions that are used (as "condition_hashes") throughout # this survey. In the reduced representation of this task (nearly always, for db i/o, in global_vars) # this field will be null. - conditions: Optional[Dict[str, PrecisionCondition]] = Field(default=None) + conditions: dict[str, PrecisionCondition] | None = Field(default=None) # This comes from the API - expected_end_date: Optional[AwareDatetimeISO] = Field( + expected_end_date: AwareDatetimeISO | None = Field( default=None, validation_alias="end_date" ) # This does not come from the API. We set it when we update this in the db. - created: Optional[AwareDatetimeISO] = Field(default=None) - updated: Optional[AwareDatetimeISO] = Field(default=None) + created: AwareDatetimeISO | None = Field(default=None) + updated: AwareDatetimeISO | None = Field(default=None) @property def internal_id(self) -> str: @@ -199,7 +198,7 @@ class PrecisionSurvey(MarketplaceTask): @computed_field @cached_property - def all_hashes(self) -> Set[str]: + def all_hashes(self) -> set[str]: s = set() for q in self.quotas: s.update(set(q.condition_hashes)) @@ -219,7 +218,7 @@ class PrecisionSurvey(MarketplaceTask): return data @property - def condition_model(self) -> Type[MarketplaceCondition]: + def condition_model(self) -> type[MarketplaceCondition]: return PrecisionCondition @property @@ -227,7 +226,7 @@ class PrecisionSurvey(MarketplaceTask): return "age" @property - def marketplace_genders(self) -> Dict[Gender, Optional[MarketplaceCondition]]: + def marketplace_genders(self) -> dict[Gender, MarketplaceCondition | None]: return { Gender.MALE: PrecisionCondition( question_id="gender", @@ -262,7 +261,7 @@ class PrecisionSurvey(MarketplaceTask): exclude={"updated", "conditions", "created"} ) == other.model_dump(exclude={"updated", "conditions", "created"}) - def to_mysql(self) -> Dict[str, Any]: + def to_mysql(self) -> dict[str, Any]: d = self.model_dump( mode="json", exclude={ @@ -283,11 +282,11 @@ class PrecisionSurvey(MarketplaceTask): return d @classmethod - def from_db(cls, d: Dict[str, Any]) -> Self: - d["created"] = d["created"].replace(tzinfo=timezone.utc) - d["updated"] = d["updated"].replace(tzinfo=timezone.utc) + def from_db(cls, d: dict[str, Any]) -> Self: + d["created"] = d["created"].replace(tzinfo=UTC) + d["updated"] = d["updated"].replace(tzinfo=UTC) d["expected_end_date"] = ( - d["expected_end_date"].replace(tzinfo=timezone.utc) + d["expected_end_date"].replace(tzinfo=UTC) if d["expected_end_date"] else None ) @@ -295,7 +294,7 @@ class PrecisionSurvey(MarketplaceTask): d["used_question_ids"] = json.loads(d["used_question_ids"]) return cls.model_validate(d) - def passes_quotas(self, criteria_evaluation: Dict[str, Optional[bool]]) -> bool: + def passes_quotas(self, criteria_evaluation: dict[str, bool | None]) -> bool: # We have to match 1 or more quota. # Quotas are exclusionary: they can NOT match a quota where currently_open=0 any_pass = False @@ -308,8 +307,8 @@ class PrecisionSurvey(MarketplaceTask): return any_pass def passes_quotas_soft( - self, criteria_evaluation: Dict[str, Optional[bool]] - ) -> Tuple[Optional[bool], Set[str]]: + self, criteria_evaluation: dict[str, bool | None] + ) -> tuple[bool | None, set[str]]: # Quotas are exclusionary. They can NOT match a quota where currently_open=0 quota_eval = { quota: quota.matches_soft(criteria_evaluation) for quota in self.quotas @@ -345,19 +344,19 @@ class PrecisionSurvey(MarketplaceTask): return False, set() def determine_eligibility( - self, criteria_evaluation: Dict[str, Optional[bool]] + self, criteria_evaluation: dict[str, bool | None] ) -> bool: return self.is_open and self.passes_quotas(criteria_evaluation) def determine_eligibility_soft( - self, criteria_evaluation: Dict[str, Optional[bool]] - ) -> Tuple[Optional[bool], Optional[Set[str]]]: + self, criteria_evaluation: dict[str, bool | None] + ) -> tuple[bool | None, set[str] | None]: if not self.is_open: return False, None return self.passes_quotas_soft(criteria_evaluation) def participation_allowed( - self, att_survey_ids: Set[str], att_group_ids: Set[str] + self, att_survey_ids: set[str], att_group_ids: set[str] ) -> bool: """ Checks if this user can participate in this survey diff --git a/generalresearch/models/precision/task_collection.py b/generalresearch/models/precision/task_collection.py index 233d329..daea448 100644 --- a/generalresearch/models/precision/task_collection.py +++ b/generalresearch/models/precision/task_collection.py @@ -36,8 +36,8 @@ PrecisionTaskCollectionSchema = DataFrameSchema( "expected_end_date": Column(dtype=pd.DatetimeTZDtype(tz="UTC"), nullable=True), "created": Column(dtype=pd.DatetimeTZDtype(tz="UTC")), "updated": Column(dtype=pd.DatetimeTZDtype(tz="UTC")), - "used_question_ids": Column(List[str]), - "all_hashes": Column(List[str]), # set >> list for column support + "used_question_ids": Column(list[str]), + "all_hashes": Column(list[str]), # set >> list for column support }, checks=[], index=Index( @@ -53,10 +53,10 @@ PrecisionTaskCollectionSchema = DataFrameSchema( class PrecisionTaskCollection(TaskCollection): - items: List[PrecisionSurvey] + items: list[PrecisionSurvey] _schema = PrecisionTaskCollectionSchema - def to_row(self, s: PrecisionSurvey) -> Dict[str, Any]: + def to_row(self, s: PrecisionSurvey) -> dict[str, Any]: d = s.model_dump( mode="json", exclude={ diff --git a/generalresearch/models/prodege/__init__.py b/generalresearch/models/prodege/__init__.py index d419c0c..5c6659a 100644 --- a/generalresearch/models/prodege/__init__.py +++ b/generalresearch/models/prodege/__init__.py @@ -1,8 +1,7 @@ from enum import Enum -from typing import Literal +from typing import Annotated, Literal from pydantic import Field -from typing_extensions import Annotated ProdegeQuestionIdType = Annotated[ str, Field(min_length=1, max_length=16, pattern=r"^[0-9]+$") diff --git a/generalresearch/models/prodege/question.py b/generalresearch/models/prodege/question.py index 1c61ab9..3ef4772 100644 --- a/generalresearch/models/prodege/question.py +++ b/generalresearch/models/prodege/question.py @@ -3,7 +3,7 @@ from __future__ import annotations import json import logging -from datetime import datetime, timezone +from datetime import UTC, datetime, timezone from enum import Enum from functools import cached_property from typing import TYPE_CHECKING, Any, Literal @@ -43,9 +43,7 @@ class ProdegeUserQuestionAnswer(BaseModel): # This may be a pipe-separated string if the question_type is multi. regex means any chars except capital letters option_id: str = Field(pattern=r"^[^A-Z]*$") - created: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(tz=timezone.utc) - ) + created: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC)) # ISO 3166-1 alpha-2 (two-letter codes, lowercase) country_iso: str = Field( diff --git a/generalresearch/models/prodege/survey.py b/generalresearch/models/prodege/survey.py index c12f130..1601fa3 100644 --- a/generalresearch/models/prodege/survey.py +++ b/generalresearch/models/prodege/survey.py @@ -4,7 +4,7 @@ from __future__ import annotations import json import logging from collections import defaultdict -from datetime import datetime, timezone +from datetime import UTC, datetime, timezone from decimal import Decimal from functools import cached_property from typing import Any, Literal, Type @@ -127,7 +127,7 @@ class ProdegeQuota(BaseModel): return self.remaining_count >= min_open_spots @property - def condition_model(self) -> Type[MarketplaceCondition]: + def condition_model(self) -> type[MarketplaceCondition]: return ProdegeCondition @property @@ -276,7 +276,7 @@ class ProdegeUserPastParticipation(BaseModel): raise ValueError(f"Unknown ext_status_code_1: {self.ext_status_code_1}") def days_ago(self) -> float: - now = datetime.now(timezone.utc) + now = datetime.now(UTC) return (now - self.started).total_seconds() / (3600 * 24) @@ -486,7 +486,7 @@ class ProdegeSurvey(MarketplaceTask): return data @property - def condition_model(self) -> Type[MarketplaceCondition]: + def condition_model(self) -> type[MarketplaceCondition]: return ProdegeCondition @property @@ -656,8 +656,8 @@ class ProdegeSurvey(MarketplaceTask): @classmethod def from_db(cls, d: dict[str, Any]) -> ProdegeSurvey: - d["created"] = d["created"].replace(tzinfo=timezone.utc) - d["updated"] = d["updated"].replace(tzinfo=timezone.utc) + d["created"] = d["created"].replace(tzinfo=UTC) + d["updated"] = d["updated"].replace(tzinfo=UTC) d["quotas"] = json.loads(d["quotas"]) for k in [ "max_clicks_settings", diff --git a/generalresearch/models/prodege/task_collection.py b/generalresearch/models/prodege/task_collection.py index 19e594f..d3e4a20 100644 --- a/generalresearch/models/prodege/task_collection.py +++ b/generalresearch/models/prodege/task_collection.py @@ -30,8 +30,8 @@ ProdegeTaskCollectionSchema = DataFrameSchema( "conversion_rate": Column(float, Check.between(0, 1), nullable=True), "created": Column(dtype=pd.DatetimeTZDtype(tz="UTC")), "updated": Column(dtype=pd.DatetimeTZDtype(tz="UTC")), - "used_question_ids": Column(List[str]), - "all_hashes": Column(List[str]), # set >> list for column support + "used_question_ids": Column(list[str]), + "all_hashes": Column(list[str]), # set >> list for column support "is_recontact": Column(bool), # Not including here: entrance_url, max_clicks_settings, past_participation, include_psids, exclude_psids, # quotas, source, conditions diff --git a/generalresearch/models/repdata/survey.py b/generalresearch/models/repdata/survey.py index 2290ca6..fa71c04 100644 --- a/generalresearch/models/repdata/survey.py +++ b/generalresearch/models/repdata/survey.py @@ -3,10 +3,10 @@ from __future__ import annotations import json import logging -from datetime import datetime, timezone +from datetime import UTC, datetime, timezone from decimal import Decimal from functools import cached_property -from typing import Any, Literal, Type +from typing import Any, Literal, Self, Type from uuid import UUID from pydantic import ( @@ -17,7 +17,6 @@ from pydantic import ( field_validator, model_validator, ) -from typing_extensions import Self from generalresearch.grpc import timestamp_from_datetime from generalresearch.locales import Localelator @@ -304,7 +303,7 @@ class RepDataStream(MarketplaceTask): return self.stream_status == RepDataStatus.LIVE @property - def condition_model(self) -> Type[MarketplaceCondition]: + def condition_model(self) -> type[MarketplaceCondition]: return RepDataCondition @property @@ -538,8 +537,8 @@ class RepDataSurveyHashed(RepDataSurvey): DeviceType(int(x)) for x in res["allowed_devices"].split(",") ] if res["created"] is not None: - res["created"] = res["created"].replace(tzinfo=timezone.utc) - res["last_updated"] = res["last_updated"].replace(tzinfo=timezone.utc) + res["created"] = res["created"].replace(tzinfo=UTC) + res["last_updated"] = res["last_updated"].replace(tzinfo=UTC) return cls.model_validate(res) def to_mysql(self) -> dict[str, Any]: @@ -553,7 +552,7 @@ class RepDataSurveyHashed(RepDataSurvey): return d def to_grpc(self, repdata_pb2): - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) timestamp = timestamp_from_datetime(now) return repdata_pb2.RepDataOpportunity( diff --git a/generalresearch/models/sago/__init__.py b/generalresearch/models/sago/__init__.py index 292f0f2..19e7b6d 100644 --- a/generalresearch/models/sago/__init__.py +++ b/generalresearch/models/sago/__init__.py @@ -1,7 +1,7 @@ from enum import Enum +from typing import Annotated from pydantic import Field -from typing_extensions import Annotated SagoQuestionIdType = Annotated[ str, Field(min_length=1, max_length=16, pattern=r"^[0-9]+$") diff --git a/generalresearch/models/sago/survey.py b/generalresearch/models/sago/survey.py index e73ddee..0b5830b 100644 --- a/generalresearch/models/sago/survey.py +++ b/generalresearch/models/sago/survey.py @@ -2,14 +2,13 @@ from __future__ import annotations import json import logging -from datetime import timezone +from datetime import UTC, timezone from decimal import Decimal from functools import cached_property -from typing import Annotated, Any, Literal, Type +from typing import Annotated, Any, Literal, Self, Type from more_itertools import flatten from pydantic import BaseModel, ConfigDict, Field, computed_field, model_validator -from typing_extensions import Self from generalresearch.locales import Localelator from generalresearch.models import LogicalOperator, Source @@ -235,7 +234,7 @@ class SagoSurvey(MarketplaceTask): return data @property - def condition_model(self) -> Type[MarketplaceCondition]: + def condition_model(self) -> type[MarketplaceCondition]: return SagoCondition @property @@ -314,9 +313,9 @@ class SagoSurvey(MarketplaceTask): @classmethod def from_db(cls, d: dict[str, Any]): - d["created"] = d["created"].replace(tzinfo=timezone.utc) - d["updated"] = d["updated"].replace(tzinfo=timezone.utc) - d["modified_api"] = d["modified_api"].replace(tzinfo=timezone.utc) + d["created"] = d["created"].replace(tzinfo=UTC) + d["updated"] = d["updated"].replace(tzinfo=UTC) + d["modified_api"] = d["modified_api"].replace(tzinfo=UTC) d["qualifications"] = json.loads(d["qualifications"]) d["used_question_ids"] = json.loads(d["used_question_ids"]) d["quotas"] = json.loads(d["quotas"]) diff --git a/generalresearch/models/spectrum/__init__.py b/generalresearch/models/spectrum/__init__.py index b62c089..0040551 100644 --- a/generalresearch/models/spectrum/__init__.py +++ b/generalresearch/models/spectrum/__init__.py @@ -1,7 +1,7 @@ from enum import Enum +from typing import Annotated from pydantic import Field -from typing_extensions import Annotated SpectrumQuestionIdType = Annotated[ str, Field(min_length=1, max_length=16, pattern=r"^[0-9]+$") diff --git a/generalresearch/models/spectrum/question.py b/generalresearch/models/spectrum/question.py index db8a55d..81f8655 100644 --- a/generalresearch/models/spectrum/question.py +++ b/generalresearch/models/spectrum/question.py @@ -3,10 +3,10 @@ from __future__ import annotations import json import logging -from datetime import datetime, timezone +from datetime import UTC, datetime, timezone from enum import Enum from functools import cached_property -from typing import TYPE_CHECKING, Any, Literal +from typing import TYPE_CHECKING, Any, Literal, Self from uuid import UUID from pydantic import ( @@ -16,7 +16,6 @@ from pydantic import ( field_validator, model_validator, ) -from typing_extensions import Self from generalresearch.models import MAX_INT32, Source, string_utils from generalresearch.models.custom_types import AwareDatetimeISO @@ -50,9 +49,7 @@ class SpectrumUserQuestionAnswer(BaseModel): # This may be a pipe-separated string if the question_type is multi. regex # means any chars except capital letters option_id: str = Field(pattern=r"^[^A-Z]*$") - created: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(tz=timezone.utc) - ) + created: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC)) # ISO 3166-1 alpha-2 (two-letter codes, lowercase) country_iso: str = Field( max_length=2, min_length=2, pattern=r"^[a-z]{2}$", frozen=True @@ -283,7 +280,7 @@ class SpectrumQuestion(MarketplaceQuestion): ] created = ( - datetime.utcfromtimestamp(d["crtd_on"] / 1000).replace(tzinfo=timezone.utc) + datetime.utcfromtimestamp(d["crtd_on"] / 1000).replace(tzinfo=UTC) if d.get("crtd_on") else None ) @@ -308,9 +305,7 @@ class SpectrumQuestion(MarketplaceQuestion): SpectrumQuestionOption(id=r["id"], text=r["text"], order=r["order"]) for r in d["options"] ] - d["created"] = ( - d["created"].replace(tzinfo=timezone.utc) if d["created"] else None - ) + d["created"] = d["created"].replace(tzinfo=UTC) if d["created"] else None return cls( question_id=d["question_id"], diff --git a/generalresearch/models/spectrum/survey.py b/generalresearch/models/spectrum/survey.py index a591445..f842f46 100644 --- a/generalresearch/models/spectrum/survey.py +++ b/generalresearch/models/spectrum/survey.py @@ -2,13 +2,12 @@ from __future__ import annotations import json import logging -from datetime import timezone +from datetime import UTC, timezone from decimal import Decimal -from typing import Any, Literal, Type +from typing import Any, Literal, Self, Type from more_itertools import flatten from pydantic import BaseModel, ConfigDict, Field, computed_field, model_validator -from typing_extensions import Self from generalresearch.locales import Localelator from generalresearch.models import Source, TaskCalculationType @@ -297,7 +296,7 @@ class SpectrumSurvey(MarketplaceTask): return data @property - def condition_model(self) -> Type[MarketplaceCondition]: + def condition_model(self) -> type[MarketplaceCondition]: return SpectrumCondition @property @@ -389,16 +388,14 @@ class SpectrumSurvey(MarketplaceTask): @classmethod def from_db(cls, d: dict[str, Any]) -> Self: - d["created_api"] = d["created_api"].replace(tzinfo=timezone.utc) - d["updated"] = d["updated"].replace(tzinfo=timezone.utc) - d["modified_api"] = d["modified_api"].replace(tzinfo=timezone.utc) + d["created_api"] = d["created_api"].replace(tzinfo=UTC) + d["updated"] = d["updated"].replace(tzinfo=UTC) + d["modified_api"] = d["modified_api"].replace(tzinfo=UTC) d["field_end_date"] = ( - d["field_end_date"].replace(tzinfo=timezone.utc) - if d["field_end_date"] - else None + d["field_end_date"].replace(tzinfo=UTC) if d["field_end_date"] else None ) d["project_last_complete_date"] = ( - d["project_last_complete_date"].replace(tzinfo=timezone.utc) + d["project_last_complete_date"].replace(tzinfo=UTC) if d["project_last_complete_date"] else None ) diff --git a/generalresearch/models/string_utils.py b/generalresearch/models/string_utils.py index 23c1017..d76456f 100644 --- a/generalresearch/models/string_utils.py +++ b/generalresearch/models/string_utils.py @@ -2,7 +2,7 @@ import unicodedata from typing import Optional -def remove_nbsp(s: Optional[str]) -> Optional[str]: +def remove_nbsp(s: str | None) -> str | None: # Some text comes back from the API with lots of (copied from excel or # something), and random unicode... if s: diff --git a/generalresearch/models/thl/category.py b/generalresearch/models/thl/category.py index 1ed436a..ebfc840 100644 --- a/generalresearch/models/thl/category.py +++ b/generalresearch/models/thl/category.py @@ -1,10 +1,9 @@ from __future__ import annotations -from typing import Any +from typing import Any, Self from uuid import uuid4 from pydantic import BaseModel, Field, PositiveInt, model_validator -from typing_extensions import Self from generalresearch.models.custom_types import UUIDStr diff --git a/generalresearch/models/thl/contest/__init__.py b/generalresearch/models/thl/contest/__init__.py index 363c8c0..0444586 100644 --- a/generalresearch/models/thl/contest/__init__.py +++ b/generalresearch/models/thl/contest/__init__.py @@ -1,6 +1,7 @@ from __future__ import annotations -from datetime import datetime, timezone +from datetime import UTC, datetime, timezone +from typing import Self from uuid import uuid4 from pydantic import ( @@ -10,7 +11,6 @@ from pydantic import ( computed_field, model_validator, ) -from typing_extensions import Self from generalresearch.currency import USDCent from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr @@ -106,7 +106,7 @@ class ContestWinner(BaseModel): uuid: UUIDStr = Field(default_factory=lambda: uuid4().hex) created_at: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(tz=timezone.utc), + default_factory=lambda: datetime.now(tz=UTC), description="When this user won this prize", ) diff --git a/generalresearch/models/thl/contest/contest.py b/generalresearch/models/thl/contest/contest.py index 6889dcb..8c45f13 100644 --- a/generalresearch/models/thl/contest/contest.py +++ b/generalresearch/models/thl/contest/contest.py @@ -2,8 +2,8 @@ from __future__ import annotations import json from abc import ABC, abstractmethod -from datetime import datetime, timezone -from typing import Any +from datetime import UTC, datetime, timezone +from typing import Any, Self from uuid import uuid4 from pydantic import ( @@ -14,7 +14,6 @@ from pydantic import ( NonNegativeInt, model_validator, ) -from typing_extensions import Self from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.thl.contest import ( @@ -57,7 +56,7 @@ class ContestBase(BaseModel, ABC): starts_at: AwareDatetimeISO = Field( description="When the contest starts", - default_factory=lambda: datetime.now(tz=timezone.utc), + default_factory=lambda: datetime.now(tz=UTC), ) terms_and_conditions: HttpUrl | None = Field(default=None) @@ -91,11 +90,11 @@ class Contest(ContestBase): product_id: UUIDStr = Field(description="Contest applies only to a single BP") created_at: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(tz=timezone.utc), + default_factory=lambda: datetime.now(tz=UTC), description="When this contest was created", ) updated_at: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(tz=timezone.utc), + default_factory=lambda: datetime.now(tz=UTC), description="When this contest was last modified. Does not include " "entries being created/modified", ) @@ -139,7 +138,7 @@ class Contest(ContestBase): def should_end(self) -> tuple[bool, ContestEndReason | None]: if self.status == ContestStatus.ACTIVE: if self.end_condition.ends_at: - if datetime.now(tz=timezone.utc) >= self.end_condition.ends_at: + if datetime.now(tz=UTC) >= self.end_condition.ends_at: return True, ContestEndReason.ENDS_AT return False, None @@ -158,14 +157,14 @@ class Contest(ContestBase): if winners is not None: self.update( status=ContestStatus.COMPLETED, - ended_at=datetime.now(tz=timezone.utc), + ended_at=datetime.now(tz=UTC), end_reason=reason, all_winners=winners, ) else: self.update( status=ContestStatus.COMPLETED, - ended_at=datetime.now(tz=timezone.utc), + ended_at=datetime.now(tz=UTC), end_reason=reason, ) return None @@ -211,7 +210,7 @@ class ContestUserView(Contest): ) def is_user_eligible(self, country_iso: str) -> tuple[bool, str]: - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) assert country_iso.lower() == country_iso if now < self.starts_at: diff --git a/generalresearch/models/thl/contest/contest_entry.py b/generalresearch/models/thl/contest/contest_entry.py index cddae14..31ef317 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 datetime, timezone +from datetime import UTC, datetime, timezone from uuid import uuid4 from pydantic import ( @@ -39,12 +39,8 @@ class ContestEntryCreate(BaseModel): class ContestEntry(BaseModel): uuid: UUIDStr = Field(default_factory=lambda: uuid4().hex) - created_at: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(timezone.utc) - ) - updated_at: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(timezone.utc) - ) + created_at: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(UTC)) + updated_at: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(UTC)) # entry_type and amount are the same as on ContestEntryCreate entry_type: ContestEntryType = Field() diff --git a/generalresearch/models/thl/contest/io.py b/generalresearch/models/thl/contest/io.py index e68f76e..c6af719 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 datetime, timezone +from datetime import UTC, datetime, timezone from uuid import uuid4 from generalresearch.models.thl.contest.definitions import ContestType @@ -37,7 +37,7 @@ from generalresearch.models.thl.contest.contest import Contest def contest_create_to_contest( product_id: str, contest_create: ContestCreate ) -> Contest: - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) d = contest_create.model_dump(mode="json") d["uuid"] = uuid4().hex d["product_id"] = product_id diff --git a/generalresearch/models/thl/contest/leaderboard.py b/generalresearch/models/thl/contest/leaderboard.py index 8167f46..c5e0626 100644 --- a/generalresearch/models/thl/contest/leaderboard.py +++ b/generalresearch/models/thl/contest/leaderboard.py @@ -1,7 +1,7 @@ from __future__ import annotations -from datetime import datetime, timedelta, timezone -from typing import Any, Literal +from datetime import UTC, datetime, timedelta, timezone +from typing import Any, Literal, Self from pydantic import ( ConfigDict, @@ -11,7 +11,6 @@ from pydantic import ( model_validator, ) from redis import Redis -from typing_extensions import Self from generalresearch.decorators import LOG from generalresearch.managers.leaderboard import country_timezone @@ -195,7 +194,7 @@ class LeaderboardContest(LeaderboardContestCreate, Contest): def should_end(self) -> tuple[bool, ContestEndReason | None]: if self.status == ContestStatus.ACTIVE: if self.end_condition.ends_at: - if datetime.now(tz=timezone.utc) >= self.end_condition.ends_at: + if datetime.now(tz=UTC) >= self.end_condition.ends_at: return True, ContestEndReason.ENDS_AT return False, None @@ -276,7 +275,7 @@ class LeaderboardContestUserView(LeaderboardContest, ContestUserView): if self.user_winnings: return False, "User already won" - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) if self.leaderboard_model.period_end_utc < now: return False, "Contest is over" if self.leaderboard_model.period_start_utc > now: diff --git a/generalresearch/models/thl/contest/milestone.py b/generalresearch/models/thl/contest/milestone.py index f62be8f..8b74d50 100644 --- a/generalresearch/models/thl/contest/milestone.py +++ b/generalresearch/models/thl/contest/milestone.py @@ -10,7 +10,7 @@ from pydantic import ( Field, PositiveInt, ) -from typing_extensions import Self +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 d3157e3..b497a44 100644 --- a/generalresearch/models/thl/contest/raffle.py +++ b/generalresearch/models/thl/contest/raffle.py @@ -3,8 +3,8 @@ from __future__ import annotations import logging import random from collections import defaultdict -from datetime import datetime, timezone -from typing import Any, Literal +from datetime import UTC, datetime, timezone +from typing import Any, Literal, Self from pydantic import ( ConfigDict, @@ -14,7 +14,6 @@ from pydantic import ( model_validator, ) from scipy.stats import hypergeom -from typing_extensions import Self from generalresearch.currency import USDCent from generalresearch.models.thl.contest import ( @@ -202,7 +201,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=timezone.utc) >= c.ends_at: + if c.ends_at and datetime.now(tz=UTC) >= c.ends_at: return True return False diff --git a/generalresearch/models/thl/finance.py b/generalresearch/models/thl/finance.py index b72ecf6..a992f78 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 timezone +from datetime import UTC, timezone from typing import TYPE_CHECKING from uuid import uuid4 @@ -16,8 +16,8 @@ from pydantic import ( model_validator, ) from pydantic.json_schema import SkipJsonSchema -from generalresearch.config import is_debug +from generalresearch.config import is_debug from generalresearch.currency import USDCent from generalresearch.decorators import LOG from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr @@ -150,7 +150,7 @@ class POPFinancial(BaseModel): index: int # Not useful, just a RangeIndex row: pd.DataFrame - row["time_idx"] = row.time_idx.to_pydatetime().replace(tzinfo=timezone.utc) + row["time_idx"] = row.time_idx.to_pydatetime().replace(tzinfo=UTC) instance = ProductBalances.from_pandas(row) res.append( @@ -842,12 +842,12 @@ class BusinessBalances(BaseModel): 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 - from generalresearch.managers.thl.product import ProductManager # Validate the input accounts assert len(accounts) > 0, "Must provide accounts" diff --git a/generalresearch/models/thl/ipinfo.py b/generalresearch/models/thl/ipinfo.py index 0d254e2..3f212cf 100644 --- a/generalresearch/models/thl/ipinfo.py +++ b/generalresearch/models/thl/ipinfo.py @@ -1,8 +1,8 @@ from __future__ import annotations import ipaddress -from datetime import datetime, timezone -from typing import Any, Literal +from datetime import UTC, datetime, timezone +from typing import Any, Literal, Self from faker import Faker from pydantic import ( @@ -13,7 +13,6 @@ from pydantic import ( PrivateAttr, field_validator, ) -from typing_extensions import Self from generalresearch.models.custom_types import ( AwareDatetimeISO, @@ -95,7 +94,7 @@ class IPGeoname(BaseModel): is_in_european_union: bool | None = Field(default=None) updated: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(tz=timezone.utc), + default_factory=lambda: datetime.now(tz=UTC), ) @field_validator( @@ -119,7 +118,7 @@ class IPGeoname(BaseModel): @classmethod def from_mysql(cls, d: dict[str, Any]) -> Self: - d["updated"] = d["updated"].replace(tzinfo=timezone.utc) + d["updated"] = d["updated"].replace(tzinfo=UTC) return cls.model_validate(d) @@ -205,7 +204,7 @@ class IPInformation(BaseModel): ) updated: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(tz=timezone.utc), + default_factory=lambda: datetime.now(tz=UTC), ) _geoname: IPGeoname | None = PrivateAttr(default=None) @@ -255,7 +254,7 @@ class IPInformation(BaseModel): @classmethod def from_mysql(cls, d: dict) -> Self: - d["updated"] = d["updated"].replace(tzinfo=timezone.utc) + d["updated"] = d["updated"].replace(tzinfo=UTC) return cls.model_validate(d) diff --git a/generalresearch/models/thl/leaderboard.py b/generalresearch/models/thl/leaderboard.py index 399a906..dce3280 100644 --- a/generalresearch/models/thl/leaderboard.py +++ b/generalresearch/models/thl/leaderboard.py @@ -2,10 +2,11 @@ from __future__ import annotations import logging import math -from datetime import datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta, timezone from enum import Enum from typing import Literal from uuid import UUID, uuid3 +from zoneinfo import ZoneInfo import pandas as pd from pydantic import ( @@ -17,7 +18,6 @@ from pydantic import ( field_validator, model_validator, ) -from zoneinfo import ZoneInfo from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.legacy.api_status import StatusResponse @@ -167,13 +167,13 @@ class Leaderboard(BaseModel): def period_start_utc(self) -> datetime: # The start of the time period covered by this board in UTC, tz-aware # e.g. datetime(2024, 7, 12, 4, 0, 0, 0, tzinfo=timezone.utc) - return self.period_start_local.astimezone(timezone.utc) + return self.period_start_local.astimezone(UTC) @property def period_end_utc(self) -> datetime: # The end of the time period covered by this board in UTC, tz-aware # e.g. datetime(2024, 7, 13, 3, 59, 59, 999999, tzinfo=timezone.utc) - return self.period_end_local.astimezone(timezone.utc) + return self.period_end_local.astimezone(UTC) @computed_field( description="(unix timestamp) The start time of the time range this leaderboard covers.", diff --git a/generalresearch/models/thl/ledger.py b/generalresearch/models/thl/ledger.py index 3f8b123..dd37d98 100644 --- a/generalresearch/models/thl/ledger.py +++ b/generalresearch/models/thl/ledger.py @@ -1,8 +1,8 @@ from __future__ import annotations -from datetime import datetime, timezone +from datetime import UTC, datetime, timezone from enum import Enum -from typing import Annotated, Any, Literal, Union +from typing import Annotated, Any, Literal, Self, Union from uuid import uuid4 from pydantic import ( @@ -15,7 +15,6 @@ from pydantic import ( field_validator, model_validator, ) -from typing_extensions import Self from generalresearch.models.custom_types import ( AwareDatetimeISO, @@ -253,7 +252,7 @@ class LedgerTransaction(BaseModel): id: int | None = Field(default=None) created: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(tz=timezone.utc), + default_factory=lambda: datetime.now(tz=UTC), description="When the Transaction (TX) was created into the database." "This does not represent the exact time for any action" "which may be responsible for this Transaction (TX), and " @@ -283,9 +282,7 @@ class LedgerTransaction(BaseModel): """Created should not be in the future. This will mess up LedgerAccountStatement / groupby rollups. """ - assert ( - datetime.now(tz=timezone.utc) > created - ), "created cannot be in the future" + assert datetime.now(tz=UTC) > created, "created cannot be in the future" return created @field_validator("entries", mode="after") @@ -536,12 +533,10 @@ class UserLedgerTransactionTaskAdjustment(UserLedgerTransaction): UserLedgerTransactionType = Annotated[ - Union[ - UserLedgerTransactionUserPayout, - UserLedgerTransactionUserBonus, - UserLedgerTransactionTaskAdjustment, - UserLedgerTransactionTaskComplete, - ], + UserLedgerTransactionUserPayout + | UserLedgerTransactionUserBonus + | UserLedgerTransactionTaskAdjustment + | UserLedgerTransactionTaskComplete, Field(discriminator="tx_type"), ] diff --git a/generalresearch/models/thl/ledger_example.py b/generalresearch/models/thl/ledger_example.py index 767be85..92ad83d 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 datetime, timezone +from datetime import UTC, datetime, timezone from typing import Any from uuid import uuid4 @@ -16,7 +16,7 @@ def _example_user_tx_payout(schema: dict[str, Any]) -> None: amount=-5, description="HIT Reward", payout_format="${payout/100:.2f}", - created=datetime.now(tz=timezone.utc), + created=datetime.now(tz=UTC), ).model_dump(mode="json") @@ -30,7 +30,7 @@ def _example_user_tx_bonus(schema: dict[str, Any]) -> None: amount=100, description="Compensation Bonus", payout_format="${payout/100:.2f}", - created=datetime.now(tz=timezone.utc), + created=datetime.now(tz=UTC), ).model_dump(mode="json") @@ -44,7 +44,7 @@ def _example_user_tx_complete(schema: dict[str, Any]) -> None: amount=38, description="Task Complete", payout_format="${payout/100:.2f}", - created=datetime.now(tz=timezone.utc), + created=datetime.now(tz=UTC), tsid=uuid4().hex, ).model_dump(mode="json") @@ -59,6 +59,6 @@ def _example_user_tx_adjustment(schema: dict[str, Any]) -> None: amount=-38, description="Task Adjustment", payout_format="${payout/100:.2f}", - created=datetime.now(tz=timezone.utc), + created=datetime.now(tz=UTC), tsid=uuid4().hex, ).model_dump(mode="json") diff --git a/generalresearch/models/thl/offerwall/__init__.py b/generalresearch/models/thl/offerwall/__init__.py index d2d7d36..3acc592 100644 --- a/generalresearch/models/thl/offerwall/__init__.py +++ b/generalresearch/models/thl/offerwall/__init__.py @@ -4,7 +4,7 @@ import hashlib import json from decimal import Decimal from enum import Enum -from typing import Any, Literal +from typing import Any, Literal, Self from pydantic import ( BaseModel, @@ -13,7 +13,6 @@ from pydantic import ( computed_field, model_validator, ) -from typing_extensions import Self from generalresearch.models import Source from generalresearch.models.custom_types import IPvAnyAddressStr diff --git a/generalresearch/models/thl/offerwall/base.py b/generalresearch/models/thl/offerwall/base.py index 3a867b6..33489df 100644 --- a/generalresearch/models/thl/offerwall/base.py +++ b/generalresearch/models/thl/offerwall/base.py @@ -4,7 +4,7 @@ import statistics from datetime import timedelta from decimal import Decimal from string import Formatter -from typing import Any +from typing import Annotated, Any, Self from uuid import uuid4 import numpy as np @@ -18,7 +18,6 @@ from pydantic import ( field_validator, model_validator, ) -from typing_extensions import Annotated, Self from generalresearch.models import Source from generalresearch.models.custom_types import HttpsUrl, UUIDStr diff --git a/generalresearch/models/thl/offerwall/cache.py b/generalresearch/models/thl/offerwall/cache.py index c36568e..82ab36d 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 datetime, timezone +from datetime import UTC, datetime, timezone from typing import Any from pydantic import BaseModel, Field @@ -26,9 +26,7 @@ class GetOfferWallCache(BaseModel): request_id: str = Field() offerwall: OfferwallBase = Field() all_sids: list[str] = Field() - timestamp: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(timezone.utc) - ) + timestamp: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(UTC)) latest_ip_info: dict[str, Any] = Field( description="So we can easily check if user's IP info has changed" ) @@ -51,9 +49,7 @@ class SessionInfoCache(BaseModel): # will get pruned as tasks are attempted tasks: list[ScoredTaskResult] = Field() - started: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(tz=timezone.utc) - ) + started: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC)) # The count of attempts per marketplace mp_retry_count: dict[Source, int] = Field(default_factory=dict) diff --git a/generalresearch/models/thl/payout_format.py b/generalresearch/models/thl/payout_format.py index 9f22ace..a7fc4fb 100644 --- a/generalresearch/models/thl/payout_format.py +++ b/generalresearch/models/thl/payout_format.py @@ -2,9 +2,9 @@ from __future__ import annotations import decimal import re +from typing import Annotated from pydantic import AfterValidator, Field -from typing_extensions import Annotated # Matches only digits, parenthesis, + , -, *, / and the string payout. xform_format_re = re.compile(pattern=r"^[\d()+\-*/.]*payout[\d()+\-*/.]*$") diff --git a/generalresearch/models/thl/product.py b/generalresearch/models/thl/product.py index 9b7d66a..8377e45 100644 --- a/generalresearch/models/thl/product.py +++ b/generalresearch/models/thl/product.py @@ -6,14 +6,15 @@ import json import math import warnings from collections import defaultdict +from collections.abc import Callable from decimal import Decimal from enum import Enum from functools import cached_property, partial from typing import ( TYPE_CHECKING, Any, - Callable, Literal, + Self, ) from urllib.parse import parse_qs, urlencode, urlsplit, urlunsplit from uuid import uuid4 @@ -34,7 +35,6 @@ from pydantic import ( model_validator, ) from pydantic.json_schema import SkipJsonSchema -from typing_extensions import Self from generalresearch.currency import USDCent from generalresearch.decorators import LOG @@ -941,9 +941,7 @@ class Product(BaseModel, validate_assignment=True): # Initialization is deferred until unless it's called # (see .prebuild_***()) - balance: ProductBalances | None = Field( - default=None, description="Product Balance" - ) + balance: ProductBalances | None = Field(default=None, description="Product Balance") payouts_total_str: str | None = Field(default=None) payouts_total: USDCent | None = Field(default=None) diff --git a/generalresearch/models/thl/profiling/marketplace.py b/generalresearch/models/thl/profiling/marketplace.py index 027aa4c..9038cf6 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 datetime, timezone +from datetime import UTC, datetime, timezone from functools import cached_property from typing import Any @@ -111,9 +111,7 @@ class MarketplaceUserQuestionAnswer(BaseModel): # This may be a pipe-separated string if the question_type is multi. Regex # means any chars except capital letters option_id: str = Field(pattern=r"^[^A-Z]*$") - created: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(tz=timezone.utc) - ) + created: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC)) country_iso: CountryISO = Field(frozen=True) language_iso: LanguageISO = Field(frozen=True) diff --git a/generalresearch/models/thl/profiling/other_option.py b/generalresearch/models/thl/profiling/other_option.py index ae58e90..6d789e5 100644 --- a/generalresearch/models/thl/profiling/other_option.py +++ b/generalresearch/models/thl/profiling/other_option.py @@ -40,7 +40,7 @@ texts_in = { } -def option_is_catch_all(c: "UpkQuestionChoice") -> bool: +def option_is_catch_all(c: UpkQuestionChoice) -> bool: """ Exclusive not specifically in the sense that it is a multi-select question and if this option is selected no others can be selected. But also in the diff --git a/generalresearch/models/thl/profiling/upk_question.py b/generalresearch/models/thl/profiling/upk_question.py index 307bc33..78f9511 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 Enum from functools import cached_property -from typing import Any, List, Literal, Union +from typing import Annotated, Any, List, Literal, Union from pydantic import ( BaseModel, @@ -16,7 +16,6 @@ from pydantic import ( field_validator, model_validator, ) -from typing_extensions import Annotated from generalresearch.models import Source from generalresearch.models.custom_types import UUIDStr @@ -215,11 +214,9 @@ SelectorType = ( | UpkQuestionSelectorHIDDEN ) Configuration = Annotated[ - Union[ - UpkQuestionConfigurationMC, - UpkQuestionConfigurationTE, - UpkQuestionConfigurationSLIDER, - ], + UpkQuestionConfigurationMC + | UpkQuestionConfigurationTE + | UpkQuestionConfigurationSLIDER, Field(discriminator="type"), ] @@ -433,7 +430,7 @@ class UpkQuestion(BaseModel): @field_validator("choices") @classmethod - def order_choices(cls, choices: List): + def order_choices(cls, choices: list): if choices: choices.sort(key=lambda x: x.order) return choices diff --git a/generalresearch/models/thl/profiling/upk_question_answer.py b/generalresearch/models/thl/profiling/upk_question_answer.py index 0024e68..c59d99d 100644 --- a/generalresearch/models/thl/profiling/upk_question_answer.py +++ b/generalresearch/models/thl/profiling/upk_question_answer.py @@ -1,7 +1,7 @@ from __future__ import annotations -from datetime import datetime, timezone -from typing import Any +from datetime import UTC, datetime, timezone +from typing import Any, Self from uuid import uuid4 from pydantic import ( @@ -12,7 +12,6 @@ from pydantic import ( computed_field, model_validator, ) -from typing_extensions import Self from generalresearch.models import MAX_INT32 from generalresearch.models.custom_types import ( @@ -60,9 +59,7 @@ class UpkQuestionAnswer(BaseModel): # ISO 3166-1 alpha-2 (two-letter codes, lowercase) country_iso: CountryISOLike = Field() - created: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(tz=timezone.utc) - ) + created: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC)) # If the property is PropertyType.UPK_ITEM, it should have an item (and no value). # If the property is UPK_NUMERICAL or UPK_TEXT, it'll have a value (and no item). diff --git a/generalresearch/models/thl/profiling/user_question_answer.py b/generalresearch/models/thl/profiling/user_question_answer.py index 8248623..a55b205 100644 --- a/generalresearch/models/thl/profiling/user_question_answer.py +++ b/generalresearch/models/thl/profiling/user_question_answer.py @@ -1,8 +1,9 @@ from __future__ import annotations import json -from datetime import datetime, timedelta, timezone -from typing import Any, Iterator, Literal +from collections.abc import Iterator +from datetime import UTC, datetime, timedelta, timezone +from typing import Any, Literal, Self from pydantic import ( BaseModel, @@ -12,7 +13,6 @@ from pydantic import ( field_validator, model_validator, ) -from typing_extensions import Self from generalresearch.grpc import timestamp_to_datetime from generalresearch.models import MAX_INT32, Source @@ -28,9 +28,7 @@ class UserQuestionAnswer(BaseModel): user_id: PositiveInt | None = Field(lt=MAX_INT32, default=None) question_id: UUIDStr = Field() answer: tuple[str, ...] = Field() - timestamp: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(tz=timezone.utc) - ) + timestamp: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC)) country_iso: CountryISO | Literal["xx"] = Field() language_iso: LanguageISO | Literal["xxx"] = Field() @@ -105,7 +103,7 @@ class UserQuestionAnswer(BaseModel): return True, "" def is_stale(self) -> bool: - return self.timestamp < datetime.now(tz=timezone.utc) - timedelta(days=30) + return self.timestamp < datetime.now(tz=UTC) - timedelta(days=30) @classmethod def from_grpc(cls, msg, default_timestamp: datetime) -> Self: @@ -129,7 +127,7 @@ class UserQuestionAnswer(BaseModel): DUMMY_UQA = UserQuestionAnswer( question_id="f118edd01cf1476ba7200a175fb4351d", answer=("0",), - timestamp=datetime(2020, 1, 1, tzinfo=timezone.utc), + timestamp=datetime(2020, 1, 1, tzinfo=UTC), country_iso="xx", language_iso="xxx", property_code="dummy", diff --git a/generalresearch/models/thl/session.py b/generalresearch/models/thl/session.py index 17142f3..c8e681c 100644 --- a/generalresearch/models/thl/session.py +++ b/generalresearch/models/thl/session.py @@ -2,9 +2,9 @@ from __future__ import annotations import json import logging -from datetime import datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta, timezone from decimal import Decimal -from typing import TYPE_CHECKING, Annotated, Any +from typing import TYPE_CHECKING, Annotated, Any, Self from uuid import uuid4 from pydantic import ( @@ -17,7 +17,6 @@ from pydantic import ( field_validator, model_validator, ) -from typing_extensions import Self from generalresearch.models import DeviceType, Source from generalresearch.models.custom_types import ( @@ -69,9 +68,7 @@ class WallBase(BaseModel): buyer_id: str | None = Field(default=None, max_length=32) req_survey_id: str = Field(max_length=32) req_cpi: Decimal = Field(decimal_places=5, lt=1000, ge=0) - started: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(tz=timezone.utc) - ) + started: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC)) # These get set on creation, or updated when the wall event is finished. So # they shouldn't really ever be NULL, but you don't have to pass them in @@ -158,9 +155,7 @@ class WallBase(BaseModel): @model_validator(mode="after") def check_timestamps(self): - assert self.started <= datetime.now( - tz=timezone.utc - ), "Started must not be in the future" + assert self.started <= datetime.now(tz=UTC), "Started must not be in the future" if self.finished: assert self.finished > self.started, "Finished must be after started" assert self.finished - self.started <= timedelta( @@ -278,7 +273,7 @@ class WallBase(BaseModel): # This is just used in tests at the moment. This needs to be adjusted. if finished is None: - finished = datetime.now(tz=timezone.utc) + finished = datetime.now(tz=UTC) self.update( status=status, @@ -313,7 +308,7 @@ class WallBase(BaseModel): ext_status_code_3, ) if finished is None: - finished = datetime.now(tz=timezone.utc) + finished = datetime.now(tz=UTC) self.update( status=status, status_code_1=status_code_1, @@ -386,7 +381,7 @@ class WallBase(BaseModel): TODO: Transition this over to use the ReportTask pydantic model. """ report_timestamp = ( - report_timestamp if report_timestamp else datetime.now(tz=timezone.utc) + report_timestamp if report_timestamp else datetime.now(tz=UTC) ) if self.status is None and self.finished is None: self.status = Status.ABANDON @@ -587,9 +582,7 @@ class Session(BaseModel): id: int | None = None uuid: UUIDStr = Field(default_factory=lambda: uuid4().hex) user: User - started: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(tz=timezone.utc) - ) + started: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC)) # This is the "bucket" the user clicked on to start this session. We only # store the 4 fields: loi_min, loi_max, user_payout_min, user_payout_max @@ -881,7 +874,7 @@ class Session(BaseModel): if ( last_wall.status is None and self.status is None - and datetime.now(tz=timezone.utc) + and datetime.now(tz=UTC) > self.started + timedelta(seconds=task_timeout_seconds) ): last_wall.status = Status.TIMEOUT @@ -962,7 +955,7 @@ class Session(BaseModel): self, max_session_len: timedelta, max_session_hard_retry: int ) -> bool: - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) last_wall = self.get_last_visible_wall() if last_wall and last_wall.status == Status.COMPLETE: diff --git a/generalresearch/models/thl/survey/__init__.py b/generalresearch/models/thl/survey/__init__.py index b6ac740..d6ba895 100644 --- a/generalresearch/models/thl/survey/__init__.py +++ b/generalresearch/models/thl/survey/__init__.py @@ -108,7 +108,7 @@ class MarketplaceTask(BaseModel, ABC): @property @abstractmethod - def condition_model(self) -> Type[MarketplaceCondition]: + def condition_model(self) -> type[MarketplaceCondition]: """ The Condition Model for this survey class """ diff --git a/generalresearch/models/thl/survey/buyer.py b/generalresearch/models/thl/survey/buyer.py index 6d4d7a1..384bab4 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 datetime, timezone +from datetime import UTC, datetime, timezone from decimal import Decimal from math import log from typing import Annotated @@ -46,7 +46,7 @@ class Buyer(BaseModel): ) label: str | None = Field(default=None, max_length=255) created: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(tz=timezone.utc), + default_factory=lambda: datetime.now(tz=UTC), description="When this entry was made, or when the buyer was first seen", ) diff --git a/generalresearch/models/thl/survey/condition.py b/generalresearch/models/thl/survey/condition.py index 927b7e1..a85073c 100644 --- a/generalresearch/models/thl/survey/condition.py +++ b/generalresearch/models/thl/survey/condition.py @@ -4,7 +4,7 @@ import hashlib from abc import ABC from enum import Enum from functools import cached_property -from typing import Any +from typing import Annotated, Any, Self from pydantic import ( BaseModel, @@ -16,7 +16,6 @@ from pydantic import ( field_validator, model_validator, ) -from typing_extensions import Annotated, Self from generalresearch.models import LogicalOperator diff --git a/generalresearch/models/thl/survey/model.py b/generalresearch/models/thl/survey/model.py index 3794c00..57bcbe2 100644 --- a/generalresearch/models/thl/survey/model.py +++ b/generalresearch/models/thl/survey/model.py @@ -1,8 +1,8 @@ from __future__ import annotations -from datetime import datetime, timezone +from datetime import UTC, datetime, timezone from decimal import Decimal -from typing import Any +from typing import Annotated, Any from pydantic import ( BaseModel, @@ -15,7 +15,6 @@ from pydantic import ( field_validator, model_validator, ) -from typing_extensions import Annotated from generalresearch.managers.thl.buyer import Buyer from generalresearch.models import Source @@ -71,12 +70,8 @@ class Survey(BaseModel): min_length=1, max_length=128, default=None, examples=["124"] ) - created_at: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(tz=timezone.utc) - ) - updated_at: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(tz=timezone.utc) - ) + created_at: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC)) + updated_at: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC)) is_live: bool = Field(default=True) is_recontact: bool = Field(default=False) @@ -188,9 +183,7 @@ class SurveyStat(BaseModel): # ---- Metadata ---- - updated_at: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(tz=timezone.utc) - ) + updated_at: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC)) @property def natural_key(self) -> str: diff --git a/generalresearch/models/thl/survey/penalty.py b/generalresearch/models/thl/survey/penalty.py index e9515d4..05153fe 100644 --- a/generalresearch/models/thl/survey/penalty.py +++ b/generalresearch/models/thl/survey/penalty.py @@ -1,11 +1,10 @@ from __future__ import annotations import abc -from datetime import datetime, timezone -from typing import Literal +from datetime import UTC, datetime, timezone +from typing import Annotated, Literal from pydantic import BaseModel, ConfigDict, Field, TypeAdapter -from typing_extensions import Annotated from generalresearch.models import Source from generalresearch.models.custom_types import ( @@ -29,9 +28,7 @@ class SurveyPenalty(BaseModel, abc.ABC): penalty: float = Field(ge=0, le=1) - created: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(tz=timezone.utc) - ) + created: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC)) @property def sid(self): diff --git a/generalresearch/models/thl/task_adjustment.py b/generalresearch/models/thl/task_adjustment.py index 89a3873..1834898 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 datetime, timezone +from datetime import UTC, datetime, timezone from decimal import Decimal from uuid import uuid4 @@ -26,11 +26,11 @@ class TaskAdjustmentEvent(BaseModel): uuid: UUIDStr = Field(default_factory=lambda: uuid4().hex) created: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(tz=timezone.utc), + default_factory=lambda: datetime.now(tz=UTC), description="When this event was created in the db", ) alerted: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(tz=timezone.utc), + default_factory=lambda: datetime.now(tz=UTC), description="When we were notified about this change", ) diff --git a/generalresearch/models/thl/task_status.py b/generalresearch/models/thl/task_status.py index ee713a5..011d743 100644 --- a/generalresearch/models/thl/task_status.py +++ b/generalresearch/models/thl/task_status.py @@ -1,7 +1,7 @@ from __future__ import annotations from datetime import datetime -from typing import Annotated, Any, Literal +from typing import Annotated, Any, Literal, Self from pydantic import ( BaseModel, @@ -12,7 +12,6 @@ from pydantic import ( field_validator, model_validator, ) -from typing_extensions import Self from generalresearch.models.custom_types import ( AwareDatetimeISO, diff --git a/generalresearch/models/thl/user.py b/generalresearch/models/thl/user.py index 355e331..55bbd18 100644 --- a/generalresearch/models/thl/user.py +++ b/generalresearch/models/thl/user.py @@ -3,8 +3,8 @@ from __future__ import annotations import json import logging import re -from datetime import datetime, timezone -from typing import TYPE_CHECKING +from datetime import UTC, datetime, timezone +from typing import TYPE_CHECKING, Annotated, Self from uuid import UUID, uuid4 from pydantic import ( @@ -19,7 +19,6 @@ from pydantic import ( model_validator, ) from sentry_sdk import set_tag, set_user -from typing_extensions import Annotated, Self from generalresearch.models import MAX_INT32 from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr @@ -126,7 +125,7 @@ class User(BaseModel): def check_not_in_future(cls, v: AwareDatetime) -> AwareDatetime: if v is not None: try: - assert v < datetime.now(tz=timezone.utc) + assert v < datetime.now(tz=UTC) except Exception: raise ValueError("Input is in the future") return v @@ -137,7 +136,7 @@ class User(BaseModel): def check_after_anno_domini(cls, v: AwareDatetime) -> AwareDatetime: if v is not None: try: - assert v > datetime(year=2016, month=7, day=13, tzinfo=timezone.utc) + assert v > datetime(year=2016, month=7, day=13, tzinfo=UTC) except Exception: raise ValueError("Input is before Anno Domini") return v @@ -294,9 +293,9 @@ class User(BaseModel): @classmethod def from_db(cls, res) -> Self: if res["created"]: - res["created"] = res["created"].replace(tzinfo=timezone.utc) + res["created"] = res["created"].replace(tzinfo=UTC) if res["last_seen"]: - res["last_seen"] = res["last_seen"].replace(tzinfo=timezone.utc) + res["last_seen"] = res["last_seen"].replace(tzinfo=UTC) res["product_id"] = UUID(res["product_id"]).hex res["uuid"] = UUID(res["uuid"]).hex return cls( diff --git a/generalresearch/models/thl/user_iphistory.py b/generalresearch/models/thl/user_iphistory.py index 5892d41..469f8ba 100644 --- a/generalresearch/models/thl/user_iphistory.py +++ b/generalresearch/models/thl/user_iphistory.py @@ -1,7 +1,8 @@ from __future__ import annotations import ipaddress -from datetime import datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta, timezone +from typing import Self from faker import Faker from pydantic import ( @@ -11,7 +12,6 @@ from pydantic import ( PositiveInt, field_validator, ) -from typing_extensions import Self from generalresearch.models.custom_types import ( AwareDatetimeISO, @@ -113,7 +113,7 @@ class IPRecord(BaseModel): # --- ORM --- @classmethod def from_mysql(cls, d: dict) -> Self: - created = d["created"].replace(tzinfo=timezone.utc) + created = d["created"].replace(tzinfo=UTC) d["created"] = created d["forwarded_ip_records"] = [] @@ -169,7 +169,7 @@ class UserIPHistory(BaseModel): def ips_timestamp(cls, ips): if ips is None: return None - cutoff = datetime.now(tz=timezone.utc) - timedelta(days=28) + cutoff = datetime.now(tz=UTC) - timedelta(days=28) return sorted( [x for x in ips if x.created > cutoff], key=lambda x: x.created, diff --git a/generalresearch/models/thl/user_profile.py b/generalresearch/models/thl/user_profile.py index e96266a..0ec605a 100644 --- a/generalresearch/models/thl/user_profile.py +++ b/generalresearch/models/thl/user_profile.py @@ -1,7 +1,7 @@ from __future__ import annotations import hashlib -from typing import Any +from typing import Annotated, Any, Self from pydantic import ( BaseModel, @@ -12,7 +12,6 @@ from pydantic import ( computed_field, ) from pydantic.json_schema import SkipJsonSchema -from typing_extensions import Annotated, Self from generalresearch.models import MAX_INT32, Source 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 d6ebddc..52903e3 100644 --- a/generalresearch/models/thl/user_quality_event.py +++ b/generalresearch/models/thl/user_quality_event.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import datetime, timezone +from datetime import UTC, datetime, timezone from decimal import Decimal from enum import Enum from typing import Literal @@ -61,9 +61,7 @@ class TaskAdjustmentEvent(BaseModel): mid: UUIDStr = Field() source: Source = Field() status: WallAdjustedStatus = Field() - alert_time: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(tz=timezone.utc) - ) + alert_time: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC)) quality_event_type: Literal[QualityEventType.task_adjustment] = Field( default=QualityEventType.task_adjustment ) diff --git a/generalresearch/models/thl/userhealth.py b/generalresearch/models/thl/userhealth.py index e556dc8..fb15572 100644 --- a/generalresearch/models/thl/userhealth.py +++ b/generalresearch/models/thl/userhealth.py @@ -1,11 +1,10 @@ from __future__ import annotations -from datetime import datetime, timezone +from datetime import UTC, datetime, timezone from enum import Enum -from typing import Dict, Optional +from typing import Dict, Optional, Self from pydantic import BaseModel, Field, NonNegativeFloat, PositiveInt -from typing_extensions import Self from generalresearch.models.custom_types import AwareDatetimeISO @@ -26,12 +25,12 @@ class AuditLog(BaseModel): are related to a User """ - id: Optional[PositiveInt] = Field(default=None) + id: PositiveInt | None = Field(default=None) user_id: PositiveInt = Field() created: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(tz=timezone.utc), - examples=[datetime.now(tz=timezone.utc)], + default_factory=lambda: datetime.now(tz=UTC), + examples=[datetime.now(tz=UTC)], description="When did this event occur", ) @@ -51,14 +50,14 @@ class AuditLog(BaseModel): # e.g. "upk-audit", "ip-audit", "entrance-limit" event_type: str = Field(max_length=64, examples=["entrance-limit"]) - event_msg: Optional[str] = Field( + event_msg: str | None = Field( default=None, min_length=3, max_length=256, description="The event message. Could be displayed on user's page", ) - event_value: Optional[NonNegativeFloat] = Field( + event_value: NonNegativeFloat | None = Field( default=None, description="Optionally store a numeric value associated with this " "event. For e.g. if we recalculate the user's normalized " @@ -68,12 +67,12 @@ class AuditLog(BaseModel): examples=[0.42], ) - def model_dump_mysql(self, **kwargs) -> Dict: + def model_dump_mysql(self, **kwargs) -> dict: d = self.model_dump(mode="json", **kwargs) d["created"] = self.created.replace(tzinfo=None) return d @classmethod - def from_mysql(cls, d: Dict) -> Self: - d["created"] = d["created"].replace(tzinfo=timezone.utc) + def from_mysql(cls, d: dict) -> Self: + d["created"] = d["created"].replace(tzinfo=UTC) return AuditLog.model_validate(d) diff --git a/generalresearch/models/thl/wallet/cashout_method.py b/generalresearch/models/thl/wallet/cashout_method.py index 59cf721..4757eba 100644 --- a/generalresearch/models/thl/wallet/cashout_method.py +++ b/generalresearch/models/thl/wallet/cashout_method.py @@ -2,7 +2,7 @@ from __future__ import annotations import hashlib import logging -from datetime import datetime, timezone +from datetime import datetime, timezone, UTC from enum import Enum from typing import Any, Literal @@ -16,7 +16,7 @@ from pydantic import ( field_validator, model_validator, ) -from typing_extensions import Self +from typing import Self from generalresearch.currency import USDCent from generalresearch.models.custom_types import ( @@ -145,7 +145,7 @@ class CashoutMethod(CashoutMethodBase): "email associated.", ) last_updated: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(tz=timezone.utc) + default_factory=lambda: datetime.now(tz=UTC) ) is_live: bool = Field(default=True) diff --git a/generalresearch/models/thl/wallet/payout.py b/generalresearch/models/thl/wallet/payout.py index 8c78bef..cbb37fe 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 datetime, timezone +from datetime import UTC, datetime, timezone from typing import Any from uuid import uuid4 @@ -50,9 +50,7 @@ class PayoutEvent(BaseModel, validate_assignment=True): # populated from the db and so does not need to be set (there is no # `description` field in event_payout) description: str | None = Field(default=None) - created: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(tz=timezone.utc) - ) + created: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC)) # In the smallest unit of the currency being transacted. For USD, this # is cents. @@ -159,7 +157,7 @@ class BPPayoutEvent(BaseModel): created: AwareDatetimeISO = Field( description="When the Brokerage Product was paid out", - default_factory=lambda: datetime.now(tz=timezone.utc), + default_factory=lambda: datetime.now(tz=UTC), ) amount: USDCent = Field( diff --git a/generalresearch/pg_helper.py b/generalresearch/pg_helper.py index b9a7d79..b5e124a 100644 --- a/generalresearch/pg_helper.py +++ b/generalresearch/pg_helper.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import timezone +from datetime import UTC, timezone import psycopg from psycopg.adapt import Buffer @@ -24,7 +24,7 @@ class UTCTimestampLoader(TimestampLoader): if dt is None: return None assert dt.tzinfo is None, "expected naive dt" - return dt.replace(tzinfo=timezone.utc) + return dt.replace(tzinfo=UTC) class BPCharLoader(TextLoader): diff --git a/generalresearch/sql_helper.py b/generalresearch/sql_helper.py index 2c526f3..2bbdbe5 100644 --- a/generalresearch/sql_helper.py +++ b/generalresearch/sql_helper.py @@ -130,7 +130,7 @@ def decode_uuids(row: dict[str, Any]) -> dict[str, Any]: class SqlHelper(SqlConnector): - def __init__(self, dsn: Optional[DataBaseDsn] = None, **kwargs): + def __init__(self, dsn: DataBaseDsn | None = None, **kwargs): super().__init__(dsn, **kwargs) def execute_sql_query( @@ -281,7 +281,7 @@ class SqlHelper(SqlConnector): cursor=None, commit=True, primary_key=None, - ) -> Optional[int]: + ) -> int | None: """ Create the item in table `table_name`. In postgresql, `primary_key` needs to be given in order to return the diff --git a/generalresearch/utils/aggregation.py b/generalresearch/utils/aggregation.py index 4023dc9..bd962f1 100644 --- a/generalresearch/utils/aggregation.py +++ b/generalresearch/utils/aggregation.py @@ -2,7 +2,7 @@ from collections import defaultdict from typing import Any, Dict, List -def group_by_year(records: List[Dict], datetime_field: str) -> Dict[int, List[Any]]: +def group_by_year(records: list[dict], datetime_field: str) -> dict[int, list[Any]]: """Memory efficient - processes records one at a time""" by_year = defaultdict(list) diff --git a/generalresearch/utils/copying_cache.py b/generalresearch/utils/copying_cache.py index ea13f69..a1cb37c 100644 --- a/generalresearch/utils/copying_cache.py +++ b/generalresearch/utils/copying_cache.py @@ -1,6 +1,6 @@ +from collections.abc import Callable from copy import deepcopy from functools import wraps -from typing import Callable def deepcopy_return(fn: Callable) -> Callable: diff --git a/generalresearch/utils/enum.py b/generalresearch/utils/enum.py index 14a31de..56706ba 100644 --- a/generalresearch/utils/enum.py +++ b/generalresearch/utils/enum.py @@ -41,7 +41,7 @@ class ReprEnumMeta(EnumMeta): ) -def get_enum_comments(enum_class) -> Dict: +def get_enum_comments(enum_class) -> dict: source = inspect.getsource(enum_class) # Regular expression to match multi-line comments and enum values pattern = re.compile(r"((?:\s*#.*?\n)+)\s*(\w+)\s*=") diff --git a/generalresearch/wall_status_codes/__init__.py b/generalresearch/wall_status_codes/__init__.py index 3a0abb8..1d80924 100644 --- a/generalresearch/wall_status_codes/__init__.py +++ b/generalresearch/wall_status_codes/__init__.py @@ -22,9 +22,9 @@ from generalresearch.wall_status_codes import ( def annotate_status_code( source: Source, ext_status_code_1: str, - ext_status_code_2: Optional[str] = None, - ext_status_code_3: Optional[str] = None, -) -> Tuple[Status, Optional[StatusCode1], Optional[str]]: + ext_status_code_2: str | None = None, + ext_status_code_3: str | None = None, +) -> tuple[Status, StatusCode1 | None, str | None]: """ :params ext_status_code_1: marketplace-dependent code :params ext_status_code_2: marketplace-dependent code diff --git a/generalresearch/wall_status_codes/cint.py b/generalresearch/wall_status_codes/cint.py index 8042cd2..ecb6219 100644 --- a/generalresearch/wall_status_codes/cint.py +++ b/generalresearch/wall_status_codes/cint.py @@ -6,9 +6,9 @@ from generalresearch.wall_status_codes import lucid def annotate_status_code( ext_status_code_1: str, - ext_status_code_2: Optional[str] = None, - ext_status_code_3: Optional[str] = None, -) -> Tuple[Status, StatusCode1, Optional[Any]]: + ext_status_code_2: str | None = None, + ext_status_code_3: str | None = None, +) -> tuple[Status, StatusCode1, Any | None]: return lucid.annotate_status_code( ext_status_code_1=ext_status_code_1, ext_status_code_2=ext_status_code_2, diff --git a/generalresearch/wall_status_codes/dynata.py b/generalresearch/wall_status_codes/dynata.py index 958e06f..69f3e74 100644 --- a/generalresearch/wall_status_codes/dynata.py +++ b/generalresearch/wall_status_codes/dynata.py @@ -8,7 +8,7 @@ from typing import Any, Dict, List, Optional, Tuple from generalresearch.models.thl.definitions import Status, StatusCode1 -status_codes_name: Dict[str, str] = { +status_codes_name: dict[str, str] = { "0.0": "Unknown", "0.1": "Missing Language", "0.2": "Missing Respondent ID", @@ -51,10 +51,10 @@ status_codes_name: Dict[str, str] = { "5.10": "Daily Limit", } -status_map: Dict[str, Status] = defaultdict( +status_map: dict[str, Status] = defaultdict( lambda: Status.FAIL, **{"1.0": Status.COMPLETE, "1.1": Status.COMPLETE} ) -status_codes_ext_map: Dict[StatusCode1, List[str]] = { +status_codes_ext_map: dict[StatusCode1, list[str]] = { StatusCode1.COMPLETE: ["1.0", "1.1"], StatusCode1.BUYER_FAIL: ["2.2", "3.2"], StatusCode1.BUYER_QUALITY_FAIL: ["5.1", "5.2"], @@ -88,10 +88,10 @@ status_codes_ext_map: Dict[StatusCode1, List[str]] = { "5.10", ], } -ext_status_code_map: Dict[str, StatusCode1] = dict() +ext_status_code_map: dict[str, StatusCode1] = dict() for k, v in status_codes_ext_map.items(): k: StatusCode1 - v: List[str] + v: list[str] for vv in v: vv: str @@ -100,9 +100,9 @@ for k, v in status_codes_ext_map.items(): def annotate_status_code( ext_status_code_1: str, - ext_status_code_2: Optional[str] = None, - ext_status_code_3: Optional[str] = None, -) -> Tuple[Status, StatusCode1, Optional[Any]]: + ext_status_code_2: str | None = None, + ext_status_code_3: str | None = None, +) -> tuple[Status, StatusCode1, Any | None]: """ :params ext_status_code_1: this is from the callback url params: disposition and status, '.'-joined diff --git a/generalresearch/wall_status_codes/fullcircle.py b/generalresearch/wall_status_codes/fullcircle.py index eda4d9c..f37ffda 100644 --- a/generalresearch/wall_status_codes/fullcircle.py +++ b/generalresearch/wall_status_codes/fullcircle.py @@ -11,7 +11,7 @@ from typing import Any, Dict, List, Optional, Tuple from generalresearch.models.thl.definitions import Status, StatusCode1 -status_codes_map: Dict[str, str] = { +status_codes_map: dict[str, str] = { "1": "Complete", "2": "Terminate", "3": "Over-quota", @@ -19,7 +19,7 @@ status_codes_map: Dict[str, str] = { } status_map = defaultdict(lambda: Status.FAIL, **{"1": Status.COMPLETE}) -status_codes_ext_map: Dict[StatusCode1, List[str]] = { +status_codes_ext_map: dict[StatusCode1, list[str]] = { StatusCode1.COMPLETE: ["1"], StatusCode1.BUYER_FAIL: ["2", "3"], StatusCode1.BUYER_QUALITY_FAIL: ["4"], @@ -29,10 +29,10 @@ 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] = dict() for k, v in status_codes_ext_map.items(): k: StatusCode1 - v: List[str] + v: list[str] for vv in v: vv: str @@ -41,9 +41,9 @@ for k, v in status_codes_ext_map.items(): def annotate_status_code( ext_status_code_1: str, - ext_status_code_2: Optional[str] = None, - ext_status_code_3: Optional[str] = None, -) -> Tuple[Status, StatusCode1, Optional[Any]]: + ext_status_code_2: str | None = None, + ext_status_code_3: str | None = None, +) -> tuple[Status, StatusCode1, Any | None]: """ :params ext_status_code_1: this is from the callback url param 's' :params ext_status_code_2: not used diff --git a/generalresearch/wall_status_codes/innovate.py b/generalresearch/wall_status_codes/innovate.py index b650ade..fba3f54 100644 --- a/generalresearch/wall_status_codes/innovate.py +++ b/generalresearch/wall_status_codes/innovate.py @@ -13,7 +13,7 @@ from typing import Any, Dict, List, Optional, Tuple from generalresearch.models.thl.definitions import Status, StatusCode1 -status_codes_innovate: Dict[str, str] = { +status_codes_innovate: dict[str, str] = { "1": "Complete", "2": "Buyer Fail", "3": "Buyer Over Quota", @@ -29,7 +29,7 @@ status_map = defaultdict( lambda: Status.FAIL, **{"1": Status.COMPLETE, "0": Status.ABANDON, "6": Status.ABANDON}, ) -status_codes_ext_map: Dict[StatusCode1, List[str]] = { +status_codes_ext_map: dict[StatusCode1, list[str]] = { StatusCode1.BUYER_FAIL: ["2", "3"], StatusCode1.BUYER_QUALITY_FAIL: ["4"], StatusCode1.PS_BLOCKED: [], @@ -43,7 +43,7 @@ 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 -category_innovate: Dict[str, StatusCode1] = { +category_innovate: dict[str, StatusCode1] = { "Selected threat potential score at joblevel not allow the survey": StatusCode1.PS_QUALITY, "OE Validation": StatusCode1.PS_QUALITY, "Unique IP": StatusCode1.PS_DUPLICATE, @@ -78,9 +78,9 @@ category_innovate: Dict[str, StatusCode1] = { def annotate_status_code( ext_status_code_1: str, - ext_status_code_2: Optional[str] = None, - ext_status_code_3: Optional[str] = None, -) -> Tuple[Status, StatusCode1, Optional[Any]]: + ext_status_code_2: str | None = None, + ext_status_code_3: str | None = None, +) -> tuple[Status, StatusCode1, Any | None]: """ Only quality terminate (4 and 8), and PS term (5) return a term_reason (af=). diff --git a/generalresearch/wall_status_codes/lucid.py b/generalresearch/wall_status_codes/lucid.py index c0098e6..3ce89b0 100644 --- a/generalresearch/wall_status_codes/lucid.py +++ b/generalresearch/wall_status_codes/lucid.py @@ -9,7 +9,7 @@ from typing import Any, Dict, List, Optional, Tuple from generalresearch.models.thl.definitions import Status, StatusCode1 -mp_codes: Dict[str, str] = { +mp_codes: dict[str, str] = { "-6": "Pre-Client Intermediary Page Drop Off", "-5": "Failure in the Post Answer Behavior", "-1": "Failure to Load the Lucid Marketplace", @@ -54,7 +54,7 @@ mp_codes: Dict[str, str] = { } # todo: finish, there's a bunch more -client_status_map: Dict[str, StatusCode1] = { +client_status_map: dict[str, StatusCode1] = { "30": StatusCode1.BUYER_QUALITY_FAIL, "33": StatusCode1.BUYER_QUALITY_FAIL, "34": StatusCode1.BUYER_QUALITY_FAIL, @@ -62,7 +62,7 @@ client_status_map: Dict[str, StatusCode1] = { } status_map = defaultdict(lambda: Status.FAIL, **{"s": Status.COMPLETE}) -status_codes_ext_map: Dict[StatusCode1, List[str]] = { +status_codes_ext_map: dict[StatusCode1, list[str]] = { StatusCode1.COMPLETE: [], StatusCode1.BUYER_FAIL: ["3"], StatusCode1.BUYER_QUALITY_FAIL: [], @@ -102,10 +102,10 @@ 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] = dict() for k, v in status_codes_ext_map.items(): k: StatusCode1 - v: List[str] + v: list[str] for vv in v: vv: str @@ -115,9 +115,9 @@ for k, v in status_codes_ext_map.items(): def annotate_status_code( ext_status_code_1: str, - ext_status_code_2: Optional[str] = None, - ext_status_code_3: Optional[str] = None, -) -> Tuple[Status, StatusCode1, Optional[Any]]: + ext_status_code_2: str | None = None, + ext_status_code_3: str | None = None, +) -> tuple[Status, StatusCode1, Any | None]: """ :params ext_status_code_1: this indicates which callback url was hit. possible values {'s', *anything else*} :params ext_status_code_2: this is from the callback url params: InitialStatus diff --git a/generalresearch/wall_status_codes/morning.py b/generalresearch/wall_status_codes/morning.py index 318d1b2..42e34b3 100644 --- a/generalresearch/wall_status_codes/morning.py +++ b/generalresearch/wall_status_codes/morning.py @@ -17,7 +17,7 @@ timeout: The respondent completed the survey after the timeout period had expire in_progress: The respondent interview session is still in progress, such as in the prescreener or survey. """ -short_code_to_status_codes_morning: Dict[str, str] = { +short_code_to_status_codes_morning: dict[str, str] = { "att_che": "attention_check", "banned": "banned", "bid_clo": "bid_closed", @@ -54,7 +54,7 @@ short_code_to_status_codes_morning: Dict[str, str] = { } status_map = defaultdict(lambda: Status.FAIL, **{"complete": Status.COMPLETE}) -status_codes_ext_map: Dict[StatusCode1, List[str]] = { +status_codes_ext_map: dict[StatusCode1, list[str]] = { StatusCode1.COMPLETE: ["complete"], StatusCode1.BUYER_FAIL: [ "in_survey_failure", @@ -97,10 +97,10 @@ 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] = dict() for k, v in status_codes_ext_map.items(): k: StatusCode1 - v: List[str] + v: list[str] for vv in v: vv: str @@ -109,9 +109,9 @@ for k, v in status_codes_ext_map.items(): def annotate_status_code( ext_status_code_1: str, - ext_status_code_2: Optional[str] = None, - ext_status_code_3: Optional[str] = None, -) -> Tuple[Status, StatusCode1, Optional[Any]]: + ext_status_code_2: str | None = None, + ext_status_code_3: str | None = None, +) -> tuple[Status, StatusCode1, Any | None]: """ :params ext_status_code_1: from callback url params: &sti={{status_id}} :params ext_status_code_2: from callback url params: &sdi={{status_detail_id}} diff --git a/generalresearch/wall_status_codes/pollfish.py b/generalresearch/wall_status_codes/pollfish.py index 361785d..1ae732d 100644 --- a/generalresearch/wall_status_codes/pollfish.py +++ b/generalresearch/wall_status_codes/pollfish.py @@ -1,9 +1,9 @@ from collections import defaultdict -from typing import Any, Dict, List, Optional, Tuple +from typing import Any from generalresearch.models.thl.definitions import Status, StatusCode1 -status_codes_map: Dict[str, str] = { +status_codes_map: dict[str, str] = { "quo_ful": "quota_full", "sur_clo": "survey_closed", "profilin": "profiling", @@ -29,7 +29,7 @@ status_codes_map: Dict[str, str] = { "complete": "complete", } status_map = defaultdict(lambda: Status.FAIL, **{"complete": Status.COMPLETE}) -status_codes_ext_map: Dict[StatusCode1, List[str]] = { +status_codes_ext_map: dict[StatusCode1, list[str]] = { StatusCode1.COMPLETE: ["complete"], StatusCode1.BUYER_FAIL: ["third_party_termination", "screenout"], StatusCode1.BUYER_QUALITY_FAIL: [ @@ -61,7 +61,7 @@ status_codes_ext_map: Dict[StatusCode1, List[str]] = { ext_status_code_map = dict() for k, v in status_codes_ext_map.items(): k: StatusCode1 - v: List[str] + v: list[str] for vv in v: vv: str @@ -70,9 +70,9 @@ for k, v in status_codes_ext_map.items(): def annotate_status_code( ext_status_code_1: str, - ext_status_code_2: Optional[str] = None, - ext_status_code_3: Optional[str] = None, -) -> Tuple[Status, StatusCode1, Optional[Any]]: + ext_status_code_2: str | None = None, + ext_status_code_3: str | None = None, +) -> tuple[Status, StatusCode1, Any | None]: """ :params ext_status_code_1: from callback url params: &sti={{status_id}} :params ext_status_code_2: from callback url params: &sdi={{status_detail_id}} diff --git a/generalresearch/wall_status_codes/precision.py b/generalresearch/wall_status_codes/precision.py index ffeaeca..d3c4471 100644 --- a/generalresearch/wall_status_codes/precision.py +++ b/generalresearch/wall_status_codes/precision.py @@ -11,7 +11,7 @@ from typing import Any, Dict, List, Optional, Tuple from generalresearch.models.thl.definitions import Status, StatusCode1 -status_codes_precision: Dict[str, str] = { +status_codes_precision: dict[str, str] = { "10": "Complete", "20": "Client Terminate", "21": "PS Terminate", @@ -46,7 +46,7 @@ status_codes_precision: Dict[str, str] = { "80": "Final Complete", } status_map = defaultdict(lambda: Status.FAIL, **{"s": Status.COMPLETE}) -status_codes_ext_map: Dict[StatusCode1, List[str]] = { +status_codes_ext_map: dict[StatusCode1, list[str]] = { StatusCode1.COMPLETE: ["10"], StatusCode1.BUYER_FAIL: ["20", "30"], StatusCode1.BUYER_QUALITY_FAIL: ["60"], @@ -76,7 +76,7 @@ status_codes_ext_map: Dict[StatusCode1, List[str]] = { ext_status_code_map = dict() for k, v in status_codes_ext_map.items(): k: StatusCode1 - v: List[str] + v: list[str] for vv in v: vv: str @@ -85,9 +85,9 @@ for k, v in status_codes_ext_map.items(): def annotate_status_code( ext_status_code_1: str, - ext_status_code_2: Optional[str] = None, - ext_status_code_3: Optional[str] = None, -) -> Tuple[Status, StatusCode1, Optional[Any]]: + ext_status_code_2: str | None = None, + ext_status_code_3: str | None = None, +) -> tuple[Status, StatusCode1, Any | None]: """ :params ext_status_code_1: from callback url params: status :params ext_status_code_2: from callback url params: code diff --git a/generalresearch/wall_status_codes/prodege.py b/generalresearch/wall_status_codes/prodege.py index a1e25ea..aaac376 100644 --- a/generalresearch/wall_status_codes/prodege.py +++ b/generalresearch/wall_status_codes/prodege.py @@ -8,7 +8,7 @@ from typing import Any, Dict, List, Optional, Tuple from generalresearch.models.thl.definitions import Status, StatusCode1 status_map = defaultdict(lambda: Status.FAIL, **{"1": Status.COMPLETE}) -status_code_map: Dict[StatusCode1, List[str]] = { +status_code_map: dict[StatusCode1, list[str]] = { StatusCode1.COMPLETE: [], StatusCode1.BUYER_FAIL: ["1", "2"], StatusCode1.BUYER_QUALITY_FAIL: ["10", "12"], @@ -34,7 +34,7 @@ status_code_map: Dict[StatusCode1, List[str]] = { status_class = dict() for k, v in status_code_map.items(): k: StatusCode1 - v: List[str] + v: list[str] for vv in v: vv: str @@ -43,9 +43,9 @@ for k, v in status_code_map.items(): def annotate_status_code( ext_status_code_1: str, - ext_status_code_2: Optional[str] = None, - ext_status_code_3: Optional[str] = None, -) -> Tuple[Status, StatusCode1, Optional[Any]]: + ext_status_code_2: str | None = None, + ext_status_code_3: str | None = None, +) -> tuple[Status, StatusCode1, Any | None]: """ :params ext_status_code_1: status from redirect url :params ext_status_code_2: termreason from redirect url diff --git a/generalresearch/wall_status_codes/repdata.py b/generalresearch/wall_status_codes/repdata.py index c502532..8b690b5 100644 --- a/generalresearch/wall_status_codes/repdata.py +++ b/generalresearch/wall_status_codes/repdata.py @@ -7,7 +7,7 @@ from typing import Any, Dict, List, Optional, Tuple from generalresearch.models.thl.definitions import Status, StatusCode1 -status_codes_name: Dict[str, str] = { +status_codes_name: dict[str, str] = { "2": "Search Failed", "3": "Activity Failed", "4": "Review Failed", @@ -26,7 +26,7 @@ status_codes_name: Dict[str, str] = { "6003": "In-Survey maximum exceeded (Research Desk)", } # See: 02, and 13 are de-dupes -rd_threat_name: Dict[str, str] = { +rd_threat_name: dict[str, str] = { "02": "Duplicate entrant into survey", "03": "Emulator Usage", "04": "VPN usage detected", @@ -47,7 +47,7 @@ rd_threat_name: Dict[str, str] = { } status_map = defaultdict(lambda: Status.FAIL, **{"complete": Status.COMPLETE}) -status_code_map: Dict[StatusCode1, List[str]] = { +status_code_map: dict[StatusCode1, list[str]] = { StatusCode1.COMPLETE: ["1000"], StatusCode1.BUYER_FAIL: ["2000", "4000"], StatusCode1.BUYER_QUALITY_FAIL: ["3000"], @@ -61,7 +61,7 @@ status_code_map: Dict[StatusCode1, List[str]] = { status_class = dict() for k, v in status_code_map.items(): k: StatusCode1 - v: List[str] + v: list[str] for vv in v: vv: str @@ -70,9 +70,9 @@ for k, v in status_code_map.items(): def annotate_status_code( ext_status_code_1: str, - ext_status_code_2: Optional[str] = None, - ext_status_code_3: Optional[str] = None, -) -> Tuple[Status, StatusCode1, Optional[Any]]: + ext_status_code_2: str | None = None, + ext_status_code_3: str | None = None, +) -> tuple[Status, StatusCode1, Any | None]: """ :params ext_status_code_1: the redirect urls category (as defined in url param 549f3710b) {'term', 'overquota', 'fraud', 'complete'} diff --git a/generalresearch/wall_status_codes/sago.py b/generalresearch/wall_status_codes/sago.py index 9c8710f..028192b 100644 --- a/generalresearch/wall_status_codes/sago.py +++ b/generalresearch/wall_status_codes/sago.py @@ -7,7 +7,7 @@ from typing import Any, Dict, List, Optional, Tuple from generalresearch.models.thl.definitions import Status, StatusCode1 -status_codes_schlesinger: Dict[str, str] = { +status_codes_schlesinger: dict[str, str] = { "1": "Complete", "2": "Buyer Fail", "3": "Buyer Fail", @@ -20,7 +20,7 @@ status_codes_schlesinger: Dict[str, str] = { "11": "Abandon", # really it is "Buyer Abandon" } -status_reason_name: Dict[str, str] = { +status_reason_name: dict[str, str] = { "1": "Not a Unique Sample Cube User", "4": "GeoIP - wrong country", "7": "Duplicate - not a unique IP", @@ -121,7 +121,7 @@ status_map = defaultdict( lambda: Status.FAIL, **{"1": Status.COMPLETE, "0": Status.ABANDON} ) -status_codes_ext_map: Dict[StatusCode1, List[str]] = { +status_codes_ext_map: dict[StatusCode1, list[str]] = { StatusCode1.COMPLETE: ["48"], StatusCode1.BUYER_FAIL: ["16", "29", "49", "50", "78", "114", "110", "114"], StatusCode1.BUYER_QUALITY_FAIL: ["26", "52", "68", "81", "84"], @@ -167,10 +167,10 @@ status_codes_ext_map: Dict[StatusCode1, List[str]] = { StatusCode1.PS_FAIL: ["7", "29", "36", "47", "56", "58", "64"], StatusCode1.PS_OVERQUOTA: ["29", "46", "33", "31"], } -ext_status_code_map: Dict[str, StatusCode1] = dict() +ext_status_code_map: dict[str, StatusCode1] = dict() for k, v in status_codes_ext_map.items(): k: StatusCode1 - v: List[str] + v: list[str] for vv in v: vv: str @@ -179,9 +179,9 @@ for k, v in status_codes_ext_map.items(): def annotate_status_code( ext_status_code_1: str, - ext_status_code_2: Optional[str] = None, - ext_status_code_3: Optional[str] = None, -) -> Tuple[Status, StatusCode1, Optional[Any]]: + ext_status_code_2: str | None = None, + ext_status_code_3: str | None = None, +) -> tuple[Status, StatusCode1, Any | None]: """ :params ext_status_code_1: from callback url params: scstatus :params ext_status_code_2: from callback url params: scsecuritystatus diff --git a/generalresearch/wall_status_codes/spectrum.py b/generalresearch/wall_status_codes/spectrum.py index 9e9e0a2..cf200a9 100644 --- a/generalresearch/wall_status_codes/spectrum.py +++ b/generalresearch/wall_status_codes/spectrum.py @@ -7,7 +7,7 @@ from typing import Any, Dict, List, Optional, Tuple from generalresearch.models.thl.definitions import Status, StatusCode1 -status_codes_spectrum: Dict[str, str] = { +status_codes_spectrum: dict[str, str] = { "11": "PS Drop", "12": "PS Quota Full Core", "13": "PS Termination Core", @@ -80,7 +80,7 @@ status_codes_spectrum: Dict[str, str] = { "88": "PS_Supplier_Allocation_Throttle", } status_map = defaultdict(lambda: Status.FAIL, **{"21": Status.COMPLETE}) -status_codes_ext_map: Dict[StatusCode1, List[str]] = { +status_codes_ext_map: dict[StatusCode1, list[str]] = { StatusCode1.COMPLETE: ["21"], StatusCode1.BUYER_FAIL: ["16", "17", "18", "19", "30", "59", "84"], StatusCode1.BUYER_QUALITY_FAIL: ["20", "31"], @@ -143,7 +143,7 @@ status_codes_ext_map: Dict[StatusCode1, List[str]] = { ext_status_code_map = dict() for k, v in status_codes_ext_map.items(): k: StatusCode1 - v: List[str] + v: list[str] for vv in v: vv: str @@ -152,9 +152,9 @@ for k, v in status_codes_ext_map.items(): def annotate_status_code( ext_status_code_1: str, - ext_status_code_2: Optional[str] = None, - ext_status_code_3: Optional[str] = None, -) -> Tuple[Status, StatusCode1, Optional[Any]]: + ext_status_code_2: str | None = None, + ext_status_code_3: str | None = None, +) -> tuple[Status, StatusCode1, Any | None]: """ :params ext_status_code_1: from url params: ps_rstatus https://purespectrum.atlassian.net/wiki/spaces/PA/pages/33613201/Minimizing+Clickwaste+with+ps+rstatus diff --git a/generalresearch/wall_status_codes/wxet.py b/generalresearch/wall_status_codes/wxet.py index e7cf67d..1b7c514 100644 --- a/generalresearch/wall_status_codes/wxet.py +++ b/generalresearch/wall_status_codes/wxet.py @@ -8,10 +8,10 @@ from generalresearch.wxet.models.definitions import ( WXETStatusCode2, ) -status_map: Dict[WXETStatus, Status] = defaultdict( +status_map: dict[WXETStatus, Status] = defaultdict( lambda: Status.FAIL, **{WXETStatus.COMPLETE: Status.COMPLETE} ) -status_codes_ext_map: Dict[StatusCode1, List[WXETStatusCode1]] = { +status_codes_ext_map: dict[StatusCode1, list[WXETStatusCode1]] = { StatusCode1.COMPLETE: [WXETStatusCode1.COMPLETE], StatusCode1.BUYER_FAIL: [ WXETStatusCode1.BUYER_DUPLICATE, @@ -33,13 +33,13 @@ status_codes_ext_map: Dict[StatusCode1, List[WXETStatusCode1]] = { ext_status_code_map = dict() for k, v in status_codes_ext_map.items(): k: StatusCode1 - v: List[WXETStatusCode1] + v: list[WXETStatusCode1] for vv in v: vv: WXETStatusCode1 ext_status_code_map[vv] = k -status_code2_map: Dict[StatusCode1, List[WXETStatusCode2]] = { +status_code2_map: dict[StatusCode1, list[WXETStatusCode2]] = { StatusCode1.PS_QUALITY: [], StatusCode1.PS_DUPLICATE: [ WXETStatusCode2.WORKER_INELIGIBLE, @@ -67,9 +67,9 @@ for k, v in status_code2_map.items(): def annotate_status_code( ext_status_code_1: str, - ext_status_code_2: Optional[str] = None, - ext_status_code_3: Optional[str] = None, -) -> Tuple[Status, StatusCode1, Optional[WXETStatusCode2]]: + ext_status_code_2: str | None = None, + ext_status_code_3: str | None = None, +) -> tuple[Status, StatusCode1, WXETStatusCode2 | None]: """ :params ext_status_code_1: WXETStatus :params ext_status_code_2: WXETStatusCode1 diff --git a/generalresearch/wxet/models/definitions.py b/generalresearch/wxet/models/definitions.py index 8d0d6b8..b08e178 100644 --- a/generalresearch/wxet/models/definitions.py +++ b/generalresearch/wxet/models/definitions.py @@ -166,8 +166,8 @@ class WXETStatusCode2(int, Enum, metaclass=ReprEnumMeta): def check_wxet_status_consistent( status: WXETStatus, - status_code_1: Optional[WXETStatusCode1] = None, - status_code_2: Optional[WXETStatusCode2] = None, + status_code_1: WXETStatusCode1 | None = None, + status_code_2: WXETStatusCode2 | None = None, ) -> bool: """ Raises an AssertionError if inconsistent @@ -203,13 +203,13 @@ def check_wxet_status_consistent( def check_wxet_adjusted_status_attempt_consistent( status: WXETStatus, - status_code_1: Optional[WXETStatusCode1] = None, - cpi: Optional[USDMill] = None, - adjusted_status: Optional[WXETAdjustedStatus] = None, - adjusted_cpi: Optional[USDMill] = None, - new_adjusted_status: Optional[WXETAdjustedStatus] = None, - new_adjusted_cpi: Optional[USDMill] = None, -) -> Tuple[bool, str]: + status_code_1: WXETStatusCode1 | None = None, + cpi: USDMill | None = None, + adjusted_status: WXETAdjustedStatus | None = None, + adjusted_cpi: USDMill | None = None, + new_adjusted_status: WXETAdjustedStatus | None = None, + new_adjusted_cpi: USDMill | None = None, +) -> tuple[bool, str]: """ Raises an AssertionError if inconsistent. - status, status_code_1, adjusted_status, adjusted_cpi, cpi are the attempt's CURRENT values @@ -233,12 +233,12 @@ def check_wxet_adjusted_status_attempt_consistent( def _check_wxet_adjusted_status_attempt_consistent( status: WXETStatus, - status_code_1: Optional[WXETStatusCode1] = None, - cpi: Optional[USDMill] = None, - adjusted_status: Optional[WXETAdjustedStatus] = None, - adjusted_cpi: Optional[USDMill] = None, - new_adjusted_status: Optional[WXETAdjustedStatus] = None, - new_adjusted_cpi: Optional[USDMill] = None, + status_code_1: WXETStatusCode1 | None = None, + cpi: USDMill | None = None, + adjusted_status: WXETAdjustedStatus | None = None, + adjusted_cpi: USDMill | None = None, + new_adjusted_status: WXETAdjustedStatus | None = None, + new_adjusted_cpi: USDMill | None = None, ) -> None: """ Raises an AssertionError if inconsistent. @@ -297,8 +297,8 @@ def _check_wxet_adjusted_status_attempt_consistent( def _check_wxet_adjusted_status_consistent( - adjusted_status: Optional[WXETAdjustedStatus] = None, - adjusted_cpi: Optional[USDMill] = None, + adjusted_status: WXETAdjustedStatus | None = None, + adjusted_cpi: USDMill | None = None, ) -> None: """ Raises an AssertionError if inconsistent. diff --git a/generalresearch/wxet/models/finish_type.py b/generalresearch/wxet/models/finish_type.py index af60fe6..a57dce8 100644 --- a/generalresearch/wxet/models/finish_type.py +++ b/generalresearch/wxet/models/finish_type.py @@ -33,7 +33,7 @@ class FinishType(str, Enum, metaclass=ReprEnumMeta): FAIL = "fail" @property - def finish_statuses(self) -> Set[Optional[WXETStatus]]: + def finish_statuses(self) -> set[WXETStatus | None]: """For this particular FinishType, what are the different WXETStatus values that are consider """ @@ -64,9 +64,9 @@ class FinishType(str, Enum, metaclass=ReprEnumMeta): def is_a_finish( - status: Optional[WXETStatus], - status_code_1: Optional[WXETStatusCode1], - finish_type: Optional[FinishType], + status: WXETStatus | None, + status_code_1: WXETStatusCode1 | None, + finish_type: FinishType | None, ) -> bool: """Determines if a wall event should be considered a finish or not. diff --git a/test_utils/conftest.py b/test_utils/conftest.py index 378b9cc..ffe458c 100644 --- a/test_utils/conftest.py +++ b/test_utils/conftest.py @@ -6,16 +6,17 @@ import stat import subprocess import sys import tempfile -from datetime import datetime, timedelta, timezone +from collections.abc import Callable, Generator +from datetime import UTC, datetime, timedelta, timezone from os.path import join as pjoin from pathlib import Path -from typing import Callable, Generator from uuid import uuid4 import pytest from _pytest.config import Config 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 @@ -93,7 +94,7 @@ def postgres_instance(settings: GRLBaseSettings) -> Generator[PostgresDsn]: from psycopg import connect from psycopg.sql import SQL, Identifier - now = datetime.now(timezone.utc) + now = datetime.now(UTC) ts: str = now.strftime("%Y-%m-%d") db_name = f"unittest-{ts}-{uuid4().hex[:6]}" @@ -152,38 +153,48 @@ def postgres_instance_host( yield value -# @pytest.fixture(scope="session") -# def git_key_path(settings: GRLBaseSettings) -> Path: -# return Path('/tmp/') - - @pytest.fixture(scope="session") def git_key_path( + tmp_path_factory: TempPathFactory, settings: GRLBaseSettings, ) -> Generator[Path]: + # We are using the tmp_path_factory because unlike the tmp_path (which + # is function scoped), this is session scoped. - assert settings.git_creds - with tempfile.NamedTemporaryFile(mode="w", delete=False, suffix="_id_rsa") as f: - f.write(settings.git_creds) - key_path = f.name - - os.chmod(key_path, stat.S_IRUSR | stat.S_IWUSR) + assert settings.git_creds, "Must define key to download alternative models" + fn = tmp_path_factory.mktemp("keys") / "git_creds" + fn.write_text(settings.git_creds, encoding="utf-8") + os.chmod(fn, stat.S_IRUSR | stat.S_IWUSR) - yield Path(key_path) + yield Path(fn) - os.unlink(key_path) + os.unlink(fn) @pytest.fixture(scope="session") -def gr_repo(git_key_path: Path) -> Callable[..., Path]: +def gr_repo( + git_key_path: Path, + tmp_path_factory: TempPathFactory, +) -> Callable[..., Path | None]: repo_url = "ssh://code.g-r-l.com/general-research/gr-carer.git" - repo_path = Path("/tmp/gr-carer") + + _ran = {} + if _ran.get(repo_url, False): + print(f"Already ran django_db_factory.{repo_url}") + return + + _ran[repo_url] = True + + fn = tmp_path_factory.mktemp("repos") + repo_path = fn / "gr-carer" + repo_path.mkdir(parents=True, exist_ok=True) def _inner() -> Path: + ssh_cmd = ( f"ssh -i {git_key_path} " "-o IdentitiesOnly=yes " - "-o StrictHostKeyChecking=no " # or accept-new, see note below + "-o StrictHostKeyChecking=no " ) env = {"GIT_SSH_COMMAND": ssh_cmd} @@ -196,6 +207,11 @@ def gr_repo(git_key_path: Path) -> Callable[..., Path]: env=env, ) + result = subprocess.run( + ["cat", git_key_path], capture_output=True, text=True, check=False + ) + print(repr(result.stdout)) + return repo_path return _inner @@ -206,21 +222,29 @@ def django_db_factory( postgres_instance: PostgresDsn, postgres_instance_dict: PostgresDict, gr_repo: Callable[..., Path], -) -> Callable[..., PostgresDsn]: +) -> Callable[..., PostgresDsn | None]: + + _ran = {} import django + from django.apps import apps from django.conf import settings as django_settings from django.core.management import call_command - def _inner(django_project: str = "generalresearch.thl_django"): + def _inner( + django_project: str = "generalresearch.thl_django", + ) -> PostgresDsn | None: + + if _ran.get(django_project, False): + print(f"Already ran django_db_factory.{django_project}") + return + _ran[django_project] = True if "gr" in django_project: # We need model files that are NOT in this repo. gr_path = gr_repo() sys.path.insert(0, str(gr_path)) - print(sys.path) - # 1. Bootstrapping Django settings if not django_settings.configured: django_settings.configure( @@ -242,10 +266,11 @@ def django_db_factory( ) django.setup() - # for model in apps.get_models(): - # print(f"Discovered model: {model._meta.label}") + for model in apps.get_models(): + print(f"Discovered model: {model._meta.label}") # 2. Run migrations directly during fixture activation + call_command("makemigrations", "gr", interactive=False) call_command("migrate") # 3. Return the Dsn so the factory gives a way to connect @@ -276,37 +301,37 @@ def spectrum_rw(settings: GRLBaseSettings) -> SqlHelper: @pytest.fixture def start() -> datetime: - return datetime(year=1900, month=1, day=1, tzinfo=timezone.utc) + return datetime(year=1900, month=1, day=1, tzinfo=UTC) @pytest.fixture def utc_now() -> datetime: - return datetime.now(tz=timezone.utc) + return datetime.now(tz=UTC) @pytest.fixture def utc_hour_ago() -> datetime: - return datetime.now(tz=timezone.utc) - timedelta(hours=1) + return datetime.now(tz=UTC) - timedelta(hours=1) @pytest.fixture def utc_day_ago() -> datetime: - return datetime.now(tz=timezone.utc) - timedelta(hours=24) + return datetime.now(tz=UTC) - timedelta(hours=24) @pytest.fixture def utc_90days_ago() -> datetime: - return datetime.now(tz=timezone.utc) - timedelta(days=90) + return datetime.now(tz=UTC) - timedelta(days=90) @pytest.fixture def utc_60days_ago() -> datetime: - return datetime.now(tz=timezone.utc) - timedelta(days=60) + return datetime.now(tz=UTC) - timedelta(days=60) @pytest.fixture def utc_30days_ago() -> datetime: - return datetime.now(tz=timezone.utc) - timedelta(days=30) + return datetime.now(tz=UTC) - timedelta(days=30) # === Clean up === @@ -322,7 +347,7 @@ def delete_df_collection( DFCollectionType, ) - def _inner(coll: "DFCollection"): + def _inner(coll: DFCollection): match coll.data_type: case DFCollectionType.LEDGER: for table in [ diff --git a/test_utils/grliq/conftest.py b/test_utils/grliq/conftest.py index e8175a5..7665b52 100644 --- a/test_utils/grliq/conftest.py +++ b/test_utils/grliq/conftest.py @@ -1,7 +1,7 @@ from __future__ import annotations -from datetime import datetime, timedelta, timezone -from typing import Callable +from collections.abc import Callable +from datetime import UTC, datetime, timedelta, timezone from uuid import uuid4 import pytest @@ -83,7 +83,7 @@ def grliq_data() -> GrlIqData: g.id = None g.uuid = uuid4().hex - g.created_at = datetime.now(tz=timezone.utc) + g.created_at = datetime.now(tz=UTC) g.timestamp = g.created_at - timedelta(seconds=10) return g @@ -117,7 +117,7 @@ def grliq_data_factory(grliq_dm: GrlIqDataManager) -> Callable[..., GrlIqData]: product_user_id = product_user_id or uuid4().hex uuid = uuid or uuid4().hex mid = mid or uuid4().hex - created_at = created_at or datetime.now(tz=timezone.utc) + created_at = created_at or datetime.now(tz=UTC) res["data"].product_id = product_id res["data"].product_user_id = product_user_id diff --git a/test_utils/incite/collections/conftest.py b/test_utils/incite/collections/conftest.py index 88eef72..631bb7b 100644 --- a/test_utils/incite/collections/conftest.py +++ b/test_utils/incite/collections/conftest.py @@ -1,7 +1,8 @@ from __future__ import annotations +from collections.abc import Callable from datetime import datetime, timedelta -from typing import TYPE_CHECKING, Callable +from typing import TYPE_CHECKING import pytest diff --git a/test_utils/incite/conftest.py b/test_utils/incite/conftest.py index 12e57c5..87ea7ae 100644 --- a/test_utils/incite/conftest.py +++ b/test_utils/incite/conftest.py @@ -1,11 +1,12 @@ from __future__ import annotations -from datetime import datetime, timedelta, timezone +from collections.abc import Callable +from datetime import UTC, datetime, timedelta, timezone from os.path import join as pjoin from pathlib import Path from random import choice as randchoice from shutil import rmtree -from typing import TYPE_CHECKING, Callable +from typing import TYPE_CHECKING from uuid import uuid4 import pytest @@ -166,7 +167,7 @@ def incite_item_factory( for _ in range(5): item_time = fake.date_time_between( - start_date=item.start, end_date=item.finish, tzinfo=timezone.utc + start_date=item.start, end_date=item.finish, tzinfo=UTC ) match data_type: diff --git a/test_utils/incite/mergers/conftest.py b/test_utils/incite/mergers/conftest.py index e9970c2..c0f0bcf 100644 --- a/test_utils/incite/mergers/conftest.py +++ b/test_utils/incite/mergers/conftest.py @@ -1,7 +1,7 @@ from __future__ import annotations +from collections.abc import Callable from datetime import datetime, timedelta -from typing import Callable import pytest diff --git a/test_utils/managers/conftest.py b/test_utils/managers/conftest.py index d2e5d20..4dacb29 100644 --- a/test_utils/managers/conftest.py +++ b/test_utils/managers/conftest.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import Callable +from collections.abc import Callable import pytest diff --git a/test_utils/managers/gr/conftest.py b/test_utils/managers/gr/conftest.py index 37da164..4da8fe3 100644 --- a/test_utils/managers/gr/conftest.py +++ b/test_utils/managers/gr/conftest.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import Callable +from collections.abc import Callable import pytest import redis.asyncio as redis_async diff --git a/test_utils/managers/thl/conftest.py b/test_utils/managers/thl/conftest.py index 5b70961..21b2007 100644 --- a/test_utils/managers/thl/conftest.py +++ b/test_utils/managers/thl/conftest.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import Callable +from collections.abc import Callable import pytest from pydantic import PostgresDsn diff --git a/test_utils/managers/upk/conftest.py b/test_utils/managers/upk/conftest.py index d8f956c..7eabee1 100644 --- a/test_utils/managers/upk/conftest.py +++ b/test_utils/managers/upk/conftest.py @@ -1,4 +1,4 @@ -from typing import Callable, Generator +from collections.abc import Callable, Generator import pytest diff --git a/test_utils/models/conftest.py b/test_utils/models/conftest.py index 93e2f44..3a9e45c 100644 --- a/test_utils/models/conftest.py +++ b/test_utils/models/conftest.py @@ -1,10 +1,11 @@ from __future__ import annotations -from datetime import datetime, timedelta, timezone +from collections.abc import Callable +from datetime import UTC, datetime, timedelta, timezone from decimal import Decimal from random import choice as randchoice from random import randint -from typing import TYPE_CHECKING, Callable +from typing import TYPE_CHECKING from uuid import uuid4 import pytest @@ -124,7 +125,7 @@ def wall_factory(wall_manager: WallManager) -> Callable[..., Wall]: ) -> Wall: assert session.started <= datetime.now( - tz=timezone.utc + tz=UTC ), "Session can't start in the future" if session.wall_events: diff --git a/test_utils/models/contest/conftest.py b/test_utils/models/contest/conftest.py index bfbe9f8..0060946 100644 --- a/test_utils/models/contest/conftest.py +++ b/test_utils/models/contest/conftest.py @@ -1,14 +1,13 @@ from __future__ import annotations -from datetime import datetime, timezone +from collections.abc import Callable +from datetime import UTC, datetime, timezone from decimal import Decimal -from typing import Callable 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 @@ -136,7 +135,7 @@ def milestone_contest_create() -> MilestoneContestCreate: ), ], end_condition=MilestoneContestEndCondition( - ends_at=datetime(year=2030, month=1, day=1, tzinfo=timezone.utc), + ends_at=datetime(year=2030, month=1, day=1, tzinfo=UTC), max_winners=5, ), entry_trigger=ContestEntryTrigger.TASK_COMPLETE, diff --git a/test_utils/models/gr/conftest.py b/test_utils/models/gr/conftest.py index df97306..90b86aa 100644 --- a/test_utils/models/gr/conftest.py +++ b/test_utils/models/gr/conftest.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import Callable +from collections.abc import Callable from uuid import uuid4 import pytest diff --git a/test_utils/models/ledger/conftest.py b/test_utils/models/ledger/conftest.py index 14a7465..b428468 100644 --- a/test_utils/models/ledger/conftest.py +++ b/test_utils/models/ledger/conftest.py @@ -1,9 +1,10 @@ from __future__ import annotations +from collections.abc import Callable from datetime import datetime from decimal import Decimal from random import randint -from typing import TYPE_CHECKING, Callable +from typing import TYPE_CHECKING from uuid import uuid4 import pytest diff --git a/test_utils/models/network/conftest.py b/test_utils/models/network/conftest.py index bebc691..cabd8dc 100644 --- a/test_utils/models/network/conftest.py +++ b/test_utils/models/network/conftest.py @@ -1,5 +1,5 @@ import os -from datetime import datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta, timezone from uuid import uuid4 import pytest @@ -90,7 +90,7 @@ def rdns_result(dig_raw_output: str) -> RDNSResult: def rdns_run(rdns_result: RDNSResult, scan_group_id: str): r = rdns_result ip = "45.33.32.156" - utc_now = datetime.now(tz=timezone.utc) + utc_now = datetime.now(tz=UTC) config = RDNSRunCommand(command="dig", options=RDNSRunCommandOptions(ip=ip)) return RDNSRun( tool_version="1.2.3", @@ -121,7 +121,7 @@ def mtr_result(mtr_raw_output: str) -> MTRResult: @pytest.fixture(scope="session") def mtr_run(mtr_result: MTRResult, scan_group_id: str): r = mtr_result - utc_now = datetime.now(tz=timezone.utc) + utc_now = datetime.now(tz=UTC) config = MTRRunCommand( command="mtr", options=MTRRunCommandOptions( diff --git a/test_utils/models/thl/conftest.py b/test_utils/models/thl/conftest.py index cf8d2fa..a2adcce 100644 --- a/test_utils/models/thl/conftest.py +++ b/test_utils/models/thl/conftest.py @@ -1,11 +1,12 @@ from __future__ import annotations -from datetime import datetime, timezone +from collections.abc import Callable +from datetime import UTC, datetime, timezone from decimal import ROUND_DOWN, Decimal from random import choice as rand_choice from random import choice as rchoice from random import randint, random -from typing import Any, Callable +from typing import Any from uuid import uuid4 import faker @@ -105,9 +106,9 @@ def wall_factory( user_id = user_id or fake.random_int(min=1, max=2_147_483_648) started = started or fake.date_time_between( - start_date=datetime(year=1900, month=1, day=1, tzinfo=timezone.utc), - end_date=datetime.now(tz=timezone.utc), - tzinfo=timezone.utc, + start_date=datetime(year=1900, month=1, day=1, tzinfo=UTC), + end_date=datetime.now(tz=UTC), + tzinfo=UTC, ) if session_id is None: @@ -199,9 +200,9 @@ def session_factory(session_manager: SessionManager): ) -> Session: """To be used in tests, where we don't care about certain fields""" started = started or fake.date_time_between( - start_date=datetime(year=1900, month=1, day=1, tzinfo=timezone.utc), - end_date=datetime(year=2000, month=1, day=1, tzinfo=timezone.utc), - tzinfo=timezone.utc, + start_date=datetime(year=1900, month=1, day=1, tzinfo=UTC), + end_date=datetime(year=2000, month=1, day=1, tzinfo=UTC), + tzinfo=UTC, ) user = user or User( user_id=fake.random_int(min=1, max=2_147_483_648), uuid=uuid4().hex diff --git a/test_utils/spectrum/conftest.py b/test_utils/spectrum/conftest.py index 0afc3f5..9c067d3 100644 --- a/test_utils/spectrum/conftest.py +++ b/test_utils/spectrum/conftest.py @@ -1,6 +1,6 @@ import logging import time -from datetime import datetime, timezone +from datetime import UTC, datetime, timezone from typing import TYPE_CHECKING import pytest @@ -19,7 +19,7 @@ if TYPE_CHECKING: @pytest.fixture(scope="session") -def spectrum_rw(settings: "GRLBaseSettings") -> SqlHelper: +def spectrum_rw(settings: GRLBaseSettings) -> SqlHelper: logging.info(f"{settings.spectrum_rw_db=}") assert settings.spectrum_rw_db is not None @@ -49,11 +49,11 @@ def spectrum_survey_manager(spectrum_rw: SqlHelper) -> SpectrumSurveyManager: def setup_spectrum_surveys( spectrum_rw: SqlHelper, spectrum_survey_manager, spectrum_criteria_manager ) -> None: - now = datetime.now(timezone.utc) + now = datetime.now(UTC) # make sure these example surveys exist in db surveys = [SpectrumSurvey.model_validate_json(x) for x in SURVEYS_JSON] for s in surveys: - s.modified_api = datetime.now(tz=timezone.utc) + s.modified_api = datetime.now(tz=UTC) spectrum_survey_manager.create_or_update(surveys) spectrum_criteria_manager.update(CONDITIONS) @@ -66,10 +66,10 @@ def setup_spectrum_surveys( ["687", "GRL", "x", "x", "x", "x"], commit=True, ) - supplier687_pk = spectrum_rw.execute_sql_query( - f""" - select id from `{spectrum_rw.db}`.spectrum_supplier where supplier_id = '687'""" - )[0]["id"] + supplier687_pk = spectrum_rw.execute_sql_query(f""" + select id from `{spectrum_rw.db}`.spectrum_supplier where supplier_id = '687'""")[ + 0 + ]["id"] conn = spectrum_rw.make_connection() c = conn.cursor() c.executemany( diff --git a/tests/grliq/models/test_forensic_data.py b/tests/grliq/models/test_forensic_data.py index 4fbf962..a901dc3 100644 --- a/tests/grliq/models/test_forensic_data.py +++ b/tests/grliq/models/test_forensic_data.py @@ -9,16 +9,16 @@ if TYPE_CHECKING: class TestGrlIqData: - def test_supported_fonts(self, grliq_data: "GrlIqData"): + def test_supported_fonts(self, grliq_data: GrlIqData): s = grliq_data.supported_fonts_binary assert len(s) == 1043 assert "Ubuntu" in grliq_data.supported_fonts - def test_battery(self, grliq_data: "GrlIqData"): + def test_battery(self, grliq_data: GrlIqData): assert not grliq_data.battery_charging assert grliq_data.battery_level == 0.41 - def test_base(self, grliq_data: "GrlIqData"): + def test_base(self, grliq_data: GrlIqData): from generalresearch.grliq.models.forensic_data import Platform assert grliq_data.timezone == "America/Los_Angeles" @@ -41,7 +41,7 @@ class TestGrlIqData: # Testing things that will cause a validation error, should only be # because something is "corrupt", not b/c the user is a baddie - def test_corrupt(self, grliq_data: "GrlIqData"): + def test_corrupt(self, grliq_data: GrlIqData): """Test for timestamp and timezone offset mismatch validation.""" from generalresearch.grliq.models.forensic_data import GrlIqData 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, diff --git a/tests/incite/mergers/foundations/test_enriched_session.py b/tests/incite/mergers/foundations/test_enriched_session.py index 47f243e..a0ae01e 100644 --- a/tests/incite/mergers/foundations/test_enriched_session.py +++ b/tests/incite/mergers/foundations/test_enriched_session.py @@ -1,4 +1,4 @@ -from datetime import datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta, timezone from decimal import Decimal from itertools import product from typing import Optional @@ -77,15 +77,15 @@ class TestEnrichedSession: class TestEnrichedSessionAdmin: @pytest.fixture - def start(self) -> "datetime": - return datetime(year=2020, month=3, day=14, tzinfo=timezone.utc) + def start(self) -> datetime: + return datetime(year=2020, month=3, day=14, tzinfo=UTC) @pytest.fixture def offset(self) -> str: return "1d" @pytest.fixture - def duration(self) -> Optional["timedelta"]: + def duration(self) -> timedelta | None: return timedelta(days=5) def test_to_admin_response( diff --git a/tests/incite/mergers/foundations/test_enriched_wall.py b/tests/incite/mergers/foundations/test_enriched_wall.py index 8f4995b..b421df8 100644 --- a/tests/incite/mergers/foundations/test_enriched_wall.py +++ b/tests/incite/mergers/foundations/test_enriched_wall.py @@ -1,4 +1,4 @@ -from datetime import timedelta, timezone, datetime +from datetime import UTC, datetime, timedelta, timezone from decimal import Decimal from itertools import product as iter_product from typing import Optional @@ -9,13 +9,13 @@ import pytest # noinspection PyUnresolvedReferences from distributed.utils_test import ( - gen_cluster, + cleanup, + client, client_no_amm, + cluster_fixture, + gen_cluster, loop, loop_in_thread, - cleanup, - cluster_fixture, - client, ) from generalresearch.incite.mergers.foundations.enriched_wall import ( @@ -158,15 +158,15 @@ class TestEnrichedWall: class TestEnrichedWallToAdmin: @pytest.fixture - def start(self) -> "datetime": - return datetime(year=2020, month=3, day=14, tzinfo=timezone.utc) + def start(self) -> datetime: + return datetime(year=2020, month=3, day=14, tzinfo=UTC) @pytest.fixture def offset(self) -> str: return "1d" @pytest.fixture - def duration(self) -> Optional["timedelta"]: + def duration(self) -> timedelta | None: return timedelta(days=5) def test_empty(self, enriched_wall_merge, client_no_amm, start): diff --git a/tests/incite/mergers/foundations/test_user_id_product.py b/tests/incite/mergers/foundations/test_user_id_product.py index f96bfb4..a696b45 100644 --- a/tests/incite/mergers/foundations/test_user_id_product.py +++ b/tests/incite/mergers/foundations/test_user_id_product.py @@ -1,4 +1,4 @@ -from datetime import timedelta, datetime, timezone +from datetime import UTC, datetime, timedelta, timezone from itertools import product import pandas as pd @@ -6,13 +6,13 @@ import pytest # noinspection PyUnresolvedReferences from distributed.utils_test import ( - gen_cluster, + cleanup, + client, client_no_amm, + cluster_fixture, + gen_cluster, loop, loop_in_thread, - cleanup, - cluster_fixture, - client, ) from generalresearch.incite.mergers.foundations.user_id_product import ( @@ -27,11 +27,7 @@ from test_utils.incite.mergers.conftest import user_id_product_merge product( ["12h", "3D"], [timedelta(days=5)], - [ - (datetime.now(tz=timezone.utc) - timedelta(days=35)).replace( - microsecond=0 - ) - ], + [(datetime.now(tz=UTC) - timedelta(days=35)).replace(microsecond=0)], ) ), ) diff --git a/tests/incite/mergers/test_merge_collection.py b/tests/incite/mergers/test_merge_collection.py index ec507bc..77fa8c7 100644 --- a/tests/incite/mergers/test_merge_collection.py +++ b/tests/incite/mergers/test_merge_collection.py @@ -1,4 +1,4 @@ -from datetime import datetime, timezone, timedelta +from datetime import UTC, datetime, timedelta, timezone from itertools import product import pandas as pd @@ -21,11 +21,7 @@ merge_types = list(e for e in MergeType if e != MergeType.TEST) merge_types, ["5min", "6h", "14D"], [timedelta(days=30)], - [ - (datetime.now(tz=timezone.utc) - timedelta(days=35)).replace( - microsecond=0 - ) - ], + [(datetime.now(tz=UTC) - timedelta(days=35)).replace(microsecond=0)], ) ), ) diff --git a/tests/incite/mergers/test_pop_ledger.py b/tests/incite/mergers/test_pop_ledger.py index 6f96108..7583faf 100644 --- a/tests/incite/mergers/test_pop_ledger.py +++ b/tests/incite/mergers/test_pop_ledger.py @@ -1,4 +1,4 @@ -from datetime import timedelta, datetime, timezone +from datetime import UTC, datetime, timedelta, timezone from itertools import product as iter_product from typing import Optional @@ -10,7 +10,7 @@ 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 mnt_filepath, incite_item_factory +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 @@ -27,11 +27,11 @@ from test_utils.managers.ledger.conftest import create_main_accounts class TestMergePOPLedger: @pytest.fixture - def start(self) -> "datetime": - return datetime(year=2020, month=3, day=14, tzinfo=timezone.utc) + def start(self) -> datetime: + return datetime(year=2020, month=3, day=14, tzinfo=UTC) @pytest.fixture - def duration(self) -> Optional["timedelta"]: + def duration(self) -> timedelta | None: return timedelta(days=5) def test_base( @@ -145,9 +145,9 @@ class TestMergePOPLedger: delete_ledger_db, session_collection, ): + from generalresearch.models.thl.finance import ProductBalances from generalresearch.models.thl.ledger import LedgerAccount from generalresearch.models.thl.product import Product - from generalresearch.models.thl.finance import ProductBalances u = user_factory(product=product, created=session_collection.start) diff --git a/tests/incite/mergers/test_ym_survey_merge.py b/tests/incite/mergers/test_ym_survey_merge.py index 4c2df6b..9107f21 100644 --- a/tests/incite/mergers/test_ym_survey_merge.py +++ b/tests/incite/mergers/test_ym_survey_merge.py @@ -1,4 +1,4 @@ -from datetime import timedelta, timezone, datetime +from datetime import UTC, datetime, timedelta, timezone from itertools import product import pandas as pd @@ -6,16 +6,16 @@ import pytest # noinspection PyUnresolvedReferences from distributed.utils_test import ( - gen_cluster, + cleanup, + client, client_no_amm, + cluster_fixture, + gen_cluster, loop, loop_in_thread, - cleanup, - cluster_fixture, - client, ) -from test_utils.incite.collections.conftest import wall_collection, session_collection +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, @@ -28,11 +28,7 @@ from test_utils.incite.mergers.conftest import ( product( ["12h", "3D"], [timedelta(days=30)], - [ - (datetime.now(tz=timezone.utc) - timedelta(days=35)).replace( - microsecond=0 - ) - ], + [(datetime.now(tz=UTC) - timedelta(days=35)).replace(microsecond=0)], ) ), ) diff --git a/tests/incite/schemas/test_admin_responses.py b/tests/incite/schemas/test_admin_responses.py index 43aa399..29d93fe 100644 --- a/tests/incite/schemas/test_admin_responses.py +++ b/tests/incite/schemas/test_admin_responses.py @@ -1,4 +1,4 @@ -from datetime import datetime, timezone, timedelta +from datetime import UTC, datetime, timedelta, timezone from random import sample from typing import List @@ -8,8 +8,8 @@ import pytest from generalresearch.incite.schemas import empty_dataframe_from_schema from generalresearch.incite.schemas.admin_responses import ( - AdminPOPSchema, SIX_HOUR_SECONDS, + AdminPOPSchema, ) from generalresearch.locales import Localelator @@ -72,8 +72,7 @@ class TestAdminPOPSchema: def test_index_tz_parser(self): tz_dates = [ - datetime(year=2024, month=1, day=i, tzinfo=timezone.utc) - for i in range(1, 10) + datetime(year=2024, month=1, day=i, tzinfo=UTC) for i in range(1, 10) ] df = pd.DataFrame( @@ -85,16 +84,16 @@ class TestAdminPOPSchema: df = self.assign_valid_vals(df) # 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 == timezone.utc for ts in timestmaps]) + timestmaps: list[pd.Timestamp] = [i for i in df.index.get_level_values(0)] + 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)] + timestmaps: list[pd.Timestamp] = [i for i in df.index.get_level_values(0)] assert all([ts.tz is None for ts in timestmaps]) def test_index_tz_no_future_beyond_one_year(self): - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) tz_dates = [now + timedelta(days=i * 365) for i in range(1, 10)] df = pd.DataFrame( @@ -157,9 +156,7 @@ class TestAdminPOPSchema: def test_invalid_parsing(self): # (1) Timezones AND as strings will still parse correctly tz_str_dates = [ - datetime( - year=2024, month=1, day=1, minute=i, tzinfo=timezone.utc - ).isoformat() + datetime(year=2024, month=1, day=1, minute=i, tzinfo=UTC).isoformat() for i in range(1, 10) ] df = pd.DataFrame( @@ -173,12 +170,12 @@ class TestAdminPOPSchema: df = AdminPOPSchema.validate(df, lazy=True) assert isinstance(df, pd.DataFrame) - timestmaps: List[pd.Timestamp] = [i for i in df.index.get_level_values(0)] + timestmaps: list[pd.Timestamp] = [i for i in df.index.get_level_values(0)] assert all([ts.tz is None for ts in timestmaps]) # (2) Timezones are removed dates = [ - datetime(year=2024, month=1, day=1, minute=i, tzinfo=timezone.utc) + datetime(year=2024, month=1, day=1, minute=i, tzinfo=UTC) for i in range(1, 10) ] df = pd.DataFrame( @@ -190,12 +187,12 @@ class TestAdminPOPSchema: df = self.assign_valid_vals(df) # 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 timezone.utc for ts in timestmaps]) + timestmaps: list[pd.Timestamp] = [i for i in df.index.get_level_values(0)] + 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)] + timestmaps: list[pd.Timestamp] = [i for i in df.index.get_level_values(0)] assert all([ts.tz is None for ts in timestmaps]) def test_clipping(self): diff --git a/tests/incite/test_collection_base.py b/tests/incite/test_collection_base.py index 7e6605f..5a63019 100644 --- a/tests/incite/test_collection_base.py +++ b/tests/incite/test_collection_base.py @@ -1,4 +1,4 @@ -from datetime import datetime, timedelta, timezone +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 @@ -12,11 +12,9 @@ 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=timezone.utc) - timedelta(minutes=15)).replace( - microsecond=0 -) -AGO_1HR = (datetime.now(tz=timezone.utc) - timedelta(hours=1)).replace(microsecond=0) -AGO_2HR = (datetime.now(tz=timezone.utc) - timedelta(hours=2)).replace(microsecond=0) +AGO_15min = (datetime.now(tz=UTC) - timedelta(minutes=15)).replace(microsecond=0) +AGO_1HR = (datetime.now(tz=UTC) - timedelta(hours=1)).replace(microsecond=0) +AGO_2HR = (datetime.now(tz=UTC) - timedelta(hours=2)).replace(microsecond=0) class TestCollectionBase: @@ -50,7 +48,7 @@ class TestCollectionBase: with pytest.raises(expected_exception=ValueError) as cm: cm: ExceptionInfo CollectionBase( - start=datetime.now(tz=timezone.utc) - timedelta(days=10), + start=datetime.now(tz=UTC) - timedelta(days=10), archive_path=mnt_filepath.data_src, ) assert "Collection.start must not have microseconds" in str(cm.value) @@ -66,9 +64,7 @@ class TestCollectionBase: assert "Timezone is not UTC" in str(cm.value) instance = CollectionBase(archive_path=mnt_filepath.data_src) - assert instance.start == datetime( - year=2018, month=1, day=1, tzinfo=timezone.utc - ) + assert instance.start == datetime(year=2018, month=1, day=1, tzinfo=UTC) with pytest.raises(expected_exception=ValueError) as cm: cm: ExceptionInfo @@ -145,7 +141,7 @@ class TestCollectionBaseProperties: instance._interval_range(end=datetime.now(tz=tz)) assert "Timezones must match" in str(cm.value) - res = instance._interval_range(end=datetime.now(tz=timezone.utc)) + res = instance._interval_range(end=datetime.now(tz=UTC)) assert isinstance(res, pd.IntervalIndex) assert res.closed_left assert res.is_non_overlapping_monotonic @@ -282,7 +278,7 @@ class TestCollectionBaseMethodsSourceTiming: def test_get_item_start(self, mnt_filepath): instance = CollectionBase(archive_path=mnt_filepath.data_src) - dt = datetime.now(tz=timezone.utc) + dt = datetime.now(tz=UTC) start = pd.Timestamp(dt) with pytest.raises(expected_exception=NotImplementedError) as cm: @@ -292,7 +288,7 @@ class TestCollectionBaseMethodsSourceTiming: def test_get_items(self, mnt_filepath): instance = CollectionBase(archive_path=mnt_filepath.data_src) - dt = datetime.now(tz=timezone.utc) + dt = datetime.now(tz=UTC) with pytest.raises(expected_exception=NotImplementedError) as cm: instance.get_items(since=dt) diff --git a/tests/incite/test_collection_base_item.py b/tests/incite/test_collection_base_item.py index e5d1d02..3f4d023 100644 --- a/tests/incite/test_collection_base_item.py +++ b/tests/incite/test_collection_base_item.py @@ -1,4 +1,4 @@ -from datetime import datetime, timezone +from datetime import UTC, datetime, timezone from os.path import join as pjoin from pathlib import Path from uuid import uuid4 @@ -13,7 +13,7 @@ from generalresearch.incite.base import CollectionItemBase class TestCollectionItemBase: def test_init(self): - dt = datetime.now(tz=timezone.utc).replace(microsecond=0) + dt = datetime.now(tz=UTC).replace(microsecond=0) instance = CollectionItemBase() instance2 = CollectionItemBase(start=dt) @@ -25,7 +25,7 @@ class TestCollectionItemBase: assert 0 == instance.start.microsecond == instance2.start.microsecond def test_init_start(self): - dt = datetime.now(tz=timezone.utc) + dt = datetime.now(tz=UTC) with pytest.raises(expected_exception=ValidationError) as cm: CollectionItemBase(start=dt) diff --git a/tests/managers/leaderboard.py b/tests/managers/leaderboard.py index 4d32dd0..149bdbb 100644 --- a/tests/managers/leaderboard.py +++ b/tests/managers/leaderboard.py @@ -1,7 +1,7 @@ import os import time import zoneinfo -from datetime import datetime, timezone +from datetime import UTC, datetime, timezone from decimal import Decimal from uuid import uuid4 @@ -10,9 +10,6 @@ import pytest from generalresearch.managers.leaderboard.manager import LeaderboardManager from generalresearch.managers.leaderboard.tasks import hit_leaderboards from generalresearch.models.thl.definitions import Status -from generalresearch.models.thl.user import User -from generalresearch.models.thl.product import Product -from generalresearch.models.thl.session import Session from generalresearch.models.thl.leaderboard import ( LeaderboardCode, LeaderboardFrequency, @@ -22,7 +19,10 @@ from generalresearch.models.thl.product import ( PayoutConfig, PayoutTransformation, PayoutTransformationPercentArgs, + Product, ) +from generalresearch.models.thl.session import Session +from generalresearch.models.thl.user import User # random uuid for leaderboard tests product_id = uuid4().hex @@ -63,7 +63,7 @@ def _create_session( ) session = Session( user=user, - started=datetime(2025, 2, 5, 6, tzinfo=timezone.utc), + started=datetime(2025, 2, 5, 6, tzinfo=UTC), id=1, country_iso=country_iso, status=Status.COMPLETE, @@ -152,7 +152,7 @@ class TestLeaderboards: 999999, tzinfo=zoneinfo.ZoneInfo(key="America/New_York"), ) - assert lb.period_start_utc == datetime(2025, 2, 5, 5, tzinfo=timezone.utc) + assert lb.period_start_utc == datetime(2025, 2, 5, 5, tzinfo=UTC) assert lb.row_count == 7 assert lb.rows == [ LeaderboardRow(bpuid="aaa", rank=1, value=10), @@ -270,5 +270,5 @@ class TestLeaderboards: ) assert lb.local_start_time == "2025-02-01T00:00:00+09:00" assert lb.local_end_time == "2025-02-01T23:59:59.999999+09:00" - assert lb.period_start_utc == datetime(2025, 1, 31, 15, tzinfo=timezone.utc) + assert lb.period_start_utc == datetime(2025, 1, 31, 15, tzinfo=UTC) print(lb.model_dump(mode="json")) diff --git a/tests/managers/test_events.py b/tests/managers/test_events.py index a0fab38..6941c00 100644 --- a/tests/managers/test_events.py +++ b/tests/managers/test_events.py @@ -1,22 +1,22 @@ +import math import random import time -from datetime import timedelta, datetime, timezone +from datetime import UTC, datetime, timedelta, timezone from decimal import Decimal from functools import partial +from math import floor from typing import Optional from uuid import uuid4 -import math import pytest -from math import floor from generalresearch.managers.events import EventSubscriber from generalresearch.models import Source from generalresearch.models.events import ( - MessageKind, - EventType, AggregateBySource, + EventType, MaxGaugeBySource, + MessageKind, ) from generalresearch.models.legacy.bucket import Bucket from generalresearch.models.thl.definitions import Status, StatusCode1 @@ -41,13 +41,13 @@ def event_subscriber(thl_redis_config, product_id): def create_dummy( - product_id: Optional[str] = None, product_user_id: Optional[str] = None + product_id: str | None = None, product_user_id: str | None = None ) -> User: return User( product_id=product_id, product_user_id=product_user_id or uuid4().hex, uuid=uuid4().hex, - created=datetime.now(tz=timezone.utc), + created=datetime.now(tz=UTC), user_id=random.randint(0, floor(2**32 / 2)), ) @@ -496,7 +496,7 @@ class TestChannelsSubscriptions: wall.update( status=Status.COMPLETE, status_code_1=StatusCode1.COMPLETE, - finished=datetime.now(tz=timezone.utc), + finished=datetime.now(tz=UTC), cpi=Decimal("1"), ) event_manager.handle_task_finish(wall, session, user) diff --git a/tests/managers/thl/test_contest/test_leaderboard.py b/tests/managers/thl/test_contest/test_leaderboard.py index 80a88a5..1a52f83 100644 --- a/tests/managers/thl/test_contest/test_leaderboard.py +++ b/tests/managers/thl/test_contest/test_leaderboard.py @@ -1,10 +1,10 @@ -from datetime import datetime, timezone, timedelta +from datetime import UTC, datetime, timedelta, timezone from zoneinfo import ZoneInfo from generalresearch.currency import USDCent from generalresearch.models.thl.contest.definitions import ( - ContestStatus, ContestEndReason, + ContestStatus, ) from generalresearch.models.thl.contest.leaderboard import ( LeaderboardContest, @@ -13,9 +13,11 @@ 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_in_db as contest_in_db, leaderboard_contest_create as contest_create, ) +from test_utils.managers.contest.conftest import ( + leaderboard_contest_in_db as contest_in_db, +) class TestLeaderboardContestCRUD: @@ -39,7 +41,7 @@ class TestLeaderboardContestCRUD: # We have it set in the fixture as the daily contest for 2025-01-01 assert c.end_condition.ends_at == datetime( 2025, 1, 1, 23, 59, 59, 999999, tzinfo=ZoneInfo("America/New_York") - ).astimezone(tz=timezone.utc) + timedelta(minutes=90) + ).astimezone(tz=UTC) + timedelta(minutes=90) def test_enter( self, diff --git a/tests/managers/thl/test_contest/test_milestone.py b/tests/managers/thl/test_contest/test_milestone.py index 7312a64..66c5dc4 100644 --- a/tests/managers/thl/test_contest/test_milestone.py +++ b/tests/managers/thl/test_contest/test_milestone.py @@ -1,23 +1,29 @@ -from datetime import datetime, timezone +from datetime import UTC, datetime, timezone from generalresearch.models.thl.contest.definitions import ( - ContestStatus, ContestEndReason, + ContestStatus, ) from generalresearch.models.thl.contest.milestone import ( + ContestEntryTrigger, MilestoneContest, MilestoneContestCreate, MilestoneUserView, - ContestEntryTrigger, ) 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, - milestone_contest_in_db as contest_in_db, +) +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: @@ -28,7 +34,7 @@ class TestMilestoneContest: assert not should, msg # Change so that the contest ends now - contest.end_condition.ends_at = datetime.now(tz=timezone.utc) + contest.end_condition.ends_at = datetime.now(tz=UTC) should, msg = contest.should_end() assert should assert msg == ContestEndReason.ENDS_AT diff --git a/tests/managers/thl/test_contest/test_raffle.py b/tests/managers/thl/test_contest/test_raffle.py index 060055a..5804ea3 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 datetime, timezone +from datetime import UTC, datetime, timezone import pytest from pydantic import ValidationError @@ -9,21 +9,19 @@ from generalresearch.managers.thl.ledger_manager.exceptions import ( LedgerTransactionConditionFailedError, ) from generalresearch.models.thl.contest import ( - ContestPrize, - ContestEntryRule, ContestEndCondition, + ContestEntryRule, + ContestPrize, ) from generalresearch.models.thl.contest.definitions import ( - ContestStatus, - ContestPrizeKind, ContestEndReason, + ContestPrizeKind, + ContestStatus, ) from generalresearch.models.thl.contest.exceptions import ContestError from generalresearch.models.thl.contest.raffle import ( ContestEntry, ContestEntryType, -) -from generalresearch.models.thl.contest.raffle import ( RaffleContest, RaffleContestCreate, RaffleUserView, @@ -32,10 +30,16 @@ 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, - raffle_contest_in_db as contest_in_db, +) +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: @@ -46,7 +50,7 @@ class TestRaffleContest: assert not should, msg # Change so that the contest ends now - contest.end_condition.ends_at = datetime.now(tz=timezone.utc) + contest.end_condition.ends_at = datetime.now(tz=UTC) should, msg = contest.should_end() assert should assert msg == ContestEndReason.ENDS_AT diff --git a/tests/managers/thl/test_harmonized_uqa.py b/tests/managers/thl/test_harmonized_uqa.py index 6bbbbe1..3b6df48 100644 --- a/tests/managers/thl/test_harmonized_uqa.py +++ b/tests/managers/thl/test_harmonized_uqa.py @@ -1,11 +1,11 @@ -from datetime import datetime, timezone +from datetime import UTC, datetime, timezone import pytest from generalresearch.managers.thl.profiling.uqa import UQAManager from generalresearch.models.thl.profiling.user_question_answer import ( - UserQuestionAnswer, DUMMY_UQA, + UserQuestionAnswer, ) from generalresearch.models.thl.user import User @@ -18,7 +18,7 @@ class TestUQAManager: assert len(uqas) == 0 def test_create(self, uqa_manager: UQAManager, user: User): - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) uqas = [ UserQuestionAnswer( user_id=user.user_id, @@ -38,7 +38,7 @@ class TestUQAManager: assert res[0] == uqas[0] # Same question, so this gets updated - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) uqas_update = [ UserQuestionAnswer( user_id=user.user_id, @@ -57,7 +57,7 @@ class TestUQAManager: assert res[0] == uqas_update[0] # Add a new answer - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) uqas_new = [ UserQuestionAnswer( user_id=user.user_id, @@ -103,7 +103,7 @@ class TestUQAManagerCache: UserQuestionAnswer( question_id="5d6d9f3c03bb40bf9d0a24f306387d7c", answer=("1",), - timestamp=datetime.now(tz=timezone.utc), + timestamp=datetime.now(tz=UTC), country_iso="us", language_iso="eng", property_code="gr:gender", diff --git a/tests/managers/thl/test_ledger/test_lm_accounts.py b/tests/managers/thl/test_ledger/test_lm_accounts.py index 5cfaac1..faef5fb 100644 --- a/tests/managers/thl/test_ledger/test_lm_accounts.py +++ b/tests/managers/thl/test_ledger/test_lm_accounts.py @@ -1,6 +1,7 @@ +from collections.abc import Callable from itertools import product as iproduct from random import randint -from typing import TYPE_CHECKING, Callable +from typing import TYPE_CHECKING from uuid import uuid4 import pytest @@ -51,10 +52,10 @@ class TestLedgerAccountManagerNoResults: def test_get_account_no_results( self, - currency: "LedgerCurrency", + currency: LedgerCurrency, kind: str, - acct_id: "UUIDStr", - lm: "LedgerManager", + acct_id: UUIDStr, + lm: 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 @@ -74,10 +75,10 @@ class TestLedgerAccountManagerNoResults: def test_get_account_no_results_many( self, - currency: "LedgerCurrency", + currency: LedgerCurrency, kind: str, - acct_id: "UUIDStr", - lm: "LedgerManager", + acct_id: UUIDStr, + lm: LedgerManager, ): qn = ":".join([currency, kind, acct_id]) @@ -114,10 +115,10 @@ class TestLedgerAccountManagerCreate: def test_create_account_error_permission( self, - currency: "LedgerCurrency", - account_type: "AccountType", - direction: "Direction", - lm: "LedgerManager", + currency: LedgerCurrency, + account_type: AccountType, + direction: Direction, + lm: LedgerManager, ): """Confirm that the Permission values that are set on the Ledger Manger allow the Creation action to occur. @@ -164,10 +165,10 @@ class TestLedgerAccountManagerCreate: def test_create( self, - currency: "LedgerCurrency", - account_type: "AccountType", - direction: "Direction", - lm: "LedgerManager", + currency: LedgerCurrency, + account_type: AccountType, + direction: Direction, + lm: LedgerManager, ): """Confirm that the Permission values that are set on the Ledger Manger allow the Creation action to occur. @@ -194,10 +195,10 @@ class TestLedgerAccountManagerCreate: def test_get_or_create( self, - currency: "LedgerCurrency", - account_type: "AccountType", - direction: "Direction", - lm: "LedgerManager", + currency: LedgerCurrency, + account_type: AccountType, + direction: Direction, + lm: LedgerManager, ): """Confirm that the Permission values that are set on the Ledger Manger allow the Creation action to occur. @@ -225,7 +226,7 @@ class TestLedgerAccountManagerCreate: class TestLedgerAccountManagerGet: - def test_get(self, ledger_account: "LedgerAccount", lm: "LedgerManager"): + def test_get(self, ledger_account: LedgerAccount, lm: LedgerManager): res = lm.get_account(qualified_name=ledger_account.qualified_name) assert res is not None assert res.uuid == ledger_account.uuid @@ -243,11 +244,11 @@ class TestLedgerAccountManagerGet: def test_get_balance_empty( self, - ledger_account: "LedgerAccount", - ledger_account_credit: "LedgerAccount", - ledger_account_debit: "LedgerAccount", - ledger_tx: "LedgerTransaction", - lm: "LedgerManager", + ledger_account: LedgerAccount, + ledger_account_credit: LedgerAccount, + ledger_account_debit: LedgerAccount, + ledger_tx: LedgerTransaction, + lm: LedgerManager, ): res = lm.get_account_balance(account=ledger_account) assert res == 0 @@ -261,12 +262,12 @@ class TestLedgerAccountManagerGet: @pytest.mark.parametrize("n_times", range(5)) def test_get_account_filtered_balance( self, - ledger_account: "LedgerAccount", - ledger_account_credit: "LedgerAccount", - ledger_account_debit: "LedgerAccount", - ledger_tx: "LedgerTransaction", - n_times: "PositiveInt", - lm: "LedgerManager", + ledger_account: LedgerAccount, + ledger_account_credit: LedgerAccount, + ledger_account_debit: LedgerAccount, + ledger_tx: LedgerTransaction, + n_times: PositiveInt, + lm: LedgerManager, ): """Try searching for random metadata and confirm it's always 0 because Tx can be found. @@ -320,7 +321,7 @@ class TestLedgerAccountManagerGet: ) def test_get_balance_timerange_empty( - self, ledger_account: "LedgerAccount", lm: "LedgerManager" + self, ledger_account: LedgerAccount, lm: LedgerManager ): res = lm.get_account_balance_timerange(account=ledger_account) assert res == 0 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 df2611b..07c3712 100644 --- a/tests/managers/thl/test_ledger/test_lm_tx_locks.py +++ b/tests/managers/thl/test_ledger/test_lm_tx_locks.py @@ -1,7 +1,7 @@ import logging -from datetime import datetime, timezone, timedelta +from collections.abc import Callable +from datetime import UTC, datetime, timedelta, timezone from decimal import Decimal -from typing import Callable import pytest @@ -9,21 +9,21 @@ from generalresearch.managers.thl.ledger_manager.conditions import ( generate_condition_mp_payment, ) from generalresearch.managers.thl.ledger_manager.exceptions import ( + LedgerTransactionCreateError, LedgerTransactionCreateLockError, LedgerTransactionFlagAlreadyExistsError, - LedgerTransactionCreateError, ) from generalresearch.models import Source from generalresearch.models.thl.ledger import LedgerTransaction from generalresearch.models.thl.session import ( - Wall, + Session, Status, StatusCode1, - Session, + Wall, WallAdjustedStatus, ) from generalresearch.models.thl.user import User -from test_utils.models.conftest import user_factory, session, product_user_wallet_no +from test_utils.models.conftest import product_user_wallet_no, session, user_factory logger = logging.getLogger("LedgerManager") @@ -139,7 +139,7 @@ class TestLedgerLocks: delete_ledger_db() create_main_accounts() - now = datetime.now(timezone.utc) - timedelta(hours=1) + now = datetime.now(UTC) - timedelta(hours=1) user: User = user_factory(product=product_user_wallet_no) # A User does a Wall complete on Session.id=1 and the transaction is @@ -283,8 +283,8 @@ class TestLedgerLocks: session_id=3, status=Status.COMPLETE, status_code_1=StatusCode1.COMPLETE, - started=datetime.now(timezone.utc), - finished=datetime.now(timezone.utc) + timedelta(seconds=1), + started=datetime.now(UTC), + finished=datetime.now(UTC) + timedelta(seconds=1), ) thl_lm.create_tx_task_complete(wall=wall1, user=user, created=wall1.started) @@ -327,8 +327,8 @@ class TestLedgerLocks: session_id=3, status=Status.COMPLETE, status_code_1=StatusCode1.COMPLETE, - started=datetime.now(timezone.utc), - finished=datetime.now(timezone.utc) + timedelta(seconds=1), + started=datetime.now(UTC), + finished=datetime.now(UTC) + timedelta(seconds=1), ) thl_lm.create_tx_task_complete(wall1, user, created=wall1.started) 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 294d092..1fb9c01 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 datetime, timezone, timedelta +from datetime import UTC, datetime, timedelta, timezone from decimal import Decimal from random import randint from uuid import uuid4 @@ -11,22 +11,22 @@ from redis.lock import Lock from generalresearch.currency import USDCent from generalresearch.managers.base import Permission -from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager from generalresearch.managers.thl.ledger_manager.exceptions import ( - LedgerTransactionFlagAlreadyExistsError, LedgerTransactionConditionFailedError, - LedgerTransactionReleaseLockError, LedgerTransactionCreateError, + LedgerTransactionFlagAlreadyExistsError, + LedgerTransactionReleaseLockError, ) from generalresearch.managers.thl.ledger_manager.ledger import LedgerTransaction +from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager from generalresearch.models import Source from generalresearch.models.thl.definitions import PayoutStatus from generalresearch.models.thl.ledger import Direction, TransactionType from generalresearch.models.thl.session import ( - Wall, + Session, Status, StatusCode1, - Session, + Wall, ) from generalresearch.models.thl.user import User from generalresearch.models.thl.wallet import PayoutType @@ -55,7 +55,7 @@ class TestThlLedgerManagerBPPayout: delete_ledger_db() create_main_accounts() - now = datetime.now(timezone.utc) - timedelta(hours=1) + now = datetime.now(UTC) - timedelta(hours=1) user: User = user_factory(product=product_user_wallet_no) wall1 = Wall( @@ -158,7 +158,7 @@ class TestThlLedgerManagerBPPayout: product=product, amount=rand_amount, payoutevent_uuid=payoutevent_uuid, - created=datetime.now(tz=timezone.utc), + created=datetime.now(tz=UTC), skip_wallet_balance_check=True, skip_one_per_day_check=True, skip_flag_check=True, @@ -189,7 +189,7 @@ class TestThlLedgerManagerBPPayout: product=product, amount=rand_amount, payoutevent_uuid=uuid4().hex, - created=datetime.now(tz=timezone.utc), + created=datetime.now(tz=UTC), skip_wallet_balance_check=False, skip_one_per_day_check=False, skip_flag_check=False, @@ -199,7 +199,7 @@ class TestThlLedgerManagerBPPayout: def test_create_tx_redis_failure(self, product, thl_web_rw, thl_lm): rand_amount: USDCent = USDCent(randint(100, 1_000)) payoutevent_uuid = uuid4().hex - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) thl_lm.create_tx_plug_bp_wallet( product, rand_amount, now, direction=Direction.CREDIT @@ -226,7 +226,7 @@ class TestThlLedgerManagerBPPayout: product=product, amount=rand_amount, payoutevent_uuid=payoutevent_uuid, - created=datetime.now(tz=timezone.utc), + created=datetime.now(tz=UTC), ) assert e.type is redis.exceptions.TimeoutError # No txs were created @@ -238,7 +238,7 @@ class TestThlLedgerManagerBPPayout: def test_create_tx_multiple_per_day(self, product, thl_lm): rand_amount: USDCent = USDCent(randint(100, 1_000)) payoutevent_uuid = uuid4().hex - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) thl_lm.create_tx_plug_bp_wallet( product, rand_amount * USDCent(2), now, direction=Direction.CREDIT @@ -248,7 +248,7 @@ class TestThlLedgerManagerBPPayout: product=product, amount=rand_amount, payoutevent_uuid=payoutevent_uuid, - created=datetime.now(tz=timezone.utc), + created=datetime.now(tz=UTC), ) # Try to create another @@ -258,7 +258,7 @@ class TestThlLedgerManagerBPPayout: product=product, amount=rand_amount, payoutevent_uuid=payoutevent_uuid, - created=datetime.now(tz=timezone.utc), + created=datetime.now(tz=UTC), ) assert e.type is LedgerTransactionFlagAlreadyExistsError @@ -270,7 +270,7 @@ class TestThlLedgerManagerBPPayout: product=product, amount=rand_amount, payoutevent_uuid=payoutevent_uuid2, - created=datetime.now(tz=timezone.utc), + created=datetime.now(tz=UTC), ) assert e.type is LedgerTransactionConditionFailedError assert str(e.value) == ">1 tx per day" @@ -280,14 +280,14 @@ class TestThlLedgerManagerBPPayout: product=product, amount=rand_amount, payoutevent_uuid=payoutevent_uuid2, - created=datetime.now(tz=timezone.utc), + created=datetime.now(tz=UTC), skip_one_per_day_check=True, ) def test_create_tx_redis_lock_release_error(self, product, thl_lm): rand_amount: USDCent = USDCent(randint(100, 1_000)) payoutevent_uuid = uuid4().hex - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(product=product) thl_lm.create_tx_plug_bp_wallet( @@ -304,7 +304,7 @@ class TestThlLedgerManagerBPPayout: product=product, amount=rand_amount, payoutevent_uuid=payoutevent_uuid, - created=datetime.now(tz=timezone.utc), + created=datetime.now(tz=UTC), ) assert e.type is LedgerTransactionCreateError assert str(e.value) == "Redis error: Simulated timeout during acquire" @@ -321,7 +321,7 @@ class TestThlLedgerManagerBPPayout: product=product, amount=rand_amount, payoutevent_uuid=payoutevent_uuid, - created=datetime.now(tz=timezone.utc), + created=datetime.now(tz=UTC), ) assert e.type is LedgerTransactionReleaseLockError assert str(e.value) == "Redis error: Simulated timeout during release" @@ -337,7 +337,7 @@ class TestPayoutEventManagerBPPayout: def test_create(self, product, thl_lm, brokerage_product_payout_event_manager): rand_amount: USDCent = USDCent(randint(100, 1_000)) - now = datetime.now(tz=timezone.utc) + 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( @@ -369,7 +369,7 @@ class TestPayoutEventManagerBPPayout: original_release = Lock.release rand_amount: USDCent = USDCent(randint(100, 1_000)) - now = datetime.now(tz=timezone.utc) + 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( @@ -435,7 +435,7 @@ class TestPayoutEventManagerBPPayout: # We wouldn't do this in practice, because this is paying out the BP again, but # we can if want to. # Change the timestamp so it'll create a new payout event - now = datetime.now(tz=timezone.utc) + 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, @@ -450,7 +450,7 @@ class TestPayoutEventManagerBPPayout: assert pe.status == PayoutStatus.FAILED # And if we really want to, we can make it again - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) pe = brokerage_product_payout_event_manager.create_bp_payout_event( thl_ledger_manager=thl_lm, product=product, @@ -478,7 +478,7 @@ class TestPayoutEventManagerBPPayout: original_release = Lock.release rand_amount: USDCent = USDCent(randint(100, 1_000)) - now = datetime.now(tz=timezone.utc) + 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) 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 31c7107..be988a1 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 datetime, timezone, timedelta +from datetime import UTC, datetime, timedelta, timezone from decimal import Decimal from random import randint from uuid import uuid4 @@ -14,23 +14,22 @@ from generalresearch.models import Source from generalresearch.models.thl.definitions import ( WALL_ALLOWED_STATUS_STATUS_CODE, ) -from generalresearch.models.thl.ledger import Direction -from generalresearch.models.thl.ledger import TransactionType +from generalresearch.models.thl.ledger import Direction, TransactionType +from generalresearch.models.thl.payout import UserPayoutEvent from generalresearch.models.thl.product import ( PayoutConfig, PayoutTransformation, UserWalletConfig, ) from generalresearch.models.thl.session import ( - Wall, + Session, Status, StatusCode1, - Session, + Wall, WallAdjustedStatus, ) from generalresearch.models.thl.user import User from generalresearch.models.thl.wallet import PayoutType -from generalresearch.models.thl.payout import UserPayoutEvent logger = logging.getLogger("LedgerManager") @@ -82,7 +81,7 @@ class TestThlLedgerTxManager: session=s1, status=Status.COMPLETE, status_code_1=status_code_1, - finished=datetime.now(tz=timezone.utc) + timedelta(minutes=10), + finished=datetime.now(tz=UTC) + timedelta(minutes=10), payout=bp_pay, user_payout=user_pay, ) @@ -127,7 +126,7 @@ class TestThlLedgerTxManager: session=s1, status=Status.COMPLETE, status_code_1=status_code_1, - finished=datetime.now(tz=timezone.utc) + timedelta(minutes=10), + finished=datetime.now(tz=UTC) + timedelta(minutes=10), payout=bp_pay, user_payout=user_pay, ) @@ -209,7 +208,7 @@ class TestThlLedgerTxManager: # there is no financial changes needed session.update( **{ - "finished": datetime.now(tz=timezone.utc) + timedelta(minutes=10), + "finished": datetime.now(tz=UTC) + timedelta(minutes=10), } ) assert session.finished @@ -229,7 +228,7 @@ class TestThlLedgerTxManager: product=product, amount=rand_amount, payoutevent_uuid=payoutevent_uuid, - created=datetime.now(tz=timezone.utc), + created=datetime.now(tz=UTC), skip_wallet_balance_check=True, skip_one_per_day_check=True, skip_flag_check=True, @@ -260,7 +259,7 @@ class TestThlLedgerTxManager: product=product, amount=rand_amount, payoutevent_uuid=uuid4().hex, - created=datetime.now(tz=timezone.utc), + created=datetime.now(tz=UTC), skip_wallet_balance_check=False, skip_one_per_day_check=False, skip_flag_check=False, @@ -276,7 +275,7 @@ class TestThlLedgerTxManager: product=product, amount=rand_amount, payoutevent_uuid=payoutevent_uuid, - created=datetime.now(tz=timezone.utc), + created=datetime.now(tz=UTC), ) # Check the basic attributes @@ -300,7 +299,7 @@ class TestThlLedgerTxManager: tx = thl_lm.create_tx_plug_bp_wallet( product=product, amount=rand_amount, - created=datetime.now(tz=timezone.utc), + created=datetime.now(tz=UTC), direction=Direction.DEBIT, skip_flag_check=False, ) @@ -328,7 +327,7 @@ class TestThlLedgerTxManager: tx = thl_lm.create_tx_plug_bp_wallet_( product=product, amount=rand_amount, - created=datetime.now(tz=timezone.utc), + created=datetime.now(tz=UTC), direction=Direction.DEBIT, ) @@ -345,7 +344,7 @@ class TestThlLedgerTxManager: thl_lm.create_tx_plug_bp_wallet_( product=product, amount=rand_amount + rand_amount, - created=datetime.now(tz=timezone.utc), + created=datetime.now(tz=UTC), direction=Direction.CREDIT, ) balance = thl_lm.get_account_balance( @@ -727,8 +726,8 @@ class TestThlLedgerTxManagerFlows: session_id=1, status=Status.COMPLETE, status_code_1=StatusCode1.COMPLETE, - started=datetime.now(timezone.utc), - finished=datetime.now(timezone.utc) + timedelta(seconds=1), + started=datetime.now(UTC), + finished=datetime.now(UTC) + timedelta(seconds=1), ) thl_lm.create_tx_task_complete(wall=wall1, user=user, created=wall1.started) @@ -740,8 +739,8 @@ class TestThlLedgerTxManagerFlows: session_id=1, status=Status.COMPLETE, status_code_1=StatusCode1.COMPLETE, - started=datetime.now(timezone.utc), - finished=datetime.now(timezone.utc) + timedelta(seconds=1), + started=datetime.now(UTC), + finished=datetime.now(UTC) + timedelta(seconds=1), ) thl_lm.create_tx_task_complete(wall=wall2, user=user, created=wall2.started) @@ -793,8 +792,8 @@ class TestThlLedgerTxManagerFlows: session_id=1, status=Status.COMPLETE, status_code_1=StatusCode1.COMPLETE, - started=datetime.now(timezone.utc), - finished=datetime.now(timezone.utc) + timedelta(seconds=1), + started=datetime.now(UTC), + finished=datetime.now(UTC) + timedelta(seconds=1), ) tx = thl_lm.create_tx_task_complete( wall=wall1, user=user, created=wall1.started @@ -880,8 +879,8 @@ class TestThlLedgerTxManagerFlows: session_id=3, status=Status.COMPLETE, status_code_1=StatusCode1.COMPLETE, - started=datetime.now(timezone.utc), - finished=datetime.now(timezone.utc) + timedelta(seconds=1), + started=datetime.now(UTC), + finished=datetime.now(UTC) + timedelta(seconds=1), ) tx = thl_lm.create_tx_task_complete( @@ -922,8 +921,8 @@ class TestThlLedgerTxManagerFlows: session_id=3, status=Status.COMPLETE, status_code_1=StatusCode1.COMPLETE, - started=datetime.now(timezone.utc), - finished=datetime.now(timezone.utc) + timedelta(seconds=1), + started=datetime.now(UTC), + finished=datetime.now(UTC) + timedelta(seconds=1), ) thl_lm.create_tx_task_complete(wall=wall1, user=user, created=wall1.started) @@ -963,8 +962,8 @@ class TestThlLedgerTxManagerFlows: session_id=3, status=Status.COMPLETE, status_code_1=StatusCode1.COMPLETE, - started=datetime.now(timezone.utc), - finished=datetime.now(timezone.utc) + timedelta(seconds=1), + started=datetime.now(UTC), + finished=datetime.now(UTC) + timedelta(seconds=1), ) thl_lm.create_tx_task_complete(wall=wall1, user=user, created=wall1.started) @@ -1416,7 +1415,7 @@ class TestThlLedgerManagerAdj: delete_ledger_db() create_main_accounts() - now = datetime.now(timezone.utc) - timedelta(days=1) + now = datetime.now(UTC) - timedelta(days=1) user: User = user_factory(product=product_user_wallet_yes) # Create 2 Wall completes and create the respective transaction for 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 1e7146a..9253ff0 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,17 +1,17 @@ import logging -from datetime import datetime, timezone, timedelta +from datetime import UTC, datetime, timedelta, timezone from decimal import Decimal from uuid import uuid4 import pytest from generalresearch.managers.thl.ledger_manager.exceptions import ( - LedgerTransactionFlagAlreadyExistsError, LedgerTransactionConditionFailedError, + LedgerTransactionFlagAlreadyExistsError, ) +from generalresearch.models.thl.payout import UserPayoutEvent from generalresearch.models.thl.user import User from generalresearch.models.thl.wallet import PayoutType -from generalresearch.models.thl.payout import UserPayoutEvent from test_utils.managers.ledger.conftest import create_main_accounts @@ -243,7 +243,7 @@ class TestLedgerManagerAMT: delete_ledger_db() create_main_accounts() - now = datetime.now(timezone.utc) - timedelta(hours=1) + now = datetime.now(UTC) - timedelta(hours=1) user: User = user_factory(product=product_amt_true) pe = UserPayoutEvent( @@ -394,7 +394,7 @@ class TestLedgerManagerPaypal: delete_ledger_db() create_main_accounts() - now = datetime.now(tz=timezone.utc) - timedelta(hours=1) + 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? diff --git a/tests/managers/thl/test_ledger/test_user_txs.py b/tests/managers/thl/test_ledger/test_user_txs.py index ecf146f..b4b0437 100644 --- a/tests/managers/thl/test_ledger/test_user_txs.py +++ b/tests/managers/thl/test_ledger/test_user_txs.py @@ -1,6 +1,7 @@ -from datetime import datetime, timedelta, timezone +from collections.abc import Callable +from datetime import UTC, datetime, timedelta, timezone from decimal import Decimal -from typing import TYPE_CHECKING, Callable +from typing import TYPE_CHECKING from uuid import uuid4 from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager @@ -24,8 +25,8 @@ if TYPE_CHECKING: def test_user_txs( - user_factory: Callable[..., "User"], - product_amt_true: "Product", + user_factory: Callable[..., User], + product_amt_true: Product, create_main_accounts: Callable[..., None], thl_lm: ThlLedgerManager, lm, @@ -36,7 +37,7 @@ def test_user_txs( session_factory, user_payout_event_manager, utc_now: datetime, - settings: "GRLSettings", + settings: GRLSettings, ): delete_ledger_db() create_main_accounts() @@ -136,13 +137,13 @@ def test_user_txs( def test_user_txs_pagination( - user_factory: Callable[..., "User"], - product_amt_true: "Product", + user_factory: Callable[..., User], + product_amt_true: Product, create_main_accounts: Callable[..., None], - thl_lm: "ThlLedgerManager", - lm: "LedgerManager", + thl_lm: ThlLedgerManager, + lm: LedgerManager, delete_ledger_db: Callable[..., None], - session_with_tx_factory: Callable[..., "Session"], + session_with_tx_factory: Callable[..., Session], adj_to_fail_with_tx_factory, user_payout_event_manager, utc_now: datetime, @@ -187,7 +188,7 @@ def test_user_txs_pagination( assert txs.summary.user_bonus.entry_count == 12 # Test filtering. We should pull back only this one - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) user_compensate( ledger_manager=thl_lm, user=user, @@ -203,7 +204,7 @@ def test_user_txs_pagination( assert txs.summary.user_bonus.entry_count == 1 # And filtering with 0 results - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) txs = thl_lm.get_user_txs(user, page=1, size=5, time_start=now) assert len(txs.transactions) == 0 assert txs.total == 0 @@ -215,8 +216,8 @@ def test_user_txs_pagination( def test_user_txs_rolling_balance( - user_factory: Callable[..., "User"], - product_amt_true: "Product", + user_factory: Callable[..., User], + product_amt_true: Product, create_main_accounts, thl_lm, lm, @@ -224,7 +225,7 @@ def test_user_txs_rolling_balance( session_with_tx_factory, adj_to_fail_with_tx_factory, user_payout_event_manager, - settings: "GRLSettings", + settings: GRLSettings, ): """ Creates 3 $1.00 bonuses (postive), diff --git a/tests/managers/thl/test_maxmind.py b/tests/managers/thl/test_maxmind.py index c588c58..75bf0e9 100644 --- a/tests/managers/thl/test_maxmind.py +++ b/tests/managers/thl/test_maxmind.py @@ -1,17 +1,12 @@ import json import logging -from typing import Callable +from collections.abc import Callable -import geoip2.models import pytest from faker import Faker from faker.providers.address.en_US import Provider as USAddressProvider from generalresearch.managers.thl.ipinfo import GeoIpInfoManager -from generalresearch.managers.thl.maxmind import MaxmindManager -from generalresearch.managers.thl.maxmind.basic import ( - MaxmindBasicManager, -) from generalresearch.models.thl.ipinfo import ( GeoIPInformation, normalize_ip, diff --git a/tests/managers/thl/test_profiling/test_user_upk.py b/tests/managers/thl/test_profiling/test_user_upk.py index 53bb8fe..491e2b1 100644 --- a/tests/managers/thl/test_profiling/test_user_upk.py +++ b/tests/managers/thl/test_profiling/test_user_upk.py @@ -1,8 +1,8 @@ -from datetime import datetime, timezone +from datetime import UTC, datetime, timezone from generalresearch.managers.thl.profiling.user_upk import UserUpkManager -now = datetime.now(tz=timezone.utc) +now = datetime.now(tz=UTC) base = { "country_iso": "us", "language_iso": "eng", diff --git a/tests/managers/thl/test_survey.py b/tests/managers/thl/test_survey.py index 58c4577..4b4a579 100644 --- a/tests/managers/thl/test_survey.py +++ b/tests/managers/thl/test_survey.py @@ -1,24 +1,24 @@ import uuid -from datetime import datetime, timezone +from datetime import UTC, datetime, timezone from decimal import Decimal import pytest from generalresearch.models import Source from generalresearch.models.legacy.bucket import ( - SurveyEligibilityCriterion, - TopNPlusBucket, DurationSummary, PayoutSummary, + SurveyEligibilityCriterion, + TopNPlusBucket, ) from generalresearch.models.thl.profiling.user_question_answer import ( UserQuestionAnswer, ) from generalresearch.models.thl.survey.model import ( Survey, - SurveyStat, SurveyCategoryModel, SurveyEligibilityDefinition, + SurveyStat, ) @@ -258,7 +258,7 @@ class TestSurveyStat: return # 1,000 of the 20,000 are "new" - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) for s in ss[:1000]: s.survey__survey_id = "b" s.updated_at = now @@ -298,7 +298,7 @@ class TestSurveyStat: source=source, surveys=surveys, survey_stats=survey_stats ) # UPDATE ------- - since = datetime.now(tz=timezone.utc) + since = datetime.now(tz=UTC) print(f"{since=}") # 10 survey disappear diff --git a/tests/managers/thl/test_task_adjustment.py b/tests/managers/thl/test_task_adjustment.py index 839bbe1..71e3535 100644 --- a/tests/managers/thl/test_task_adjustment.py +++ b/tests/managers/thl/test_task_adjustment.py @@ -1,9 +1,9 @@ import logging +from datetime import UTC, datetime, timedelta, timezone +from decimal import Decimal from random import randint import pytest -from datetime import datetime, timezone, timedelta -from decimal import Decimal from generalresearch.models import Source from generalresearch.models.thl.definitions import ( @@ -31,16 +31,14 @@ def session_complete_with_wallet(session_with_tx_factory, user_with_wallet): @pytest.fixture() def session_fail(user, session_manager, wall_manager): - session = session_manager.create_dummy( - started=datetime.now(timezone.utc), user=user - ) + session = session_manager.create_dummy(started=datetime.now(UTC), user=user) wall1 = wall_manager.create_dummy( session_id=session.id, user_id=user.user_id, source=Source.DYNATA, req_survey_id="72723", req_cpi=Decimal("3.22"), - started=datetime.now(timezone.utc), + started=datetime.now(UTC), ) wall_manager.finish( wall=wall1, @@ -109,7 +107,7 @@ class TestHandleRecons: assert ledger_manager.get_account_balance(commission_account) == 0 # Now, say we get the exact same *adjust to incomplete* msg again. It should do nothing! - adjusted_timestamp = datetime.now(tz=timezone.utc) + adjusted_timestamp = datetime.now(tz=UTC) wall = wall_manager.get_from_uuid(wall_uuid=wall_uuid) with pytest.raises(match=" is already "): wall_manager.adjust_status( diff --git a/tests/managers/thl/test_task_status.py b/tests/managers/thl/test_task_status.py index 55c89c0..468fd5e 100644 --- a/tests/managers/thl/test_task_status.py +++ b/tests/managers/thl/test_task_status.py @@ -1,31 +1,31 @@ -import pytest -from datetime import datetime, timezone, timedelta +from datetime import UTC, datetime, timedelta, timezone from decimal import Decimal +import pytest + from generalresearch.managers.thl.session import SessionManager from generalresearch.models import Source from generalresearch.models.thl.definitions import ( Status, - WallAdjustedStatus, StatusCode1, + WallAdjustedStatus, ) from generalresearch.models.thl.product import ( PayoutConfig, - UserWalletConfig, PayoutTransformation, PayoutTransformationPercentArgs, + UserWalletConfig, ) from generalresearch.models.thl.session import Session, WallOut from generalresearch.models.thl.task_status import TaskStatusResponse from generalresearch.models.thl.user import User - -start1 = datetime(2023, 2, 1, tzinfo=timezone.utc) +start1 = datetime(2023, 2, 1, tzinfo=UTC) finish1 = start1 + timedelta(minutes=5) recon1 = start1 + timedelta(days=20) -start2 = datetime(2023, 2, 2, tzinfo=timezone.utc) +start2 = datetime(2023, 2, 2, tzinfo=UTC) finish2 = start2 + timedelta(minutes=5) -start3 = datetime(2023, 2, 3, tzinfo=timezone.utc) +start3 = datetime(2023, 2, 3, tzinfo=UTC) finish3 = start3 + 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 0d7ffef..2704490 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 datetime, timezone +from datetime import UTC, datetime, timezone from random import randint from uuid import uuid4 @@ -118,7 +118,7 @@ class TestBlockUserManager: ) assert not user.blocked - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) # Adds user to whitelist thl_web_rw.execute_write( """ diff --git a/tests/managers/thl/test_user_streak.py b/tests/managers/thl/test_user_streak.py index 7728f9f..ef25e2b 100644 --- a/tests/managers/thl/test_user_streak.py +++ b/tests/managers/thl/test_user_streak.py @@ -1,17 +1,17 @@ import copy -from datetime import datetime, timezone, timedelta, date +from datetime import UTC, date, datetime, timedelta, timezone from decimal import Decimal from zoneinfo import ZoneInfo import pytest from generalresearch.managers.thl.user_streak import compute_streaks_from_days -from generalresearch.models.thl.definitions import StatusCode1, Status +from generalresearch.models.thl.definitions import Status, StatusCode1 from generalresearch.models.thl.user_streak import ( - UserStreak, - StreakState, - StreakPeriod, StreakFulfillment, + StreakPeriod, + StreakState, + UserStreak, ) @@ -126,7 +126,7 @@ def test_user_streaks_active_broken( user_streak_manager, user, session_manager, broken_active_streak ): # Testing active streak, but broken (not today or yesterday) - start1 = datetime(2025, 2, 12, tzinfo=timezone.utc) + start1 = datetime(2025, 2, 12, tzinfo=UTC) end1 = start1 + timedelta(minutes=1) # abandon counts as inactive @@ -176,7 +176,7 @@ def test_user_streak_complete_active(user_streak_manager, user, session_manager) # They completed yesterday NY time. Today isn't over so streak is pending start1 = datetime.now(tz=ZoneInfo("America/New_York")) - timedelta(days=1) - create_session_complete(session_manager, start1.astimezone(tz=timezone.utc), user) + create_session_complete(session_manager, start1.astimezone(tz=UTC), user) last_complete_day = start1.date() expected_streak = UserStreak( @@ -201,7 +201,7 @@ def test_user_streak_complete_active(user_streak_manager, user, session_manager) # And now they complete today start2 = datetime.now(tz=ZoneInfo("America/New_York")) - create_session_complete(session_manager, start2.astimezone(tz=timezone.utc), user) + create_session_complete(session_manager, start2.astimezone(tz=UTC), user) last_complete_day = start2.date() expected_streak = UserStreak( longest_streak=2, diff --git a/tests/managers/thl/test_userhealth.py b/tests/managers/thl/test_userhealth.py index 1cda8de..e2ea4a8 100644 --- a/tests/managers/thl/test_userhealth.py +++ b/tests/managers/thl/test_userhealth.py @@ -1,4 +1,4 @@ -from datetime import timezone, datetime +from datetime import UTC, datetime, timezone from uuid import uuid4 import faker @@ -12,7 +12,7 @@ from generalresearch.models.thl.ipinfo import GeoIPInformation from generalresearch.models.thl.user_iphistory import ( IPRecord, ) -from generalresearch.models.thl.userhealth import AuditLogLevel, AuditLog +from generalresearch.models.thl.userhealth import AuditLog, AuditLogLevel fake = faker.Faker() @@ -51,7 +51,7 @@ class TestAuditLog: res = audit_log_manager.get_by_id(auditlog_id=audit_log.id) assert isinstance(res, AuditLog) assert res.id == audit_log.id - assert res.created.tzinfo == timezone.utc + assert res.created.tzinfo == UTC def test_filter_by_product( self, @@ -179,7 +179,7 @@ class TestAuditLog: res = audit_log_manager.filter_count( user_ids=[u1.user_id, u2.user_id, u3.user_id], - created_after=datetime.now(tz=timezone.utc), + created_after=datetime.now(tz=UTC), ) assert isinstance(res, int) assert res == 0 diff --git a/tests/managers/thl/test_wall_manager.py b/tests/managers/thl/test_wall_manager.py index ee44e23..46d199f 100644 --- a/tests/managers/thl/test_wall_manager.py +++ b/tests/managers/thl/test_wall_manager.py @@ -1,4 +1,4 @@ -from datetime import datetime, timezone, timedelta +from datetime import UTC, datetime, timedelta, timezone from decimal import Decimal from uuid import uuid4 @@ -10,7 +10,7 @@ from generalresearch.models.thl.session import ( Status, StatusCode1, ) -from test_utils.models.conftest import user, session +from test_utils.models.conftest import session, user class TestWallManager: @@ -88,7 +88,7 @@ class TestWallManager: session_id=session.id, user_id=user.user_id, uuid_id=uuid4().hex, - started=datetime.now(tz=timezone.utc), + started=datetime.now(tz=UTC), source=Source.DYNATA, buyer_id="123", req_survey_id="456", @@ -217,9 +217,9 @@ class TestWallCacheManager: def test_get_wall_events( self, wall_cache_manager, wall_manager, session_manager, user ): - start1 = datetime.now(timezone.utc) - timedelta(hours=3) - start2 = datetime.now(timezone.utc) - timedelta(hours=2) - start3 = datetime.now(timezone.utc) - timedelta(hours=1) + start1 = datetime.now(UTC) - timedelta(hours=3) + start2 = datetime.now(UTC) - timedelta(hours=2) + start3 = datetime.now(UTC) - timedelta(hours=1) session = session_manager.create_dummy(started=start1, user=user) wall1 = wall_manager.create_dummy( diff --git a/tests/models/admin/test_report_request.py b/tests/models/admin/test_report_request.py index a80afbe..4626ab4 100644 --- a/tests/models/admin/test_report_request.py +++ b/tests/models/admin/test_report_request.py @@ -1,4 +1,4 @@ -from datetime import datetime, timezone +from datetime import UTC, datetime, timezone import pandas as pd import pytest @@ -19,19 +19,19 @@ class TestReportRequest: assert rr.report_type == ReportType.POP_SESSION assert rr.start != rr.start_floor, "rr.start != rr.start_floor" - assert rr.start_floor.tzinfo == timezone.utc, "rr.start_floor.tzinfo not utc" + assert rr.start_floor.tzinfo == UTC, "rr.start_floor.tzinfo not utc" rr1 = ReportRequest.model_validate( { "start": datetime( - year=datetime.now(tz=timezone.utc).year, + year=datetime.now(tz=UTC).year, month=1, day=1, hour=0, minute=30, second=25, microsecond=35, - tzinfo=timezone.utc, + tzinfo=UTC, ), "interval": "1h", } @@ -43,14 +43,14 @@ class TestReportRequest: rr2 = ReportRequest.model_validate( { "start": datetime( - year=datetime.now(tz=timezone.utc).year, + year=datetime.now(tz=UTC).year, month=1, day=1, hour=6, minute=30, second=25, microsecond=35, - tzinfo=timezone.utc, + tzinfo=UTC, ), "interval": "1d", } @@ -92,8 +92,8 @@ class TestReportRequest: with pytest.raises(expected_exception=ValidationError): ReportRequest.model_validate( { - "start": datetime(year=1990, month=1, day=1, tzinfo=timezone.utc), - "end": datetime(year=1950, month=1, day=1, tzinfo=timezone.utc), + "start": datetime(year=1990, month=1, day=1, tzinfo=UTC), + "end": datetime(year=1950, month=1, day=1, tzinfo=UTC), } ) @@ -156,8 +156,8 @@ class TestReportRequest: rr = ReportRequest.model_validate( { "interval": "1d", - "start": datetime(year=2000, month=1, day=1, tzinfo=timezone.utc), - "end": datetime(year=2000, month=1, day=10, tzinfo=timezone.utc), + "start": datetime(year=2000, month=1, day=1, tzinfo=UTC), + "end": datetime(year=2000, month=1, day=10, tzinfo=UTC), } ) diff --git a/tests/models/custom_types/test_aware_datetime.py b/tests/models/custom_types/test_aware_datetime.py index 530142e..043fba0 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 datetime, timezone +from datetime import UTC, datetime, timezone import pytest import pytz @@ -27,14 +27,14 @@ class TestAwareDatetimeISO: AwareDatetimeISOModel.model_validate_json(t.model_dump_json()) def test_dt(self): - dt = datetime(2023, 10, 10, 1, 1, 1, tzinfo=timezone.utc) + dt = datetime(2023, 10, 10, 1, 1, 1, tzinfo=UTC) t = AwareDatetimeISOModel(dt=dt, dt_optional=dt) AwareDatetimeISOModel.model_validate_json(t.model_dump_json()) t = AwareDatetimeISOModel(dt=dt, dt_optional=None) AwareDatetimeISOModel.model_validate_json(t.model_dump_json()) - dt = datetime(2023, 10, 10, 1, 1, 1, microsecond=123, tzinfo=timezone.utc) + dt = datetime(2023, 10, 10, 1, 1, 1, microsecond=123, tzinfo=UTC) t = AwareDatetimeISOModel(dt=dt, dt_optional=dt) AwareDatetimeISOModel.model_validate_json(t.model_dump_json()) diff --git a/tests/models/custom_types/test_dsn.py b/tests/models/custom_types/test_dsn.py index 16e1f83..050976e 100644 --- a/tests/models/custom_types/test_dsn.py +++ b/tests/models/custom_types/test_dsn.py @@ -11,9 +11,9 @@ from generalresearch.models.custom_types import DaskDsn, SentryDsn class SettingsModel(BaseModel): - dask: Optional["DaskDsn"] = Field(default=None) - sentry: Optional["SentryDsn"] = Field(default=None) - db: Optional["MySQLDsn"] = Field(default=None) + dask: DaskDsn | None = Field(default=None) + sentry: SentryDsn | None = Field(default=None) + db: MySQLDsn | None = Field(default=None) # --- Pytest themselves --- diff --git a/tests/models/dynata/test_eligbility.py b/tests/models/dynata/test_eligbility.py index 736c971..16cad26 100644 --- a/tests/models/dynata/test_eligbility.py +++ b/tests/models/dynata/test_eligbility.py @@ -1,4 +1,4 @@ -from datetime import datetime, timezone +from datetime import UTC, datetime, timezone class TestEligibility: @@ -40,7 +40,7 @@ class TestEligibility: "project_id": "p1", "status": "OPEN", "project_exclusions": set(), - "created": datetime.now(tz=timezone.utc), + "created": datetime.now(tz=UTC), "category_exclusions": set(), "category_ids": set(), "cpi": 1, @@ -172,7 +172,7 @@ class TestEligibility: "project_id": "p1", "status": "OPEN", "project_exclusions": set(), - "created": datetime.now(tz=timezone.utc), + "created": datetime.now(tz=UTC), "category_exclusions": set(), "category_ids": set(), "cpi": 1, diff --git a/tests/models/gr/test_authentication.py b/tests/models/gr/test_authentication.py index 6c84a5d..51595a7 100644 --- a/tests/models/gr/test_authentication.py +++ b/tests/models/gr/test_authentication.py @@ -1,9 +1,9 @@ import binascii import json import os -from datetime import datetime, timezone +from collections.abc import Callable +from datetime import UTC, datetime, timezone from random import randint -from typing import Callable from uuid import uuid4 import pytest @@ -251,7 +251,7 @@ class TestGRToken: def gr_token(self, gr_user): from generalresearch.models.gr.authentication import GRToken - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) token = binascii.hexlify(os.urandom(20)).decode() gr_token = GRToken(key=token, created=now, user_id=gr_user.id) diff --git a/tests/models/gr/test_base.py b/tests/models/gr/test_base.py index a9f01a8..8da28d3 100644 --- a/tests/models/gr/test_base.py +++ b/tests/models/gr/test_base.py @@ -1,6 +1,6 @@ import subprocess +from collections.abc import Callable from pathlib import Path -from typing import Callable import pytest from pydantic import PostgresDsn @@ -10,9 +10,11 @@ from generalresearch.pg_helper import PostgresConfig class TestGRPostgresDjangoCreation: - def test_git(self, git_key_path: Path, gr_repo: Callable[..., Path]): + def test_git(self, gr_repo: Callable[..., Path]): repo_path = gr_repo() + print("test_git.PATH:", repo_path) + try: # Run the git command inside the target directory result = subprocess.run( @@ -36,11 +38,13 @@ class TestGRPostgresDjangoCreation: dsn = django_db_factory("gr") assert isinstance(dsn, PostgresDsn) - # def test_django_tables(self, thl_web_rw: PostgresConfig): - # res = thl_web_rw.execute_sql_query(query=""" - # SELECT COUNT(*) - # FROM information_schema.tables - # WHERE table_schema = 'public'; - # """) - # assert len(res) == 1 - # assert res[0]["count"] == 56 + def test_django_tables(self, gr_db: PostgresConfig): + res = gr_db.execute_sql_query(query=""" + SELECT COUNT(*) + FROM information_schema.tables + WHERE table_schema = 'public'; + """) + print(res) + assert len(res) == 1 + assert res[0]["count"] == 56 + assert True diff --git a/tests/models/gr/test_business.py b/tests/models/gr/test_business.py index 7a84f23..716ec75 100644 --- a/tests/models/gr/test_business.py +++ b/tests/models/gr/test_business.py @@ -1,5 +1,5 @@ import os -from datetime import datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta, timezone from decimal import Decimal from typing import Optional from uuid import uuid4 @@ -82,15 +82,15 @@ class TestBusinessContact: class TestBusiness: @pytest.fixture - def start(self) -> "datetime": - return datetime(year=2018, month=3, day=14, hour=0, tzinfo=timezone.utc) + def start(self) -> datetime: + return datetime(year=2018, month=3, day=14, hour=0, tzinfo=UTC) @pytest.fixture def offset(self) -> str: return "30d" @pytest.fixture - def duration(self) -> Optional["timedelta"]: + def duration(self) -> timedelta | None: return None def test_init(self, business): @@ -413,15 +413,15 @@ class TestBusiness: class TestBusinessBalance: @pytest.fixture - def start(self) -> "datetime": - return datetime(year=2018, month=3, day=14, hour=0, tzinfo=timezone.utc) + def start(self) -> datetime: + return datetime(year=2018, month=3, day=14, hour=0, tzinfo=UTC) @pytest.fixture def offset(self) -> str: return "30d" @pytest.fixture - def duration(self) -> Optional["timedelta"]: + def duration(self) -> timedelta | None: return None @pytest.mark.skip @@ -1138,7 +1138,7 @@ class TestBusinessBalance: class TestBusinessMethods: @pytest.fixture(scope="function") - def start(self, utc_90days_ago) -> "datetime": + def start(self, utc_90days_ago) -> datetime: s = utc_90days_ago.replace(microsecond=0) return s @@ -1149,7 +1149,7 @@ class TestBusinessMethods: @pytest.fixture(scope="function") def duration( self, - ) -> Optional["timedelta"]: + ) -> timedelta | None: return None def test_cache_key(self, business, gr_redis): @@ -1212,7 +1212,7 @@ class TestBusinessMethods: # We're going to pull only a specific year, but make sure that # it's being assigned to the field regardless - year = datetime.now(tz=timezone.utc).year + year = datetime.now(tz=UTC).year res = Business.from_redis( uuid=business.uuid, fields=[f"pop_financial:{year}"], diff --git a/tests/models/morning/test.py b/tests/models/morning/test.py index bedf9c2..222cb93 100644 --- a/tests/models/morning/test.py +++ b/tests/models/morning/test.py @@ -1,4 +1,4 @@ -from datetime import datetime, timezone +from datetime import UTC, datetime, timezone from generalresearch.models.morning.question import MorningQuestion @@ -163,8 +163,8 @@ bid = { # what gets run in MorningAPI._format_bid bid["language_isos"] = ("eng",) bid["country_iso"] = "us" -bid["end_date"] = datetime(2024, 7, 19, 9, 1, 13, 520243, tzinfo=timezone.utc) -bid["published_at"] = datetime(2024, 6, 19, 9, 1, 13, 520243, tzinfo=timezone.utc) +bid["end_date"] = datetime(2024, 7, 19, 9, 1, 13, 520243, tzinfo=UTC) +bid["published_at"] = datetime(2024, 6, 19, 9, 1, 13, 520243, tzinfo=UTC) bid.update(bid["statistics"]) bid["qualified_conversion"] /= 100 bid["system_conversion"] /= 100 diff --git a/tests/models/prodege/test_survey_participation.py b/tests/models/prodege/test_survey_participation.py index 68d7838..3b35d0c 100644 --- a/tests/models/prodege/test_survey_participation.py +++ b/tests/models/prodege/test_survey_participation.py @@ -1,4 +1,4 @@ -from datetime import datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta, timezone class TestProdegeParticipation: @@ -10,7 +10,7 @@ class TestProdegeParticipation: ProdegeUserPastParticipation, ) - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) pp = ProdegePastParticipation.from_api( { "participation_project_ids": [152677146, 152803285], @@ -89,7 +89,7 @@ class TestProdegeParticipation: ProdegeUserPastParticipation, ) - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) pp = ProdegePastParticipation.from_api( { "participation_project_ids": [152677146, 152803285], diff --git a/tests/models/spectrum/test_question.py b/tests/models/spectrum/test_question.py index ba118d7..4f92961 100644 --- a/tests/models/spectrum/test_question.py +++ b/tests/models/spectrum/test_question.py @@ -1,17 +1,17 @@ -from datetime import datetime, timezone +from datetime import UTC, datetime, timezone from generalresearch.models import Source from generalresearch.models.spectrum.question import ( - SpectrumQuestionOption, SpectrumQuestion, - SpectrumQuestionType, SpectrumQuestionClass, + SpectrumQuestionOption, + SpectrumQuestionType, ) from generalresearch.models.thl.profiling.upk_question import ( UpkQuestion, + UpkQuestionChoice, UpkQuestionSelectorMC, UpkQuestionType, - UpkQuestionChoice, ) @@ -43,7 +43,7 @@ class TestSpectrumQuestion: tags=None, options=None, class_num=SpectrumQuestionClass.CORE, - created=datetime(2017, 8, 16, 7, 52, 7, 688000, tzinfo=timezone.utc), + created=datetime(2017, 8, 16, 7, 52, 7, 688000, tzinfo=UTC), is_live=True, source=Source.SPECTRUM, category_id=None, @@ -85,7 +85,7 @@ class TestSpectrumQuestion: SpectrumQuestionOption(id="112", text="Female", order=1), ], class_num=SpectrumQuestionClass.CORE, - created=datetime(2017, 8, 16, 7, 52, 7, 688000, tzinfo=timezone.utc), + created=datetime(2017, 8, 16, 7, 52, 7, 688000, tzinfo=UTC), is_live=True, source=Source.SPECTRUM, category_id=None, @@ -160,7 +160,7 @@ class TestSpectrumQuestion: SpectrumQuestionOption(id="999", text="None of the above", order=3), ], class_num=SpectrumQuestionClass.EXTENDED, - created=datetime(2017, 8, 16, 7, 52, 7, 688000, tzinfo=timezone.utc), + created=datetime(2017, 8, 16, 7, 52, 7, 688000, tzinfo=UTC), is_live=True, source=Source.SPECTRUM, category_id=None, diff --git a/tests/models/spectrum/test_survey.py b/tests/models/spectrum/test_survey.py index b612a63..5e095a3 100644 --- a/tests/models/spectrum/test_survey.py +++ b/tests/models/spectrum/test_survey.py @@ -1,4 +1,4 @@ -from datetime import datetime, timezone +from datetime import UTC, datetime, timezone from decimal import Decimal @@ -140,11 +140,11 @@ class TestSpectrumSurvey: "survey_id": 29333264, "survey_name": "Exciting New Survey #29333264", "survey_status": 22, - "field_end_date": datetime(2024, 5, 23, 18, 18, 31, tzinfo=timezone.utc), + "field_end_date": datetime(2024, 5, 23, 18, 18, 31, tzinfo=UTC), "category": "Exciting New", "category_code": 232, - "crtd_on": datetime(2024, 5, 20, 17, 48, 13, tzinfo=timezone.utc), - "mod_on": datetime(2024, 5, 20, 18, 18, 31, tzinfo=timezone.utc), + "crtd_on": datetime(2024, 5, 20, 17, 48, 13, tzinfo=UTC), + "mod_on": datetime(2024, 5, 20, 18, 18, 31, tzinfo=UTC), "soft_launch": False, "click_balancing": 0, "price_type": 1, @@ -212,7 +212,7 @@ class TestSpectrumSurvey: survey_id="29333264", survey_name="Exciting New Survey #29333264", status=SpectrumStatus.LIVE, - field_end_date=datetime(2024, 5, 23, 18, 18, 31, tzinfo=timezone.utc), + field_end_date=datetime(2024, 5, 23, 18, 18, 31, tzinfo=UTC), category_code="232", calculation_type=TaskCalculationType.COMPLETES, requires_pii=False, @@ -240,8 +240,8 @@ class TestSpectrumSurvey: values=["18-64"], ) }, - created_api=datetime(2024, 5, 20, 17, 48, 13, tzinfo=timezone.utc), - modified_api=datetime(2024, 5, 20, 18, 18, 31, tzinfo=timezone.utc), + created_api=datetime(2024, 5, 20, 17, 48, 13, tzinfo=UTC), + modified_api=datetime(2024, 5, 20, 18, 18, 31, tzinfo=UTC), updated=None, ) assert expected_survey.model_dump_json() == s.model_dump_json() @@ -255,11 +255,11 @@ class TestSpectrumSurvey: "survey_id": 29333264, "survey_name": "#29333264", "survey_status": 22, - "field_end_date": datetime(2024, 5, 23, 18, 18, 31, tzinfo=timezone.utc), + "field_end_date": datetime(2024, 5, 23, 18, 18, 31, tzinfo=UTC), "category": "Exciting New", "category_code": 232, - "crtd_on": datetime(2024, 5, 20, 17, 48, 13, tzinfo=timezone.utc), - "mod_on": datetime(2024, 5, 20, 18, 18, 31, tzinfo=timezone.utc), + "crtd_on": datetime(2024, 5, 20, 17, 48, 13, tzinfo=UTC), + "mod_on": datetime(2024, 5, 20, 18, 18, 31, tzinfo=UTC), "soft_launch": False, "click_balancing": 0, "price_type": 1, @@ -318,11 +318,11 @@ class TestSpectrumSurvey: "survey_id": 29333264, "survey_name": "#29333264", "survey_status": 22, - "field_end_date": datetime(2024, 5, 23, 18, 18, 31, tzinfo=timezone.utc), + "field_end_date": datetime(2024, 5, 23, 18, 18, 31, tzinfo=UTC), "category": "Exciting New", "category_code": 232, - "crtd_on": datetime(2024, 5, 20, 17, 48, 13, tzinfo=timezone.utc), - "mod_on": datetime(2024, 5, 20, 18, 18, 31, tzinfo=timezone.utc), + "crtd_on": datetime(2024, 5, 20, 17, 48, 13, tzinfo=UTC), + "mod_on": datetime(2024, 5, 20, 18, 18, 31, tzinfo=UTC), "soft_launch": False, "click_balancing": 0, "price_type": 1, diff --git a/tests/models/spectrum/test_survey_manager.py b/tests/models/spectrum/test_survey_manager.py index 582093c..11970bf 100644 --- a/tests/models/spectrum/test_survey_manager.py +++ b/tests/models/spectrum/test_survey_manager.py @@ -1,22 +1,21 @@ import copy import logging -from datetime import timezone, datetime +from datetime import UTC, datetime, timezone from decimal import Decimal from pymysql import IntegrityError - logger = logging.getLogger() example_survey_api_response = { "survey_id": 29333264, "survey_name": "#29333264", "survey_status": 22, - "field_end_date": datetime(2024, 5, 23, 18, 18, 31, tzinfo=timezone.utc), + "field_end_date": datetime(2024, 5, 23, 18, 18, 31, tzinfo=UTC), "category": "Exciting New", "category_code": 232, - "crtd_on": datetime(2024, 5, 20, 17, 48, 13, tzinfo=timezone.utc), - "mod_on": datetime(2024, 5, 20, 18, 18, 31, tzinfo=timezone.utc), + "crtd_on": datetime(2024, 5, 20, 17, 48, 13, tzinfo=UTC), + "mod_on": datetime(2024, 5, 20, 18, 18, 31, tzinfo=UTC), "soft_launch": False, "click_balancing": 0, "price_type": 1, @@ -66,7 +65,7 @@ class TestSpectrumSurvey: assert settings.debug, "CRITICAL: Do not run this on production." - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) spectrum_rw.execute_sql_query( query=f""" DELETE FROM `{spectrum_rw.db}`.spectrum_survey @@ -93,7 +92,7 @@ class TestSpectrumSurvey: assert settings.debug, "CRITICAL: Do not run this on production." - now = datetime.now(tz=timezone.utc) + now = datetime.now(tz=UTC) spectrum_rw.execute_sql_query( query=f""" DELETE FROM `{spectrum_rw.db}`.spectrum_survey diff --git a/tests/models/test_finance.py b/tests/models/test_finance.py index bd548b3..6dcd441 100644 --- a/tests/models/test_finance.py +++ b/tests/models/test_finance.py @@ -1,7 +1,7 @@ -from datetime import datetime, timedelta, timezone +from collections.abc import Callable +from datetime import UTC, datetime, timedelta, timezone from itertools import product as iter_product from random import randint -from typing import Callable from uuid import uuid4 import pandas as pd @@ -683,7 +683,7 @@ class TestProductFinanceData: rand_item_time = fake.date_time_between( start_date=item.start, end_date=item.finish, - tzinfo=timezone.utc, + tzinfo=UTC, ) session_with_tx_factory(started=rand_item_time, user=u) @@ -773,7 +773,7 @@ class TestPOPFinancialData: rand_item_time = fake.date_time_between( start_date=item.start, end_date=item.finish, - tzinfo=timezone.utc, + tzinfo=UTC, ) session_with_tx_factory(started=rand_item_time, user=u) @@ -870,7 +870,7 @@ class TestBusinessBalanceData: item_time = fake.date_time_between( start_date=item.start, end_date=item.finish, - tzinfo=timezone.utc, + tzinfo=UTC, ) session_with_tx_factory(started=item_time, user=u) item.initial_load(overwrite=True) diff --git a/tests/models/thl/test_adjustments.py b/tests/models/thl/test_adjustments.py index 27091bb..c2c035d 100644 --- a/tests/models/thl/test_adjustments.py +++ b/tests/models/thl/test_adjustments.py @@ -1,6 +1,6 @@ -from datetime import datetime, timedelta, timezone +from collections.abc import Callable +from datetime import UTC, datetime, timedelta, timezone from decimal import Decimal -from typing import Callable import pytest @@ -16,14 +16,14 @@ from generalresearch.models.thl.session import ( ) from generalresearch.models.thl.user import User -started1 = datetime(2023, 1, 1, tzinfo=timezone.utc) -started2 = datetime(2023, 1, 1, 0, 10, 0, tzinfo=timezone.utc) +started1 = datetime(2023, 1, 1, tzinfo=UTC) +started2 = datetime(2023, 1, 1, 0, 10, 0, tzinfo=UTC) finished1 = started1 + timedelta(minutes=10) finished2 = started2 + timedelta(minutes=10) -adj_ts = datetime(2023, 2, 2, tzinfo=timezone.utc) -adj_ts2 = datetime(2023, 2, 3, tzinfo=timezone.utc) -adj_ts3 = datetime(2023, 2, 4, tzinfo=timezone.utc) +adj_ts = datetime(2023, 2, 2, tzinfo=UTC) +adj_ts2 = datetime(2023, 2, 3, tzinfo=UTC) +adj_ts3 = datetime(2023, 2, 4, tzinfo=UTC) class TestProductAdjustments: diff --git a/tests/models/thl/test_contest/test_contest.py b/tests/models/thl/test_contest/test_contest.py index 0fbd4cc..acb501c 100644 --- a/tests/models/thl/test_contest/test_contest.py +++ b/tests/models/thl/test_contest/test_contest.py @@ -1,4 +1,4 @@ -from typing import Callable +from collections.abc import Callable import pytest diff --git a/tests/models/thl/test_contest/test_leaderboard_contest.py b/tests/models/thl/test_contest/test_leaderboard_contest.py index 8b714ee..3efcf2f 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 timezone +from datetime import UTC, timezone from uuid import uuid4 import pytest @@ -26,7 +26,7 @@ class TestLeaderboardContest(TestContest): @pytest.fixture def leaderboard_contest( self, product: Product, thl_redis, user_manager - ) -> "LeaderboardContest": + ) -> LeaderboardContest: board_key = f"leaderboard:{product.uuid}:us:weekly:2025-05-26:complete_count" c = LeaderboardContest( @@ -91,7 +91,7 @@ class TestLeaderboardContest(TestContest): country_iso=model.country_iso, freq=model.freq, product_id=leaderboard_contest.product_id, - within_time=model.period_start_local.astimezone(tz=timezone.utc), + within_time=model.period_start_local.astimezone(tz=UTC), ) lbm.hit_complete_count(product_user_id=user_1.product_user_id) diff --git a/tests/models/thl/test_ledger.py b/tests/models/thl/test_ledger.py index 257de3c..5edcc9d 100644 --- a/tests/models/thl/test_ledger.py +++ b/tests/models/thl/test_ledger.py @@ -1,4 +1,4 @@ -from datetime import datetime, timezone +from datetime import UTC, datetime, timezone from uuid import uuid4 import pytest @@ -21,7 +21,7 @@ class TestLedgerTransaction: assert [] == t.entries assert {} == t.metadata t = LedgerTransaction( - created=datetime.now(tz=timezone.utc), + created=datetime.now(tz=UTC), metadata={"a": "b", "user": "1234"}, ext_description="foo", ) diff --git a/tests/models/thl/test_product.py b/tests/models/thl/test_product.py index 39469dc..78bc10a 100644 --- a/tests/models/thl/test_product.py +++ b/tests/models/thl/test_product.py @@ -2,9 +2,9 @@ from __future__ import annotations import os import shutil -from datetime import datetime, timedelta, timezone +from collections.abc import Callable +from datetime import UTC, datetime, timedelta, timezone from decimal import Decimal -from typing import Callable from uuid import uuid4 import pytest @@ -586,7 +586,7 @@ class TestProductFinancials: @pytest.fixture def start(self) -> datetime: - return datetime(year=2018, month=3, day=14, hour=0, tzinfo=timezone.utc) + return datetime(year=2018, month=3, day=14, hour=0, tzinfo=UTC) @pytest.fixture def offset(self) -> str: @@ -769,7 +769,7 @@ class TestProductBalance: @pytest.fixture def start(self) -> datetime: - return datetime(year=2018, month=3, day=14, hour=0, tzinfo=timezone.utc) + return datetime(year=2018, month=3, day=14, hour=0, tzinfo=UTC) @pytest.fixture def offset(self) -> str: @@ -877,7 +877,7 @@ class TestProductBalance: product=product, amount=USDCent(71), ext_ref_id=uuid4().hex, - created=datetime.now(tz=timezone.utc), + created=datetime.now(tz=UTC), skip_wallet_balance_check=True, skip_one_per_day_check=True, ) @@ -892,7 +892,7 @@ class TestProductPOPFinancial: @pytest.fixture def start(self) -> datetime: - return datetime(year=2018, month=3, day=14, hour=0, tzinfo=timezone.utc) + return datetime(year=2018, month=3, day=14, hour=0, tzinfo=UTC) @pytest.fixture def offset(self) -> str: @@ -965,7 +965,7 @@ class TestProductCache: @pytest.fixture def start(self) -> datetime: - return datetime(year=2018, month=3, day=14, hour=0, tzinfo=timezone.utc) + return datetime(year=2018, month=3, day=14, hour=0, tzinfo=UTC) @pytest.fixture def offset(self) -> str: diff --git a/tests/models/thl/test_user.py b/tests/models/thl/test_user.py index 943ae8e..a4f331a 100644 --- a/tests/models/thl/test_user.py +++ b/tests/models/thl/test_user.py @@ -1,5 +1,5 @@ import json -from datetime import datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta, timezone from decimal import Decimal from random import choice as rand_choice from random import randint @@ -383,7 +383,7 @@ class TestUserCreated: from generalresearch.models.thl.user import User user = User(user_id=self.user_id) - dt = datetime.now(tz=timezone.utc) + dt = datetime.now(tz=UTC) user.created = dt assert user.created == dt @@ -419,7 +419,7 @@ class TestUserCreated: def test_not_in_future(self): from generalresearch.models.thl.user import User - the_future = datetime.now(tz=timezone.utc) + timedelta(minutes=1) + the_future = datetime.now(tz=UTC) + timedelta(minutes=1) with pytest.raises(ValueError) as cm: User(user_id=self.user_id, created=the_future) assert "1 validation error for User" in str(cm.value) @@ -428,9 +428,9 @@ class TestUserCreated: def test_after_anno_domini(self): from generalresearch.models.thl.user import User - before_ad = datetime( - year=2015, month=1, day=1, tzinfo=timezone.utc - ) + timedelta(minutes=1) + before_ad = datetime(year=2015, month=1, day=1, tzinfo=UTC) + timedelta( + minutes=1 + ) with pytest.raises(ValueError) as cm: User(user_id=self.user_id, created=before_ad) assert "1 validation error for User" in str(cm.value) @@ -444,7 +444,7 @@ class TestUserLastSeen: from generalresearch.models.thl.user import User user = User(user_id=self.user_id) - dt = datetime.now(tz=timezone.utc) + dt = datetime.now(tz=UTC) user.last_seen = dt assert user.last_seen == dt @@ -480,7 +480,7 @@ class TestUserLastSeen: def test_not_in_future(self): from generalresearch.models.thl.user import User - the_future = datetime.now(tz=timezone.utc) + timedelta(minutes=1) + the_future = datetime.now(tz=UTC) + timedelta(minutes=1) with pytest.raises(ValueError) as cm: User(user_id=self.user_id, last_seen=the_future) assert "1 validation error for User" in str(cm.value) @@ -489,9 +489,9 @@ class TestUserLastSeen: def test_after_anno_domini(self): from generalresearch.models.thl.user import User - before_ad = datetime( - year=2015, month=1, day=1, tzinfo=timezone.utc - ) + timedelta(minutes=1) + before_ad = datetime(year=2015, month=1, day=1, tzinfo=UTC) + timedelta( + minutes=1 + ) with pytest.raises(ValueError) as cm: User(user_id=self.user_id, last_seen=before_ad) assert "1 validation error for User" in str(cm.value) @@ -549,8 +549,8 @@ class TestUserTiming: def test_valid(self): from generalresearch.models.thl.user import User - created = datetime.now(tz=timezone.utc) - timedelta(minutes=60) - last_seen = datetime.now(tz=timezone.utc) - timedelta(minutes=59) + created = datetime.now(tz=UTC) - timedelta(minutes=60) + last_seen = datetime.now(tz=UTC) - timedelta(minutes=59) user = User(user_id=self.user_id, created=created, last_seen=last_seen) assert user.created == created @@ -559,8 +559,8 @@ class TestUserTiming: def test_created_first(self): from generalresearch.models.thl.user import User - created = datetime.now(tz=timezone.utc) - timedelta(minutes=60) - last_seen = datetime.now(tz=timezone.utc) - timedelta(minutes=59) + created = datetime.now(tz=UTC) - timedelta(minutes=60) + last_seen = datetime.now(tz=UTC) - timedelta(minutes=59) with pytest.raises(ValueError) as cm: User(user_id=self.user_id, created=last_seen, last_seen=created) @@ -602,7 +602,7 @@ class TestUserSerialization: user = User( product_id=product_id, product_user_id=product_user_id, - created=datetime.now(tz=timezone.utc), + created=datetime.now(tz=UTC), blocked=False, ) @@ -623,7 +623,7 @@ class TestUserSerialization: user = User( product_id=product_id, product_user_id=product_user_id, - created=datetime.now(tz=timezone.utc), + created=datetime.now(tz=UTC), blocked=False, ) @@ -633,7 +633,7 @@ class TestUserSerialization: assert not d.get("blocked") assert d.get("product") is None - assert d.get("created").tzinfo == timezone.utc + assert d.get("created").tzinfo == UTC def test_from_json(self): from generalresearch.models.thl.user import User @@ -644,14 +644,14 @@ class TestUserSerialization: user = User( product_id=product_id, product_user_id=product_user_id, - created=datetime.now(tz=timezone.utc), + created=datetime.now(tz=UTC), blocked=False, ) u = User.model_validate_json(user.to_json()) assert u.product_id == product_id assert u.product is None - assert u.created.tzinfo == timezone.utc + assert u.created.tzinfo == UTC class TestUserMethods: diff --git a/tests/models/thl/test_user_iphistory.py b/tests/models/thl/test_user_iphistory.py index 596849c..0f050b0 100644 --- a/tests/models/thl/test_user_iphistory.py +++ b/tests/models/thl/test_user_iphistory.py @@ -1,4 +1,4 @@ -from datetime import datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta, timezone from generalresearch.models.thl.user_iphistory import ( UserIPHistory, @@ -8,7 +8,7 @@ from generalresearch.models.thl.user_iphistory import ( def test_collapse_ip_records(): # This does not exist in a db, so we do not need fixtures/ real user ids, whatever - now = datetime.now(tz=timezone.utc) - timedelta(days=1) + now = datetime.now(tz=UTC) - timedelta(days=1) # Gets stored most recent first. This is reversed, but the validator will order it records = [ UserIPRecord(ip="1.2.3.5", created=now + timedelta(minutes=1)), diff --git a/tests/models/thl/test_wall.py b/tests/models/thl/test_wall.py index 8398c81..9e9483b 100644 --- a/tests/models/thl/test_wall.py +++ b/tests/models/thl/test_wall.py @@ -1,4 +1,4 @@ -from datetime import datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta, timezone from decimal import Decimal from uuid import uuid4 @@ -27,8 +27,8 @@ class TestWall: ext_status_code_1="1.0", status=Status.FAIL, status_code_1=StatusCode1.BUYER_FAIL, - started=datetime(2023, 1, 1, 0, 0, 1, tzinfo=timezone.utc), - finished=datetime(2023, 1, 1, 0, 10, 1, tzinfo=timezone.utc), + started=datetime(2023, 1, 1, 0, 0, 1, tzinfo=UTC), + finished=datetime(2023, 1, 1, 0, 10, 1, tzinfo=UTC), ) s = w.to_json() w2 = Wall.from_json(s) @@ -45,8 +45,8 @@ class TestWall: survey_id="yyy", status=Status.FAIL, status_code_1=StatusCode1.BUYER_FAIL, - started=datetime.now(timezone.utc), - finished=datetime.now(timezone.utc) + timedelta(seconds=1), + started=datetime.now(UTC), + finished=datetime.now(UTC) + timedelta(seconds=1), ) Wall( user_id=1, @@ -58,8 +58,8 @@ class TestWall: status=Status.FAIL, status_code_1=StatusCode1.MARKETPLACE_FAIL, status_code_2=WallStatusCode2.COMPLETE_TOO_FAST, - started=datetime.now(timezone.utc), - finished=datetime.now(timezone.utc) + timedelta(seconds=1), + started=datetime.now(UTC), + finished=datetime.now(UTC) + timedelta(seconds=1), ) with pytest.raises(expected_exception=ValidationError) as e: Wall( @@ -71,8 +71,8 @@ class TestWall: survey_id="yyy", status=Status.FAIL, status_code_1=StatusCode1.GRS_ABANDON, - started=datetime.now(timezone.utc), - finished=datetime.now(timezone.utc) + timedelta(seconds=1), + started=datetime.now(UTC), + finished=datetime.now(UTC) + timedelta(seconds=1), ) assert "If status is f, status_code_1 should be in" in str(e.value) @@ -87,8 +87,8 @@ class TestWall: status=Status.FAIL, status_code_1=StatusCode1.GRS_ABANDON, status_code_2=WallStatusCode2.COMPLETE_TOO_FAST, - started=datetime.now(timezone.utc), - finished=datetime.now(timezone.utc) + timedelta(seconds=1), + started=datetime.now(UTC), + finished=datetime.now(UTC) + timedelta(seconds=1), ) assert "If status is f, status_code_1 should be in" in str(e.value) @@ -104,8 +104,8 @@ class TestWall: status=Status.FAIL, status_code_1=StatusCode1.MARKETPLACE_FAIL, status_code_2=WallStatusCode2.COMPLETE_TOO_FAST, - started=datetime.now(timezone.utc), - finished=datetime.now(timezone.utc) + timedelta(seconds=1), + started=datetime.now(UTC), + finished=datetime.now(UTC) + timedelta(seconds=1), ) Wall( user_id=1, @@ -117,8 +117,8 @@ class TestWall: status=Status.FAIL, status_code_1=StatusCode1.BUYER_FAIL, status_code_2=None, - started=datetime.now(timezone.utc), - finished=datetime.now(timezone.utc) + timedelta(seconds=1), + started=datetime.now(UTC), + finished=datetime.now(UTC) + timedelta(seconds=1), ) Wall( user_id=1, @@ -130,8 +130,8 @@ class TestWall: status=Status.COMPLETE, status_code_1=StatusCode1.COMPLETE, status_code_2=None, - started=datetime.now(timezone.utc), - finished=datetime.now(timezone.utc) + timedelta(seconds=1), + started=datetime.now(UTC), + finished=datetime.now(UTC) + timedelta(seconds=1), ) with pytest.raises(expected_exception=ValidationError) as e: @@ -145,8 +145,8 @@ class TestWall: status=Status.FAIL, status_code_1=StatusCode1.BUYER_FAIL, status_code_2=WallStatusCode2.COMPLETE_TOO_FAST, - started=datetime.now(timezone.utc), - finished=datetime.now(timezone.utc) + timedelta(seconds=1), + started=datetime.now(UTC), + finished=datetime.now(UTC) + timedelta(seconds=1), ) assert "If status_code_1 is 1, status_code_2 should be in" in str(e.value) diff --git a/tests/models/thl/test_wall_session.py b/tests/models/thl/test_wall_session.py index 1208c56..10f3cba 100644 --- a/tests/models/thl/test_wall_session.py +++ b/tests/models/thl/test_wall_session.py @@ -1,4 +1,4 @@ -from datetime import datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta, timezone from decimal import Decimal import pytest @@ -12,7 +12,7 @@ from generalresearch.models.thl.user import User class TestWallSession: def test_session_with_no_wall_events(self): - started = datetime(2023, 1, 1, tzinfo=timezone.utc) + started = datetime(2023, 1, 1, tzinfo=UTC) s = Session(user=User(user_id=1), started=started) assert s.status is None assert s.status_code_1 is None @@ -24,7 +24,7 @@ class TestWallSession: # assert s.status_code_1 == StatusCode1.SESSION_START_FAIL def test_session_timeout_with_only_grs(self): - started = datetime(2023, 1, 1, tzinfo=timezone.utc) + started = datetime(2023, 1, 1, tzinfo=UTC) s = Session(user=User(user_id=1), started=started) w = Wall( user_id=1, @@ -53,7 +53,7 @@ class TestWallSession: # assert s.status_code_1 == StatusCode1.GRS_FAIL def test_session_with_only_grs_complete(self): - started = datetime(year=2023, month=1, day=1, tzinfo=timezone.utc) + started = datetime(year=2023, month=1, day=1, tzinfo=UTC) # A Session is started s = Session(user=User(user_id=1), started=started) @@ -98,7 +98,7 @@ class TestWallSession: # assert s.status_code_1 is None def test_session_with_only_non_grs_fail(self): - started = datetime(year=2023, month=1, day=1, tzinfo=timezone.utc) + started = datetime(year=2023, month=1, day=1, tzinfo=UTC) s = Session(user=User(user_id=1), started=started) w = Wall( @@ -119,7 +119,7 @@ class TestWallSession: assert s.payout is None def test_session_with_only_non_grs_timeout(self): - started = datetime(year=2023, month=1, day=1, tzinfo=timezone.utc) + started = datetime(year=2023, month=1, day=1, tzinfo=UTC) s = Session(user=User(user_id=1), started=started) w = Wall( @@ -139,7 +139,7 @@ class TestWallSession: assert s.payout is None def test_session_with_grs_and_external(self): - started = datetime(year=2023, month=1, day=1, tzinfo=timezone.utc) + started = datetime(year=2023, month=1, day=1, tzinfo=UTC) s = Session(user=User(user_id=1), started=started) w = Wall( @@ -168,7 +168,7 @@ class TestWallSession: s.append_wall_event(w) w.finish( status=Status.ABANDON, - finished=datetime.now(tz=timezone.utc) + timedelta(minutes=10), + finished=datetime.now(tz=UTC) + timedelta(minutes=10), status_code_1=StatusCode1.BUYER_ABANDON, ) status, status_code_1 = s.determine_session_status() @@ -206,7 +206,7 @@ class TestWallSession: assert s.payout is None def test_session_marketplace_fail(self): - started = datetime(2023, 1, 1, tzinfo=timezone.utc) + started = datetime(2023, 1, 1, tzinfo=UTC) s = Session(user=User(user_id=1), started=started) w = Wall( @@ -229,7 +229,7 @@ class TestWallSession: assert StatusCode1.SESSION_CONTINUE_QUALITY_FAIL == s.status_code_1 def test_session_unknown(self): - started = datetime(2023, 1, 1, tzinfo=timezone.utc) + started = datetime(2023, 1, 1, tzinfo=UTC) s = Session(user=User(user_id=1), started=started) w = Wall( diff --git a/tests/test_postgres.py b/tests/test_postgres.py index 3b3ddd0..ed5a7ae 100644 --- a/tests/test_postgres.py +++ b/tests/test_postgres.py @@ -1,6 +1,6 @@ import socket import subprocess -from typing import Callable +from collections.abc import Callable from pydantic import PostgresDsn @@ -12,7 +12,7 @@ def is_port_open(host: InternalHostname, port: int = 5432, timeout: int = 3): try: with socket.create_connection((host, port), timeout=timeout): return True - except (socket.timeout, ConnectionRefusedError, OSError): + except (TimeoutError, ConnectionRefusedError, OSError): return False -- 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 'test_utils/models/ledger/conftest.py') 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 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 'test_utils/models/ledger/conftest.py') 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 'test_utils/models/ledger/conftest.py') 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 1151b332279425e4e088bd3499c76e582f7f045d Mon Sep 17 00:00:00 2001 From: Max Nanis Date: Thu, 3 Sep 2026 09:26:17 -0700 Subject: Test cleanup all morning. mangers/thl = 78fail, 417passed --- generalresearch/thl_django/app/test_settings.py | 2 +- test_utils/models/conftest.py | 151 +--------- test_utils/models/gr/conftest.py | 7 - test_utils/models/ledger/conftest.py | 54 ++-- test_utils/models/thl/conftest.py | 382 +++++++++++++++--------- tests/grliq/managers/test_forensic_data.py | 2 +- tests/managers/test_events.py | 20 -- tests/managers/thl/test_ledger/test_thl_pem.py | 14 +- tests/managers/thl/test_payout.py | 38 ++- tests/managers/thl/test_task_adjustment.py | 12 +- tests/managers/thl/test_user_streak.py | 24 +- tests/managers/thl/test_userhealth.py | 5 +- tests/managers/thl/test_wall_manager.py | 14 +- tests/models/gr/test_business.py | 54 ++-- tests/models/thl/test_product.py | 26 +- 15 files changed, 404 insertions(+), 401 deletions(-) (limited to 'test_utils/models/ledger/conftest.py') diff --git a/generalresearch/thl_django/app/test_settings.py b/generalresearch/thl_django/app/test_settings.py index d6ab124..57cb9b9 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-44c0b4', + "NAME": 'unittest-2026-09-03-a0a584', "USER": 'jenkins', "PASSWORD": '123456789', "HOST": 'unittest-postgresql.fmt2.grl.internal', diff --git a/test_utils/models/conftest.py b/test_utils/models/conftest.py index d5c9a71..9edadd3 100644 --- a/test_utils/models/conftest.py +++ b/test_utils/models/conftest.py @@ -51,119 +51,17 @@ if TYPE_CHECKING: ) from generalresearch.models.thl.session import Session, Wall from generalresearch.models.thl.user import User - from generalresearch.pg_helper import PostgresConfig # === THL === -@pytest.fixture -def user( - request: FixtureRequest, - user_manager: UserManager, - thl_web_rr: PostgresConfig, - product_factory: Callable[..., Product], -) -> User: - product = getattr(request, "product", None) - - if product is None: - product = product_factory() - - u = user_manager.create_dummy(product_id=product.id) - u.prefetch_product(pg_config=thl_web_rr) - - return u - - -@pytest.fixture -def user_with_wallet( - user_factory: Callable[..., User], - product_user_wallet_yes: Product, -) -> User: - # A user on a product with user wallet enabled, but they have no money - return user_factory(product=product_user_wallet_yes) - - -@pytest.fixture -def user_with_wallet_amt( - user_factory: Callable[..., User], product_amt_true: Product -) -> User: - # A user on a product with user wallet enabled, on AMT, but they have no money - return user_factory(product=product_amt_true) - - -@pytest.fixture(scope="function") -def user_factory( - user_manager: UserManager, thl_web_rr: PostgresConfig -) -> Callable[..., User]: - - def _inner(product: Product, created: datetime | None = None) -> User: - u = user_manager.create_dummy(product=product, created=created) - u.prefetch_product(pg_config=thl_web_rr) - - return u - - return _inner - - -@pytest.fixture -def wall_factory(wall_manager: WallManager) -> Callable[..., Wall]: - - def _inner( - session: Session, wall_status: Status, req_cpi: Decimal | None = None - ) -> Wall: - - assert session.started <= datetime.now( - tz=UTC - ), "Session can't start in the future" - - if session.wall_events: - # Subsequent Wall events - wall = session.wall_events[-1] - assert not wall.finished, "Can't add new Walls until prior finishes" - # wall_started = last_wall.started + timedelta(milliseconds=1) - else: - # First Wall Event in a session - wall_started = session.started + timedelta(milliseconds=1) - - wall = wall_manager.create_dummy( - session_id=session.id, - user_id=session.user_id, - started=wall_started, - req_cpi=req_cpi, - ) - session.append_wall_event(w=wall) - - options = list(WALL_ALLOWED_STATUS_STATUS_CODE.get(wall_status, {})) - wall.finish( - finished=wall.started + timedelta(seconds=randint(a=60 * 2, b=60 * 10)), - status=wall_status, - status_code_1=randchoice(options), - ) - - return wall - - return _inner - - -@pytest.fixture -def wall(session: Session, user: User, wall_manager: WallManager) -> Wall | None: - from generalresearch.models.thl.task_status import StatusCode1 - - wall = wall_manager.create_dummy(session_id=session.id, user_id=user.user_id) - # thl_session.append_wall_event(wall) - wall.finish( - finished=wall.started + timedelta(seconds=randint(a=60 * 2, b=60 * 10)), - status=Status.COMPLETE, - status_code_1=StatusCode1.COMPLETE, - ) - return wall - - @pytest.fixture def session_factory( session_manager: SessionManager, wall_manager: WallManager, utc_hour_ago: datetime, + session_factory: Callable[..., Session], + wall_factory: Callable[..., Wall], ) -> Callable[..., Session]: from generalresearch.models.thl.session import Source @@ -184,7 +82,7 @@ def session_factory( if wall_statuses: assert len(wall_statuses) == wall_count - s = session_manager.create_dummy(started=started, user=user, country_iso="us") + s = session_factory(started=started, user=user, country_iso="us") for idx in range(wall_count): if idx == 0: # First Wall Event in a session @@ -195,7 +93,7 @@ def session_factory( assert last_wall.finished, "Can't add new Walls until prior finishes" wall_started = last_wall.started + timedelta(milliseconds=1) - w = wall_manager.create_dummy( + w = wall_factory( session_id=s.id, source=wall_source, user_id=s.user_id, @@ -271,11 +169,15 @@ def finished_session_factory( @pytest.fixture def session( - user: User, session_manager: SessionManager, wall_manager: WallManager + user: User, + session_manager: SessionManager, + wall_manager: WallManager, + session_factory: Callable[..., Session], + wall_factory: Callable[..., Wall], ) -> Session: - session: Session = session_manager.create_dummy(user=user, country_iso="us") - wall: Wall = wall_manager.create_dummy( + session: Session = session_factory(user=user, country_iso="us") + wall: Wall = wall_factory( session_id=session.id, user_id=session.user_id, started=session.started, @@ -342,37 +244,6 @@ def product_amt_true( ) -@pytest.fixture -def bp_payout_factory( - thl_ledger_manager: ThlLedgerManager, - product_manager: ProductManager, - business_payout_event_manager: BusinessPayoutEventManager, -) -> Callable[..., BrokerageProductPayoutEvent]: - - def _inner( - product: Product | None = None, - amount: USDCent | None = None, - ext_ref_id: str | None = None, - created: AwareDatetime | None = None, - skip_wallet_balance_check: bool = False, - skip_one_per_day_check: bool = False, - ) -> BrokerageProductPayoutEvent: - from generalresearch.currency import USDCent - - product = product or product_manager.create_dummy() - amount = amount or USDCent(randint(1, 99_99)) - - return business_payout_event_manager.create_bp_payout_event( - thl_ledger_manager=thl_ledger_manager, - product=product, - amount=amount, - ext_ref_id=ext_ref_id or uuid4().hex, - created=created, - ) - - return _inner - - @pytest.fixture(scope="session") def buyer(buyer_manager: BuyerManager) -> Buyer: buyer_code = uuid4().hex diff --git a/test_utils/models/gr/conftest.py b/test_utils/models/gr/conftest.py index 3dd73a1..1dbea0c 100644 --- a/test_utils/models/gr/conftest.py +++ b/test_utils/models/gr/conftest.py @@ -132,13 +132,6 @@ def gr_business_address_factory( return _inner -# @pytest.fixture -# def business_address( -# gr_business: Business, business_address_manager: BusinessAddressManager -# ) -> : -# return business_address_manager.create_dummy(business_id=gr_business.id) - - @pytest.fixture def gr_business_address( gr_business_address_factory: Callable[..., BusinessAddress], diff --git a/test_utils/models/ledger/conftest.py b/test_utils/models/ledger/conftest.py index 31e5eb4..9ee0df2 100644 --- a/test_utils/models/ledger/conftest.py +++ b/test_utils/models/ledger/conftest.py @@ -11,36 +11,29 @@ import pytest from pytest import FixtureRequest as Request from generalresearch.currency import USDCent -from test_utils.models.conftest import ( - payout_config, - product_amt_true, - product_user_wallet_no, - product_user_wallet_yes, - session, - session_factory, - user_factory, - wall, - wall_factory, -) -if TYPE_CHECKING: - from generalresearch.managers.base import PostgresManager - -_ = ( - user_factory, - product_user_wallet_no, - wall, - product_amt_true, - product_user_wallet_yes, - session_factory, - session, - wall_factory, - payout_config, -) +# from test_utils.models.conftest import ( +# payout_config, +# product_amt_true, +# product_user_wallet_no, +# product_user_wallet_yes, +# ) + +# _ = ( +# user_factory, +# product_user_wallet_no, +# wall, +# product_amt_true, +# product_user_wallet_yes, +# session_factory, +# session, +# wall_factory, +# payout_config, +# ) if TYPE_CHECKING: - from generalresearch.currency import LedgerCurrency + from generalresearch.managers.base import PostgresManager from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager from generalresearch.managers.thl.ledger_manager.thl_ledger import ( ThlLedgerManager, @@ -193,16 +186,17 @@ def usd_cent(request: Request) -> USDCent: def bp_payout_event( product: Product, usd_cent: USDCent, - business_payout_event_manager: BusinessPayoutEventManager, + brokerage_product_payout_event_manager: BrokerageProductPayoutEvent, thl_ledger_manager: ThlLedgerManager, ) -> BrokerageProductPayoutEvent: - return business_payout_event_manager.create_bp_payout_event( + _ext_ref_id = f"tx-{uuid4().hex[:7]}" + + return brokerage_product_payout_event_manager.create_bp_payout_event( thl_ledger_manager=thl_ledger_manager, + ext_ref_id=_ext_ref_id, product=product, amount=usd_cent, - skip_wallet_balance_check=True, - skip_one_per_day_check=True, ) diff --git a/test_utils/models/thl/conftest.py b/test_utils/models/thl/conftest.py index 14f8f36..5dc46cd 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 +from datetime import UTC, datetime, timedelta from decimal import ROUND_DOWN, Decimal from random import choice as rand_choice from random import randint, random @@ -13,37 +13,49 @@ import pytest from grip_client.enums import AccessType from pydantic import PositiveInt +from generalresearch.managers.thl.payout import UserPayoutEventManager from generalresearch.models.custom_types import ( AwareDatetimeISO, IPvAnyAddressStr, UUIDStr, ) -from generalresearch.models.thl.definitions import PayoutStatus +from generalresearch.models.thl.definitions import ( + WALL_ALLOWED_STATUS_STATUS_CODE, + PayoutStatus, +) +from generalresearch.models.thl.payout import UserPayoutEvent from generalresearch.models.thl.session import ( Source, Status, ) from generalresearch.models.thl.user import User +from generalresearch.models.thl.user_iphistory import IPRecord from generalresearch.models.thl.userhealth import AuditLogLevel from generalresearch.models.thl.wallet.definitions import PayoutType if TYPE_CHECKING: + from generalresearch.currency import USDCent from generalresearch.managers.thl.ipinfo import ( IPGeonameManager, IPInformationManager, ) - from generalresearch.managers.thl.payout import UserPayoutEventManager + from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager + from generalresearch.managers.thl.payout import ( + BrokerageProductPayoutEventManager, + BusinessPayoutEventManager, + ) 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 AwareDatetime from generalresearch.models.definitions import DeviceType from generalresearch.models.gr.business import Business from generalresearch.models.gr.team import Team from generalresearch.models.legacy.bucket import Bucket from generalresearch.models.thl.ipinfo import IPGeoname, IPInformation - from generalresearch.models.thl.payout import UserPayoutEvent + from generalresearch.models.thl.payout import BrokerageProductPayoutEvent from generalresearch.models.thl.product import ( PayoutConfig, Product, @@ -66,19 +78,31 @@ if TYPE_CHECKING: fake = faker.Faker() +# --- Wall --- -@pytest.fixture -def wall_status() -> Status: - return Status.COMPLETE + +# from generalresearch.models.thl.task_status import StatusCode1 +# # thl_session.append_wall_event(wall) +# wall.finish( +# finished=wall.started + timedelta(seconds=randint(a=60 * 2, b=60 * 10)), +# status=Status.COMPLETE, +# status_code_1=StatusCode1.COMPLETE, +# ) +# return wall @pytest.fixture def wall_factory( - wall_manager: WallManager, session_factory: Session + wall_manager: WallManager, + session_factory: Callable[..., Session], + session_manager: SessionManager, ) -> Callable[..., Wall]: def _inner( - session_id: int | None = None, + wall_status: Status, + save: bool = True, + session: Session | None = None, + session_id: PositiveInt | None = None, user_id: int | None = None, started: datetime | None = None, source: Source | None = None, @@ -86,43 +110,157 @@ def wall_factory( req_cpi: Decimal | None = None, buyer_id: str | None = None, uuid_id: str | None = None, - ): + ) -> Wall: """To be used in tests, where we don't care about certain fields""" - user_id = user_id or fake.random_int(min=1, max=2_147_483_648) - started = started or fake.date_time_between( - start_date=datetime(year=1900, month=1, day=1, tzinfo=UTC), - end_date=datetime.now(tz=UTC), - tzinfo=UTC, - ) + if save: - if session_id is None: - # session = SessionManager(pg_config=self.pg_config).create_dummy( - # started=started - # ) - session = session_factory() - session_id = session.id + user_id = user_id or fake.random_int(min=1, max=2_147_483_648) + _wall_started = started or fake.date_time_between( + start_date=datetime(year=1900, month=1, day=1, tzinfo=UTC), + end_date=datetime.now(tz=UTC), + tzinfo=UTC, + ) - source = source or rand_choice(list(Source)) - req_survey_id = req_survey_id or uuid4().hex - req_cpi = req_cpi or Decimal(fake.random_int(min=1, max=150) / 100).quantize( - Decimal(".01"), rounding=ROUND_DOWN - ) + if session: + # If an existing Session was provided, we want to do some + # additional validation. + + if session.wall_events: + # Subsequent Wall events + _last_wall = session.wall_events[-1] + assert ( + not _last_wall.finished + ), "Can't add new Walls until prior finishes" + _wall_started = _last_wall.started + timedelta(milliseconds=1) + else: + # First Wall Event in a session + _wall_started = session.started + timedelta(milliseconds=1) + else: + # If a Session was NOT provided, either (1) try to retrieve it + # from an optionally provided session_id int, or (2) proceed + # forward and make one + session = ( + session_manager.get_from_id(session_id=session_id) + if session_id + else None + ) or session_factory(save=True, user_id=user_id) + + assert session, "Wall factory requires Session" + + source = source or rand_choice(list(Source)) + req_survey_id = req_survey_id or uuid4().hex + req_cpi = req_cpi or Decimal( + fake.random_int(min=1, max=150) / 100 + ).quantize(Decimal(".01"), rounding=ROUND_DOWN) + + w = wall_manager.create( + session_id=session.id, + user_id=session.user_id, + started=_wall_started, + source=source, + req_survey_id=req_survey_id, + req_cpi=req_cpi, + buyer_id=buyer_id, + uuid_id=uuid_id, + ) - return wall_manager.create( - session_id=session_id, - user_id=user_id, - started=started, - source=source, - req_survey_id=req_survey_id, - req_cpi=req_cpi, - buyer_id=buyer_id, - uuid_id=uuid_id, - ) + _status_code_options = list( + WALL_ALLOWED_STATUS_STATUS_CODE.get(wall_status, {}) + ) + w.finish( + finished=w.started + timedelta(seconds=randint(a=60 * 2, b=60 * 10)), + status=wall_status, + status_code_1=rand_choice(_status_code_options), + ) + + session.append_wall_event(w=w) + + return w + + else: + raise ValueError("Unsaved Wall not yet supported") return _inner +@pytest.fixture +def wall(wall_factory: Callable[..., Wall]) -> Wall: + return wall_factory(save=True) + + +@pytest.fixture() +def unsaved_wall(wall_factory: Callable[..., Wall]) -> Wall: + return wall_factory(save=False) + + +# --- Wall: Enum(s) --- + + +@pytest.fixture +def wall_status() -> Status: + return Status.COMPLETE + + +# --- Session --- + + +@pytest.fixture +def session_factory(session_manager: SessionManager, user_factory: Callable[..., User]): + + def _inner( + save: bool = True, + # -- Create Dummy "optional" -- # + started: datetime | None = None, + user: User | None = None, + # -- Optional -- # + country_iso: str | None = None, + device_type: DeviceType | None = None, + ip: str | None = None, + bucket: Bucket | None = None, + url_metadata: dict[str, str] | None = None, + uuid_id: str | None = None, + ) -> Session: + + if save: + """To be used in tests, where we don't care about certain fields""" + started = started or fake.date_time_between( + start_date=datetime(year=1900, month=1, day=1, tzinfo=UTC), + end_date=datetime(year=2000, month=1, day=1, tzinfo=UTC), + tzinfo=UTC, + ) + user = user or user_factory(save=True) + assert user.user_id, "Provided User must be saved to the database" + + return session_manager.create( + started=started, + user=user, + country_iso=country_iso, + device_type=device_type, + ip=ip, + bucket=bucket, + url_metadata=url_metadata, + uuid_id=uuid_id, + ) + else: + # user = User( + # user_id=fake.random_int(min=1, max=2_147_483_648), uuid=uuid4().hex + # ) + raise ValueError("Unsaved Session not yet supported") + + return _inner + + +@pytest.fixture() +def session(session_factory: Callable[..., Session]) -> Session: + return session_factory(save=True) + + +@pytest.fixture() +def unsaved_session(session_factory: Callable[..., Session]) -> Session: + return session_factory(save=False) + + # --- Product --- @@ -193,46 +331,7 @@ def unsaved_product(product_factory: Callable[..., Product]) -> Product: return product_factory(save=False) -# --- Session --- - - -@pytest.fixture -def session_factory(session_manager: SessionManager): - - def _inner( - # -- Create Dummy "optional" -- # - started: datetime | None = None, - user: User | None = None, - # -- Optional -- # - country_iso: str | None = None, - device_type: DeviceType | None = None, - ip: str | None = None, - bucket: Bucket | None = None, - url_metadata: dict[str, str] | None = None, - uuid_id: str | None = None, - ) -> Session: - """To be used in tests, where we don't care about certain fields""" - started = started or fake.date_time_between( - start_date=datetime(year=1900, month=1, day=1, tzinfo=UTC), - end_date=datetime(year=2000, month=1, day=1, tzinfo=UTC), - tzinfo=UTC, - ) - user = user or User( - user_id=fake.random_int(min=1, max=2_147_483_648), uuid=uuid4().hex - ) - - return session_manager.create( - started=started, - user=user, - country_iso=country_iso, - device_type=device_type, - ip=ip, - bucket=bucket, - url_metadata=url_metadata, - uuid_id=uuid_id, - ) - - return _inner +# --- IP Geoname --- @pytest.fixture @@ -363,7 +462,7 @@ def ip_information( return ip_information_factory(save=True) -@pytest.fixture +@pytest.fixture() def unsaved_ip_information( ip_information_factory: Callable[..., IPInformation], ) -> IPInformation: @@ -373,41 +472,36 @@ def unsaved_ip_information( # --- IP Record --- -@pytest.fixture -def ip_record_factory( - ip_record_manager: IPRecordManager, user: User -) -> Callable[..., IPRecord]: - # return ip_record_manager.create_dummy(user_id=user.user_id) - - # def create_dummy( - # self, - # user_id: PositiveInt, - # ip: IPvAnyAddressStr | None = None, - # forwarded_ip1: IPvAnyAddressStr | None = None, - # forwarded_ip2: IPvAnyAddressStr | None = None, - # forwarded_ip3: IPvAnyAddressStr | None = None, - # forwarded_ip4: IPvAnyAddressStr | None = None, - # forwarded_ip5: IPvAnyAddressStr | None = None, - # forwarded_ip6: IPvAnyAddressStr | None = None, - # ) -> IPRecord: - # return self.create( - # user_id=user_id, - # ip=ip or fake.ipv4_public(), - # forwarded_ip1=(forwarded_ip1 or fake.ipv4_public()), - # forwarded_ip2=(forwarded_ip2 or fake.ipv6() if random() < 0.5 else None), - # forwarded_ip3=( - # forwarded_ip3 or fake.ipv4_public() if random() < 0.25 else None - # ), - # forwarded_ip4=forwarded_ip4, - # forwarded_ip5=forwarded_ip5, - # forwarded_ip6=forwarded_ip6, - # ) +@pytest.fixture() +def ip_record_factory(ip_record_manager: IPRecordManager) -> Callable[..., IPRecord]: def _inner( - user_id: PositiveInt, save: bool = True, ip: str | None = None + user_id: PositiveInt, + save: bool = True, + ip: IPvAnyAddressStr | None = None, + forwarded_ip1: IPvAnyAddressStr | None = None, + forwarded_ip2: IPvAnyAddressStr | None = None, + forwarded_ip3: IPvAnyAddressStr | None = None, + forwarded_ip4: IPvAnyAddressStr | None = None, + forwarded_ip5: IPvAnyAddressStr | None = None, + forwarded_ip6: IPvAnyAddressStr | None = None, ) -> IPRecord: + if save: - return ip_record_manager.create_dummy(user_id=user_id, ip=ip) + return ip_record_manager.create( + user_id=user_id, + ip=ip or fake.ipv4_public(), + forwarded_ip1=(forwarded_ip1 or fake.ipv4_public()), + forwarded_ip2=( + forwarded_ip2 or fake.ipv6() if random() < 0.5 else None + ), + forwarded_ip3=( + forwarded_ip3 or fake.ipv4_public() if random() < 0.25 else None + ), + forwarded_ip4=forwarded_ip4, + forwarded_ip5=forwarded_ip5, + forwarded_ip6=forwarded_ip6, + ) else: raise ValueError("Unsaved IP Record not supported") @@ -415,9 +509,7 @@ def ip_record_factory( @pytest.fixture() -def ip_record( - ip_record_manager: IPRecordManager, ip_geoname: IPGeoname, user: User -) -> IPRecord: +def ip_record(ip_record_factory: Callable[..., IPRecord]) -> IPRecord: return ip_record_factory(save=True) @@ -431,7 +523,8 @@ def unsaved_ip_record(ip_record_factory: Callable[..., IPRecord]) -> IPRecord: @pytest.fixture() def user_factory( - user_manager: UserManager, thl_web_rr: PostgresConfig + user_manager: UserManager, + thl_web_rr: PostgresConfig, ) -> Callable[..., User]: def _inner( @@ -456,8 +549,6 @@ def user_factory( created=created, ) - u = user_manager.create_dummy(product=product, created=created) - u.prefetch_product(pg_config=thl_web_rr) return u @@ -498,7 +589,7 @@ def user_with_wallet_amt( return user_factory(save=True, product=product_amt_true) -# --- User Payout --- +# --- User Payout Event --- @pytest.fixture @@ -555,30 +646,47 @@ def user_payout_event_factory( return _inner +@pytest.fixture() +def user_payout_event( + user_payout_event_factory: Callable[..., UserPayoutEvent], +) -> UserPayoutEvent: + return user_payout_event_factory(save=True) + + +@pytest.fixture() +def unsaved_user_payout_event( + user_payout_event_factory: Callable[..., UserPayoutEvent], +) -> UserPayoutEvent: + return user_payout_event_factory(save=True) + + +# -- Brokerage Product Payout Event + + @pytest.fixture -def iprecord_factory(iprecord_manager: IPRecordManager) -> Callable[..., IPRecord]: +def brokerage_product_payout_event_factory( + thl_ledger_manager: ThlLedgerManager, + brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager, + product_factory: Callable[..., Product], +) -> Callable[..., BrokerageProductPayoutEvent]: def _inner( - user_id: PositiveInt, - ip: IPvAnyAddressStr | None = None, - forwarded_ip1: IPvAnyAddressStr | None = None, - forwarded_ip2: IPvAnyAddressStr | None = None, - forwarded_ip3: IPvAnyAddressStr | None = None, - forwarded_ip4: IPvAnyAddressStr | None = None, - forwarded_ip5: IPvAnyAddressStr | None = None, - forwarded_ip6: IPvAnyAddressStr | None = None, - ) -> IPRecord: - return iprecord_manager.create( - user_id=user_id, - ip=ip or fake.ipv4_public(), - forwarded_ip1=(forwarded_ip1 or fake.ipv4_public()), - forwarded_ip2=(forwarded_ip2 or fake.ipv6() if random() < 0.5 else None), - forwarded_ip3=( - forwarded_ip3 or fake.ipv4_public() if random() < 0.25 else None - ), - forwarded_ip4=forwarded_ip4, - forwarded_ip5=forwarded_ip5, - forwarded_ip6=forwarded_ip6, + product: Product | None = None, + amount: USDCent | None = None, + ext_ref_id: str | None = None, + created: AwareDatetime | None = None, + ) -> BrokerageProductPayoutEvent: + from generalresearch.currency import USDCent + + product = product or product_factory() + amount = amount or USDCent(randint(1, 99_99)) + + return brokerage_product_payout_event_manager.create_bp_payout_event( + thl_ledger_manager=thl_ledger_manager, + product=product, + amount=amount, + ext_ref_id=ext_ref_id or uuid4().hex, + created=created, ) return _inner diff --git a/tests/grliq/managers/test_forensic_data.py b/tests/grliq/managers/test_forensic_data.py index 1b83757..2254829 100644 --- a/tests/grliq/managers/test_forensic_data.py +++ b/tests/grliq/managers/test_forensic_data.py @@ -31,7 +31,7 @@ except ImportError: class TestGrlIqDataManager: - def test_create_dummy( + def test_factory( self, grliq_data_factory: Callable[..., GrlIqData], ): diff --git a/tests/managers/test_events.py b/tests/managers/test_events.py index 8745126..e256876 100644 --- a/tests/managers/test_events.py +++ b/tests/managers/test_events.py @@ -1,13 +1,10 @@ from __future__ import annotations import math -import random import time from collections.abc import Callable from datetime import UTC, datetime, timedelta from decimal import Decimal -from functools import partial -from math import floor from typing import TYPE_CHECKING from uuid import uuid4 @@ -38,28 +35,11 @@ def product_id(product_manager: ProductManager) -> str: return uuid4().hex -@pytest.fixture(scope="function") -def user_factory(product_id: str): - return partial(create_dummy, product_id=product_id) - - @pytest.fixture(scope="function") def event_subscriber(thl_redis_config: RedisConfig, product_id: str) -> EventSubscriber: return EventSubscriber(redis_config=thl_redis_config, product_id=product_id) -def create_dummy( - product_id: str | None = None, product_user_id: str | None = None -) -> User: - return User( - product_id=product_id, - product_user_id=product_user_id or uuid4().hex, - uuid=uuid4().hex, - created=datetime.now(tz=UTC), - user_id=random.randint(0, floor(2**32 / 2)), - ) - - class TestActiveUsers: def test_run_empty(self, event_manager: EventManager, product_id: str): diff --git a/tests/managers/thl/test_ledger/test_thl_pem.py b/tests/managers/thl/test_ledger/test_thl_pem.py index 9dbec48..18102c3 100644 --- a/tests/managers/thl/test_ledger/test_thl_pem.py +++ b/tests/managers/thl/test_ledger/test_thl_pem.py @@ -25,6 +25,7 @@ if TYPE_CHECKING: BrokerageProductPayoutEventManager, UserPayoutEventManager, ) + from generalresearch.models.thl.payout import UserPayoutEvent from generalresearch.models.thl.product import Product @@ -111,7 +112,7 @@ class TestThlPayoutEventManager: # 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_ledger_manager, product_uuids=[product.id] + product_uuids=[product.id] ) assert len(res) == N_PAYOUT_EVENTS @@ -120,7 +121,6 @@ class TestThlPayoutEventManager: # ahead and query for them res = ( brokerage_product_payout_event_manager.get_bp_bp_payout_events_for_products( - thl_ledger_manager=thl_ledger_manager, product_uuids=[i.uuid for i in products], ) ) @@ -160,11 +160,15 @@ class TestThlPayoutEventManager: # def test_filter_by(self): # raise NotImplementedError - def test_create(self, user_payout_event_manager: UserPayoutEventManager): + def test_create( + self, + user_payout_event_factory: Callable[..., UserPayoutEvent], + user_payout_event_manager: UserPayoutEventManager, + ): from generalresearch.models.thl.payout import UserPayoutEvent # Confirm the creation method returns back an instance. - pe = user_payout_event_manager.create_dummy() + pe = user_payout_event_factory() assert isinstance(pe, UserPayoutEvent) # Now query the DB for that PayoutEvent to confirm it was actually @@ -260,7 +264,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_ledger_manager, product_uuids=[product.uuid] + product_uuids=[product.uuid] ) ) assert isinstance(bp_bp_res, list) diff --git a/tests/managers/thl/test_payout.py b/tests/managers/thl/test_payout.py index ad101a4..52bbbec 100644 --- a/tests/managers/thl/test_payout.py +++ b/tests/managers/thl/test_payout.py @@ -282,7 +282,9 @@ class TestBusinessPayoutEventManager: create_main_accounts: Callable[..., None], thl_ledger_manager: ThlLedgerManager, product_factory: Callable[..., Product], - bp_payout_factory: Callable[..., BrokerageProductPayoutEvent], + brokerage_product_payout_event_factory: Callable[ + ..., BrokerageProductPayoutEvent + ], gr_business: Business, ): delete_ledger_db() @@ -295,16 +297,24 @@ class TestBusinessPayoutEventManager: ach_id2 = uuid4().hex # ext_ref_id is required now - bp_payout_factory(product=p1, amount=USDCent(1), ext_ref_id="none") + brokerage_product_payout_event_factory( + product=p1, amount=USDCent(1), ext_ref_id="none" + ) - bp_payout_factory(product=p1, amount=USDCent(1), ext_ref_id=ach_id1) + brokerage_product_payout_event_factory( + product=p1, amount=USDCent(1), ext_ref_id=ach_id1 + ) with pytest.raises( expected_exception=ValueError, match="Cannot create a BusinessPayoutEvent with an existing transaction_id", ): - bp_payout_factory(product=p1, amount=USDCent(25), ext_ref_id=ach_id1) + brokerage_product_payout_event_factory( + product=p1, amount=USDCent(25), ext_ref_id=ach_id1 + ) - bp_payout_factory(product=p1, amount=USDCent(50), ext_ref_id=ach_id2) + brokerage_product_payout_event_factory( + product=p1, amount=USDCent(50), ext_ref_id=ach_id2 + ) gr_business.prebuild_payouts( bpem=business_payout_event_manager, @@ -562,9 +572,9 @@ class TestBusinessPayoutEventManager: 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, + brokerage_product_payout_event_factory: Callable[ + ..., BrokerageProductPayoutEvent + ], ledger_manager: LedgerManager, product_manager: ProductManager, ): @@ -593,7 +603,7 @@ class TestBusinessPayoutEventManager: wall_req_cpi=Decimal("5.00"), started=start + timedelta(days=6), ) - bp_payout_factory( + brokerage_product_payout_event_factory( product=u1.product, amount=USDCent(475), # 95% of $5.00 created=start + timedelta(days=1, minutes=1), @@ -602,7 +612,7 @@ class TestBusinessPayoutEventManager: ledger_collection.initial_load(client=None, sync=True) pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) gr_business.prebuild_balance( - thl_pg_config=thl_web_rr, + product_manager=product_manager, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -747,7 +757,9 @@ class TestBusinessPayoutEventManager: session_with_tx_factory: Callable[..., None], pop_ledger_merge: PopLedgerMerge, start: datetime, - bp_payout_factory: Callable[..., BrokerageProductPayoutEvent], + brokerage_product_payout_event_factory: Callable[ + ..., BrokerageProductPayoutEvent + ], adj_to_fail_with_tx_factory: Callable[..., None], thl_web_rr: PostgresConfig, ledger_manager: LedgerManager, @@ -784,7 +796,7 @@ class TestBusinessPayoutEventManager: wall_req_cpi=Decimal("5.00"), started=start + timedelta(days=1), ) - bp_payout_factory( + brokerage_product_payout_event_factory( product=u1.product, amount=USDCent(475), # 95% of $5.00 ext_ref_id=ach_id1, @@ -815,7 +827,7 @@ class TestBusinessPayoutEventManager: ledger_collection.initial_load(client=None, sync=True) pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) gr_business.prebuild_balance( - thl_pg_config=thl_web_rr, + product_manager=product_manager, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, diff --git a/tests/managers/thl/test_task_adjustment.py b/tests/managers/thl/test_task_adjustment.py index a14401e..323d6db 100644 --- a/tests/managers/thl/test_task_adjustment.py +++ b/tests/managers/thl/test_task_adjustment.py @@ -23,7 +23,7 @@ if TYPE_CHECKING: TaskAdjustmentManager, ) from generalresearch.managers.thl.wall import WallManager - from generalresearch.models.thl.session import Session + from generalresearch.models.thl.session import Session, Wall from generalresearch.models.thl.user import User @@ -47,10 +47,14 @@ def session_complete_with_wallet( @pytest.fixture() def session_fail( - user: User, session_manager: SessionManager, wall_manager: WallManager + user: User, + session_manager: SessionManager, + wall_manager: WallManager, + session_factory: Callable[..., Session], + wall_factory: Callable[..., Wall], ) -> Session: - session = session_manager.create_dummy(started=datetime.now(UTC), user=user) - wall1 = wall_manager.create_dummy( + session = session_factory(started=datetime.now(UTC), user=user) + wall1 = wall_factory( session_id=session.id, user_id=user.user_id, source=Source.DYNATA, diff --git a/tests/managers/thl/test_user_streak.py b/tests/managers/thl/test_user_streak.py index 564a142..59dee2d 100644 --- a/tests/managers/thl/test_user_streak.py +++ b/tests/managers/thl/test_user_streak.py @@ -1,6 +1,7 @@ from __future__ import annotations import copy +from collections.abc import Callable from datetime import UTC, date, datetime, timedelta from decimal import Decimal from typing import TYPE_CHECKING @@ -24,6 +25,7 @@ if TYPE_CHECKING: from generalresearch.managers.thl.user_streak import ( UserStreakManager, ) + from generalresearch.models.thl.session import Session, Wall from generalresearch.models.thl.user import User @@ -106,8 +108,14 @@ def broken_active_streak(user: User) -> list[UserStreak]: ] -def create_session_fail(session_manager: SessionManager, start: datetime, user: User): - session = session_manager.create_dummy(started=start, country_iso="us", user=user) +def create_session_fail( + session_manager: SessionManager, + start: datetime, + user: User, + session_factory: Callable[..., Session], + wall_factory: Callable[..., Wall], +): + session = session_factory(started=start, country_iso="us", user=user) session_manager.finish_with_status( session, finished=start + timedelta(minutes=1), @@ -117,9 +125,13 @@ def create_session_fail(session_manager: SessionManager, start: datetime, user: def create_session_complete( - session_manager: SessionManager, start: datetime, user: User + session_manager: SessionManager, + start: datetime, + user: User, + session_factory: Callable[..., Session], + wall_factory: Callable[..., Wall], ): - session = session_manager.create_dummy(started=start, country_iso="us", user=user) + session = session_factory(started=start, country_iso="us", user=user) session_manager.finish_with_status( session, finished=start + timedelta(minutes=1), @@ -141,13 +153,15 @@ def test_user_streaks_active_broken( user: User, session_manager: SessionManager, broken_active_streak: list[UserStreak], + session_factory: Callable[..., Session], + wall_factory: Callable[..., Wall], ): # Testing active streak, but broken (not today or yesterday) start1 = datetime(2025, 2, 12, tzinfo=UTC) end1 = start1 + timedelta(minutes=1) # abandon counts as inactive - session = session_manager.create_dummy(started=start1, country_iso="us", user=user) + session = session_factory(started=start1, country_iso="us", user=user) streak = user_streak_manager.get_user_streaks(user_id=user.user_id) assert streak == [] diff --git a/tests/managers/thl/test_userhealth.py b/tests/managers/thl/test_userhealth.py index ce6c221..a86361a 100644 --- a/tests/managers/thl/test_userhealth.py +++ b/tests/managers/thl/test_userhealth.py @@ -241,10 +241,9 @@ class TestIPRecordManager: ip_record_manager: IPRecordManager, user: User, ip_information: IPInformation, + ip_record_factory: Callable[..., IPRecord], ): - instance = ip_record_manager.create_dummy( - user_id=user.user_id, ip=ip_information.ip - ) + instance = ip_record_factory(user_id=user.user_id, ip=ip_information.ip) assert isinstance(instance, IPRecord) assert isinstance(instance.forwarded_ips, list) diff --git a/tests/managers/thl/test_wall_manager.py b/tests/managers/thl/test_wall_manager.py index 58de7a2..70db71e 100644 --- a/tests/managers/thl/test_wall_manager.py +++ b/tests/managers/thl/test_wall_manager.py @@ -19,7 +19,7 @@ from generalresearch.models.thl.definitions import ( 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.session import Session, Wall from generalresearch.models.thl.user import User @@ -250,13 +250,15 @@ class TestWallCacheManager: wall_manager: WallManager, session_manager: SessionManager, user: User, + session_factory: Callable[..., Session], + wall_factory: Callable[..., Wall], ): start1 = datetime.now(UTC) - timedelta(hours=3) start2 = datetime.now(UTC) - timedelta(hours=2) start3 = datetime.now(UTC) - timedelta(hours=1) - session = session_manager.create_dummy(started=start1, user=user) - wall_manager.create_dummy( + session = session_factory(started=start1, user=user) + wall_factory( session_id=session.id, user_id=session.user_id, started=start1, @@ -272,7 +274,7 @@ class TestWallCacheManager: attempts = wall_cache_manager.get_attempts(user_id=user.user_id) assert len(attempts) == 1 - wall_manager.create_dummy( + wall_factory( session_id=session.id, user_id=session.user_id, started=start2, @@ -298,8 +300,8 @@ class TestWallCacheManager: attempts10000 = [attempts[0]] * 6000 wall_cache_manager.update_attempts_redis_(attempts10000, user_id=user.user_id) - session = session_manager.create_dummy(started=start3, user=user) - wall_manager.create_dummy( + session = session_factory(started=start3, user=user) + wall_factory( session_id=session.id, user_id=session.user_id, started=start3, diff --git a/tests/models/gr/test_business.py b/tests/models/gr/test_business.py index e942be5..030a214 100644 --- a/tests/models/gr/test_business.py +++ b/tests/models/gr/test_business.py @@ -134,7 +134,9 @@ class TestBusiness: thl_ledger_manager: ThlLedgerManager, product_manager: ProductManager, business_payout_event_manager: BusinessPayoutEventManager, - bp_payout_factory: Callable[..., BusinessPayoutEventManager], + brokerage_product_payout_event_factory: Callable[ + ..., BusinessPayoutEventManager + ], start: datetime, user_factory: Callable[..., User], session_with_tx_factory: Callable[..., Session], @@ -179,7 +181,7 @@ class TestBusiness: wall_req_cpi=Decimal("2.50"), started=start + timedelta(days=5), ) - bp_payout_factory( + brokerage_product_payout_event_factory( product=p1, amount=USDCent(50), created=start + timedelta(days=4), @@ -329,7 +331,9 @@ class TestBusiness: self, gr_business: Business, product_factory: Callable[..., Product], - bp_payout_factory: Callable[..., BrokerageProductPayoutEvent], + brokerage_product_payout_event_factory: Callable[ + ..., BrokerageProductPayoutEvent + ], thl_ledger_manager: ThlLedgerManager, business_payout_event_manager: BusinessPayoutEventManager, create_main_accounts: Callable[..., None], @@ -341,7 +345,7 @@ class TestBusiness: thl_lm=thl_ledger_manager ) - bp_payout_factory( + brokerage_product_payout_event_factory( product=p, amount=USDCent(123), skip_wallet_balance_check=True ) @@ -352,7 +356,7 @@ class TestBusiness: assert sum([p.amount for p in gr_business.payouts]) == 123 # Add another! - bp_payout_factory( + brokerage_product_payout_event_factory( product=p, amount=USDCent(123), skip_wallet_balance_check=True, @@ -373,7 +377,9 @@ class TestBusiness: self, gr_business: Business, product_factory: Callable[..., Product], - bp_payout_factory: Callable[..., BrokerageProductPayoutEvent], + brokerage_product_payout_event_factory: Callable[ + ..., BrokerageProductPayoutEvent + ], thl_ledger_manager: ThlLedgerManager, thl_web_rr: PostgresConfig, business_payout_event_manager: BusinessPayoutEventManager, @@ -388,21 +394,21 @@ class TestBusiness: thl_lm=thl_ledger_manager ) - bp_payout_factory( + brokerage_product_payout_event_factory( product=p1, amount=USDCent(1), skip_wallet_balance_check=True, skip_one_per_day_check=True, ) - bp_payout_factory( + brokerage_product_payout_event_factory( product=p1, amount=USDCent(25), skip_wallet_balance_check=True, skip_one_per_day_check=True, ) - bp_payout_factory( + brokerage_product_payout_event_factory( product=p1, amount=USDCent(50), skip_wallet_balance_check=True, @@ -633,7 +639,9 @@ class TestBusinessBalance: user_factory: Callable[..., User], product_manager: ProductManager, mnt_filepath: GRLDatasets, - bp_payout_factory: Callable[..., BrokerageProductPayoutEvent], + brokerage_product_payout_event_factory: Callable[ + ..., BrokerageProductPayoutEvent + ], thl_ledger_manager: ThlLedgerManager, ledger_manager: LedgerManager, start: datetime, @@ -668,7 +676,7 @@ class TestBusinessBalance: payout_event_manager.set_account_lookup_table(thl_lm=thl_ledger_manager) - bp_payout_factory( + brokerage_product_payout_event_factory( product=u1.product, amount=USDCent(5), created=start + timedelta(days=4), @@ -676,7 +684,7 @@ class TestBusinessBalance: skip_one_per_day_check=True, ) - bp_payout_factory( + brokerage_product_payout_event_factory( product=u2.product, amount=USDCent(50), created=start + timedelta(days=4), @@ -707,7 +715,9 @@ class TestBusinessBalance: product_factory: Callable[..., Product], user_factory: Callable[..., User], mnt_filepath: GRLDatasets, - bp_payout_factory: Callable[..., BrokerageProductPayoutEvent], + brokerage_product_payout_event_factory: Callable[ + ..., BrokerageProductPayoutEvent + ], ledger_manager: LedgerManager, thl_ledger_manager: ThlLedgerManager, start: datetime, @@ -762,7 +772,7 @@ class TestBusinessBalance: ) payout_event_manager.set_account_lookup_table(thl_lm=thl_ledger_manager) - bp_payout_factory( + brokerage_product_payout_event_factory( product=u1.product, amount=USDCent(250), created=start + timedelta(days=3), @@ -770,7 +780,7 @@ class TestBusinessBalance: skip_one_per_day_check=True, ) - bp_payout_factory( + brokerage_product_payout_event_factory( product=u2.product, amount=USDCent(50), created=start + timedelta(days=4), @@ -846,7 +856,9 @@ class TestBusinessBalance: session_with_tx_factory: Callable[..., Session], pop_ledger_merge: PopLedgerMerge, start: datetime, - bp_payout_factory: Callable[..., BrokerageProductPayoutEvent], + brokerage_product_payout_event_factory: Callable[ + ..., BrokerageProductPayoutEvent + ], payout_event_manager, product_manager: ProductManager, adj_to_fail_with_tx_factory: Callable[..., None], @@ -876,7 +888,7 @@ class TestBusinessBalance: started=start + timedelta(days=1), ) payout_event_manager.set_account_lookup_table(thl_lm=thl_ledger_manager) - bp_payout_factory( + brokerage_product_payout_event_factory( product=u1.product, amount=USDCent(71), ext_ref_id=uuid4().hex, @@ -958,7 +970,9 @@ class TestBusinessBalance: product_factory: Callable[..., Product], user_factory: Callable[..., User], mnt_filepath: GRLDatasets, - bp_payout_factory: Callable[..., BrokerageProductPayoutEvent], + brokerage_product_payout_event_factory: Callable[ + ..., BrokerageProductPayoutEvent + ], thl_ledger_manager: ThlLedgerManager, ledger_manager: LedgerManager, product_manager: ProductManager, @@ -1029,7 +1043,7 @@ class TestBusinessBalance: ) payout_event_manager.set_account_lookup_table(thl_lm=thl_ledger_manager) - bp_payout_factory( + brokerage_product_payout_event_factory( product=u1.product, amount=USDCent(250), created=start + timedelta(days=3), @@ -1037,7 +1051,7 @@ class TestBusinessBalance: skip_one_per_day_check=True, ) - bp_payout_factory( + brokerage_product_payout_event_factory( product=u2.product, amount=USDCent(50), created=start + timedelta(days=4), diff --git a/tests/models/thl/test_product.py b/tests/models/thl/test_product.py index f1050bb..223430f 100644 --- a/tests/models/thl/test_product.py +++ b/tests/models/thl/test_product.py @@ -617,7 +617,9 @@ class TestProductFinancials: product_factory: Callable[..., Product], user_factory: Callable[..., User], mnt_filepath: GRLDatasets, - bp_payout_factory: Callable[..., BrokerageProductPayoutEvent], + brokerage_product_payout_event_factory: Callable[ + ..., BrokerageProductPayoutEvent + ], thl_ledger_manager: ThlLedgerManager, start: datetime, brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager, @@ -716,7 +718,7 @@ class TestProductFinancials: from generalresearch.currency import USDCent - bp_payout_factory( + brokerage_product_payout_event_factory( product=p1, amount=USDCent(50), created=start + timedelta(days=3), @@ -766,7 +768,7 @@ class TestProductFinancials: # -- Now pay ou another!. - bp_payout_factory( + brokerage_product_payout_event_factory( product=p1, amount=USDCent(5), created=start + timedelta(days=4), @@ -843,7 +845,9 @@ class TestProductBalance: session_with_tx_factory: Callable[..., Session], pop_ledger_merge: PopLedgerMerge, start: datetime, - bp_payout_factory: Callable[..., BrokerageProductPayoutEvent], + brokerage_product_payout_event_factory: Callable[ + ..., BrokerageProductPayoutEvent + ], payout_event_manager: PayoutEventManager, ): # Now let's load it up and actually test some things @@ -864,7 +868,7 @@ class TestProductBalance: # 2. Payout and build Parquets 2nd time payout_event_manager.set_account_lookup_table(thl_lm=thl_ledger_manager) - bp_payout_factory( + brokerage_product_payout_event_factory( product=product, amount=USDCent(71), ext_ref_id=uuid4().hex, @@ -895,7 +899,9 @@ class TestProductBalance: session_with_tx_factory: Callable[..., None], pop_ledger_merge: PopLedgerMerge, start: datetime, - bp_payout_factory: Callable[..., BrokerageProductPayoutEvent], + brokerage_product_payout_event_factory: Callable[ + ..., BrokerageProductPayoutEvent + ], payout_event_manager: PayoutEventManager, ): # This is very similar to the test_complete_payout_pq_inconsistent @@ -923,7 +929,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_ledger_manager) - bp_payout_factory( + brokerage_product_payout_event_factory( product=product, amount=USDCent(71), ext_ref_id=uuid4().hex, @@ -1114,7 +1120,9 @@ class TestProductCache: session_with_tx_factory: Callable[..., None], pop_ledger_merge: PopLedgerMerge, start: datetime, - bp_payout_factory: Callable[..., BrokerageProductPayoutEvent], + brokerage_product_payout_event_factory: Callable[ + ..., BrokerageProductPayoutEvent + ], payout_event_manager: PayoutEventManager, adj_to_fail_with_tx_factory: Callable[..., None], ): @@ -1136,7 +1144,7 @@ class TestProductCache: # 2. Payout payout_event_manager.set_account_lookup_table(thl_lm=thl_ledger_manager) - bp_payout_factory( + brokerage_product_payout_event_factory( product=product, amount=USDCent(71), ext_ref_id=uuid4().hex, -- cgit v1.2.3